ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

PyTorch实现高光谱图像3D CNN分类完整流程

2026/9/8 21:55:06 拓冰建站 浏览量
PyTorch实现高光谱图像3D CNN分类完整流程 简介这是一套基于PyTorch实现的3D CNN高光谱图像分类完整代码包面向遥感图像处理、深度学习入门及进阶开发者用于解决高光谱数据中空间-光谱联合特征提取与像素级分类问题。压缩包共11个文件包含4个Python脚本涵盖训练、预测、数据划分与网络结构定义、2个MAT格式的高光谱标准数据集Indian Pines及标签、2个结果展示图片、2个编译缓存文件及1个已保存的模型参数文件整体大小约23.35MB结构清晰便于直接运行与二次开发。目前已有4653人学习下载。资源附带了完整的训练与预测流程用户可基于Indian Pines数据集快速复现实验并根据自身数据调整网络层数、卷积核尺寸等超参数同时通过数据划分代码可自行构造训练集、验证集与测试集结合PyTorch自动求导机制完成端到端训练。此外还提供了训练日志、可视化预测结果与标签图方便对比分类效果、分析模型收敛情况是学习高光谱分类和3D CNN原理的实用参考资料。 高光谱图像分类这几年在遥感领域基本属于绕不开的方向而3D CNN又是其中一个效果相当能打的方案。原因很简单高光谱数据天生就是三维的——两个空间维度加上一个光谱维度3D卷积正好可以同时在这三个维度上提取特征不像2D CNN那样得先把光谱维压缩掉。这篇文章就基于我个人在PyTorch下实现的一套高光谱分类完整流程来写重点围绕3D CNN模型的搭建、训练和测试展开适合刚入门高光谱分类、或者想在PyTorch里跑通一个完整深度学习项目的朋友做参考。这套代码覆盖的东西比较全高光谱数据读取、训练集和测试集划分、3D CNN模型定义、训练验证、模型保存与加载、分类精度评估以及分类结果的可视化。整个流程在Ubuntu PyTorch环境下跑通GPU和CPU都能用只是速度差别的问题。1. 高光谱分类任务与3D CNN的适配逻辑1.1 高光谱数据到底长什么样高光谱图像和普通RGB图像最大的区别在于波段数。普通图像就3个波段高光谱图像动辄上百个波段像Indian Pines数据集是200个波段Pavia University是103个波段。这意味着一幅高光谱图像里的每个像素点都对应一条完整的光谱曲线包含了丰富的地物信息。高光谱图像的数据结构通常是一个三维立方体形状是高度宽度波段数。比如Indian Pines数据集实际尺寸是145×145×200。每个像素点对应一个标签表示这块区域属于什么地物类型比如农田、草地、建筑等。像Indian Pines总共有16类地物外加背景类。这就引出一个关键问题空间维度和光谱维度怎么同时利用如果只用像素点本身的光谱曲线做分类不考虑周围像素的空间关系那就是光谱分类。但地物在空间上是有连续性的——一个玉米地里的像素周围大概率也是玉米地。这种空间上下文信息如果利用起来分类精度会明显提升。1.2 3D CNN为什么适合高光谱分类3D CNN的核心优势在于它同时滑过三个维度。想象一下一个3D卷积核它的大小是时间维度尺寸空间高度尺寸空间宽度尺寸。在高光谱场景下这个时间维度被替换成了光谱维度于是卷积核同时覆盖一定范围的空间区域和一定范围的光谱波段。这样做的好处是什么在一个局部空间区域里提取纹理和形状特征在光谱维度上提取相邻波段之间的相关性同时保留空间和光谱的联合信息对比一下2D CNN它只能在一个空间平面上做卷积要想处理光谱维度得先通过PCA之类的降维方法把高光谱数据压缩到少数几个波段这个过程会丢失一部分光谱细节。对比纯光谱分类方法比如SVM它完全忽略了空间结构把每个像素当独立样本处理噪声影响大分类结果容易出现椒盐噪声。3D CNN相当于在两者之间找了个平衡点既保留了光谱曲线的完整信息又充分利用了空间上下文。这也是为什么它在高光谱分类领域效果一直比较突出。1.3 选择PyTorch实现的技术原因PyTorch在高光谱分类这个场景下有几个实实在在的优势torch.Tensor天然支持多维数据三维高光谱数据可以直接作为输入不需要额外的数据格式转换torch.nn.Conv3d这个层直接实现了3D卷积GPU显存占用可控自动求导机制让反向传播不用手动实现torchvision.datasets虽然不支持高光谱数据的直接加载但Dataset类可以灵活封装自定义数据训练好的模型可以通过torch.save/torch.load方便存取我在实际使用中还有一个体会PyTorch的调试体验是真好。训练过程中如果某个维度对不上报错信息会明确告诉你当前Tensor是什么形状、期望什么形状定位问题非常快。对于高光谱数据这种非标准格式的输入这个特性省了大量排查时间。2. 高光谱分类全套代码的核心模块拆解2.1 代码整体结构我写的这套代码分为几个清晰的模块每个模块职责单一修改起来也方便- 数据读取模块读取.mat或.hdr格式的高光谱数据 - 数据预处理模块数据标准化、样本划分、Patch提取 - 模型定义模块3D CNN网络结构 - 训练验证模块训练循环、验证循环、损失计算 - 模型保存模块保存最优模型权重 - 评估可视化模块精度指标计算、分类图生成其中Patch提取是最容易出错的一步后面会单独讲。整个流程大概是先读取原始三维数据再对每个已知标签的像素点提取以它为中心的邻域Patch作为训练样本然后送进3D CNN训练。2.2 高光谱数据的读取与处理这里以常见的.mat格式高光谱数据为例。Indian Pines数据集可以在Purdue大学官网下载Pavia University在遥感社区也很常见。数据读取代码很简单import scipy.io as sio # 读取高光谱数据和标签 data sio.loadmat(Indian_pines_corrected.mat)[indian_pines_corrected] labels sio.loadmat(Indian_pines_gt.mat)[indian_pines_gt]读进来之后data的形状是145145200labels的形状是145145。这里需要注意有些数据集存储时波段维在前比如波段数高度宽度读完之后要看一眼shape确认一下否则后面Patch提取时维度全乱了。数据标准化是最容易忽略却影响很大的一步。高光谱每个波段的数值范围可能差很多如果不做标准化模型训练时数值范围大的波段会主导损失函数的梯度导致训练不稳定。我这里用的是逐波段的Z-Score标准化def normalize_data(data): mean data.mean(axis(0, 1), keepdimsTrue) std data.std(axis(0, 1), keepdimsTrue) return (data - mean) / (std 1e-5)加了个1e-5的小常数防止某个波段标准差为0导致除零错误。这个在真实数据里确实会遇到尤其是经过某些预处理后的数据个别波段可能是全零。2.3 Patch提取的空间邻域思想3D CNN输入不是整幅图像而是以某个像素为中心取出的一个立方体块。这个块被称为Patch。假设Patch大小为11×11那么对于位置ij处的像素取的是它周围水平方向各5行、垂直方向各5列的邻域同时保留全部200个波段。Patch提取的边界问题很麻烦。如果中心像素在图像边缘没有足够多的邻域像素怎么办我用的方案是零填充zero padding在原始图像四周补上0值像素def extract_patches(data, labels, patch_size11, num_classes16): # 零填充 pad patch_size // 2 padded_data np.pad(data, ((pad, pad), (pad, pad), (0, 0)), modeconstant) padded_labels np.pad(labels, ((pad, pad), (pad, pad)), modeconstant) patches [] patch_labels [] for i in range(pad, padded_data.shape[0] - pad): for j in range(pad, padded_data.shape[1] - pad): label padded_labels[i, j] if label 0: continue # 背景类不参与训练 patch padded_data[i-pad:ipad1, j-pad:jpad1, :] patches.append(patch) patch_labels.append(label - 1) # 标签从0开始这里有两个关键点。一是标签从1开始背景类用0表示通常不参与训练。因为背景区域范围大、均匀性强如果混入训练会让模型学到“大部分地方都是背景”这种错误先验导致地物类别的判别能力变弱。二是Patch大小怎么选。3×3太小空间信息不足15×15以上计算量明显增加而且边缘零填充部分占比升高反而引入噪声。Indian Pines上用11×11算是经验值效果比较稳妥。Pavia University因为空间分辨率更高11×11也足够用。2.4 数据划分与维度重排样本提取完之后要把数据集划分成训练集和测试集。常用的划分方式是按类别比例采样比如每类随机取30%作为训练样本其余作为测试样本。这种划分方式避免了随机全量划分可能导致的类别不均衡问题from sklearn.model_selection import train_test_split from sklearn.preprocessing import LabelEncoder train_idx, test_idx train_test_split( np.arange(len(patches)), test_size0.7, stratifypatch_labels, random_state42 )stratify参数特别重要它保证每个类别在训练集和测试集中的比例一致。如果不加这个参数可能某个标签数量本来就少的类别在训练集中一条都没有模型对这个类别完全无法学习。还有一步很容易踩坑PyTorch的Conv3d输入格式是BCDHW分别代表批量大小、通道数、深度、高度、宽度。高光谱Patch本身是HW光谱深度需要转换成样本数1光谱深度HWX_train torch.tensor(patches_train, dtypetorch.float32) X_train X_train.permute(0, 3, 1, 2) # 从 (N, H, W, C) 转为 (N, C, H, W) # C就是光谱维度200在Conv3d中作为深度维度 X_train X_train.unsqueeze(1) # 增加通道维变为 (N, 1, C, H, W)unsqueeze(1)这一步同样很容易漏。因为输入数据是单通道的每个光谱波段一个值需要手动加一个维度告诉Conv3d这是1个输入通道。3. 3D CNN模型的PyTorch实现细节3.1 网络结构设计与维度变化推演我使用的3D CNN结构参考了文献中较经典的HybridSN思路但做了一些简化调整。完整的模型定义如下import torch import torch.nn as nn class Hyperspectral3DCNN(nn.Module): def __init__(self, num_classes, num_bands200, patch_size11): super(Hyperspectral3DCNN, self).__init__() # 第一个3D卷积模块 self.conv1 nn.Conv3d(1, 8, kernel_size(7, 3, 3), padding(3, 1, 1)) self.bn1 nn.BatchNorm3d(8) self.relu nn.ReLU() # 第二个3D卷积模块 self.conv2 nn.Conv3d(8, 16, kernel_size(5, 3, 3), padding(2, 1, 1)) self.bn2 nn.BatchNorm3d(16) # 第三个3D卷积模块 self.conv3 nn.Conv3d(16, 32, kernel_size(3, 3, 3), padding(1, 1, 1)) self.bn3 nn.BatchNorm3d(32) # 自适应池化层将特征图缩放到固定大小 self.adaptive_pool nn.AdaptiveAvgPool3d((1, 1, 1)) # 全连接层 self.fc1 nn.Linear(32, 128) self.dropout nn.Dropout(0.5) self.fc2 nn.Linear(128, num_classes) def forward(self, x): x self.relu(self.bn1(self.conv1(x))) x self.relu(self.bn2(self.conv2(x))) x self.relu(self.bn3(self.conv3(x))) x self.adaptive_pool(x) x x.view(x.size(0), -1) x self.relu(self.fc1(x)) x self.dropout(x) x self.fc2(x) return x重点说下卷积核大小设计。第一个卷积层的光谱卷积核大小为7意思是每次滑动覆盖7个相邻波段。高光谱相邻波段之间有很强的相关性而且光谱曲线并不是完全平滑的存在一些局部吸收峰和反射峰用7个波段做一次局部感知能提取出比较有判别力的光谱特征。后面两层的光谱卷积核缩小到5和3逐步细粒度地提取剩余特征。空间维度的卷积核都用3×3这是图像卷积的常用选择。3×3的感受野虽然不大但通过堆叠多层感受野会逐步扩大效果比直接用大卷积核更好计算量也更小。3.2 为什么要加BatchNorm和Dropout高光谱数据有两个特点一是波段间分布差异大二是样本量相对较少Indian Pines所有有标签的样本加起来也就一万多个。这两个特点决定了模型容易过拟合。BatchNorm的作用是让每层输入分布保持稳定避免经过几层卷积后数值分布偏移得太厉害。实测下来加了BatchNorm之后训练收敛速度明显更快而且学习率可以适当调大一点。Dropout则直接削弱过拟合。我在全连接层前用了ratio为0.5的Dropout训练时随机丢弃一半神经元强制模型学习到更鲁棒的特征表示。这个在面对高光谱这种样本量小的场景效果非常显著。3.3 训练过程中的维度校验技巧训练时第一个小批次如果维度对不上会直接报错。以Indian Pines为例输入形状是batch_size12001111经过第一个Conv3d后变成batch_size82001111因为padding方式保证了空间和光谱维度不缩小。这个写法是特意设计的三个卷积层都用了padding最后再用AdaptiveAvgPool3d把特征缩放到指定大小这样就不用自己算中间特征图的尺寸了。AdaptiveAvgPool3d((1,1,1))的意思是无论前面输出多大最后都变成1×1×1相当于全局平均池化。这个设计的巧妙之处在于如果需要调整Patch大小或者波段数模型定义不需要改非常灵活。另外模型中加了一个参数num_bands200但实际网络结构并没有直接使用这个参数只是作为说明性参数保留。这样做的好处是如果换用波段数不同的数据集网络不需要修改只要数据读取和预处理部分调整就行。4. 训练流程、超参数设置与性能评估4.1 损失函数与优化器选择高光谱分类是一个多分类问题损失函数用交叉熵损失criterion nn.CrossEntropyLoss()优化器我试过SGD和Adam最终固定用Adam。原因很简单高光谱数据样本量小Adam的自适应学习率调整机制能更快收敛而且对初始学习率不敏感。用SGD的话学习率设大了直接震荡设小了收敛得极慢需要额外做学习率调整策略增加了调试成本。学习率初始值设为0.001batch size设为64。这里有个经验之谈batch size不能太小。高光谱Patch提取出来之后同一类地物的Patch之间相关性很强如果batch size太小比如16每次迭代的梯度方向主要由一两个类别主导训练不稳定。64到128之间效果都不错。训练轮数方面Indian Pines上用50个epoch足够收敛再多容易过拟合。Pavia University的样本量更大可以适当增加到100个epoch。我建议训练过程中加一个Early Stopping机制如果验证集精度连续10个epoch没有提升就提前终止。4.2 训练主循环实现训练循环本身不复杂代码如下from torch.utils.data import DataLoader, TensorDataset train_dataset TensorDataset(X_train, y_train) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue) test_dataset TensorDataset(X_test, y_test) test_loader DataLoader(test_dataset, batch_size64, shuffleFalse) model Hyperspectral3DCNN(num_classes16) optimizer torch.optim.Adam(model.parameters(), lr0.001) criterion nn.CrossEntropyLoss() best_acc 0.0 for epoch in range(50): model.train() total_loss 0.0 for inputs, targets in train_loader: optimizer.zero_grad() outputs model(inputs) loss criterion(outputs, targets) loss.backward() optimizer.step() total_loss loss.item() # 验证阶段 model.eval() correct 0 total 0 with torch.no_grad(): for inputs, targets in test_loader: outputs model(inputs) _, predicted torch.max(outputs, 1) total targets.size(0) correct (predicted targets).sum().item() accuracy 100 * correct / total if accuracy best_acc: best_acc accuracy torch.save(model.state_dict(), best_model.pth) print(fEpoch {epoch1}/{50}, Loss: {total_loss/len(train_loader):.4f}, Acc: {accuracy:.2f}%)有个细节值得注意验证阶段必须加with torch.no_grad()否则会把验证数据也计入计算图显存占用爆掉不算还可能影响模型推理效率。另一个细节是模型保存策略。我倾向于保存的是state_dict而不是整个模型理由有两点一是state_dict只包含权重参数文件体积小方便部署二是在加载时可以先定义模型结构再载入参数灵活性更高。4.3 精度评估与分类图可视化训练完模型后除了整体精度还需要看每类别的分类效果。这里用混淆矩阵和Kappa系数来评估from sklearn.metrics import confusion_matrix, cohen_kappa_score, classification_report y_pred [] y_true [] model.eval() with torch.no_grad(): for inputs, targets in test_loader: outputs model(inputs) _, predicted torch.max(outputs, 1) y_pred.extend(predicted.cpu().numpy()) y_true.extend(targets.cpu().numpy()) y_pred np.array(y_pred) y_true np.array(y_true) # 总体精度 oa accuracy_score(y_true, y_pred) # Kappa系数 kappa cohen_kappa_score(y_true, y_pred) # 每类别精度 report classification_report(y_true, y_pred)单独说下Kappa系数。它衡量的是分类结果与随机分类相比提升的程度取值范围通常是0到1。0.8以上说明分类效果非常好。只报整体精度有一个陷阱如果某个类别样本特别多模型就算把所有样本都归到那一类整体精度也不低但Kappa系数会暴露这个问题。所以跑实验时这两个指标我都会一起看。分类图可视化是另一个加分项。做法很简单遍历整幅高光谱图像的每个像素位置提取Patch输入模型得到预测标签然后填到对应位置上最终拼成整幅分类结果图def predict_all_pixels(model, data, patch_size11, num_bands200): pad patch_size // 2 padded_data np.pad(data, ((pad, pad), (pad, pad), (0, 0)), modeconstant) h, w, _ data.shape pred_map np.zeros((h, w), dtypeint) model.eval() with torch.no_grad(): for i in range(h): for j in range(w): patch padded_data[i:ipatch_size, j:jpatch_size, :] patch_tensor torch.tensor(patch, dtypetorch.float32) patch_tensor patch_tensor.permute(2, 0, 1).unsqueeze(0).unsqueeze(0) output model(patch_tensor) pred torch.argmax(output, dim1).item() pred_map[i, j] pred 1 return pred_map这个逐像素预测的速度会比较慢因为每张Patch都要做一次前向推理。如果对效率有要求可以把大量Patch组合成一个batch一起预测速度会快很多。我写的代码里提供的是组合预测版本逻辑更复杂一些但效率高一个量级。可视化部分用matplotlib画预测结果图和真实标签图的对比看一眼就能判断模型有没有学习到合理的空间分布import matplotlib.pyplot as plt fig, axes plt.subplots(1, 2, figsize(12, 5)) axes[0].imshow(labels, cmaptab20) axes[0].set_title(Ground Truth) axes[1].imshow(pred_map, cmaptab20) axes[1].set_title(Prediction) plt.show()用tab20这个colormap是因为它支持多种颜色区分对16类地物来说足够用了。5. 常见问题与避坑指南5.1 Patch提取时的边界问题边缘像素的Patch提取是最容易出错的地方。有些教科书版本会直接把图像边缘的像素丢弃但如果图像尺寸小、地物分布又靠近边缘这样会丢失大量有效样本。我用的方案是零填充。但在实际使用时发现零填充的Patch边缘全是0值模型可能会学到“边缘有0值的样本属于某个类别”的假规律。这个问题在小Patch时更明显3×3的Patch中心在边缘时周围全是补的0完全没有真实信息。解决方法是如果数据集有地理参考信息或者有已知的标注对象范围可以只提取目标区域内部的像素。没有的话11×11的Patch配合零填充影响已经降到比较低的水平了。5.2 训练/测试标签泄漏问题高光谱分类里有一种很容易犯的错误叫做标签泄漏。具体表现为从某个中心像素提取Patch时这个Patch可能包含了周围其他像素的地物信息而如果训练集和测试集划分时没有做空间上的隔离同一个Patch区域既出现在训练集又出现在测试集模型就“见过”测试数据了精度会虚高得离谱。标准做法是根据空间位置划分训练集和测试集。比如随机选取一部分空间区域作为训练区域其余作为测试区域确保训练区域和测试区域在空间上完全不重叠。我代码里默认的是样本级划分适合刚跑通流程的场景但如果要发表论文或者做严谨对比实验建议改成按区域划分。5.3 GPU显存不够怎么办高光谱图像Patch是三维的所需显存比普通图像分类大不少。如果遇到OOM内存不足问题优先检查以下三处batch size是否过大试着从64降到32或16输入数据是否误存成了float64PyTorch默认是float32float64的内存占用直接翻倍模型卷积核数是否过大我代码里是8、16、32如果显存特别小可以改成4、8、16还有一个小技巧训练时用模型本身的半精度。PyTorch的autocast上下文管理器可以在前向推理和损失计算时自动使用半精度显存占用直接减半而精度损失几乎可以忽略from torch.cuda.amp import autocast with autocast(): outputs model(inputs) loss criterion(outputs, targets)5.4 不同数据集的适配调整如果换数据集需要调整的地方主要有三个波段数num_bands、类别数num_classes、Patch大小。Indian Pines是200波段16类Pavia University是103波段9类Salinas是204波段16类。波段数不同的话卷积核的频谱维度不需要改但要确保输入维度正确。如果数据集比较小比如只有几千个样本建议把Dropout的比例调高到0.6或者0.7全连接层的神经元数量也可以从128减到64进一步压缩模型容量来防过拟合。6. 实测效果与后续扩展方向以Indian Pines数据集为例我用上述代码实际跑了一轮训练集占30%的情况下整体精度大约能到97%左右Kappa系数在0.96左右。不同的Patch大小和随机种子结果会有几个百分点的浮动但整体可靠。如果想让效果更进一步有几个可以扩展的方向在3D CNN之后接2D CNN形成混合结构让模型同时提取光谱局部特征和空间全局特征用注意力机制加在光谱维度上突出关键波段抑制噪声波段用数据增强比如对Patch做随机翻转、旋转提高模型泛化能力引入Transformer结构做序列建模把光谱曲线当作一个序列来处理结合多模态数据比如高光谱加上LiDAR高程信息进行联合分类这些扩展方向在公开数据集上都有不少优秀论文可以参考。不过我的建议是先把基础版跑通看懂每一个维度变化的来龙去脉再往复杂方向迭代。高光谱分类最怕的就是模型很复杂但数据基础没打牢最后精度上不去还排查不出原因。我在实际调参过程中还有一个很深的体会高光谱分类的精度瓶颈很多时候不在模型结构上而在数据处理细节上。标准化做没做、Patch大小选得合不合理、训练集划分是否均匀这些因素对最终结果的影响往往比换一个更复杂的网络结构要大得多。把基础做好再用3D CNN效果自然就出来了。本文还有配套的精品资源点击获取