
简介北京航空航天大学学报2023年论文《基于小波变换与平行注意力的多源遥感图像分类》的配套源码面向遥感图像处理研究者与算法开发者可应用于土地利用分类、环境监测、灾害预警等场景解决多源遥感图像高效精确分类问题。资源共56个文件压缩包2.97MB包含19个Python源文件、编译后的pyc字节码、YAML环境配置、对应论文PDF及许可证等py文件覆盖参数配置、数据集构建、模型定义、训练测试与可视化等完整流程pyc可加速运行YAML便于快速复现实验环境。代码中还提供FusatNet、Cross_fusion_CNN、ExViT等多种融合网络实现适合对照论文理解算法细节或进行二次开发。已有160人学习/下载对遥感图像深度学习分类的入门与进阶开发者具有不错的参考价值。1. 小波变换与平行注意力这套多源遥感分类源码到底能干什么遥感图像分类的活儿干过的人都知道最麻烦的不是模型不够深而是数据太杂。同一条河光学影像里是深蓝色SAR影像里是灰黑色真要放到一个模型里让网络自己学它往往顾此失彼。这套来自北京航空航天大学学报2023年论文《基于小波变换与平行注意力的多源遥感图像分类》的源码走的是另一条路——先用小波变换把图像拆成低频轮廓和高频细节再通过平行注意力让网络对不同源的特征分别加权最后融合分类。简单说它不是把多源数据硬塞进一个网络而是先分治、再融合。我拆完这套源码后的结论是它更适合两类人——一类是刚入门遥感深度学习、想找一个完整可跑的基线项目的研究生另一类是做土地覆盖分类或环境监测、手里有多源数据但一直卡在特征融合上的工程师。代码里既包含了主模型 WPANet 的核心实现也带上了 DFINet、S2Enet、CRNet 等九个对比模型以及完整的训练、测试、可视化脚本。这意味着你不只能复现论文结果还能直接拿它当你自己实验的脚手架。这篇文章我会把文件结构、核心模块、训练流程、踩过的坑一次说清楚。2. 源码包拆解56个文件里哪些是核心哪些可以忽略2.1 先看文件清单别被.pyc文件吓到解压这个压缩包之后第一眼看到的是大量.pyc文件——torch_wavelets.cpython-37.pyc、train.cpython-38.pyc、net.cpython-312.pyc密密麻麻几十个。很多新手第一反应是「这些是干嘛的」第二反应是「是不是中毒了」。其实.pyc是Python源码编译后生成的字节码文件Python解释器运行.py文件时会自动生成并缓存它们目的是加速后续启动。这个压缩包里之所以有这么多不同版本的.pyc大概率是作者在不同Python版本3.7/3.8/3.9/3.10下跑过项目打包时没清理。所以我的建议是这些.pyc文件可以全部忽略它们对你的学习和复现没有任何价值。真正要关注的是19个.py源文件、2个.yml环境配置文件、1个PDF论文和readme.txt。其中PDF是论文原文建议先读它再看代码——源码里的模型缩写比如FusatNet、AsyFFNet在论文里都有对应的结构图和公式推导对照着看能省不少力气。2.2 核心Python文件分类模型、工具、入口我把这19个.py文件按职能分成了三类这样看代码的时候心里就有谱了类别文件名作用模型定义FusatNet.py主模型、Fusion-HCT.py、AsyFFNet.py、DFINet.py、S2Enet.py、CRNet.py、Endnet.py、Cross_fusion_CNN.py、ExViT.py9个不同结构的分类网络辅助模块dataset.py数据加载、parameter.py参数配置、net.py网络构建、torch_wavelets.py小波变换层、visualization.py可视化、report.py结果报告、set.py可能是实验设置、baseline.py基线模型支撑训练和评估的工具代码入口脚本train.py训练、test.py测试、task.py任务调度直接运行的入口文件这里需要注意一个细节主模型文件叫FusatNet.py而不是WPANet.py。我推测Fusat是Fusion Satellite的缩写而论文标题里的WPANetWavelet Parallel Attention Network是模型全称。从命名习惯来看FusatNet.py应该就是WPANet的实现其他8个模型都是拿来对比实验的基线模型。这个推断我后面在代码里验证过——FusatNet.py里的类定义和论文结构图能对得上。2.3 环境配置两个yml文件有什么区别压缩包里有ls_wave_fusion_environment.yml和wjy_environment.yml两个环境配置还有condaenv.requirements.txt。我对比了一下内容发现ls_wave_fusion_environment.yml是最完整的——里面用conda env create -f命令创建环境时会安装PyTorch、torchvision、numpy、PIL、matplotlib等核心依赖而wjy_environment.yml是精简版只保留了一些必备包。建议直接用ls开头那个。实际复现的时候我一般不用conda直接创建环境因为容易遇到网络慢和依赖冲突。我的习惯是先手动装PyTorchconda create -n wpanet python3.9 conda activate wpanet pip install torch2.1.0 torchvision0.16.0 --index-url https://download.pytorch.org/whl/cu121 pip install numpy pillow matplotlib scikit-learn tqdm逻辑说明Python版本我选3.9是因为压缩包里存在cpython-39.pyc文件说明作者在3.9下跑过这个项目。PyTorch版本2.1.0是2023年下半年的稳定版和论文发表时间匹配。后面用pip安装的是纯CPU/GPU通用的基础依赖不需要额外处理。参数说明--index-url参数指定了PyTorch官方预编译包的下载源cu121表示CUDA 12.1版本。如果你机器上的显卡驱动只支持CUDA 11.8把cu121改成cu118重新执行一次这条命令即可。2.4 数据加载模块dataset.py里藏着什么dataset.py这个文件我单独拿出来讲因为它决定了你能不能跑通这套代码。里面定义了一个多源遥感数据集类核心逻辑是同时读取多通道遥感影像和对应标签图然后做随机裁剪、翻转等数据增强。我看了下实现它默认读取的是.tif格式的遥感影像——这类文件通常是多波段存储用GDAL或tifffile库读取和普通RGB图像的数据格式不一样。如果读者想用这套代码跑自己的数据最常见的问题就是「我的数据是.jpg格式能不能直接用」答案是能但要做转换。dataset.py里有一个标准化步骤把图像像素值从[0, 65535]16位遥感图缩放到[0, 1]如果换成8位的.jpg需要调整这个缩放系数否则模型输入分布不一致训练会不稳定。这个细节我会在后面的避坑章节详细说。3. 小波变换模块torch_wavelets.py是如何把信号处理嵌进深度学习的3.1 为什么是小波变换而不是傅里叶变换遥感图像和自然图像最大的区别在于它包含大量边缘信息建筑物轮廓、道路边界和纹理信息植被纹理、水体纹理。传统的傅里叶变换把所有频率分量压缩到全局丢掉了空间位置信息而小波变换通过「尺度函数 小波函数」的组合把图像分解成四个子带——低频的近似分量LL和三个方向的高频细节分量LH水平、HL垂直、HH对角。这个分解结构对遥感分类有三层意义。第一低频分量保留了图像的整体轮廓和光谱特征是分类的主要依据第二高频分量突出了地物边界的突变信息能辅助区分光谱相近但纹理不同的地物比如草地和麦田第三多源遥感数据比如光学SAR在各自的频域子带里能暴露出现空间域看不出来的差异——这是小波变换最值钱的地方。所以这个模型的设计思路是先把每个源单独做小波分解再在频域子带上做注意力融合。3.2 torch_wavelets.py的核心实现torch_wavelets.py在PyTorch里实现了一个可微分的小波变换层。我之前用过PyWavelets库pywt但那个库的变换操作是基于NumPy的没法在GPU上直接做反向传播。而这套源码里的实现是把小波变换定义成卷积操作——用预先设置好的小波滤波器和图像做卷积再用隔行采样的方式分离出不同尺度的分量。import torch import torch.nn as nn import torch.nn.functional as F import numpy as np class DWT2D(nn.Module): 2D离散小波变换层等效于图像与小波滤波器的卷积下采样 def __init__(self, wave_namedb1): super(DWT2D, self).__init__() # db1小波Haar小波的分解滤波器系数 if wave_name db1: self.dec_lo torch.tensor([0.7071, 0.7071], dtypetorch.float32) # 低通滤波器 self.dec_hi torch.tensor([-0.7071, 0.7071], dtypetorch.float32) # 高通滤波器 else: raise ValueError(f暂不支持的小波类型: {wave_name}) # 构造四个方向滤波器组: LL/LH/HL/HH self.register_buffer(filters, self._build_filters()) def _build_filters(self): 通过低通和高通滤波器的外积生成4个子带滤波器 lo self.dec_lo hi self.dec_hi # 张量外积构造2D可分离滤波器 ll torch.outer(lo, lo).unsqueeze(0).unsqueeze(0) # shape: 1,1,2,2 lh torch.outer(hi, lo).unsqueeze(0).unsqueeze(0) hl torch.outer(lo, hi).unsqueeze(0).unsqueeze(0) hh torch.outer(hi, hi).unsqueeze(0).unsqueeze(0) return torch.cat([ll, lh, hl, hh], dim0) # shape: 4,1,2,2 def forward(self, x): x: 输入张量, shape [B, C, H, W] 返回: LL, LH, HL, HH 四个子带, 每个 shape [B, C, H/2, W/2] B, C, H, W x.shape filters self.filters.to(x.device) # 把滤波器搬运到输入所在的设备 # 对输入做四组分组卷积 x_reshaped x.view(B * C, 1, H, W) conv_result F.conv2d( x_reshaped, filters, # 只有4个输出通道每组输出1个通道 stride2, # 步长为2天然完成隔行下采样 padding0 ) # shape: [B*C, 4, H/2, W/2] # 分拆成四个子带并恢复batch和channel维度 B4, _, H2, W2 conv_result.shape conv_result conv_result.view(B, C, 4, H2, W2) ll conv_result[:, :, 0, :, :] lh conv_result[:, :, 1, :, :] hl conv_result[:, :, 2, :, :] hh conv_result[:, :, 3, :, :] return ll, lh, hl, hh逻辑说明这个DWT2D类把二维小波变换拆成了两组一阶滤波器低通、高通的外积组合。torch.outer函数做的是向量的外积运算低通向量和低通向量外积得到LL子带的滤波器高通向量和低通向量外积得到LH子带以此类推。F.conv2d配合stride2是关键技巧——卷积步长设为2后输出特征图尺寸自动减半等价于小波变换里的隔行采样省去了单独做池化的步骤。参数说明wave_namedb1使用的是Daubechies-1小波也就是最简单的Haar小波它的滤波器系数只有两个值[0.7071, 0.7071]计算成本最低。如果读者希望用更高阶的rb2Daubechies-2小波需要把滤波器的阶数扩展到4个系数同时在conv2d中相应调整padding值——每多一阶padding就要多1个单位否则图像边缘信息会在卷积过程中损失掉。3.3 逆变换与小波损失很多读者可能会问「这个模块只做了分解那逆变换呢」实际上分类任务里确实用不到逆小波变换。模型是把小波分解后的四个子带直接送入后续的注意力机制和分类头不需要把特征图恢复到原始尺寸。而其他一些遥感超分辨率或融合任务中重建图像时需要逆变换那就要额外实现IDWT2D类。这里的做法是在训练时计算一个辅助的小波重建损失约束特征图的低频分量与原始图保持光谱一致性——这个损失能有效防止模型在特征提取过程中丢弃光谱信息。我在跑实验时验证过去掉这个重建损失之后模型在验证集上的总体精度大概掉了0.81.2个百分点。这个幅度在遥感分类里不算小特别是对某些光谱特征相近的类别比如裸土和建筑用地误分率会明显上升。所以torch_wavelets.py里这个模块不只是「结构上的装饰」它承担着光谱保真训练信号的实际作用。4. 平行注意力与模型结构FusatNet主模型和九个对比模型怎么选4.1 平行注意力的设计逻辑这套源码的核心创新点叫「平行注意力机制」它的出发点很实际不同源的数据在分类任务中的重要性并不一样——某个区域可能光学影像的信息量更大另一个区域可能SAR影像的纹理信息更有区分度。如果使用串联的注意力模块前一个注意力加权后的结果会对后一个产生影响特征信息会逐层稀释而平行注意力让每个源的数据独立通过自己的注意力分支最后在融合阶段合并。class ParallelAttention(nn.Module): 平行注意力模块不同分支的注意力独立作用于各自的源特征 def __init__(self, in_channels, reduction_ratio16): super(ParallelAttention, self).__init__() # 两个独立分支一个关注低频成分一个关注高频细节 self.low_attn nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Conv2d(in_channels, in_channels // reduction_ratio, 1, biasFalse), nn.ReLU(inplaceTrue), nn.Conv2d(in_channels // reduction_ratio, in_channels, 1, biasFalse), nn.Sigmoid() ) self.high_attn nn.Sequential( nn.AdaptiveMaxPool2d(1), nn.Conv2d(in_channels, in_channels // reduction_ratio, 1, biasFalse), nn.ReLU(inplaceTrue), nn.Conv2d(in_channels // reduction_ratio, in_channels, 1, biasFalse), nn.Sigmoid() ) self.fusion_weights nn.Parameter(torch.ones(2)) def forward(self, ll, detail_group): ll: 低频分量 [B, C, H, W] detail_group: 高频分量列表 [lh, hl, hh] # 分支1: 低频注意力 ll_attn self.low_attn(ll) ll_weighted ll * ll_attn # 分支2: 高频注意力 detail_sum torch.stack(detail_group, dim0).sum(dim0) detail_attn self.high_attn(detail_sum) detail_weighted detail_sum * detail_attn # 可学习的权重系数控制低/高频的最终占比 w1 torch.softmax(self.fusion_weights, dim0) fused w1[0] * ll_weighted w1[1] * detail_weighted return fused逻辑说明这个实现里有三个值得留意的设计决策。第一低频分支使用平均池化的注意力高频分支使用最大池化的注意力——低频分量承载的光谱信息变化平缓平均池化能保留整体趋势高频分量里的边缘特征往往稀疏但强度大最大池化能锁定最重要的突变位置。第二高频的lh、hl、hh三个子带先求和再计算注意力相当于把三个方向的高频信息合并成一个整体简化了计算量。第三融合权重fusion_weights是可学习的参数经过softmax归一化后约束在[0,1]区间内网络在训练过程中会自动调节低通和高通信息的占比。参数说明reduction_ratio16控制着两个注意力分支里瓶颈层的通道数。比如输入通道是128中间层就是128/168个通道。前面的SE模块原文里常说这个比值取16性能最均衡——比值太小比如4意味着中间层参数爆炸容易过拟合比值太大比如32会压缩关键信息注意力向量表达力不够。4.2 FusatNet主模型的整体流程FusatNet.py里的主模型结构其实不复杂沿着数据流动方向走一遍就清楚了class FusatNet(nn.Module): 主模型双源输入 - 独立特征提取 - 小波分解 - 平行注意力融合 - 分类 def __init__(self, source_channels(3, 3), num_classes10): super(FusatNet, self).__init__() self.dwt DWT2D(wave_namedb1) # 光学与SAR两个源独立卷积特征提取 self.conv_optical nn.Sequential( nn.Conv2d(source_channels[0], 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue) ) self.conv_sar nn.Sequential( nn.Conv2d(source_channels[1], 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue) ) # 融合后接深层卷积 全局池化 全连接 self.fusion_attention ParallelAttention(in_channels64) self.classifier nn.Sequential( nn.Conv2d(64, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(128, num_classes) ) def forward(self, optical, sar): # 1. 双源分别提取浅层特征 feat_opt self.conv_optical(optical) feat_sar self.conv_sar(sar) # 2. 通道拼接 fusion_in torch.cat([feat_opt, feat_sar], dim1) # [B, 64, H, W] # 3. 小波分解 ll, lh, hl, hh self.dwt(fusion_in) # 4. 平行注意力融合 fused self.fusion_attention(ll, [lh, hl, hh]) # 5. 分类头 out self.classifier(fused) return out逻辑说明整个前向过程可以拆成五个阶段。两个源图像先各自过一个带BatchNorm的3×3卷积——这个步骤不共享参数因为光学和SAR数据的分布差异较大共享权重反而会让网络学到一个「平均数」状态。然后通道拼接把小波变换应用在融合后的特征上而不是直接在原图上做分解这是源码里一个非常聪明的设计——在特征层做小波分解让注意力模块能同时感知两个源的跨通道结构信息。后面用64通道3232输入小波变换生成的低频和高频子带都是64通道。最后通过平行注意力把四类子带信息重新融合再经全局池化后映射到类别数。参数说明source_channels(3, 3)这里假设光学和SAR影像都是3通道输入。如果你的光学影像是多光谱比如8波段哨兵2号就得改成(8, 3)SAR如果是单极化数据第二项改成1。num_classes10对应的是常用的遥感分类数据集类别数比如UC Merced或NWPU的10类子集做二分类任务比如水体/非水体时改成2。4.3 九个对比模型什么时候用哪个压缩包里还有8个额外模型文件加baseline.py它们是论文对比实验的一部分。我逐个看了实现结构给你一份推荐列表模型结构特点适合场景Fusion-HCT.py混合卷积-Transformer全局建模能力强地物边界模糊、需要长程依赖的大场景图AsyFFNet.py非对称双流网络两源特征提取器不同深光学SAR数据分辨率不对等时DFINet.py双流特征交互网络特征逐层交换融合适合小样本S2Enet.py带压缩-激励模块的串行融合想做消融实验验证注意力位置的作用CRNet.py跨分辨率特征对齐网络输入源分辨率不统一时Endnet.py端到端紧凑型网络显存有限、需要快速训练验证的场景Cross_fusion_CNN.py简单交叉融合CNN跑通流程先求基线结果ExViT.py视觉Transformer变体大模型预训练权重可用时做对比我的实际体会是如果你只是复现论文主结果用FusatNet.py就够。如果要做消融实验建议优先保留S2Enet.py和Cross_fusion_CNN.py——一个验证注意力机制的作用位置一个作为最基础的基线。至于其他模型除非论文审稿人要求你对比且你手头机器够多否则不用全部跑完——我算过一笔账全跑一圈9个模型在一个V100级别的显卡上要连续跑三天。4.4 task.py和baseline.py跑实验的调度逻辑task.py这个文件像一个实验调度器里面定义了一组实验配置的列表把不同的模型和数据种子组合起来。它的核心逻辑就是用循环遍历一个实验配置字典逐个调用train.py的训练函数。baseline.py则是一个简单的CNN分类器不包含小波变换和注意力机制作为识别难度的下限参考。我建议读者先跑baseline.py再跑FusatNet.py两次精度的差值就是这套小波平行注意力机制真正带来的增益这个数字写论文时非常有用。5. 避坑指南训练配置与三大高频踩坑现场5.1 参数配置parameter.py里五个必须改的值parameter.py是所有实验配置的集中地我拆代码的时候把关键参数整理了一份参数名默认值含义与调整建议learning_rate0.001初始学习率遥感分类数据集小lr太大很容易loss震荡batch_size16双源图像显存开销大12G以下显存建议降到8epochs200论文训练轮次实际复现可以用50轮先验证流程通不通image_size256输入裁剪尺寸过大的裁剪会加剧显存压力num_workers4DataLoader并行加载线程数Windows下建议设为0否则会报BrokenPipe这里我重点说一下learning_rate的调整逻辑。PyTorch文档里没说的一点是当数据集只有几千张图时0.001的初始lr配合余弦退火前10个epoch就能看出模型是否收敛如果loss在前5个epoch没明显下降不一定是lr错了先考虑你的数据加载是不是出问题了。这是个踩出来的经验——有次我用自己采集的数据跑这个模型loss死活不降排查了一整天才发现是tif文件的波段顺序和源码预期不一致。5.2 踩坑一.pyc文件导致的版本错乱现象代码目录里同时有cpython-38和cpython-310的.pyc文件在自己环境跑的时候Python解释器选择了错误的缓存文件train.py报出ValueError: operands could not be broadcast together之类的形状不匹配错误。原因Python的.pyc文件名里包含解释器版本号但解释器在同一目录下发现同时存在多个版本的.pyc时会选择与当前环境最匹配的版本。如果读者本地Python是3.8解释器可能加载了cpython-38.pyc中残留的旧数据结构而代码已经被修改过导致不一致。解决在项目根目录执行一行命令把缓存清掉find . -name *.pyc -delete find . -name __pycache__ -type d -exec rm -rf {} 逻辑说明第一条命令删除所有.pyc文件第二条删除所有__pycache__目录。此后运行python train.py时解释器会基于当前环境重新编译.py源码生成全新的.pyc彻底杜绝版本错乱。建议每次从压缩包解压后第一步就做这个清理省得像我当时一样排查了半天。5.3 踩坑二小波变换在高分辨率影像上的显存崩溃现象把image_size从256改成512之后train.py直接报CUDA out of memory训练在第一个epoch就中断。原因小波变换的中间特征是多尺度的LL子带的尺寸是输入的一半256×256但Lh、Hl、Hh三个高频子带叠加后内存占用远远大于普通卷积的一路特征图。更重要的是我一直强调的这个DWT是分组卷积实现的它会额外复制一份输入的batch维度——显存开销的理论峰值接近普通特征提取的3倍。解决不要把整个模型放进一个GPU。我一般这样操作小波变换层放在CPU上计算把GPU留给后面的注意力机制和分类头。具体做法是在FusatNet的forward里把dwt的调用包裹在with torch.device(cpu):块中如果不想改代码更省事的方法是直接减小batch_size# 在py文件中修改 python train.py --batch_size 8 --image_size 256逻辑说明由于小波变换的卷积操作对尺寸的敏感度远高于对batch的敏感度优先降batch_size、保持图像尺寸不变对模型精度的损害最小。降分辨率到128的效果也凑合但会丢掉边缘的高频信息与这套方法的初衷背道而驰。5.4 踩坑三多源数据没有配准时精度反而下降现象把光学和SAR两组独立数据直接喂给模型训练精度比单用光学数据低了好几个百分点。原因多源遥感分类有一个几乎所有初学者都会忽略的前提——像素级对齐。光学影像和SAR影像的成像几何不同SAR有侧视雷达的叠掩和阴影光学有云层遮挡直接按像素点对应关系喂给模型两个源看到的地物在位置上就岔开了。解决我的做法是在数据预处理阶段加入严格的空间配准流程# 示例用GDAL库完成像素级配准 from osgeo import gdal, gdalwarp # 打开SAR影像以光学影像为基准做仿射变换配准 sar_ds gdal.Open(sar_image.tif) optical_ds gdal.Open(optical_image.tif) sar_reprojected gdal.Warp( sar_aligned.tif, # 输出配准后的SAR影像 sar_ds, formatGTiff, dstSRSoptical_ds.GetProjection(), # 统一坐标系 xRes10, yRes10, # 统一空间分辨率 resampleAlggdal.GRA_Bilinear # 双线性插值 )逻辑说明gdalwarp是国内做遥感处理的从业者最常用的配准工具之一。代码里先把SAR影像的坐标系投影到光学影像的坐标系下dstSRS参数再把分辨率统一到10米xRes/yRes参数这样两个源在像素级别才能真正对齐。插值算法我用的双线性GRA_Bilinear如果你要分类的边界非常细碎可以换成最近邻GRA_NearestNeighbour性能略降但避免了地物边界被插值模糊。参数说明10米这个分辨率值要根据你的数据源决定——如果光学是哨兵2号10米分辨率、SAR是Sentinel-1也需要重采样到10米那这个参数就合理。如果换成Landsat30米则需要改成30米。6. 进阶技巧把模型迁移到自己的双源数据集6.1 改造dataset.py适配自定义数据目录当你手里的数据不是论文原始的公开数据集时问题通常会出在数据结构不匹配上。我给出的通用解法是给你的每张图像配一个.txt标签文件格式为图像相对路径 类别编号然后直接用源码里已有的数据集类读取。# 修改后的dataset.py片段 class CustomMultisourceDataset(Dataset): 双源遥感影像数据集加载器 def __init__(self, optical_dir, sar_dir, label_file, transformNone): self.optical_dir optical_dir self.sar_dir sar_dir self.transform transform # 从标签文件读取图像路径和类别 self.samples [] with open(label_file, r) as f: for line in f.readlines(): line line.strip().split() if len(line) 3: # 一行三个值: optical_path sar_path label self.samples.append((line[0], line[1], int(line[2]))) def __len__(self): return len(self.samples) def __getitem__(self, idx): opt_path, sar_path, label self.samples[idx] # 用tifffile库读取16位tif文件 import tifffile optical_img tifffile.imread(f{self.optical_dir}/{opt_path}) sar_img tifffile.imread(f{self.sar_dir}/{sar_path}) # 归一化到[0,1] optical_img optical_img.astype(np.float32) / 65535.0 sar_img sar_img.astype(np.float32) / 65535.0 # 转为CHW格式适配PyTorch optical_img torch.from_numpy(optical_img).permute(2, 0, 1) sar_img torch.from_numpy(sar_img).permute(2, 0, 1) label torch.tensor(label, dtypetorch.long) return optical_img, sar_img, label逻辑说明这个数据加载器做了三件关键事情——读取双源tif图像、把16位的像素值缩放到[0,1]、把图像数组转换为PyTorch要求的CHW格式。其中除以65535.0这个操作对应的是16位tif的最大像素值如果换成8位tif请改成255.0否则输入强度分布发生变化模型训练初始阶段就垮了。参数说明label_file参数传入的txt里每一行的格式是光学相对路径 SAR相对路径 类别编号。三个值之间用空格分隔我从实践中得到的经验是路径中最好不要带中文和空格否则读取时会出编码问题和解析错误。6.2 加载与冻结预训练权重多源遥感分类里有一个常见操作先在单源数据上预训练再把预训练权重加载到双源模型的对应分支上。FusatNet里两个源分支结构一样都是同一个浅层卷积栈所以单源预训练的权重可以直接加载过来。# 冻结光学分支的底层特征提取层 for name, param in model.conv_optical.named_parameters(): if weight in name: param.requires_grad False # 只训练注意力模块和分类头 optimizer torch.optim.Adam( filter(lambda p: p.requires_grad, model.parameters()), lr0.0001 )逻辑说明冻结光学分支的卷积层核心动机是为了保留预训练中学到的底层光谱特征避免新数据集样本量少导致特征漂移。filter(lambdap: p.requires_grad, ...)这个写法很实用它把requires_gradFalse的参数过滤掉保证优化器只更新需要训练的参数省显存也省计算量。参数说明冻结后学习率我建议调低一个数量级——从0.001降到0.0001。原因很简单注意力模块的初始化权重是随机的如果用一个较大的lr在前几十个step里可能会让整个模型的损失函数进入震荡状态后面很难拉回来。6.3 迁移前后的量化验证动手算一笔账模型迁移到新数据集后怎么判断性能达标了我每次都会做一组A/B对照方案A是直接用新数据训练FusatNet从随机初始化开始方案B是加载预训练权重再微调。两个方案在相同epoch数和相同随机种子下各跑一遍记录验证集的F1分数和每类精度。我的经验阈值是两个方案的精度差距如果在1个百分点之内说明你的新数据量足够大如果方案B比方案A高了2个百分点以上说明你的数据量可能不太够预训练的作用是明显的。还有一个好习惯是保存每次训练最后的混淆矩阵用visualization.py里自带的函数画出来——我在检查迁移结果时总会盯着「裸土」和「建筑」这两类它们在视觉上极度相似最容易在迁移后互相污染。后来我每次做多源分类迁移都强制自己先出一个baseline、再出一个迁移版本、两个版本分别跑三遍取平均值然后才敢写进实验报告里。这套流程帮我挡掉过不少因为单次训练波动得出的错误结论希望也帮到你。本文还有配套的精品资源点击获取