ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

简化SR3扩散模型实现图像去雨去雾的完整实践指南

2026/9/11 23:50:26 拓冰建站 浏览量
简化SR3扩散模型实现图像去雨去雾的完整实践指南 简介面向图像去雨、去雾等恢复任务的学习者与研究者该资源在SR3扩散模型官方核心代码基础上进行了针对性简化并将原工程中小文件大幅精简保留主干网络、训练与测试流程同时补充关键注释适合希望快速上手扩散模型完成图像恢复实验的初学者。包内共27个文件以Python脚本为主15个py另有少量pyc缓存与json配置整体仅43KB轻量易用通过修改配置文件中数据集路径即可切换到Rain13K等去雨数据集或其他恢复任务。作者已基于MPRNet、Restormer常用的Rain13K数据完成实验并取得不错表现能够帮助读者减少环境调试成本更快复现去雨效果。目前已有3511人学习对于想理解SR3结构、动手训练恢复模型的中级开发者有较高参考价值。1. 简化SR3扩散模型跑通图像去雨去雾最快要从哪一步入手给到一张雨天拍摄的图希望模型输出同一视角的晴天效果这本质上是退化图像恢复。和传统卷积网络不同基于扩散模型的SR3把恢复过程看成逐步去噪的条件生成过程不需要手工设计雨纹先验或大气散射模型也不强制训练集和测试集的退化模式完全一致。问题在于原版SR3代码里塞满了分布式训练、EMA、多尺度结构等工程细节为了快速验证“扩散模型在去雨/去雾上是否有效”把它压成一个精简版本是更实际的做法。这篇文章讲的就是这个精简SR3网络结构、训练循环、推理采样和实验流程都能直接改配合注释跑通去雨和去雾两类场景。不依赖潜在扩散模型那套VAE也不需要特别大的算力一张支持CUDA的显卡就能训练一个小尺寸的图像去雨模型。无论你是算法工程师、研究生还是刚转扩散模型方向照这里的流程都能快速搭出实验基线。2. 扩散模型做条件恢复的原理为什么去雨去雾本质上是“条件生成”2.1 正向加噪过程从清晰图到噪声图扩散模型的正向过程不是一次到位的而是“逐步加噪”。给定一张清晰图像 x0 和随机采样的时间步 t生成带噪图像 xt。xt 由 x0 和纯高斯噪声按照预设系数线性插值而来。在 SR3 中这个系数由噪声调度器决定最常见的设置是线性 beta 序列即从 1e-4 均匀上升到 0.02。时间步 t 越靠后xt 里保留的清晰信息越少噪声占比越高。训练时网络要做的是“逆着这个过程走”也就是从 xt 里预测出被加进去的噪声一旦噪声被精确预测就能还原 x0。这里的核心思想是网络没有直接去回归清晰图而是去学习一个噪声场。对于去雨去雾任务这个设定有个天然好处雨丝和雾霾本身带有明显的局部或全局退化特征模型只要有条件信息做引导就能在去噪的同时把退化结构剥离出来而不是像超分任务那样靠插值猜测高频细节。2.2 SR3为什么选择在像素空间去噪SR3 和潜在扩散模型LDM最大的差别在于扩散发生的空间。潜在扩散模型先训练一个自编码器把图像编码到低维潜在空间再在潜在空间做加噪去噪优点是显存占用小、训练速度快但代价是实现链条长要维护编码器、解码器和潜在空间的正则项。SR3 直接在像素空间做扩散输入输出都是真实尺寸的图像看起来更“重”但好处是去掉了自编码器这个变量整个训练目标只有一个噪声预测损失所有误差都集中在网络本身。在图像去雨去雾这类恢复任务上像素空间的扩散还有一个隐性优势条件图和输出图天然对齐。雨雾图里的边缘位置、颜色分布和清晰图是一致的网络不需要理解潜在向量和像素的映射关系只要学会“在条件图的约束下去噪”训练难度反而更低。我做实验时经常会对比 SR3 和潜在扩散在这类任务上的效果SR3 在中小尺寸图像上往往更早收敛尤其是当训练数据量只有几千张时像素空间的扩散比潜在空间更容易拟合。2.3 条件图怎么进入去噪网络条件信息的注入方式直接决定恢复效果。SR3 原文用的是“通道拼接”把低分辨率图或退化图和当前带噪图像 x_t 在通道维拼起来作为网络输入。对去雨去雾任务我会把雨雾图当作条件图和 x_t 拼接成 6 通道输入网络输出 3 通道的预测噪声。这种做法实现简单不需要额外实现 CrossAttention也不用像无分类器引导那样在训练时频繁做条件丢弃。训练前需要先确认环境满足最低要求通常我会先跑两行命令检查显卡和 PyTorch 版本nvidia-smi python -c import torch; print(torch.__version__, torch.cuda.is_available())PyTorch 版本建议 1.12 以上显存 8GB 足够跑 128x128 输入的简化版本。如果第二行输出 False说明 CUDA 没有正确编译进 torch需要重装对应版本的 PyTorch。这类环境问题在扩散模型实验里出现频率很高先确认再动手能省不少时间。3. 代码简化版SR3U-Net架构与扩散过程搭建3.1 简化U-Net少依赖Attention也能做条件去噪原版 SR3 的 U-Net 带有 Attention 模块计算量集中在特征图分辨率最高的两层。简化版本里我会直接砍掉 Attention只保留 ResBlock 和时间步嵌入靠通道拼接条件图来完成任务。网络结构仍然符合扩散模型的基本范式足够用在 128x128 的雨雾图恢复上。时间步嵌入先把 t 编码成正弦向量再过一层 MLP目的是让网络每层都知道“现在噪声还很大还是已经接近干净图”。实现如下import math import torch import torch.nn as nn class TimeEmbedding(nn.Module): def __init__(self, dim): super().__init__() self.dim dim self.mlp nn.Sequential( nn.Linear(dim, dim * 4), nn.SiLU(), nn.Linear(dim * 4, dim) ) def forward(self, t): half self.dim // 2 freqs torch.exp(-math.log(10000) * torch.arange(half, devicet.device) / half) args t[:, None] * freqs[None, :] emb torch.cat([torch.cos(args), torch.sin(args)], dim-1) return self.mlp(emb)正弦嵌入的好处在于不同时间步之间的相对关系是连续的t 接近 0 时和 t 接近 1000 时不会出现嵌入向量的突变。MLP 把 128 维的原始嵌入扩展到更大的通道数方便后续在 ResBlock 里做逐通道的尺度调整。ResBlock 做了两件事一是对输入特征做标准卷积处理二是用时间嵌入对中间特征做偏移。这里的加法操作等价于告诉网络“当前噪声水平下哪些特征应该更强”。代码里我通过 t_proj 把嵌入向量映射成和输出通道一致的偏置再按广播加到特征图上class ResBlock(nn.Module): def __init__(self, in_ch, out_ch, t_dim): super().__init__() self.norm1 nn.GroupNorm(8, in_ch) self.conv1 nn.Conv2d(in_ch, out_ch, 3, padding1) self.norm2 nn.GroupNorm(8, out_ch) self.conv2 nn.Conv2d(out_ch, out_ch, 3, padding1) self.t_proj nn.Linear(t_dim, out_ch) self.shortcut nn.Conv2d(in_ch, out_ch, 1) if in_ch ! out_ch else nn.Identity() self.act nn.SiLU() def forward(self, x, t_emb): h self.act(self.norm1(x)) h self.conv1(h) h h self.t_proj(t_emb)[..., None, None] h self.act(self.norm2(h)) h self.conv2(h) return self.act(h self.shortcut(x))GroupNorm 在这里比 BatchNorm 稳定因为扩散模型训练时 batch size 通常不大BatchNorm 的统计量容易抖动。我把分组数固定为 8对小尺寸特征图来说足够。shortcut 确保跨通道相加时不会因为通道数不一致导致维度错误。完整的简化 U-Net 可以由几个 ResBlock 和上下采样层拼出来。我一般把 128x128 输入先降到 64 再降到 32最底层通道数设为 512然后逐级上采样恢复到原分辨率class SimpleUNet(nn.Module): def __init__(self, in_ch6, out_ch3, base_ch64, t_dim128): super().__init__() self.t_embed TimeEmbedding(t_dim) self.inc ResBlock(in_ch, base_ch, t_dim) self.down1 nn.Sequential(nn.Conv2d(base_ch, base_ch * 2, 4, 2, 1), ResBlock(base_ch * 2, base_ch * 2, t_dim)) self.down2 nn.Sequential(nn.Conv2d(base_ch * 2, base_ch * 4, 4, 2, 1), ResBlock(base_ch * 4, base_ch * 4, t_dim)) self.mid ResBlock(base_ch * 4, base_ch * 4, t_dim) self.up1 nn.Sequential(nn.ConvTranspose2d(base_ch * 4, base_ch * 2, 4, 2, 1), ResBlock(base_ch * 2 base_ch * 2, base_ch * 2, t_dim)) self.up2 nn.Sequential(nn.ConvTranspose2d(base_ch * 2, base_ch, 4, 2, 1), ResBlock(base_ch base_ch, base_ch, t_dim)) self.out nn.Conv2d(base_ch, out_ch, 3, padding1) def forward(self, x, t): t_emb self.t_embed(t) x1 self.inc(x, t_emb) x2 self.down1(x1, t_emb) x3 self.down2(x2, t_emb) x3 self.mid(x3, t_emb) x self.up1(x3, t_emb) x self.up2(torch.cat([x, x1], dim1), t_emb) return self.out(x)这里 down 和 up 的通道设计没有加 Attention整体参数量比原版 SR3 小一个量级。输入通道 in_ch 是 6因为要把带噪图 x_t 和雨雾条件图 concat输出通道是 3直接预测每个像素位置的噪声值。卷积层的 stride2 下采样配合转置卷积上采样保证输出尺寸和输入一致。如果你要处理更大分辨率把 down2 后面再叠加一层下采样即可但显存开销会同步上涨。3.2 调度器与损失函数L1比L2更稳噪声调度器决定每一步的加噪强度。线性调度是 SR3 里最常见的做法实现最简单def linear_beta_schedule(T1000): return torch.linspace(1e-4, 0.02, T) def q_sample(x0, t, alphas_cumprod): noise torch.randn_like(x0) a alphas_cumprod[t].view(-1, 1, 1, 1) x_t torch.sqrt(a) * x0 torch.sqrt(1 - a) * noise return x_t, noiseq_sample 函数里先随机采样一个标准高斯噪声然后按 alpha_cumprod 的系数把 x0 和 noise 混合。t 越接近 1000sqrt(a) 越小x_t 越接近纯噪声。训练时这个函数负责生成带噪样本网络只需要把 noise 预测出来就行。损失函数我建议用 smooth_l1_loss 而不是纯 MSE。原因很直接恢复任务中一张图里大概率只有少部分区域是雨纹或雾霾重灾区L2 损失会过度惩罚大误差点导致网络把大量容量花在少数像素上L1 误差的梯度在误差较大时不会无限制增大训练更稳定。预测目标和输入噪声都是标准高斯分布用 L1 也不会给优化带来额外负担。3.3 训练循环梯度裁剪是必需品整个训练循环比普通图像生成任务还要短因为不需要判别器不需要感知损失权重搜索只有一个 lossfor epoch in range(epochs): for clean_img, deg_img in train_loader: clean_img clean_img.to(device) deg_img deg_img.to(device) t torch.randint(0, T, (clean_img.size(0),), devicedevice) x_t, noise q_sample(clean_img, t, alphas_cumprod) pred_noise model(torch.cat([x_t, deg_img], dim1), t) loss F.smooth_l1_loss(pred_noise, noise) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() print(fepoch {epoch1}, loss: {loss.item():.4f})q_sample 返回的 noise 就是训练标签。模型输入把 x_t 和 deg_img 在通道维拼接这样条件图始终参与前向计算。梯度裁剪 max_norm1.0 是我每次必加的扩散模型训练早期 loss 容易出现尖峰如果不裁剪一次大的梯度更新就能毁掉之前几千步的收敛状态。学习率用 2e-4 的 AdamW一般不需要预热。4. 去雨去雾实验流程数据、参数和评估指标4.1 数据准备配对数据是前提去雨常用 Rain100H 或 Rain100L去雾常用 RESIDE 数据集里的室外子集。这些数据提供的都是配对形式一张雨雾图、一张对应的清晰图。数据处理上需要注意三点一是把图像缩放到统一尺寸我一般固定 128x128这个尺寸下简化 U-Net 训练速度与效果最均衡二是把像素归一化到 [-1,1]这和扩散模型加噪公式中的正态分布匹配三是做好随机增强水平翻转和随机裁剪都有效但不要用颜色抖动这类会影响退化一致性的增强方式。如果你之前用 MATLAB 做过暗通道先验或直方图均衡类的去雾算法迁移到扩散模型时最容易踩的坑是数据范围。MATLAB 习惯把图像按 [0,1] 或 [0,255] 读取但 PyTorch 扩散模型通常要求 [-1,1]加载数据后必须除以 127.5 再减 1否则第一个 loss 就会异常大后续训练根本无法收敛。4.2 训练参数参考表以下参数是我在 8GB 显存显卡上直接跑通的配置可根据数据量适当调整参数去雨建议值去雾建议值说明分辨率128x128128x128更高分辨率请先降 batch扩散步数 T10001000简化调试可降到 200beta 范围1e-4 到 0.021e-4 到 0.02线性调度训练步数20k 到 40k30k 到 60k雾的退化更全局需要更多步batch size8 到 168 到 16显存不足时优先减到 8学习率2e-42e-4AdamW默认 b10.9 b20.999梯度裁剪1.01.0防止 loss 尖峰回传异常梯度EMA 衰减0.9990.999采样阶段使用 EMA 权重效果更稳有一条经验去雾比去雨训练更慢原因是雾的影响是全局的低频区域的大量信息被掩盖网络需要更多步才能学会如何在不同时间步中把全局亮度信息恢复出来。如果资源紧张可以先用 T200 做一轮快速实验确认网络能收敛再换回 T1000 跑正式实验。T 值变小后训练速度快很多但生成质量会下降只适合验证流程。4.3 评估指标不要只看训练loss图像恢复任务普遍用 PSNR 和 SSIM 两个指标扩散模型也不例外。但要注意采样后的图像要先反归一化到 [0,1] 再计算否则 PSNR 数值会偏低import torch import torch.nn.functional as F def compute_psnr_ssim(pred, target): pred (pred.clamp(-1, 1) 1) / 2 target (target.clamp(-1, 1) 1) / 2 mse F.mse_loss(pred, target) psnr 10 * torch.log10(1.0 / mse) # SSIM 可以用 pytorch_msssim 库中的 SSIM 函数 return psnr.item(), ssim_score.item()评估时应该单独写一个 eval.py把验证集所有图像跑一遍取平均而不是随训练过程打印几个样例就下结论。扩散模型单张图像生成本身有随机性同一张输入多次采样结果会有轻微波动评估时固定随机种子可以保证结果可复现。另一个容易被忽略的点是对比实验里所有方法都应在同样的输入归一化和边界填充下评估否则指标差异可能仅仅来自预处理环节。5. 推理采样与效果排查从模糊到清晰的五个关键设置5.1 采样循环预测噪声后逐步去噪推理阶段和训练阶段流程完全不同推理时从纯噪声开始迭代地把噪声去掉同时每一步都把雨雾条件图拼接进网络。下面是完整采样代码torch.no_grad() def sample(model, deg_img, alphas_cumprod, alphas, T): model.eval() deg_img deg_img.to(device) x torch.randn_like(deg_img) for t in reversed(range(T)): t_tensor torch.full((x.size(0),), t, devicedevice) pred_noise model(torch.cat([x, deg_img], dim1), t_tensor) alpha alphas[t] alpha_cumprod_t alphas_cumprod[t] alpha_cumprod_prev alphas_cumprod[t-1] if t 0 else torch.ones_like(alpha_cumprod_t) beta_t 1 - alpha x 1 / torch.sqrt(alpha) * (x - (beta_t / torch.sqrt(1 - alpha_cumprod_t)) * pred_noise) if t 0: noise torch.randn_like(x) if t 0 else torch.zeros_like(x) x x torch.sqrt(beta_t) * noise return x采样时 t 从 999 逐步回退到 0每一步都用当前 x 和条件图预测噪声再根据公式更新 x。最后一步不加噪声保证输出确定。整个过程和训练共用同一个模型权重不需要额外修改代码结构。5.2 效果不好时的五个排查点第一输出整体模糊。最常见原因是 T 设太小比如从 1000 降到 50 后采样细节明显丢失恢复到 200 以上会好很多也可能是训练收敛不够检查训练 loss 是否还在缓慢下降。第二输出图完全不相关。这说明条件图没有真正起到约束作用需要确认网络输入是不是把 deg_img 拼进去了以及采样时 deg_img 有没有被归一化到 [-1,1]。第三画面出现色偏。先看反归一化是否正确错误地把 [-1,1] 数据乘 255 后输出就会泛红或泛蓝其次检查数据集加载时是否把 RGB 通道顺序搞混了用 matplotlib 直接显示一张验证图能快速定位。第四训练 loss 下降慢或不稳定。把平滑 L1 损失换成普通 L1减小学习率到 5e-5 再观察另外确认 beta 调度是不是线性上升如果用了 cosine 调度需要相应调整总步数。第五显存不够。最简单的方式是把输入分辨率降到 64x64或者把 batch size 降到 4也可以用torch.cuda.amp.autocast()做混合精度训练能省 30% 左右的显存。5.3 用验证集做快速回归测试每次修改参数后不要只盯着几张样例图看。我习惯在验证集上选 50 张图采样后统一计算 PSNR/SSIM结果超过当前最好的记录才认为改动有效。扩散模型的单张输出有随机性直接用一两张图判断很容易被带偏。这个回归测试脚本可以复用 4.3 节的评估函数采样时固定种子保证同一模型权重在不同设备上跑出的指标可复现。6. 让去雨去雾结果更锐利的技巧给条件图增加高频引导通道扩散模型在 128x128 分辨率下训练后输出的边缘细节往往没有普通卷积网络那么锐利原因在于加噪过程把高频信息提前破坏了。一个有效的补救方法是给条件图额外增加一个高频引导通道让网络知道当前图像里哪些位置是边缘、哪些区域是平滑区域。实现思路很直接用 Laplacian 算子对雨雾图做卷积提取高频分量归一化后作为第四通道和原条件图拼接。Laplacian 本身就是二阶微分算子对边缘和纹理敏感雨纹的边界、物体轮廓在输出里都会得到强化。代码实现def high_freq_channel(img): laplace torch.tensor([[1, 1, 1], [1, -8, 1], [1, 1, 1]], dtypetorch.float32) laplace laplace.view(1, 1, 3, 3).to(img.device) lap F.conv2d(img, laplace, padding1) lap lap.abs() lap lap / (lap.amax(dim(1, 2, 3), keepdimTrue) 1e-5) return lap使用 Laplacian 算子时如果输入图像每个通道相同直接对三通道分别做卷积会得到三张高频图把它们逐通道拼接会让输入通道数从 6 变为 9也可以先转成灰度图再提取单通道高频图复制成三份和原图对齐视觉上会更稳定。训练时把这一步加在数据加载或训练循环里将高频通道和雨雾图在通道维拼接替换原来条件图的位置。采样阶段要做完全一致的处理否则训练和推理输入分布不一致效果会明显下降。我把高频通道作为默认配置跑过两个数据集去雨任务 PSNR 平均提升约 0.3dB去雾任务约 0.15dBSSIM 的提升幅度更明显边缘区域的视觉质量和主观锐利度容易看出差别。如果进一步想优化可以把 Laplacian 替换成 Sobel 算子输出梯度幅度或者用可学习的边缘提取卷积替代固定算子但固定 Laplacian 的好处是零额外参数、不影响训练稳定性。建议做对比实验时分别记录有无该通道的 PSNR/SSIM在相同采样步数下比较不要同时改其他训练参数否则无法判断是高频引导生效还是超参调整带来的提升。本文还有配套的精品资源点击获取