恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
Python实战:非局部卷积神经网络SAR图像去噪从零搭建
首页
资讯中心
/
Python实战:非局部卷积神经网络SAR图像去噪从零搭建
Python实战:非局部卷积神经网络SAR图像去噪从零搭建
发布时间:2026/10/11 22:58:33
简介这份源码面向希望深入理解SAR图像去噪与深度学习结合的开发者与研究者提供了一套基于Python的非局部卷积神经网络NL-CNN完整实现。项目围绕SAR图像噪声抑制这一典型难题将非局部注意力机制与卷积网络融合适合具备一定深度学习基础、想动手复现或改进去噪算法的读者学习实践。压缩包共33个文件约18.83MB以19个Python源码文件为核心涵盖网络结构、数据集加载、训练与评估脚本另含YAML环境配置、T7与Pickle权重文件、Markdown说明文档及少量TXT、Shell脚本便于快速搭建实验环境。目录中划分了模型、数据集、工具函数与实验入口等模块结构清晰。目前已有486人学习下载。读者可从中获得完整的网络定义、训练流程、权重文件与评估指标实现用于复现实验、对比NLM与CNN等基线方法并在此基础上开展SAR去噪的二次开发与调参实践。1. 从一张斑点图说起非局部卷积为什么成了 SAR 去噪的刚需SAR 图像天生带相干斑成像机理决定了噪声不是加性的而是乘性的跟光学图像里那种高斯白噪声完全两码事。你拿一张 Sentinel-1 的强度图放大看均匀区域里全是颗粒状的明暗跳变边缘和纹理又被这层颗粒糊住。传统做法是 Lee 滤波、Frost 滤波、BM3D 这一路靠局部窗口统计或者块匹配来压噪问题是窗口一大边缘就糊窗口一小噪声压不干净。非局部卷积神经网络Non-Local CNN换了个思路它不只看像素邻域而是在整张特征图上算任意两个位置的相似度把远处相似的纹理块拉过来一起参与重建。对 SAR 这种充满重复纹理农田、海面、城区网格的图像来说这个机制天然对路。这篇要讲的就是怎么用 Python 把这套东西从零搭出来包括数据准备、非局部模块怎么写、训练怎么调、推理怎么落地源码结构也一并拆开讲。适合已经会 PyTorch 基础、想把手里的 SAR 去噪从传统滤波升级到深度学习的人。2. 非局部模块到底在算什么从自相似到注意力权重2.1 非局部操作的数学直觉非局部操作的核心公式其实就一行输出某个位置的值等于全图所有位置值的加权和权重由两个位置的相似度决定。用数学写出来是y_i (1/C(x)) * Σ_j f(x_i, x_j) * g(x_j)其中f是相似度函数g是输入变换C是归一化因子。放到卷积网络里f通常用嵌入空间的内积或者高斯函数实现g就是一次线性映射。这跟自注意力机制本质上是同一个东西只不过在去噪任务里我们关心的是相似块之间的信息聚合而不是序列建模里的长程依赖。为什么这个机制对 SAR 特别有效因为相干斑虽然看起来随机但它在同质区域内的统计特性是一致的。一个农田块里的斑点分布和另一块农田的斑点分布在特征空间里是接近的。非局部模块能把这些同质区域的统计信息聚合起来相当于在整张图上做了一次自适应的、内容感知的均值滤波但比固定窗口的均值滤波聪明得多——它不会把边缘另一侧的异质区域混进来。2.2 在 CNN 里嵌入非局部块的具体写法下面是一个可以直接用的非局部块实现我一般把它插在编码器和解码器之间的瓶颈层或者每个尺度跳跃连接之前。import torch import torch.nn as nn class NonLocalBlock(nn.Module): def __init__(self, in_channels, reduction2): super().__init__() # 把通道数压缩减少计算量这是实操里的常规做法 self.inter_channels in_channels // reduction self.theta nn.Conv2d(in_channels, self.inter_channels, 1) self.phi nn.Conv2d(in_channels, self.inter_channels, 1) self.g nn.Conv2d(in_channels, self.inter_channels, 1) self.out nn.Conv2d(self.inter_channels, in_channels, 1) # 归一化层稳定训练 self.norm nn.GroupNorm(8, in_channels) def forward(self, x): B, C, H, W x.shape N H * W # theta 和 phi 做内积得到相似度矩阵 theta self.theta(x).view(B, self.inter_channels, N).permute(0, 2, 1) # B,N,C phi self.phi(x).view(B, self.inter_channels, N) # B,C,N attn torch.softmax(torch.bmm(theta, phi), dim-1) # B,N,N # g 分支做线性变换后聚合 g self.g(x).view(B, self.inter_channels, N) # B,C,N y torch.bmm(g, attn.permute(0, 2, 1)) # B,C,N y y.view(B, self.inter_channels, H, W) y self.out(y) return self.norm(y) x # 残差连接保证原始信息不丢逻辑说明theta和phi是两个嵌入函数把每个位置的通道特征映射到一个低维空间然后做矩阵乘法得到N×N的相似度矩阵。softmax按行归一化保证每个位置的权重和为 1。g分支负责提供被聚合的值。最后用1×1卷积把通道数恢复加残差连接。参数说明reduction控制中间通道的压缩比默认 2 是我在 256×256 的 SAR 图上试出来的平衡点。压得太狠比如 8相似度矩阵的判别力会下降压得太少比如 1显存直接爆。GroupNorm的组数设 8 是经验值换成BatchNorm在小 batch 下反而不稳。softmax的维度是-1也就是对每一行做归一化这个别搞错搞错了权重就不满足概率分布。2.3 整体网络结构怎么搭单靠一个非局部块不够实际去噪网络我一般用 U-Net 骨架在三个地方插非局部块瓶颈层、最底层跳跃连接、以及倒数第二层跳跃连接。浅层特征感受野小插非局部块收益不大还费显存深层特征语义强相似度计算更有意义。class NLUNet(nn.Module): def __init__(self, base64): super().__init__() # 编码器 self.enc1 self._block(1, base) self.enc2 self._block(base, base*2) self.enc3 self._block(base*2, base*4) # 瓶颈层加非局部 self.bottleneck nn.Sequential( self._block(base*4, base*8), NonLocalBlock(base*8), self._block(base*8, base*8) ) # 解码器 self.dec3 self._block(base*8 base*4, base*4) self.dec2 self._block(base*4 base*2, base*2) self.dec1 self._block(base*2 base, base) self.nl3 NonLocalBlock(base*4) # 跳跃连接前 self.nl2 NonLocalBlock(base*2) self.final nn.Conv2d(base, 1, 1) self.pool nn.MaxPool2d(2) def _block(self, cin, cout): return nn.Sequential( nn.Conv2d(cin, cout, 3, padding1), nn.GroupNorm(8, cout), nn.LeakyReLU(0.1, inplaceTrue), nn.Conv2d(cout, cout, 3, padding1), nn.GroupNorm(8, cout), nn.LeakyReLU(0.1, inplaceTrue) ) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) b self.bottleneck(self.pool(e3)) d3 self.dec3(torch.cat([self.nl3(e3), nn.functional.interpolate(b, scale_factor2)], 1)) d2 self.dec2(torch.cat([self.nl2(e2), nn.functional.interpolate(d3, scale_factor2)], 1)) d1 self.dec1(torch.cat([e1, nn.functional.interpolate(d2, scale_factor2)], 1)) return self.final(d1)这个结构里非局部块放在跳跃连接之前是为了让编码器传过来的特征先经过一次全局相似度聚合再和解码器特征拼接。实测比放在拼接之后效果好因为拼接后的特征通道数翻倍非局部矩阵计算量直接乘 4不划算。3. 数据准备与训练从 SAR 原始格式到可训练张量3.1 SAR 数据的读取与归一化SAR 数据常见格式有 GeoTIFF、ENVI、以及一些厂商私有格式。Python 里用rasterio读 GeoTIFF 最稳读进来是浮点强度值动态范围可能从 0 到几万。直接送进网络会梯度爆炸必须先做对数变换再归一化。import rasterio import numpy as np def load_sar(path): with rasterio.open(path) as src: arr src.read(1).astype(np.float32) # 对数变换压缩动态范围加 1 防止 log(0) arr np.log1p(arr) # 按百分位裁剪去掉极端亮斑的影响 lo, hi np.percentile(arr, [1, 99]) arr np.clip(arr, lo, hi) # 归一化到 [0,1] arr (arr - lo) / (hi - lo 1e-8) return arr逻辑说明log1p是 SAR 强度图的标准预处理把乘性噪声转化成近似加性。百分位裁剪是为了排除强散射体比如建筑物角反射对归一化的干扰用 1% 和 99% 是我在多个数据集上试出来的稳健值。归一化到[0,1]之后再按均值和标准差做一次标准化这一步在训练脚本里做。参数说明如果你的 SAR 数据已经是 dB 格式跳过log1p直接做百分位裁剪和归一化。percentile的[1, 99]可以改成[0.5, 99.5]取决于图像里强散射体的比例。裁剪之后一定要检查直方图确认没有大量像素被压到 0 或 1。3.2 训练对的构造干净图从哪来SAR 去噪有个尴尬的地方你拿不到真正的干净图。常见做法有三种。第一种是用光学图像加模拟斑点噪声优点是干净图已知缺点是域不匹配。第二种是用多时相平均图当伪干净图优点是真实缺点是平均之后仍有残余噪声。第三种是自监督比如 Noise2Noise 那套用两张独立噪声图互相监督。我一般用第一种做预训练第二种做微调。模拟斑点噪声用 Gamma 分布因为单视 SAR 强度图的斑点服从 Gamma 分布。def add_speckle(clean, looks1): # 单视 Gamma 分布形状参数为 looks noise np.random.gamma(looks, 1.0/looks, clean.shape) noisy clean * noise return np.clip(noisy, 0, 1)逻辑说明np.random.gamma的形状参数设looks尺度参数设1/looks这样均值为 1方差为1/looks。乘到干净图上就得到带斑点的图像。looks1是最难的情况多视数据把looks调大就行。参数说明looks对应 SAR 的视数单视设 1四视设 4。实际训练时我会把looks在 1 到 4 之间随机采样让模型适应不同噪声强度。clip到[0,1]是防止乘性噪声把值推到归一化范围之外。3.3 损失函数的选择与训练参数去噪任务最常用的损失是 L1 和 L2。L2 对异常值敏感SAR 里的强散射体容易把训练带偏L1 更鲁棒但收敛慢。我的做法是 L1 为主加一个 SSIM 项做结构约束。class HybridLoss(nn.Module): def __init__(self, alpha0.8): super().__init__() self.alpha alpha self.l1 nn.L1Loss() def forward(self, pred, target): l1 self.l1(pred, target) # 简化版 SSIM用局部均值和方差算 mu_p pred.mean(dim[2,3], keepdimTrue) mu_t target.mean(dim[2,3], keepdimTrue) var_p pred.var(dim[2,3], keepdimTrue) var_t target.var(dim[2,3], keepdimTrue) cov ((pred - mu_p) * (target - mu_t)).mean(dim[2,3], keepdimTrue) c1, c2 0.01**2, 0.03**2 ssim ((2*mu_p*mu_t c1)*(2*cov c2)) / ((mu_p**2 mu_t**2 c1)*(var_p var_t c2)) return self.alpha * l1 (1 - self.alpha) * (1 - ssim.mean())逻辑说明alpha0.8表示 L1 占主导SSIM 做辅助。SSIM 的常数c1和c2用标准值。这个简化版 SSIM 没有做高斯加权但在去噪任务里够用而且计算快。参数说明alpha在 0.7 到 0.9 之间调低于 0.7 结构约束太强容易过平滑高于 0.9 又回到纯 L1边缘保持不够。优化器用 Adam初始学习率1e-4每 20 个 epoch 衰减 0.5。batch size 设 8patch 大小 128×128显存占用约 6GB。4. 推理与后处理把模型输出变成可用的 SAR 图4.1 滑窗推理与重叠融合SAR 图像往往很大比如 10000×10000直接送进网络显存扛不住。滑窗推理是标准做法但窗口边界会有拼接痕迹。解决办法是重叠采样然后对重叠区域做加权平均。def sliding_inference(model, img, patch256, stride128): model.eval() H, W img.shape output np.zeros((H, W), dtypenp.float32) weight np.zeros((H, W), dtypenp.float32) # 汉宁窗做加权边界权重低中心权重高 win np.hanning(patch) win2d np.outer(win, win) for i in range(0, H - patch 1, stride): for j in range(0, W - patch 1, stride): patch_img img[i:ipatch, j:jpatch] tensor torch.from_numpy(patch_img).unsqueeze(0).unsqueeze(0).cuda() with torch.no_grad(): pred model(tensor).squeeze().cpu().numpy() output[i:ipatch, j:jpatch] pred * win2d weight[i:ipatch, j:jpatch] win2d return output / (weight 1e-8)逻辑说明stride设patch/2保证 50% 重叠。汉宁窗在中心为 1边界趋近 0这样拼接时边界贡献小不会出现明显接缝。最后除以权重和做归一化。参数说明patch设 256 是显存和速度的平衡点128 更快但边界更多512 更慢但接缝更少。stride一般设patch//2再小就浪费计算。如果你的图特别大可以分块读入但要注意块之间的重叠区域也要做融合。4.2 后处理对比度恢复与伪影抑制网络输出是归一化后的值要恢复到原始动态范围。如果前面做了对数变换这里要做指数逆变换。另外非局部模块偶尔会在极亮或极暗区域产生轻微伪影可以用一个导向滤波做后处理。import cv2 def postprocess(pred, lo, hi): # 反归一化 pred pred * (hi - lo) lo # 指数逆变换 pred np.expm1(pred) # 导向滤波抑制伪影半径 4eps 1e-3 pred cv2.ximgproc.guidedFilter(pred.astype(np.float32), pred.astype(np.float32), radius4, eps1e-3) return np.clip(pred, 0, None)逻辑说明lo和hi是预处理时保存的百分位值。expm1是log1p的逆运算。导向滤波用自身做引导相当于边缘保持的平滑能把非局部模块产生的孤立亮点压下去。参数说明radius4对应约 9×9 的窗口太大边缘会糊。eps1e-3控制平滑强度SAR 去噪一般设1e-3到1e-2。如果输出没有明显伪影这一步可以跳过省时间。5. 避坑与排查非局部 SAR 去噪里最容易翻车的五个地方5.1 相似度矩阵显存爆炸现象训练到一半报CUDA out of memory而且是在非局部块那一层。原因N×N的相似度矩阵NH×W。如果特征图是 64×64N4096矩阵就是 4096×4096float32 下约 64MB看着不大。但 batch 里每个样本都有一份而且反向传播要存中间激活实际占用是前向的 3 到 4 倍。特征图再大一点比如 128×128N16384矩阵直接 1GB 起步。解决控制非局部块输入的特征图尺寸别在浅层插。如果必须在浅层用先做一次下采样再算相似度算完上采样回去。或者用reduction把通道压到 1/4 甚至 1/8减少theta和phi的维度。5.2 斑点噪声被当成纹理保留现象去噪后均匀区域还是花的噪声没压干净但边缘确实保住了。原因训练时模拟的斑点噪声分布和真实 SAR 不匹配。比如你用looks4训练实际数据是单视的噪声强度差一倍模型没见过那么强的噪声就把它当成了纹理。解决训练时把looks随机化范围覆盖实际数据的视数。另外损失函数里 L1 的权重可以调高逼模型更激进地压噪。如果还是不行在输入前先做一个轻度的 Lee 滤波把噪声方差降下来再送网络。5.3 非局部权重全图均匀现象训练 loss 降不下去推理结果和普通卷积网络没区别。原因theta和phi初始化太小内积之后所有位置的相似度都差不多softmax输出接近均匀分布非局部退化成全局平均池化。解决检查theta和phi的初始化用kaiming_normal而不是默认的。另外可以在相似度计算前加一个温度系数让分布更尖锐。attn torch.softmax(torch.bmm(theta, phi) / temperature, dim-1)temperature设 0.1 到 0.5越小分布越尖锐。这个参数可以学习也可以固定。5.4 训练集和验证集来自同一区域现象验证集指标很好换一张新图就崩。原因SAR 图像的空间相关性很强同一区域的训练块和验证块在纹理上高度相似模型相当于记住了这片区域的统计特性没有泛化能力。解决按地理区域划分训练集和验证集别按块随机分。如果数据来自多个传感器留一个传感器的数据做验证。另外数据增强里加随机旋转和翻转但别加颜色抖动SAR 是单通道的。5.5 推理速度慢到无法落地现象一张 5000×5000 的图推理要十几分钟。原因滑窗推理的窗口重叠率高而且非局部块的计算量随特征图尺寸平方增长。解决用torch.cuda.amp做半精度推理速度能快一倍精度损失很小。另外把stride从patch//2改成patch//4重叠区域少了接缝用后处理补。如果还慢把非局部块只在瓶颈层保留其他层去掉速度能再快一倍。6. 进阶技巧让非局部模块真正跑出优势的三个调参习惯第一个习惯是监控相似度矩阵的熵。在验证阶段把attn拿出来算一下每行的熵如果熵接近log(N)说明权重太均匀非局部没学到东西。正常情况熵应该在log(N)的 0.3 到 0.6 倍之间。这个指标比 loss 更早暴露问题。第二个习惯是分阶段训练。前 10 个 epoch 把非局部块的输出乘以 0.1让网络先靠卷积层学基础特征再逐步放开非局部的影响。直接端到端训容易让非局部块在早期就退化成均匀权重后面很难纠正。# 在 forward 里加一个可调的缩放因子 scale min(1.0, epoch / 10.0) return self.norm(y) * scale x第三个习惯是保存预处理参数。lo、hi、looks这些值在推理时必须和训练时一致否则归一化对不上输出会偏。我一般把它们存成 JSON和模型权重放一起。import json meta {lo: float(lo), hi: float(hi), looks: 1} with open(preprocess.json, w) as f: json.dump(meta, f)最后一个技巧是关于验证的别只看 PSNR 和 SSIM。SAR 去噪的最终评价要看等效视数ENL和边缘保持指数EPI。ENL 在均匀区域算越大说明噪声压得越干净EPI 在边缘区域算越接近 1 说明边缘保得越好。这两个指标比 PSNR 更贴近实际需求。我一般会在验证脚本里同时算这四个指标ENL 低于 100 或者 EPI 低于 0.7 就说明模型有问题得回去查。这套东西我从最早的 Lee 滤波一路踩坑踩过来最大的教训就是别迷信单一指标。PSNR 高的模型不一定好看ENL 高的模型可能糊成一片。多指标交叉验证再加上人眼看一下均匀区域和边缘区域基本就不会翻车。希望帮到你。本文还有配套的精品资源点击获取