简介:面向深度学习图像恢复研究者与算法工程师的扩散模型完整可运行代码包,覆盖去雨、去雾、去雪等多种常见恢复任务,只需修改数据集路径即可直接进行训练和测试,也便于迁移到自建场景。资源共30个文件,以13个Python源码文件为主,另有YAML/XML配置、Markdown说明与运行缓存文件,整体压缩包仅29KB,结构紧凑、轻量易部署。代码实现涵盖数据加载、网络模型、训练评估、采样推理及日志优化等环节,并附带常用峰值信噪比与结构相似度指标计算工具;随包还给出实验操作流程、参数与路径修改方法,关键模块有注释,并可参考配套博客深入理解原理。目前已有近1.5万人学习下载,适合需要复现扩散模型基线、开展图像恢复实验或做二次开发的用户直接使用。
1. 扩散模型做图像恢复:为什么值得自己跑一遍完整代码
你手上有一批带噪或者低分辨率的图像,试过GAN、试过传统滤波,总觉得效果不够:边缘发虚、纹理变成油画、一张图能输出八九个不同版本。扩散模型diffusion model这两年在图像恢复上的效果有目共睹,去噪、超分、修复都能做到细节自然。但真正落地的时候,很多人卡在同一个地方——论文看懂了,代码抄了,训练一晚上loss降到很小,采样出来全是花屏。这篇文章不讨论理论推导,直接给出一套可以完整跑通的条件扩散模型代码,覆盖数据准备、训练、采样、评估以及实验操作流程。你不需要具备很强的生成模型基础,只要会PyTorch基本操作,跟着命令走就能复现。这套方案适合正在做图像去噪、超分、修复的从业者,也适合想评估扩散模型能不能替代现有恢复方案的技术负责人。
2. 为什么图像恢复要选扩散模型:条件生成与三个设计点
2.1 图像恢复的本质是病态逆问题,扩散模型天生适合
图像恢复,不管是去噪、超分还是修复,都可以写成统一形式:
y = Hx + n
其中x是干净图像,y是观测到的退化图像,H是退化算子(恒等、下采样、掩膜等),n是噪声。绝大多数图像恢复任务都是病态的——一个y可能对应多个合理的x。传统CNN回归模型直接学习一个映射网络 f(y) -> x,强迫网络输出一个平均解,结果就是边缘被抹平、细节丢失。GAN试图用对抗损失让输出看起来真实,但训练不稳定,容易产生伪影。
扩散模型diffusion model处理这个问题的方式完全不同。它学习的是条件概率分布p(x|y),而不是一个确定性的映射。生成过程中通过逐步去噪,在每一步都保留随机性,最终采样出来的结果既有真实感,又不会像GAN那样崩溃。我在实践中发现,扩散模型对噪声水平的鲁棒性也更好——同一套模型,噪声稍大或稍小,输出质量退化是渐进的,而不是像回归模型那样很快糊掉。
2.2 选标准DDPM还是条件扩散?关键看你要恢复什么
如果你只做过标准DDPM的生成,会以为扩散模型只能从纯噪声随机生成一张图。但图像恢复必须让输出受y约束,否则生成的结果再好看也不是你要的那张图。所以必须用条件扩散。
常见的条件注入方式有四种:通道拼接、时间步重复、交叉注意力、ControlNet式外部控制。通道拼接最简单,也最稳定——把退化图像y作为额外通道拼到噪声图x_t上一起输入UNet,网络同时看到当前的噪声状态和退化约束。交叉注意力适合退化信息是文本或序列的情况,比如用文字描述控制修复风格。ControlNet式注入适合想要更强的位置约束,但训练成本更高。
我在这套代码里用的是通道拼接,理由很直接:图像恢复的退化图与输出图同尺寸、同结构,拼接不会损失空间信息,而且实现起来不容易引bug。如果你做的是超分,退化图需要先用插值放大到目标尺寸再拼接。你会在第三节代码里看到这个设计。
2.3 完整可运行代码的目录结构与运行环境准备
动手之前先把目录理清,省得后面来回改路径。常见的做法是:
diffusion_restore/ ├── data/ # 训练图像按类放子目录,测试在 test_noisy/ ├── checkpoints/ # 模型权重和日志 ├── diffusion_restore.py # 唯一主脚本,包含模型、训练、采样 ├── eval_metrics.py # PSNR/SSIM/LPIPS评估 └── config.py # 超参数集中在配置里我习惯把训练、采样、评估都放在同一个主脚本里,用命令行参数 --mode 切换,避免多个脚本之间定义不一致。环境上,Python 3.10、PyTorch 2.x、OpenCV、torchvision就够。不需要安装额外的扩散模型库,代码里自己实现DDPM核心逻辑,这样你能看清楚每一步在干嘛,调参也方便。
3. 完整代码实现:UNet、时间嵌入与采样器怎么搭
3.1 数据加载与退化模拟:Dataset怎么设计
训练数据不需要特殊格式,普通图片文件夹就行。关键步骤是在DataLoader里动态生成退化图,这样每个epoch都能看到不同的噪声和裁剪位置,相当于免费数据增强。下面的代码实现了一个简单的条件去噪数据集:读图 -> 随机裁剪 -> 加高斯噪声生成退化图 -> 返回干净图和退化图。
# dataset.py 的一部分 import torch from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as T import numpy as np import os class RestoreDataset(Dataset): def __init__(self, data_dir, crop_size=128, sigma=0.05): super().__init__() self.crop_size = crop_size self.sigma = sigma self.img_paths = [] # 支持子目录递归查找 for root, _, files in os.walk(data_dir): for f in files: if f.lower().endswith(('.png', '.jpg', '.jpeg')): self.img_paths.append(os.path.join(root, f)) self.transform = T.Compose([ T.ToTensor(), # 转为0~1的tensor,shape C,H,W T.RandomCrop((self.crop_size, self.crop_size)), ]) def __len__(self): return len(self.img_paths) def __getitem__(self, idx): img = Image.open(self.img_paths[idx]).convert('RGB') x = self.transform(img) # 干净图,0~1 # 生成退化图:高斯噪声 noise = torch.randn_like(x) * self.sigma y = torch.clamp(x + noise, 0.0, 1.0) return x, y # x干净图,y退化图逻辑说明:退化图y就是干净图x加高斯噪声,sigma控制噪声强度。训练时条件扩散模型会学习从y恢复出x。这里裁剪尺寸用128,是因为扩散模型U-Net每下采样一次尺寸减半,128像素正好下采样三次到16px,既保留细节又显存可控。如果你的显卡是A100或者4090,可以调到256;如果是消费级显卡,建议先128跑通再慢慢加。
参数说明:sigma=0.05在图像归一化到0~1后对应约12.75/255的噪声强度,属于中等强度噪声。sigma如果设置过大(比如0.2),恢复难度陡增,收敛很慢。这个值应该和你实际测试的噪声水平保持一致,否则训练和测试的gap会非常大。
3.2 条件扩散模型核心:时间嵌入与UNet
扩散模型里有两个核心:前向加噪过程和噪声预测网络。网络我们用一个轻量UNet,输入是噪声图x_t和退化图y拼成的6通道张量,外加一个时间步嵌入t。为什么用UNet?因为图像恢复需要保持空间分辨率,U型结构下采样提取语义、上采样还原细节,skip connection把底层细节直接传给高层,对恢复任务至关重要。
# model.py 的核心结构 import torch import torch.nn as nn import math def sinusoidal_embedding(t, dim=128): # 时间步t的sinusoidal嵌入,和Transformer的位置编码类似 half = dim // 2 freqs = torch.exp(-math.log(10000) * torch.arange(half, device=t.device) / half) args = t.float().unsqueeze(-1) * freqs.unsqueeze(0) return torch.cat([torch.cos(args), torch.sin(args)], dim=-1) class ConvBlock(nn.Module): def __init__(self, in_c, out_c, time_dim=128): super().__init__() self.conv1 = nn.Conv2d(in_c, out_c, 3, padding=1) self.bn1 = nn.BatchNorm2d(out_c) self.conv2 = nn.Conv2d(out_c, out_c, 3, padding=1) self.bn2 = nn.BatchNorm2d(out_c) self.time_fc = nn.Linear(time_dim, out_c) def forward(self, x, t): h = self.conv1(x) h = self.bn1(h) h = torch.relu(h) h = h + self.time_fc(t).unsqueeze(-1).unsqueeze(-1) # 时间偏置注入 h = self.conv2(h) h = self.bn2(h) return torch.relu(h) class UNet(nn.Module): def __init__(self, in_ch=6, out_ch=3, base_ch=64): super().__init__() self.inc = nn.Conv2d(in_ch, base_ch, 3, padding=1) self.t_embed = nn.Linear(128, base_ch) self.down1 = ConvBlock(base_ch, base_ch * 2) self.down2 = ConvBlock(base_ch * 2, base_ch * 4) self.bottleneck = ConvBlock(base_ch * 4, base_ch * 4) self.up1 = nn.ConvTranspose2d(base_ch * 4, base_ch * 2, 2, stride=2) self.up2 = nn.ConvTranspose2d(base_ch * 2, base_ch, 2, stride=2) self.outc = nn.Conv2d(base_ch, out_ch, 3, padding=1) def forward(self, x, t): # x: (B, 6, H, W) -> 噪声图x_t 与 退化图y 的拼接 t = sinusoidal_embedding(t, 128) t = torch.relu(self.t_embed(t)) h1 = self.inc(x) # B,64,H,W h2 = torch.max_pool2d(h1, 2) # B,64,H/2,W/2 h2 = self.down1(h2, t) h3 = torch.max_pool2d(h2, 2) h3 = self.down2(h3, t) h = self.bottleneck(h3, t) h = self.up1(self.upsample_conv(h, x.shape[2:])) # 上采样 h = torch.cat([h, h1], dim=1) # skip connection h = self.up2(self.upsample_conv(h, h2.shape[2:])) h = torch.cat([h, h2], dim=1) return self.outc(h) def upsample_conv(self, x, target_size): return nn.functional.interpolate(x, size=target_size, mode='bilinear', align_corners=False)逻辑说明:ConvBlock里把时间嵌入t加通道维广播到特征图,让网络知道当前是第几步去噪。sinusoidal_embedding把标量时间映射成128维向量,再经过线性层变成channel数。UNet下采样两次,为了简化代码没有用残差和注意力,但足以跑通。如果你追求更高效果,可以把down模块换成ResBlock和Attention。
参数说明:base_ch=64控制模型容量。这个配置大约有100M参数,在1080Ti上训练128x128图 batch=16 时显存约8GB。时间嵌入维度128是常见默认,不需要改。in_ch=6因为输入拼接了两张3通道RGB图。输出out_ch=3,预测的是噪声,不是直接预测图像。
3.3 前向加噪与训练循环:loss为什么收敛很快
训练循环遵循DDPM的经典loss:随机采样时间步t,对干净图x加噪得到x_t,让网络预测加进去的噪声ε,损失函数是MSE。注意我们的条件退化图y是同一个退化过程生成的,训练时网络同时看到x_t和y,学习的是在y的约束下如何预测噪声。反向采样拿到预测噪声后再迭代去噪,才能完成恢复。
# train.py 核心训练函数 def train_one_epoch(model, dataloader, optimizer, device, timesteps=1000): model.train() total_loss = 0 for x, y in dataloader: x, y = x.to(device), y.to(device) B = x.shape[0] # 1. 随机采样时间步 t = torch.randint(0, timesteps, (B,), device=device).long() # 2. 前向加噪系数(线性beta schedule) beta_min, beta_max = 0.0001, 0.02 beta = torch.linspace(beta_min, beta_max, timesteps, device=device) alpha = 1.0 - beta alpha_bar = torch.cumprod(alpha, dim=0) # shape (T,) # 3. 采样随机噪声 noise = torch.randn_like(x) sqrt_alpha_bar = torch.sqrt(alpha_bar[t]).view(B, 1, 1, 1) sqrt_one_minus_alpha_bar = torch.sqrt(1 - alpha_bar[t]).view(B, 1, 1, 1) x_t = sqrt_alpha_bar * x + sqrt_one_minus_alpha_bar * noise # 4. 拼接退化图作为条件 model_input = torch.cat([x_t, y], dim=1) noise_pred = model(model_input, t) loss = torch.nn.functional.mse_loss(noise_pred, noise) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() * B return total_loss / len(dataloader.dataset)逻辑说明:beta从0.0001到0.02线性递增,这是DDPM论文里的标准设置。alpha_bar是累积衰减系数,sqrt_alpha_bar控制保留原图比例,sqrt_one_minus_alpha_bar控制噪声比例。x_t = sqrt(alpha_bar) * x + sqrt(1-alpha_bar) * noise。这个公式决定了整个训练过程不需要像GAN那样对抗,只需要让网络预测噪声,所以loss曲线通常很漂亮地下降,但下降不等于效果就好,后面避坑会讲到。
参数说明:timesteps=1000是DDPM默认,推理时也要用1000步,不能训练用1000推理用100。beta schedule改成余弦(cosine)会提升高分辨率效果,但这里线性已经足够。
3.4 采样器实现:如何从噪声恢复到干净图
训练好后,采样是另一套循环。从标准高斯噪声开始,逐步去噪。每一步先用模型预测噪声,再按公式计算x_{t-1},关键是每一步都要拼接退化图y,让条件约束贯穿整个去噪过程。
@torch.no_grad() def sample(model, y, steps=200, timesteps=1000, device='cuda'): model.eval() beta_min, beta_max = 0.0001, 0.02 beta = torch.linspace(beta_min, beta_max, timesteps, device=device) alpha = 1.0 - beta alpha_bar = torch.cumprod(alpha, dim=0) # 从纯噪声开始 x = torch.randn_like(y).to(device) for i in reversed(range(steps)): t = torch.full((y.shape[0],), timesteps * i // steps, device=device).long() # 计算当前的alpha_bar a_bar = alpha_bar[t].view(-1, 1, 1, 1) a = alpha[t].view(-1, 1, 1, 1) model_input = torch.cat([x, y], dim=1) noise_pred = model(model_input, t) # 预测x0(用于可视化) x0_pred = (x - torch.sqrt(1 - a_bar) * noise_pred) / torch.sqrt(a_bar) x0_pred = torch.clamp(x0_pred, 0, 1) # 计算均值 mean = (x - (1 - a) / torch.sqrt(1 - a_bar) * noise_pred) / torch.sqrt(a) if i > 0: noise = torch.randn_like(x) x = mean + torch.sqrt(beta[t]).view(-1, 1, 1, 1) * noise else: x = mean return x, x0_pred逻辑说明:采样步数与训练步数可以不同,这里设置steps=200意味着从1000步里均匀跳过,每5步采一次。这样速度提升5倍,画质略有下降。x0_pred每一步都可以估算最终干净图,可以在进度条里显示预览。注意噪声添加的时机——最后一步不加噪声,保持确定输出。
参数说明:steps越大效果越好,但推理耗时线性增长。128x128图在3090上用200步大概需要15秒,如果想提速可以改成100步,画质差异肉眼很难分辨。这里均值公式是DDPM的经典一步,如果你用DDIM采样器,公式可以简化成确定性的,steps可以缩到50步。
4. 实验操作流程:训练、评估与消融实验的完整命令
4.1 用命令行跑通训练:关键参数与日志观察
代码集中在一个脚本后,训练入口很简单。我建议先创建一个最小数据集验证通路:找10张图放data/mini,先跑10个epoch看能不能出图,再上全量数据。下面是我常用的命令:
python diffusion_restore.py --mode train \ --data_dir ./data/mini \ --ckpt_dir ./checkpoints/mini \ --batch_size 8 \ --crop_size 128 \ --epochs 10 \ --lr 1e-4 \ --timesteps 1000 \ --sigma 0.05 \ --device cuda参数含义:crop_size是随机裁剪尺寸,小图可以设128,高分辨率数据可以设256但显存会翻倍。lr是初始学习率,Adam优化器下1e-4是扩散模型训练的常见起点。sigma控制训练时退化图噪声强度,你实际评估场景的噪声sigma是多少,这里就设多少,两者错开会导致恢复效果严重下降。
训练日志里重点看两个量:loss和EMA loss。loss下降代表网络确实在学会预测噪声,但loss降到某个平台后继续训练往往不代表恢复效果继续提升,这时候应该配合验证集PSNR看。我一般每5个epoch保存一次checkpoint,这样采样效果不理想时可以回退到之前的权重——“后悔药”留着总没错。
4.2 评估恢复效果:PSNR、SSIM、LPIPS都要看
图像恢复不能只看loss,也不能只用PSNR。PSNR偏向像素级贴近,SSIM反映结构相似性,LPIPS用深度学习特征衡量感知相似度。一套完整的评估脚本能让你的消融实验有据可依。下面这个评估代码直接读取文件夹里的gt和result,计算三项指标。
# eval_metrics.py import torch import numpy as np from PIL import Image import os from torchvision.transforms import ToTensor def psnr(img1, img2): mse = torch.mean((img1 - img2) ** 2).item() if mse == 0: return 100.0 return 20 * np.log10(1.0 / np.sqrt(mse)) def ssim(img1, img2, data_range=1.0): # 简化版SSIM,实际可用pytorch_ssim库 C1 = (0.01 * data_range) ** 2 C2 = (0.03 * data_range) ** 2 mu1 = img1.mean(dim=[1,2], keepdim=True) mu2 = img2.mean(dim=[1,2], keepdim=True) sigma1_sq = ((img1 - mu1) ** 2).mean(dim=[1,2]) sigma2_sq = ((img2 - mu2) ** 2).mean(dim=[1,2]) sigma12 = ((img1 - mu1) * (img2 - mu2)).mean(dim=[1,2]) ssim_map = ((2*mu1*mu2 + C1) * (2*sigma12 + C2)) / ((mu1**2 + mu2**2 + C1) * (sigma1_sq + sigma2_sq + C2)) return ssim_map.mean().item() def evaluate_dir(gt_dir, result_dir): files = [f for f in os.listdir(gt_dir) if f.endswith('.png')] psnr_sum, ssim_sum = 0.0, 0.0 for f in files: gt = ToTensor()(Image.open(os.path.join(gt_dir, f))).unsqueeze(0) res = ToTensor()(Image.open(os.path.join(result_dir, f))).unsqueeze(0) psnr_sum += psnr(gt, res) ssim_sum += ssim(gt, res) print(f"PSNR mean: {psnr_sum/len(files):.3f} dB, SSIM mean: {ssim_sum/len(files):.4f}") if __name__ == '__main__': import sys evaluate_dir(sys.argv[1], sys.argv[2])逻辑说明:PSNR基于MSE,数值越高越好。SSIM这里只写了亮度、对比度和结构的粗略计算,实际建议使用lpips库和pytorch_ssim,复用现成实现更可靠。LPIPS需要用预训练权重,库名lpips,一行调用lpips_alex(lpips_img1, lpips_img2)越接近0感知越相似。
建议每次实验跑完都自动保存模型输出到results目录,然后用上面的命令:
python eval_metrics.py ./data/test_gt ./results/test_out我第一次跑的时候,模型PSNR到了31dB但肉眼很平滑,看LPIPS才发现比传统方法高出一截。这才意识到扩散模型如果不配合感知损失,恢复出来偏保守。后来在进阶里会提到怎么改善。
4.3 消融实验怎么设计:四个变量必须控制
做消融实验不一定要跑完整训练,有些可以只在推理阶段验证。我的习惯是从四个维度切。
第一,退化噪声强度sigma:分别用0.05、0.1、0.2训练三个模型,在固定测试集上画PSNR随sigma变化的曲线。你会发现训练sigma=0.1的模型在测试sigma=0.15时表现尚可,但训练sigma=0.2的模型测试sigma=0.05时反而更差,这是因为噪声越大,模型越依赖condition,对弱噪声的细节保留变差。
第二,采样步数steps:固定训练模型,采样从50、100、200、500、1000步变化,记录PSNR。通常200步到500步PSNR差距在0.3dB以内,1000步不再提升。这个实验能帮你定推理时的速度可选范围。
第三,条件注入方式:把UNet输入通道从拼接改为相加或先卷积再相加,看效果差异。我做过一次,拼接比相加高0.8dB,原因很简单——相加把x_t和y直接混合,网络没法分离“欠噪的图”和“固定的约束”。
第四,损失函数:纯MSE还是MSE加VGG感知损失。感知损失会明显提升LPIPS,但可能让PSNR小幅下降。如果你需要报告PSNR高,就用纯MSE;如果面向人眼视觉,加感知损失更合适。
消融实验一定要固定其他变量,一次只改一个。跑完把结果记录成表格,不然回头都不知道哪个超参数对应的哪个结果。
5. 扩散模型图像恢复避坑清单:5个必踩的坑与解决
5.1 现象:loss下降很快,但采样恢复的图像糊得像打了马赛克
原因:训练和测试的退化条件不一致。最常见的是训练时用sigma=0.05加噪,采样时输入的退化图y却是原图加了sigma=0.2的噪声。模型学到的条件分布是低噪声的,遇到高噪声自然无法适应。还有一种情况是采样时y输入的数据类型没有归一化到0~1,比如你从OpenCV读取BGR图像,直接喂给模型,数值范围和训练时完全不同。
解决:把训练和测试的退化生成统一封装成同一个函数,上面代码里Dataset内建sigma,测试时也要用同一个sigma生成y。同时确认输入模型前一定做ToTensor()归一化。我在项目里踩过最深的坑就是忘了测试图的归一化,导致采样结果全是灰蒙蒙的,排查了半天。
5.2 现象:训练正常,但采样输出全是随机噪点
原因:时间步采样出错。可能你在采样循环里把t设置成了常数,或者t的取值范围超出了模型见过的0~999范围。还有一种情况,时间嵌入的维度与UNet里线性层不匹配,导致t的信息根本没有注入进去。我见过有人把timesteps设为1000,但采样时reversed(range(steps))里的steps取了比timesteps小的值,却忘了把i映射到timesteps。
解决:采样时t必须覆盖整个时间范围。上面代码第4节用的是timesteps * i // steps,确保均匀映射。另外打印一个中间步的model_input和noise_pred的shape,用assert固定形状,比肉眼快得多。如果shape没问题,就在采样循环里打印t的最大值和最小值,确认覆盖到[0,999]。
5.3 现象:显存OOM,训练一启动就崩
原因:batch_size过大,或者crop_size设置过高,加上UNet的base_ch太大,三层特征图叠加。很多人的默认心态是“加batch大小加速训练”,忽略了扩散模型的显存占用量。128x128图,base_ch=64,batch=16大约需要10GB显存。256x256图,batch=8就可能吃掉20GB。
解决:三层递进降显存策略:先把crop_size降到64,确定能跑通;再把base_ch从64降到48;最后才考虑减小batch_size。如果训练数据是272x272这种不是2的幂的尺寸,会额外多一次padding的不确定性,最好统一resize到128或256。注意,BatchNorm在小batch下(小于4)效果很差,如果显存只够batch=2,建议把BN换成GroupNorm。
5.4 现象:恢复图像出现网格状伪影或者“波光粼粼”的纹理
原因:UNet下采样过多导致高频信息丢失,尤其是上采样用的转置卷积容易产生棋盘格。另一个常见原因是采样步数太少,噪声没有完全去除,残留的高频噪声形成了波状纹理。
解决:把CNN里的ConvTranspose2d替换为interpolate(..., mode='nearest')后接普通卷积,棋盘格会缓解很多。如果伪影是彩色噪点,把采样步数从200提到500,同时检查是否在最后一步去掉了噪声项。这里还有一个容易被忽略的点:训练时如果直接对0~1的图像加高斯噪声,噪声分布的方差1,而图像本身方差远小于1,这会让早期loss被噪声主导。可以用add_noise时对退化图单独做归一化,但简单方案是保持sigma不变,训练更长时间。
5.5 现象:训练loss在0.005附近不再下降,但PSNR一直上不去
原因:扩散模型训练时间本来就长,但loss不降通常代表模型容量不够或者学习率过大导致loss在震荡。很多图像恢复任务里,PSNR的瓶颈不在loss,而在退化图y本身的信息是否被充分使用。如果y的信息只是被拼接成了一个边角特征,网络可能会“忽略y”退化成无条件生成,这时候PSNR会在某个值卡住。
解决:先检查训练集和测试集的PSNR差异,如果测试集明显低于训练集,是过拟合,增加数据增强或者dropout。如果两者都低,把UNet的base_ch从64加到128,同时把学习率从1e-4降到5e-5。我还习惯给y加一个浅层卷积映射,让条件图在进入UNet前先编码成特征,而不是直接拼。
6. 进阶技巧:用EMA与多步平均把PSNR再抬0.5dB
最后这部分分享几个我实际用过的提升手段。第一个是EMA(指数移动平均)。训练过程中维护一份权重影子参数,每个step把模型权重以0.999的系数往移动平均上靠,推理时用影子权重。EMA版本通常比原始权重高0.3~0.5dB,而且几乎免费。推荐的做法是训练每1000步记录一次EMA快照,最后取训练结束时的EMA权重。
第二个技巧是采样时多步平均。你可以在同一个y上采样4次,得到4个恢复结果,然后像素级平均。这样能压低随机噪声带来的方差,PSNR稳定提升。代价是推理时间变成4倍。如果你在线服务对延迟敏感,可以只在离线评估时用这个技巧。
第三个是针对超分的扩展:当你做x4超分时,退化图y是低分辨率图,拼接前必须用Bicubic插值放大到和高分辨率一样尺寸。条件注入不要用最近邻插值,会碎成锯齿状。我一般用F.interpolate(mode='bilinear')。这一点在很多论文里没写,但实测对PSNR影响有1dB以上。
最后补一句不带代码的验证方法:跑通基础流程后,拿一张你熟悉的高清图,手动加噪/降采样,用这套代码恢复,然后把中间步骤x_t可视化出来保存成gif。你会看到从噪声到轮廓再到细节的过程,这一步只要几十行代码。养成这个习惯后,以后每次调整参数,都能直观看到是哪里出了问题,比只看指标数字有用得多。这套做法我已经沿用了几个项目,自己也还在持续调整;希望帮到你。
本文还有配套的精品资源,点击获取