1. 环境配置与依赖解析
高光谱图像分类任务对计算环境有较高要求,我们需要搭建一个完整的Python深度学习环境。以下是基于CUDA 12.1和PyTorch 2.3的详细配置方案:
1.1 基础环境搭建
推荐使用conda创建隔离的Python环境,避免与其他项目产生依赖冲突。具体步骤如下:
bash复制conda create -n dtsc python=3.9 -y
conda activate dtsc
注意:虽然作者提供了Python 3.9的GDAL安装包,但实际测试发现Python 3.8-3.10版本均可兼容。如果使用其他Python版本,需要自行编译GDAL或通过conda安装。
1.2 核心依赖安装
PyTorch的安装需要严格匹配CUDA版本。对于CUDA 12.1环境,使用以下命令安装:
bash复制pip install torch==2.3.0+cu121 torchvision==0.18.0+cu121 torchaudio==2.3.0+cu121 --extra-index-url https://download.pytorch.org/whl/cu121
其他关键依赖的版本选择依据:
- MMCV 2.2.0:专为PyTorch 2.x优化,提供高效的计算机视觉操作
- timm 1.0.7:包含预训练的视觉Transformer模型
- einops 0.8.0:优化张量操作的可读性
- opencv-python 4.8.1.78:提供图像处理基础功能
1.3 环境验证
安装完成后,建议运行以下验证脚本:
python复制import torch
print(torch.__version__, torch.cuda.is_available())
import mmcv
print(mmcv.__version__)
from einops import rearrange
print(rearrange(torch.randn(1,3,224,224), 'b c h w -> b h w c').shape)
预期输出应显示PyTorch版本、CUDA可用状态、MMCV版本以及einops操作后的张量形状。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 数据集准备与预处理
2.1 WHU-OHS数据集详解
WHU-OHS是中国武汉大学发布的高光谱遥感数据集,包含丰富的城市和自然场景。其特点包括:
- 光谱范围:400-1000nm
- 空间分辨率:0.5m
- 包含16个语义类别(建筑、道路、水体等)
2.2 数据目录结构规范
必须严格按照以下结构组织数据:
code复制data/
├── tr/ # 训练集
│ ├── image/ # 原始图像(.tif)
│ └── label/ # 标注图像(.tif)
├── val/ # 验证集
│ ├── image/
│ └── label/
└── ts/ # 测试集
├── image/
└── label/
实操技巧:使用符号链接可以避免数据复制带来的存储开销:
bash复制ln -s /path/to/raw/data/WHU-OHS ./data
2.3 数据预处理要点
-
光谱归一化:对每个波段进行Z-score标准化
python复制mean = torch.mean(hsi, dim=(1,2), keepdim=True) std = torch.std(hsi, dim=(1,2), keepdim=True) hsi = (hsi - mean) / (std + 1e-6) -
空间裁剪:将大图切割为256×256的patch
python复制from torch.nn.functional import unfold patches = unfold(img, kernel_size=256, stride=256) -
类别平衡:使用加权交叉熵损失处理类别不平衡
python复制weights = 1.0 / class_counts criterion = nn.CrossEntropyLoss(weight=weights)
3. 模型架构与实现细节
3.1 双阶段处理流程解析
DTSC模型的核心创新在于两阶段处理:
-
光谱超像素生成阶段:
- 使用SLIC算法生成超像素
- 每个超像素包含相似光谱特征的像素
- 输出光谱token序列
-
分类阶段:
- 骨干网络提取空间-光谱特征
- Transformer编码器建模全局依赖
- 输出像素级分类结果
3.2 骨干网络选型对比
DTSC支持三种骨干网络,各有特点:
| 骨干网络 | 参数量 | 计算量(GFLOPs) | 适用场景 |
|---|---|---|---|
| ResNet50 | 25.5M | 4.1 | 计算资源有限时首选 |
| PVTV2 | 32.4M | 5.8 | 平衡精度与效率 |
| Swin-T | 28.3M | 4.5 | 需要长程依赖建模 |
3.3 关键代码实现
光谱超像素生成的核心逻辑:
python复制def generate_spectral_tokens(hsi, n_segments=100):
# hsi: [C, H, W]
segments = slic(hsi.permute(1,2,0),
n_segments=n_segments,
compactness=10,
sigma=1)
tokens = []
for i in np.unique(segments):
mask = (segments == i)
token = hsi[:, mask].mean(dim=1) # [C]
tokens.append(token)
return torch.stack(tokens) # [N, C]
4. 训练流程优化策略
4.1 训练脚本参数解析
train.sh中的关键参数可通过命令行覆盖:
bash复制python train.py --config=models/yamls/PVTV2.yaml \
--exp_name=my_experiment \
--batch_size=16 \
--lr=0.001 \
--weight_decay=1e-4
4.2 学习率调度策略
采用余弦退火配合线性warmup:
python复制from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001)
warmup = LinearLR(optimizer, start_factor=0.01, total_iters=500)
cosine = CosineAnnealingLR(optimizer, T_max=10000)
scheduler = SequentialLR(optimizer, [warmup, cosine], [500])
4.3 混合精度训练
使用PyTorch自动混合精度(AMP)加速训练:
python复制scaler = torch.cuda.amp.GradScaler()
with torch.amp.autocast(device_type='cuda'):
outputs = model(inputs)
loss = criterion(outputs, labels)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
5. 测试与结果分析
5.1 评估指标解读
DTSC使用以下指标评估性能:
- OA (Overall Accuracy):整体分类准确率
- AA (Average Accuracy):各类别准确率均值
- Kappa:考虑随机因素的分类一致性系数
5.2 结果可视化技巧
使用matplotlib生成对比图:
python复制def visualize(pred, label):
fig, (ax1, ax2) = plt.subplots(1, 2)
ax1.imshow(pred, cmap='jet')
ax1.set_title('Prediction')
ax2.imshow(label, cmap='jet')
ax2.set_title('Ground Truth')
plt.savefig('comparison.png')
5.3 常见问题排查
-
CUDA内存不足:
- 减小
batch_size(默认16可降至8) - 使用
--gradient_accumulation_steps=2模拟更大batch
- 减小
-
验证指标波动大:
- 检查学习率是否过高
- 增加
--patience=20参数早停
-
预测结果全为同一类别:
- 检查类别权重计算是否正确
- 验证数据标注是否平衡
6. 进阶优化方向
6.1 自定义数据集适配
对于新数据集,需要:
- 修改
datasets/whu.py中的类别映射 - 调整
models/yamls/*.yaml中的num_classes - 重新计算类别权重
6.2 模型轻量化方案
-
知识蒸馏:
python复制teacher_model = load_pretrained('PVTV2') student_model = ResNet50() loss = KLDivLoss(teacher_logits, student_logits) -
量化感知训练:
python复制
model = quantize_model(model)
6.3 多模态融合扩展
可结合LiDAR数据提升性能:
- 在数据加载器中添加LiDAR分支
- 设计跨模态特征融合模块
- 调整损失函数平衡各模态贡献
在实际部署中发现,使用PVTV2骨干时,将学习率设置为0.0005、batch_size=12可以在16GB显存显卡上取得最佳平衡。对于Swin Transformer,建议启用--use_checkpoint参数节省显存。
