简介:这是一份面向图像处理学习者和算法研究人员的Swin-Transformer与UNet结合的去噪项目。项目针对传统去噪方法细节保留不足的问题,提出融合Transformer长距离建模与UNet多尺度特征融合的神经网络,并设计了专用损失函数以提升视觉自然度和信噪比,适用于图像修复、低照度增强等场景。Swin-Transformer采用层级窗口注意力机制,在高分辨率图像处理中兼顾效率与全局感知,配合UNet跳跃连接可有效抑制噪声并保留边缘纹理。压缩包共25个文件,以17个Python源文件为主,涵盖模型定义、训练、测试与数据增强模块,同时包含3个MATLAB文件用于数据合成与评估、2份Markdown说明文档、1个YAML训练配置及若干编译缓存文件,整体仅37KB,轻量易部署。已有581人学习下载。源码结构清晰,包含完整实现、数据集生成工具、多分辨率推理demo及自定义训练配置,方便读者快速复现算法并根据需求调整网络参数,适合有一定深度学习与PyTorch基础的开发者用于算法验证和二次开发,也是深入理解Transformer在底层视觉任务中应用的优质动手项目。
1. 图像去噪选型:Swin-Transformer + UNet 到底解决什么问题
拿到这个项目压缩包的时候,我第一反应不是解压跑 train.py,而是先翻模型定义文件——因为 Swin-Transformer + UNet 这种混合架构,近几年在去噪任务里确实是效果和计算量平衡得比较好的一档。图像去噪算法做了这么多年,从 BM3D 到 DnCNN 再到 Restormer,传统方法在强噪声下细节全丢,纯 CNN 又对长距离结构关系无能为力,而这份资源走的是把 Swin-Transformer 的窗口注意力嵌进 UNet 骨架的路线,既保留 UNet 多尺度特征提取和跳跃连接的细节恢复能力,又用 Transformer 弥补 CNN 感受野不足的问题。适合手里有 PyTorch 基础、想跑通一个真实去噪项目并改到自己数据集上的开发者,也适合做低层视觉方向课程设计的学生直接复现出对比结果。
2. 架构拆解:Swin 窗口注意力如何嵌进 UNet 的跳跃连接
2.1 UNet 还是那个 UNet:编码器、瓶颈、解码器与跳跃连接
UNet 的骨架在去噪任务里至今没被淘汰,核心原因在于它天然适合“输入是图、输出也是图”的像素级预测问题。编码器逐层下采样,把分辨率从 H×W 一路压到 H/16×W/16,特征图的通道数翻倍,每一层学到的是不同尺度的图像结构;解码器再逐层上采样恢复分辨率,同时通过跳跃连接把编码器对应层的特征拼回来——这个拼接操作非常关键,它让解码器在恢复细节时,不需要完全依赖瓶颈层的信息,而是可以直接拿到浅层的边缘、纹理等高频信息。
去噪任务里,噪声通常被认为是高频成分,但图像边缘和纹理也是高频成分,两者在频域上重叠。如果只用深层特征做恢复,很容易把边缘也抹掉,UNet 的跳跃连接恰恰缓解了这个矛盾。实际代码里,跳跃连接做的是通道维度的 concat,比如编码器第 i 层输出是 B×C×H×W,解码器对应层上采样后也是 B×C×H×W,拼接后变成 B×2C×H×W,再经过卷积降维。
Swin-Transformer + UNet 的混合架构在这里的安排一般是:UNet 作为整体骨架,编码器和解码器保留卷积操作,但在瓶颈层附近插入若干个 Swin-Transformer 块;也有变体是把编码器最深层直接替换成 Swin 块堆叠。这样做的理由是瓶颈层分辨率最低、通道数最多,Swin 的窗口注意力在这里计算代价相对可控,同时又能对全局结构关系建模。
2.2 Swin-Transformer 的关键机制:窗口注意力与移位窗口
Swin-Transformer 和 ViT 最大的不同是注意力不是在整个特征图上算,而是在局部窗口内算。ViT 对一张 256×256 的图直接做 global attention,token 数量是 (256/patch_size)²,计算量随分辨率平方增长,去噪任务的输入输出都是全分辨率图,直接上 ViT 不现实。Swin 的做法是把特征图切成不重叠的窗口,比如 window_size=8,每个窗口内部做 self-attention,计算量只与窗口大小有关,与全局分辨率解耦。
更关键的是移位窗口操作(SW-MSA)。如果只做固定窗口的注意力,信息永远只在窗口内流动,跨窗口的关系学不到。Swin 的做法是在相邻层之间交替使用规则窗口和移位窗口——上一层的窗口边界在下一层被打破,token 得以跨窗口交互。去噪任务里这个特性很有用,因为噪声是逐像素破坏,但图像里的大尺度结构(比如天空渐变、墙面纹理走向)是跨区域的,移位窗口让模型有机会把远处结构信息“传递”到当前窗口内。
窗口内注意力的实现细节里有一个容易踩坑的点:移位窗口后特征图边界会多出一些不完整的窗口,常规做法是对特征图做 padding 再切窗口,同时生成一个 attention mask,把不匹配位置的权重置为负无穷,避免 padding 区域参与注意力计算。这个 mask 的生成逻辑在项目代码里通常集中在 get_attn_mask 函数中,后面第 4 章会展开。
2.3 为什么是“Swin + UNet”而不是纯 CNN 或纯 Transformer
纯 CNN 路线(比如 DnCNN)的问题在于感受野有限。虽然深层卷积堆叠能扩大感受野,但去噪需要同时理解局部纹理和全局结构,纯 CNN 对大面积光滑区域的噪声去除经常出现“花斑”伪影,因为它缺少明确的全局建模手段。纯 Transformer 路线理论上能做全局建模,但计算量和训练数据需求都很大,低层视觉任务的数据集规模往往撑不起一个完全抛弃归纳偏置的模型。
混合架构是工程上最务实的选择。UNet 提供多尺度特征和跳跃连接,保证细节恢复的下限;Swin 提供窗口注意力,拉高全局建模的上限。从参数量上看,一个典型的 Swin-T + UNet 去噪模型参数在 20M~40M 之间,比 Restormer 这类纯 Transformer 模型(60M+)小不少,训练显存占用也更友好,在单卡 11GB 的 GPU 上就能跑起来。
另外要提一点,去噪任务里很多项目会额外把噪声估计作为辅助分支,但这份资源从标题看走的是端到端路线——输入带噪图、输出干净图,不显式估计噪声方差。这样做的好处是推理时不需要知道噪声水平,坏处是如果训练时只用单一噪声强度的数据,测试时遇到更强噪声会明显力不从心。所以第 3 章我会强调数据准备时做多尺度噪声训练,这比改模型结构更有效。
3. 工程落地:环境、数据、训练参数一次跑通
3.1 拿到源码先看目录:一份完整去噪项目的文件结构
解压项目后不要急着执行,先用 tree 命令或者直接在 IDE 里展开目录,搞清楚每个文件是干什么的。我经手的去噪项目源码,常规编排大致如下:
| 文件/目录 | 职责 | 备注 |
|---|---|---|
| models/ | 模型定义,包含 UNet 骨架与 Swin 块 | 核心文件,改动最多 |
| datasets/ 或 data_loader.py | 数据读取、加噪、增强 | 决定训练效果的上限 |
| utils/ | PSNR/SSIM 计算、日志、模型保存 | 验证环节依赖这里 |
| config.py 或 options.py | 全局超参数 | 先改这里再碰代码 |
| train.py | 训练入口 | 注意断点续训逻辑 |
| test.py 或 inference.py | 推理与指标评估 | 快速验证用 |
先看 config.py 的另一个原因是你需要确认这份资源默认用的输入尺寸、batch size、噪声 sigma 和水土是否匹配你的显卡。我曾经收到过一个项目默认 batch size 32、输入 256×256 加 Swin 块堆 6 层,自己的 8GB 显卡直接 OOM,最后发现 config 里有个 “crop_size” 参数是罪魁祸首。
3.2 环境安装与依赖:PyTorch 版本、timm 与 CUDA 检查
Swin-Transformer 的实现早期大量基于微软开源的 Swin-Transformer 仓库,依赖里大概率会出现 timm 这个库,用来提供一些通用层(如 DropPath、LayerNorm2d)。安装依赖时我的建议是不要一股脑 pip install -r requirements.txt,先看 requirements 里锁定的 torch 版本和你本机 CUDA 是否匹配。
# 建议先创建干净虚拟环境,Python 3.8~3.10 均可 conda create -n swin_unet python=3.9 -y conda activate swin_unet # 安装 PyTorch,按自身 CUDA 版本调整 cu118 / cu121 pip install torch==2.0.1 torchvision==0.15.2 --index-url https://download.pytorch.org/whl/cu118 # 安装项目其他依赖 pip install timm opencv-python pillow numpy matplotlib tqdm tensorboard这里有几个参数要说明。--index-url指定了 PyTorch 官方 CUDA 11.8 的 wheel 源,如果你本机驱动支持 CUDA 12.x,也可以换成 cu121,但不要装 CPU 版——Swin 的窗口注意力在 CPU 上跑起来会慢到怀疑人生,一个 epoch 跑几个小时很正常。timm库版本不要太新,有些新版本改了 API 导致 DropPath 导入失败,如果遇到cannot import name 'DropPath' from 'timm.models.layers',直接pip install timm==0.6.13锁定旧版本。
安装完成后,先花 10 秒验证 GPU 可用性和 torch 是否正常:
import torch print("CUDA available:", torch.cuda.is_available()) print("Device name:", torch.cuda.get_device_name(0) if torch.cuda.is_available() else "CPU") print("PyTorch version:", torch.__version__)如果打印出来 CUDA available 为 False,先别急着重装,大概率是你装的 PyTorch 是 CPU 版,或者 CUDA 驱动太旧。检查nvidia-smi的 Driver 版本,再回看第一步装 torch 时选的 cu 版本是否对应。这一步跑不通,后面所有训练脚本都会在model.to(device)那行静默走 CPU,训练速度差 50 倍以上,属于最冤的翻车点。
3.3 准备数据集:自己造噪声对。
直接下载现成的去噪数据集当然可以,比如 BSD400、DIV2K,但这份资源既然叫“优质项目实战”,我更建议你用自己的图片目录走一遍全流程——这样后面第 6 章迁移到自己的数据集时就不会慌。训练数据准备的核心逻辑是:准备干净图片,然后在加载时动态加高斯噪声,形成“带噪图→干净图”的训练对。
# datasets/denoise_dataset.py 的核心逻辑 import torch import cv2 import numpy as np class DenoiseDataset(torch.utils.data.Dataset): def __init__(self, image_dir, patch_size=128, sigma_range=(15, 50)): self.image_paths = sorted(glob.glob(f"{image_dir}/*.png")) self.patch_size = patch_size self.sigma_range = sigma_range def __len__(self): return len(self.image_paths) def __getitem__(self, idx): # 读图:统一转 RGB,归一化到 [0,1] img = cv2.imread(self.image_paths[idx]) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = img.astype(np.float32) / 255.0 # 随机裁剪 patch,后面做数据增强 h, w = img.shape[:2] ih = np.random.randint(0, h - self.patch_size + 1) iw = np.random.randint(0, w - self.patch_size + 1) clean = img[ih:ih + self.patch_size, iw:iw + self.patch_size, :] # 随机水平翻转和旋转,相当于免费扩充 8 倍数据量 if np.random.rand() > 0.5: clean = clean[:, ::-1, :] if np.random.rand() > 0.5: clean = clean[::-1, :, :] # 动态加噪:每次 epoch 采样的噪声都不同,等于无限数据 sigma = np.random.uniform(*self.sigma_range) noise = np.random.randn(*clean.shape) * sigma / 255.0 noisy = np.clip(clean + noise, 0.0, 1.0) # 转成 CHW 的 tensor clean = torch.from_numpy(clean.transpose(2, 0, 1)).contiguous() noisy = torch.from_numpy(noisy.transpose(2, 0, 1)).contiguous() return noisy, clean这段代码有几个参数需要重点说明。patch_size=128决定了每张训练图的裁剪尺寸,Swin 的窗口大小要能整除 patch_size,比如 window_size=8,128 正好是 8 的倍数,不会出现窗口越界。sigma_range=(15, 50)是噪声强度的采样范围,单位是像素值域 0~255 下的标准差,15 属于轻噪,50 已经很强了,混合范围训练能让模型对噪声强度不那么敏感。注意noise = np.random.randn(...) * sigma / 255.0,因为图片已经被归一化到 [0,1],所以 sigma 也要先除以 255,这里单位换算错了的话,实际加噪强度会放大 255 倍,训练出来的模型等于在纯噪声上拟合,loss 永远降不下去。
动态加噪的做法比预生成噪声图要好,因为每个 epoch 同一个 patch 会配不同的噪声,相当于数据量膨胀了一个数量级,对防止过拟合帮助很大。这也是我拿到这份资源后第一件事就是确认它的数据加载器是不是动态加噪的原因。
3.4 训练与验证命令:核心参数怎么调
准备好数据和环境后,训练入口通常长这样:
python train.py \ --data_dir ./data/BSD400 \ --val_dir ./data/Set12 \ --patch_size 128 \ --batch_size 16 \ --epochs 200 \ --lr 1e-4 \ --sigma 25 \ --gpu 0逐个说明这些参数的用法和怎么调。--sigma 25这里如果设了固定值,训练时数据加载器就会忽略 sigma_range,每次加噪都用 25,适合想对比不同噪声强度效果的场景;但如果想让模型泛化到未知噪声,我建议在 config 里把sigma_type设为range,让数据加载器走上一小节的动态采样逻辑。--batch_size 16在 256×256 输入下大约占 8~10GB 显存,如果你的卡是 11GB 以下,先降到 8 或者把 patch_size 改成 96,不要同时硬撑。--lr 1e-4对混合架构来说是安全起点,AdamW 优化器下不建议超过 2e-4,否则训练初期 loss 容易出现锯齿。
训练过程中重点关注两个指标:loss 和验证集 PSNR。loss 用 L1 或 Charbonnier 损失,下降曲线应该是平滑递减,如果出现 loss 骤降后长期横盘,多半是学习率没配合调度器;验证集 PSNR 每 5 个 epoch 算一次,注意验证时输入的是固定噪声水平的带噪图,不要用训练时的随机噪声。训练结束后,模型权重一般会按 epoch 保存到checkpoints/目录,我习惯每 10 个 epoch 存一个,最后对比不同 epoch 的验证指标,取最优的做测试推理——这比只看最后一轮靠谱得多,因为训练后期有时会有轻微过拟合。
4. 核心代码精读:Swin 块、UNet 骨架与训练循环
4.1 Swin-Transformer 编码器块实现
这份资源的模型核心在 models/ 下的 swin_unet.py 文件里。先看 Swin 块的实现逻辑,这是整个模型理解难度最高的部分。一个完整的 Swin Transformer Block 包含窗口注意力、移位窗口注意力、MLP 和两层残差连接。
# models/swin_block.py(简化自项目源码) import torch import torch.nn as nn import torch.nn.functional as F class SwinTransformerBlock(nn.Module): def __init__(self, dim, num_heads, window_size=8, shift_size=0, mlp_ratio=4.0): 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, int(dim * mlp_ratio)), nn.GELU(), nn.Linear(int(dim * mlp_ratio), dim) ) def forward(self, x, attn_mask=None): # x: [B, H, W, C],需要先做窗口切分 B, H, W, C = x.shape shortcut = x x = self.norm1(x) # pad 到窗口的整数倍,避免窗口越界 pad_r = (self.window_size - W % self.window_size) % self.window_size pad_b = (self.window_size - H % self.window_size) % self.window_size x = F.pad(x, (0, 0, 0, pad_r, 0, pad_b)) _, Hp, Wp, _ = x.shape # 如果 shift_size > 0,做循环移位,实现跨窗口信息交换 if self.shift_size > 0: x = torch.roll(x, shifts=(-self.shift_size, -self.shift_size), dims=(1, 2)) # 把特征图切成 [B*nW, window_size, window_size, C] x = x.view(B, Hp // self.window_size, self.window_size, Wp // self.window_size, self.window_size, C) x = x.permute(0, 1, 3, 2, 4, 5).contiguous() x = x.view(-1, self.window_size * self.window_size, C) # 窗口注意力 attn_out = self.attn(x, attn_mask) attn_out = attn_out.view(B, Hp // self.window_size, Wp // self.window_size, self.window_size, self.window_size, C) attn_out = attn_out.permute(0, 1, 3, 2, 4, 5).contiguous().view(B, Hp, Wp, C) # 反向循环移位 if self.shift_size > 0: attn_out = torch.roll(attn_out, shifts=(self.shift_size, self.shift_size), dims=(1, 2)) # 裁掉 pad 的部分,恢复原始尺寸 attn_out = attn_out[:, :H, :W, :].contiguous() x = shortcut + attn_out # MLP 残差块 shortcut = x x = self.norm2(x) x = self.mlp(x) x = shortcut + x return x这段代码有几个关键参数必须理解到位。shift_size默认取window_size // 2,也就是 8 的窗口移位 4 个像素,这样规则窗口层和移位窗口层交替叠加,注意力才能跨窗口传播。mlp_ratio=4.0指 MLP 隐藏层是通道数的 4 倍,这个值偏大能涨点但会明显增加计算量。attn_mask只在 SW-MSA 层用,作用是把窗口边界外 padding 区域屏蔽掉,很多新手把注意力掩码和 Transformer 里的 padding mask 搞混,这里其实是“窗口边界 mask”,生成逻辑是判断两个位置是否属于同一窗口区域,不属于则置为-100.0,经过 softmax 后权重趋近于 0。
窗口注意力(WindowAttention)内部就是标准的 multi-head self-attention,只不过把 qkv 投影后的张量 reshape 成[B*num_windows, num_heads, window_size², dim//num_heads],然后做点积。值得注意的一点是相对位置编码 bias 的引入——每个窗口内 token 的相对位置关系是固定的,Swin 用了一个可学习的相对位置偏置表relative_position_bias_table,形状是[(2*window_size-1)², num_heads],训练中会更新。这个 bias 表对去噪效果的影响不小,别删。
4.2 UNet 骨架与跳跃连接对接
UNet 部分相对朴素,编码器每层先做两次卷积(DoubleConv),然后下采样;解码器先上采样,再和编码器对应层的输出做 concat。问题在于 Swin 块的输出是[B, H, W, C]这种排列(PyTorch 的 Transformer 惯用格式),而 UNet 卷积操作要的是[B, C, H, W],两者之间需要 permute 转换,这个转换代码如果漏了就等着 shape mismatch。
# models/swin_unet.py 中的 UNet 骨架 class UNetSwin(nn.Module): def __init__(self, in_ch=3, out_ch=3, embed_dim=96, depths=[2, 2, 6, 2]): super().__init__() # 编码器:卷积下采样 self.enc1 = DoubleConv(in_ch, 64) self.enc2 = DoubleConv(64, 128) self.enc3 = DoubleConv(128, 256) self.pool = nn.MaxPool2d(2) # 瓶颈层:换成 Swin-Transformer 块 self.swin_layers = nn.ModuleList([ SwinTransformerBlock( dim=256, num_heads=8, window_size=8, shift_size=0 if i % 2 == 0 else 4 ) for i in range(depths[2]) ]) self.norm_swin = nn.LayerNorm(256) # 解码器:上采样 + 跳跃连接 concat self.up3 = nn.ConvTranspose2d(256, 256, kernel_size=2, stride=2) self.dec3 = DoubleConv(256 + 256, 256) self.up2 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2) self.dec2 = DoubleConv(128 + 128, 128) self.up1 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2) self.dec1 = DoubleConv(64 + 64, 64) self.out_conv = nn.Conv2d(64, out_ch, kernel_size=1) def forward(self, x): # 编码器路径 e1 = self.enc1(x) # [B, 64, H, W] e2 = self.enc2(self.pool(e1)) # [B, 128, H/2, W/2] e3 = self.enc3(self.pool(e2)) # [B, 256, H/4, W/4] bottleneck = self.pool(e3) # [B, 256, H/8, W/8] # 进入 Swin 块前转成 [B, H, W, C] b, c, h, w = bottleneck.shape swin_in = bottleneck.permute(0, 2, 3, 1) # [B, H/8, W/8, 256] for blk in self.swin_layers: swin_in = blk(swin_in) swin_out = self.norm_swin(swin_in) bottleneck = swin_out.permute(0, 3, 1, 2) # 转回 [B, C, H, W] # 解码器路径 d3 = self.up3(bottleneck) d3 = self.dec3(torch.cat([d3, e3], dim=1)) d2 = self.up2(d3) d2 = self.dec2(torch.cat([d2, e2], dim=1)) d1 = self.up1(d2) d1 = self.dec1(torch.cat([d1, e1], dim=1)) return self.out_conv(d1)这段代码有几个设计点需要理解。embed_dim=96是 Swin 原文的默认通道数,UNet 的瓶颈层是 256 通道,比 Swin-T 原版 96 大不少,因为去噪任务要保留的细节信息比分类任务多。depths=[2, 2, 6, 2]表示每个 stage 里 Swin 块的个数,瓶颈层对应索引 2,也就是 6 个块,这个数字可以调小到 2~4,训练速度能快很多,但效果会略降。num_heads=8是注意力头数,需要能被通道数整除——256/8=32,每个头的维度就是 32,如果你改了瓶颈通道数,记得同步调整头数。
我特别说一下torch.cat([d3, e3], dim=1)这一步。这是 UNet 的跳跃连接,把上采样后的特征和编码器同尺度特征拼在一起,通道数翻倍后再过 DoubleConv。这里的 e3 来自编码器第三层输出,它包含的浅层细节信息是瓶颈层学不到的,这个拼接操作直接决定了去噪结果边缘是否锐利。有些改进版会把普通 concat 换成 attention gate,但那是另一个赛道了,这份资源保持经典 concat,足够用。
4.3 损失函数与训练循环:L1 还是 MSE
去噪模型训练里损失函数的选择有个经验规律:MSE(L2 损失)会让结果偏平滑,PSNR 指标好看但视觉上发糊;L1 损失训练的模型边缘更锐利,但 PSNR 略低。现在去噪论文里主流用 Charbonnier 损失,它是 L1 的平滑近似,在零点附近可导,训练更稳定。
# 训练循环片段 class CharbonnierLoss(nn.Module): def __init__(self, eps=1e-6): super().__init__() self.eps = eps def forward(self, pred, target): # 平滑 L1,eps 控制平滑程度 diff = pred - target return torch.mean(torch.sqrt(diff * diff + self.eps))训练循环里,优化器用 AdamW,学习率调度用 CosineAnnealingWarmRestarts 或 MultiStepLR。这里给出训练循环的关键片段,注意混合精度的使用——Swin 注意力在 FP16 下能省一半显存,但 LayerNorm 和注意力 softmax 建议保持 FP32,否则训练后期可能出现 NaN。
# train.py 训练循环核心 model.train() scaler = torch.cuda.amp.GradScaler() # 混合精度 for epoch in range(start_epoch, epochs): for i, (noisy, clean) in enumerate(train_loader): noisy = noisy.to(device) clean = clean.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): # 自动混合精度 pred = model(noisy) loss = criterion(pred, clean) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # 梯度裁剪 scaler.step(optimizer) scaler.update() if i % 100 == 0: print(f"Epoch {epoch} | Iter {i} | Loss {loss.item():.4f}")max_norm=1.0的梯度裁剪对 Swin+UNet 混合架构尤其重要,因为 Transformer 部分在训练初期容易出现梯度爆炸,loss 直接变成 NaN。scaler是混合精度的核心对象,backward 前用autocast()上下文自动选择 FP16/FP32,backward 后先unscale_再裁剪,顺序不能反。如果你显卡不支持混合精度或者显存充裕,也可以去掉这套逻辑,纯 FP32 训练,代码更简单。
5. 避坑手册:五个高频问题与排查路径
5.1 显存不足:训练一开始就 OOM
现象:CUDA out of memory在 train.py 启动后的第一个 forward 就爆出来。
原因:混合架构里显存大头是 Swin 块的前向激活值,它与输入分辨率平方成正比。默认 patch_size 128 + batch_size 16 + 6 个 Swin 块,在 8GB 显卡上基本必炸。
解决:按顺序尝试——先把patch_size改到 96 或 64,这是最有效的;然后把batch_size降到 8;最后把depths的中间值从 6 改成 2。注意三者联动,改完检查patch_size % window_size == 0,96/8=12 没问题。
5.2 训练 loss 下降正常,但测试 PSNR 卡在 24~25 上不去
现象:训练损失漂亮下降,但 Set12 验证集 PSNR 始终在 24~25dB 徘徊,比论文里低 5~6 个 dB。
原因:最常见的是训练噪声与测试噪声不匹配。如果你训练固定 sigma=25,测试时用了 sigma=50 的噪声,PSNR 掉 3~5dB 非常正常;第二种情况是训练集太小,UNet 的跳跃连接记住了训练图的结构模式,测试时泛化崩了。
解决:把数据加载器改成sigma_range=(15, 50)多噪声训练;如果数据集只有几百张图,先把 patch_size 调小、多做随机旋转翻转增强,再去考虑换更大的预训练模型。
5.3 推理输出出现规则网格伪影
现象:去噪结果图上能看到明显的棋盘格或网格边界,尤其在图像中央区域。
原因:这是 Swin 窗口边界信息断裂的典型表现。推理时输入尺寸如果不是 window_size 的整数倍,代码里的 pad 操作和训练时不一致,导致移位窗口的注意力 mask 错位;另外一个原因是训练和推理时的 pad 方式不同,比如训练时用反射 pad、推理时用零 pad,会放大边界伪影。
解决:推理时统一走 Python 代码里的 pad 逻辑,不要在外面手动 resize。我踩过这个坑后养成的习惯是——在 test.py 里加一个断言:
# 确保推理输入尺寸是 window_size 的整数倍 assert h % window_size == 0 and w % window_size == 0, \ f"输入尺寸必须能被窗口 {window_size} 整除,当前 {h}x{w}"如果确实无法整除,先中心裁剪到最近的整数倍尺寸,再去噪,最后把结果 resize 回原尺寸。
5.4 用彩色图训练,结果发灰或偏色
现象:模型在 RGB 图上训练,推理出来图像颜色不对,整体发灰。
原因:大概率是数据加载时归一化方式不一致。OpenCV 读图是 BGR,你需要转成 RGB;归一化时如果用了 ImageNet 的 mean/std(RGB 各自的均值和方差)做标准化,最后输出时忘了反标准化,结果就是你看到的发灰。很多分类项目迁移过来的代码都会带这个坏习惯。
解决:在数据加载器里统一用img.astype(np.float32) / 255.0,不要用 ImageNet 标准化。如果一定要用标准化,记住输出时要 denormalize:
mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) clean = pred * std + mean # 反标准化后再保存5.5 一个 epoch 要跑 4 个小时,怀疑人生
现象:batch_size 16、patch_size 128,一个 epoch 要 4 小时。
原因:三个常见因素——模型没有真正在 GPU 上跑(常见于 mixed precision 没生效);num_workers=0导致数据预处理成了瓶颈;Swin-Transformer 块堆太深。
解决:第一步跑第 3 章的 GPU 检查脚本,确认 device 正确。第二步把 DataLoader 的num_workers设为 4 或 8,pin_memory=True。第三步看模型里 Swin 块的数量,如果 depths 总和超过 12,先砍到 6 试跑一个 epoch 看时间,再决定是否加回去。训练速度本质上是时间换效果,不要指望 6 个 Swin 块的效果能和 12 个持平,但工程上首先要保证迭代速度——我一般先跑小模型确认代码没问题,再一口气跑大模型过夜。
6. 进阶验证:用 PSNR/SSIM 检验效果并迁移到自己的数据集
先建立一个观念:去噪算法跑通容易,但要证明它真的有效,需要做一组严格的对比验证。我的习惯是固定一个随机种子,在测试集上选三种典型图像——纹理密集区(草地/树叶)、平滑渐变区(天空/墙面)、边缘几何区(建筑/文字),分别计算加噪前后和去噪后的 PSNR 与 SSIM。只在测试集上算一个平均 PSNR 是不够的,因为平滑区域 PSNR 天然虚高,纹理区域才真正拉开不同算法的差距。
# utils/metrics.py import numpy as np import cv2 def psnr_between(img1, img2, max_val=255.0): # img1/img2: 0~255 的 uint8 或 0~1 的 float mse = np.mean((img1.astype(np.float64) - img2.astype(np.float64)) ** 2) if mse == 0: return float('inf') return 10 * np.log10(max_val ** 2 / mse) def ssim_between(img1, img2, win_size=11): # 直接调 skimage,比自己算稳定 from skimage.metrics import structural_similarity as ssim # 输入是 CHW 或 HWC 要提前理清 return ssim(img1, img2, channel_axis=-1)PSNR 是逐像素误差的数学指标,公式是10 * log10(MAX² / MSE),对轻微噪声很敏感,但和人眼感知并不完全一致——同一个 PSNR 下,平滑区域的噪声比纹理区域更扎眼。SSIM 考虑了局部结构相似度,取值 0~1,越接近 1 越好。做算法对比时,我一般两个指标一起报,PSNR 反映数学误差,SSIM 反映视觉感知,单独看任何一个都可能被带偏。
验证脚本写完后,做一个小实验:用固定 sigma=25 的噪声在 Set12 上分别跑这份资源和 OpenCV 的 fastNlMeansDenoising,看指标差距。这种对比不是为了证明谁更强,而是确认资源里的模型有没有正常工作——如果 Swin+UNet 的输出还不如经典非深度学习算法,说明训练出了严重问题,先排查再谈优化。
把模型迁移到自己的数据集上时,有几个联动参数要一起改。如果你的新数据集是灰度图,模型输入通道改成 1,最后一个卷积的输出通道也改成 1,同时训练数据和测试数据的通道维度保持严格一致。如果你的原始图片分辨率特别大(比如 4000×3000 的无人机航拍),推理时不要整张图直接进模型,先切成 512×512 的重叠 patch(overlap 32 像素左右),去噪后再拼回去——整张大图直接推理,Swin 的全局结构建模会因尺寸变化产生伪影,而且显存大概率扛不住。
训练自己的数据集时,我强烈建议先冻结训练好的编码器部分,只微调解码器。具体做法是给模型的 enc1/enc2/enc3 层设requires_grad_(False),只让up1~up3和最后输出层参与训练,学习率调到 5e-5。这样做的逻辑是新数据集的领域和原训练集不同,但低级特征(边缘、纹理基元)是通用的,冻结浅层能防止小数据集下把预训练权重的通用表征冲掉。等微调 20 个 epoch 后,再全部解冻,用 1e-5 的学习率对整个模型精调。
验证这一套流程跑通后,我提炼出一个血泪教训:从那以后每次拿到新的去噪数据集,我先不过夜跑大模型,而是拿 3~5 张代表图,用小 patch_size 把完整流程走一遍,确认指标能跑、推理图不花、颜色正常,才敢上全量数据和完整的 Swin 深度跑整晚。这个习惯帮我避开了至少三次数据格式不一致导致的整夜空跑。希望帮到你。
本文还有配套的精品资源,点击获取