资讯动态

Swin-Transformer+UNet图像去噪:原理、训练与避坑指南

发布时间:2026/10/2 3:06:28 来源:尧图企业网站定制
简介面向图像处理与深度学习应用的研究者与开发者项目提出结合Swin-Transformer与UNet的混合去噪网络旨在去除噪声的同时保持图像细节适用于算法对比、模型微调、毕业设计等场景。资源共25个文件其中17个Python源码覆盖模型定义、训练与推理流程、数据预处理及通用工具另有3个MATLAB脚本用于去噪效果评估2个Markdown文档提供使用说明1个YAML配置文件管理训练参数以及2个pyc缓存文件整体压缩包仅37KB轻量而清晰。当前已有581人学习/下载具有一定参考热度。借助完整源码可深入观察Swin-Transformer如何以层次方式提取特征并与UNet的跳跃连接融合进而通过调整配置参数快速复现或改进去噪实验节省从零搭建模型的时间适合作为图像重建与底层视觉方向的优质实战项目。1. Swin-TransformerUNet图像去噪为什么说这个组合是当前性价比最高的方案图像去噪是底层视觉里最“老”也是最有生命力的问题之一。传统方法如BM3D在中等噪声下依然是硬基准深度学习方法则把PSNR上限往前推了两到三个dB。最近一个绕不开的方向就是标题里的组合——Swin-TransformerUNet。这不是把两个模块缝在一起而是用UNet的编码器-解码器结构做多尺度上下文提取用Swin-Transformer的窗口自注意力在每一层做长程依赖建模两者互补后在BSD68、DND这类公开基准上往往能超过纯卷积UNet大约0.51dB。这篇文章适合两类人。一类是刚入门的同学手上有一个zip源码包但不知道从哪开始跑需要一条从数据到训练再到推理的完整路径另一类是有UNet或Transformer基础、想把自己模型改到去噪方向的工程师。我会把网络结构拆开讲清楚把训练和推理代码写给你看中间穿插参数怎么调、哪些地方最容易翻车最后单独列一节避坑记录。2. 拆解网络结构Swin-Transformer与UNet是怎么拼出图像去噪能力的2.1 窗口自注意力机制Swin-Transformer在图像去噪里到底干了什么Swin-Transformer是微软提出的视觉Transformer变体和ViT最大的区别是ViT把整张图切成固定patch然后做全局自注意力Swin改用窗口window划分把自注意力局限在每个窗口内部。对去噪任务来说这个限制反而是好事。图像去噪本质是像素级回归任务全局注意力带来的计算量是平方级增长的而窗口注意力把复杂度固定在窗口大小内允许你把网络做得更深。另一个关键设计是shifted window。Swin在每个Transformer块里交替使用W-MSA和SW-MSA前一半用规则窗口后一半把窗口偏移半个窗口大小这样信息能在窗口之间流动弥补了无全局注意力的缺陷。在去噪网络的实现中做一个两层的SwinTransformerBlock循环搭配patch merging就是基本单元。核心代码如下class SwinTransformerBlock(nn.Module): def __init__(self, dim, num_heads, window_size8, shift_size0): super().__init__() self.dim dim self.num_heads num_heads self.window_size window_size self.shift_size shift_size self.norm1 nn.LayerNorm(dim) self.attn WindowAttention(dim, num_heads, window_size) self.norm2 nn.LayerNorm(dim) self.mlp nn.Sequential( nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim) ) def forward(self, x): B, C, H, W x.shape x x.permute(0, 2, 3, 1) shortcut x x self.norm1(x) # 将图像切成window_size大小的窗口 x_windows window_partition(x, self.window_size) x_windows self.attn(x_windows) x window_reverse(x_windows, self.window_size, H, W) x shortcut x x x self.mlp(self.norm2(x)) return x.permute(0, 3, 1, 2)这里window_partition和window_reverse的细节因实现而异但逻辑是固定的把形状为(B,H,W,C)的张量重排成(BH/wsW/ws, ws, ws, C)经过注意力后还原。代码里的shift_size参数控制窗口偏移实际训练时每一个block交替设0和window_size//2。需要注意LayerNorm放在残差前面这是Pre-Norm结构也是Swin在深层网络里不炸训练的关键。2.2 三种拼接方式编码器替换、特征融合、瓶颈增强怎么选标题叫“Swin-TransformerUNet”实际工程里常见三种组合方式选错了会导致模型或慢或崩。第一种是编码器替换。把UNet的encoder每一层都换成SwinTransformerBlock堆叠decoder保持卷积上采样最典型的就是SwinUNet。这种方案在医学图像分割里表现好但在去噪里效果一般——去噪需要在每个尺度保留细节Swin的下采样patch merging太激进。第二种是特征融合最常见。保留UNet的卷积编码器在编码器输出的每个尺度上再接一个SwinTransformerBlock用窗口注意力处理特征图再传给decoder。这种做法的优点是改动小训练稳定的同时引入Transformer的建模能力。我一般推荐从这种方案上手。第三种是瓶颈增强。只在UNet最底层的低分辨率特征上比如32x32或16x16接Swin块因为这里特征图小窗口注意力几乎接近全局注意力计算量也可控。这种方式PSNR收益通常在0.20.3dB但训练速度快很多。动手改之前先要确认你拿到的源码包用的是哪种。判断方法很简单看forward里UNet的特征图有没有经过window_partition有就是第二种看encoder是卷积还是Transformer是Transformer就是第一种。如果不确定直接从第二种开始改造风险最小。2.3 关键参数的含义depth、num_heads、window_size与输入尺寸的匹配关系Swin部分需要调的参数就那几个但每个都影响训练结果。depth指SwinTransformerBlock堆叠的层数。在去噪任务上depth从2到6都测过2层偏浅感受野不足6层训练极慢。4层是通常的甜点。num_heads每层可以按通道数翻倍设计例如64通道时设4个头128通道时设8个头。window_size是最敏感的默认8在256x256的输入上很稳。如果把输入切成128x128的小patch训练window_size8依然成立但推理时要pad到8的倍数这点后面展开说。位置编码也要注意。Swin在分类任务里需要position bias但在去噪这种像素级任务里relative_position_bias_table甚至可以不初始化直接用零张量少了它反而少一层干扰。很多开源去噪实现里悄悄省掉position bias或直接冻结就是这个原因。融合阶段还有一个隐藏参数embed_dim。UNet第一层卷积通常把输入从3通道升到64通道这个64就是Transformer的嵌入维度。设成64还是128差距很大后者显存翻倍但PSNR只多0.1dB左右按你的显存决定。显存16G以下就64起步先跑通再考虑放大。3. 训练自己的去噪模型从BSD400到PSNR曲线收敛的完整落地路径3.1 数据集怎么选BSD400、BSD68、DND、RENOIR各自适合什么阶段去噪数据集的江湖地位很明确BSD400是训练集BSD68是经典测试集DND和RENOIR是真实噪声测试集。BSD400来自伯克利分割数据集取400张干净图像裁剪成128x128的小patch训练。BSD68是68张灰度图几乎所有论文都报它上面的PSNR对比。DND是真实传感器噪声不含干净参考图需要去他们的网站提交结果评测RENOIR也是真实噪声规模小一些。实操中我的建议是第一阶段只用BSD400训练patch128stride64切块到BSD68上测PSNR第二阶段追加BSD432或Flickr2K做数据增强再把模型拿到DND上验证泛化。DND因为是真实噪声模型在合成噪声上训练的PSNR在这里会明显掉这是正常现象。数据加载有一个大多数人会忽略的点归一化。图像读取后要除以255转成float噪声合成也要在浮点域做。如果用uint8做加减模型训练出来的结果会有严重的阶梯状伪影。这个坑我见的次数太多了。数据集用途典型规模说明BSD400训练400张应裁剪为patch使用BSD68验证/测试68张论文基准灰度图为主DND真实噪声测试50张需在线提交评测RENOIR真实噪声训练/测试约120张可做微调但数量有限3.2 训练脚本噪声注入、数据增强与训练循环的实现核心训练代码如下我用PyTorch写噪声选高斯噪声sigma25这是BSD68对比的标准配置。import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader import numpy as np import random class NoiseDataset(Dataset): 从干净图像中随机裁剪并注入高斯噪声 def __init__(self, img_list, patch_size128, sigma25): self.img_list img_list self.patch_size patch_size self.sigma sigma def __len__(self): return len(self.img_list) * 50 def __getitem__(self, idx): img self.img_list[idx % len(self.img_list)] h, w img.shape[:2] top random.randint(0, max(0, h - self.patch_size)) left random.randint(0, max(0, w - self.patch_size)) clean img[top:topself.patch_size, left:leftself.patch_size].copy() noise np.random.normal(0, self.sigma / 255.0, clean.shape) noisy np.clip(clean noise, 0, 1).astype(np.float32) clean torch.from_numpy(clean).permute(2, 0, 1) noisy torch.from_numpy(noisy).permute(2, 0, 1) return noisy, clean这段代码里两个细节值得说。第一sigma除以255是为了在0~1的浮点域合成噪声和图像归一化范围对齐。第二random.randint做随机裁剪每个epoch看到的patch不同相当于免费的随机平移增强。我故意没有做固定Grid采样因为随机裁剪在去噪任务上效果更好而且实现简单。训练循环部分如下def train_one_epoch(model, loader, optimizer, criterion, device): model.train() total_loss 0.0 for noisy, clean in loader: noisy, clean noisy.to(device), clean.to(device) optimizer.zero_grad() pred model(noisy) loss criterion(pred, clean) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() total_loss loss.item() return total_loss / len(loader)clip_grad_norm_是必须加的一行。Swin类模型在训练后期偶尔会遇到个别batch的梯度异常偏大不裁剪的话loss会瞬间飞掉前面的训练全部作废。max_norm设1.0比较保守一般不会影响正常收敛速度。数据增强方面我只建议水平翻转和旋转90度。颜色抖动、高斯模糊这类强度增强在去噪上收益很小反而会拖慢收敛。如果你用灰度图训练BSD400颜色增强完全没用。3.3 损失函数MSE、L1和感知损失在去噪任务中的取舍损失函数直接决定模型学到的“干净”是什么。MSEL2损失是去噪论文里的默认选择它对应PSNR指标的优化目标在合成高斯噪声下收敛平滑。但MSE的问题是它对微小结构差异的惩罚梯度很小输出容易出现过度平滑。L1损失在低噪声区域的梯度更稳收敛速度不如MSE快但最终PSNR往往接近SSIM反而更高。我在自己的实验里常用的配方是60%的L240%的L1。训练前期L2主导快速逼近正确范围后期L1主导细化边缘。如果你用SwinUNet结构这个组合损失能让验证集PSNR比纯L2高0.1dB左右。代码实现只有三行def combined_loss(pred, target, alpha0.6): mse nn.functional.mse_loss(pred, target) l1 nn.functional.l1_loss(pred, target) return alpha * mse (1 - alpha) * l1alpha控制L2比重从0.5到0.7之间调整。alpha值越大输出越平滑越小越锐利但可能残留噪声。训练到一半想改alpha也没问题直接换了继续训不会造成训练不稳。感知损失Perceptual Loss在去噪里要谨慎用。用VGG特征做感知损失会让模型偏向保留高频纹理在真实噪声数据上有帮助但在合成噪声上感知损失带来的PSNR增益通常不足0.05dB却增加了一倍GPU显存和训练时间。真实场景去噪推荐加合成数据任务不建议。3.4 评估指标PSNR与SSIM的计算方式和读数习惯PSNR和SSIM是去噪领域公认的指标。PSNR公式是10*log10(MAX^2/MSE)MAX是像素最大值归一化到0~1时就是1。计算时要注意必须在浮点域计算再乘回255如果对uint8图直接算会引入量化误差。另一个关键点是PSNR要在裁剪边界后计算还是全图计算必须和论文保持一致不然你的结果和人对不上没法比。def compute_psnr(pred, target): pred pred.clamp(0, 1) target target.clamp(0, 1) mse ((pred - target) ** 2).mean() return 10 * torch.log10(1.0 / mse)这里clamp是关键。模型输出的pred是浮点型的0~1张量必须先clamp到[0,1]再计算和保存。忘了clamp是高频翻车点pred里出现极小负值或大于1的值保存时会被截断或溢出肉眼看到的结果就是一堆奇怪的色块。SSIM用滑动窗口计算结构相似性窗口大小默认11不要改。SSIM对空间结构损失的敏感度和人眼更接近但训练中直接拿SSIM做损失函数会导致收敛不稳定。我一般每N个epoch记录PSNR和SSIM保存best PSNR的权重。4. 推理与验证拿训练好的模型处理一张真实噪图要几步4.1 推理脚本模型加载、图像预处理与结果保存推理比训练简单但实际写起来坑更多。核心流程加载权重、读图转浮点、进模型、clamp、保存。代码如下import torch import cv2 import numpy as np def denoise_image(model, img_path, output_path, devicecuda): model.eval() img cv2.imread(img_path, cv2.IMREAD_COLOR) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img img.astype(np.float32) / 255.0 # 转张量并加batch维度shape: (1, 3, H, W) x torch.from_numpy(img.transpose(2, 0, 1)).unsqueeze(0).to(device) with torch.no_grad(): pred model(x) pred pred.squeeze(0).cpu().numpy().transpose(1, 2, 0) pred np.clip(pred, 0, 1) pred (pred * 255.0).astype(np.uint8) result cv2.cvtColor(pred, cv2.COLOR_RGB2BGR) cv2.imwrite(output_path, result)这里两个容易出错的地方。第一模型是在归一化到0~1的数据上训练的推理输入必须做相同的除法255操作差一步结果就会偏暗或偏亮。第二OpenCV读图默认是BGR通道序做去噪要在RGB域最后保存再转回BGR。如果全程用BGR推理模型输出颜色会明显偏色。4.2 自集成推理旋转、翻转与多尺度投票提升PSNR测试时增强TTA是去噪实验里常见的技巧。原理很简单对输入做8次几何变换旋转0/90/180/270度各自再水平翻转分别推理后再做逆变换取平均。这样做能抑制模型对特定方向特征的偏好PSNR在白噪声条件下能稳定提高0.10.2dB。def tta_denoise(model, x): preds [] for i in range(4): img_rot torch.rot90(x, i, dims[2, 3]) p model(img_rot) preds.append(torch.rot90(p, -i, dims[2, 3])) img_flip torch.flip(img_rot, dims[3]) p model(img_flip) preds.append(torch.flip(p, dims[3])) return torch.stack(preds).mean(dim0)8次推理的代价是推理时间变为8倍。如果输出是1920x1080的大图一次推理可能已经要2~3秒TTA后接近20秒这就需要取舍。我个人的习惯是开发验证阶段不用TTA用基线快速迭代最终出效果图或提交结果时用TTA拿到最优指标。多尺度TTA是另一个方向。把输入缩小到0.8和1.2倍再推理逆变换回原尺寸取平均。注意下采样后的图像要和模型window_size对齐否则会出现网格状伪影。这个技巧能再提升约0.05dB但代码复杂度上升明显新手先做8次几何TTA就足够。4.3 模型导出TorchScript与ONNX的转换注意事项训练好的模型最终要给别人用导出成TorchScript或ONNX是常见后续工作。导出的第一个坑是动态尺寸。Swin模块里有大量reshape操作很多实现写死了输入尺寸用torch.jit.trace导出时会把batch和通道维度固化成常量换个尺寸的图就报错。解决办法有三个方向用torch.jit.script而不是trace或者在模型forward里确保所有reshape基于张量的shape属性而不是绝对值或者直接用torch.onnx.export并设置动态轴。如果导出ONNXSwin的窗口操作包含大量的reshape和transposeopset版本建议设到11以上否则部分张量操作导出不完整。导出后一定要先用onnxruntime跑一遍对比PyTorch原始输出。像素误差超过1e-3就需要检查哪个算子不被支持。如果遇到量化导出某些自定义算子还需要手动实现替换。这个步骤虽然繁琐但省掉的话模型交付后才发现跑不了返工成本更高。5. 避坑指南Swin-TransformerUNet训练里的5个踩坑记录5.1 现象训练损失平稳下降但验证集PSNR几乎不动原因模型在“复制输入”。去噪任务有个经典陷阱——如果模型学到的只是把输入直接搬到输出L2损失依然会持续下降因为干净图和带噪图的差异本就不大。这种情况下验证集PSNR最多比输入高1dB左右。解决从两个方向排查。第一确认训练时输入是加噪图、标签是干净图新手很容易在数据加载里把noisy和clean传反。第二把第一层输出可视化如果特征图看起来就是输入图说明模型通道配置可能有问题试着把嵌入维度调大或加深层次。一句话训练损失不是衡量去噪是否学好的唯一标准验证集PSNR才是。5.2 现象显存溢出16G显存也只能跑batch size4原因SwinTransformerBlock的中间特征图是(B, H*W, 4C)再加上UNet每层的卷积特征峰值显存往往是输入的十几倍。很多人不知道128x128 patch下window_size8的注意力计算量虽然小但MLP的中间线性层会把通道数扩到4倍这里才是显存大头。解决三个手段组合使用。先把输入patch从128降到96显存占用大约降一半再把batch size改为2加梯度累积最后把Swin块的mlp_ratio从4降到2PSNR损失约0.1dB。优先级从前往后排先降patch大小改动最小收益最大。5.3 现象推理结果出现规则的网格状伪影原因输入图像的宽或高不是window_size的整数倍。Swin的窗口划分要求H/ws和W/ws是整数推理时遇到非整除会paddingpadding的像素没有正确的上下文导致每ws个像素出现一条伪影带。这在1920x1080的图片上如果用了window_size10就会必定触发。解决推理前pad输入到8或16的倍数推理后crop回去。注意用reflect模式padding比zero padding效果好很多不会在边缘引入黑色边框。def pad_to_multiple(x, multiple8): _, _, h, w x.shape ph (multiple - h % multiple) % multiple pw (multiple - w % multiple) % multiple return torch.nn.functional.pad(x, [0, pw, 0, ph], modereflect)5.4 现象合成噪声上测试PSNR高真实噪图上一团糊原因合成高斯噪声的分布和真实传感器的噪声分布完全不同。真实噪声有亮度相关的泊松分量、去马赛克伪影、坏点等。SwinUNet在合成噪声上看到的pattern很单一学到的滤波器对真实噪声不鲁棒。解决两个方向。一是用真实噪声数据集做微调RENOIR或自己拍raw噪声数据二是在训练时做噪声合成增强——把高斯噪声、泊松噪声按随机比例混合甚至叠加随机条纹噪声。混合噪声训练会让模型在合成数据评测时略降约0.1dB但在真实噪声上的泛化能力显著提升这个取舍是值得的。5.5 现象同样的代码和参数第二次复现PSNR掉了0.3dB原因去噪训练的随机性很大。数据采样顺序、dropout的随机性、GPU算子的非确定性这些叠加起来两次训练的PSNR可能有0.20.4dB波动这在SwinUNet这种深度模型上尤其明显。解决在训练脚本最前面固定随机种子并设置PyTorch的确定性模式。注意固定代价是训练速度降低10%左右。random.seed(42) np.random.seed(42) torch.manual_seed(42) torch.cuda.manual_seed_all(42) torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False复现实验时batch size和patch size尽量保持完全一致。很多人改动DataLoader的num_workers后PSNR变了以为是代码bug其实是数据采样序列变化导致的正常波动不是玄学是确定性没锁住。6. 进阶技巧用噪声级别图把模型从固定噪声升级成盲去噪标题里的模型通常默认是固定噪声级别训练比如sigma25但我们实际使用时并不知道照片噪声有多大。用噪声级别图Noise Level Map可以让模型适应任意强度的噪声。做法是给模型额外输入一个与图像同尺寸的通道每个像素的值就是该位置的局部噪声强度估计。这个思路来自FFDNet。最小改动是在UNet输入层把3通道变成4通道第四个通道填充为全局sigma/255。训练时sigma从5到50随机采样测试时先用快速估计器算出sigma再填入。这样换来的是一个模型同时覆盖低中高噪声不用为每个噪声级别单独训练。验证模型是否真的在工作最快的办法不是看PSNR而是把残差图打出来。残差图等于输入减去去噪结果理论上残差就是噪声。如果残差图里还有明显的结构边缘纹理说明模型正在把细节当噪声抹掉这样的模型不能直接上线。这个诊断方法我每次训练都会截图记录几十个epoch的变化趋势比单一PSNR曲线信息量大得多。另外一个实用习惯是每次训练结束保存一份完整训练记录包含loss曲线、验证集PSNR曲线、残差缩略图以及本次训练的随机种子和所有超参数。一个月后想复现或调参时这份记录能省下大量返工时间。所谓“玄学调参”大部分时候只是因为没有记录足够信息导致无法回溯哪一步真正影响了结果。希望这条思路对你有用也希望这篇文章里从结构到代码到踩坑的完整链路能帮你真正跑通一个Swin-TransformerUNet的去噪项目。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑