简介:这份代码基于PyTorch实现了SRGAN图像超分辨率算法,并融合GaN风格的门控去噪思路,适合深度学习初学者、图像处理研究者以及希望复现GAN类生成模型的开发者,无论是课题研究还是工程实践都具备参考价值。资源配套完整的训练、测试与数据预处理工具,涵盖训练(train.py)、模型定义(model.py)、损失函数(loss.py)、数据处理(data_utils.py)及图像/视频测试(test_image.py、test_video.py)等核心脚本,并内置pytorch_ssim指标模块,可量化评估重建质量;同时提供多张PNG样例图,覆盖图像与视频超分测试场景,便于对比低分辨率、生成结果与原始高清图,进而理解感知损失、残差学习和对抗训练的关键机制,也方便修改网络层数、损失函数权重等超参数开展对比实验。压缩包共含30个文件,以Python脚本、PNG样例图、Markdown说明及License文件为主,整体约16.33MB,目录结构清晰,适合直接运行调试和二次开发。目前已有671人学习下载,代码可直接用于科研复现、课程设计或工程落地前的算法验证,是入门与进阶GAN图像增强实践的不错参考。
1. 超分模型遇到去噪任务时,SRGAN 到底在解决什么
SRGAN 让行业记住的不是 PSNR 数字,而是“超分结果更接近人眼对照片的理解”。它的核心改动在损失:用 VGG 特征空间的内容损失替代像素级均方误差,再用生成对抗网络去逼真实高频分布。问题恰恰藏在“高频分布”里——噪声也是高频。判别器在追逐纹理时,常常把传感器噪声、JPEG 块噪当成锐利细节保留甚至放大。因此“SRGAN-master 里的 srgan 算法 + gan 去噪”放在一起,真正的挑战不是跑通超分,而是给对抗训练加一道约束:只锐化结构,不锐化噪声。这篇文章按损失函数、工程骨架、训练调度、批量推理、验证指标的顺序,把这条链路上的翻车点一次讲清楚,适合做图像超分、视频增强和相关方向的人阅读。
2. SRGAN 损失函数拆解:对抗损失和内容损失如何决定噪声去留
2.1 内容损失放在 VGG 层,而不是像素层
纯像素 MSE 的问题在于它对每个像素独立惩罚,模型发现“取平均能降低整体风险”,于是输出向模糊方向收缩。对去噪来说,像素 MSE 确实能抹掉不少噪声,但同时也把发丝、毛孔、砖缝这类高频细节抹掉了。SRGAN 的做法是把相似性计算放到 VGG19 的特征空间,在 relu5_4 之后取特征图再做均方误差,模型优化的就是“结构像不像”,而不是“像素亮不亮”。
特征层的选择直接影响噪声行为。relu5_4 感知语义结构,对局部噪声相对宽容;relu2_2 或 relu4_4 更贴近边缘细节,但如果训练目标里带一点噪,它们也会把噪声模式学进损失里。所以做去噪时我会同时挂两个层:relu5_4 权重 1.0 管整体,relu4_4 权重 0.05~0.1 管局部边缘,等发现输出出现椒盐状块时再把 relu4_4 权重降到 0。
import torch import torch.nn as nn import torchvision.models as models class VGGFeatureExtractor(nn.Module): def __init__(self, feature_layer: str = "relu5_4"): super().__init__() # PyTorch 新版建议用 weights=;旧版传 pretrained=True 同样可用 vgg = models.vgg19(weights=models.VGG19_Weights.IMAGENET1K_V1).features layer_index = {"relu4_4": 28, "relu5_4": 38} # 截断到目标 ReLU 输出层,后面的层不参与计算 self.features = vgg[: layer_index[feature_layer]].eval() for p in self.features.parameters(): p.requires_grad_(False) def forward(self, x): return self.features(x) class PerceptualLoss(nn.Module): def __init__(self, weights={"relu5_4": 1.0, "relu4_4": 0.05}): super().__init__() self.extractors = {name: VGGFeatureExtractor(name) for name in weights} self.weights = weights def forward(self, fake, real): total = 0.0 for name, w in self.weights.items(): total += w * nn.functional.mse_loss( self.extractors[name](fake), self.extractors[name](real) ) return total说明:预训练 VGG 只做特征提取,不参与梯度更新,所以requires_grad_(False)是必须的;如果让 VGG 参与反向,会拖慢训练,还可能让生成器通过“破坏 VGG 特征”来钻空子。代码里的权重 dict 直接控制各层对噪声的敏感度,训练时改成{"relu5_4": 1.0}就完全忽略边缘层,适合噪声偏重的场景。
2.2 对抗损失在噪声上的双重身份
对抗损失用的是二分类交叉熵。判别器把真实高清图判定为 1,把生成器的输出判定为 0;生成器则试图让判别器对输出给出偏 1 的响应。SRGAN 常用非饱和形式:G 的对抗损失写成-E[log D(G(x))],梯度在 D 输出接近 0 时依然明显,不容易出现生成器早期梯度消失的问题。
这里有一个去噪任务绕不开的坑:判别器学习的是“真实图像长什么样”,如果训练用的高清图本身带有轻微噪声,判别器就会把这种噪声当作真实分布的一部分,等于替噪声背书。于是生成器被引导着制造和训练数据噪声相关的抖动。处理办法不是削弱判别器,而是保证训练数据的干净程度优先于数量。用干净的 4K 素材下采样当 HR,比用网络随便抓的图更可靠。
| 损失路径 | 常见权重 | 对噪声的典型作用 | 输出出问题时往哪调 |
|---|---|---|---|
| relu5_4 内容损失 | 1.0 | 保持整体结构,噪声可能被忽略 | 大体不动 |
| relu4_4 内容损失 | 0.05~0.1 | 增强局部边缘,也会带噪 | 出现雪花点就调 0 |
| 对抗损失 | 1e-3 | 补高频细节,可能放大噪声 | 伪影块明显时降一半 |
| 像素 L1 | 0.01 | 校正色调偏移,顺带压噪声 | 画面发灰时上调到 0.03 |
这个表的量级来自 SRGAN 的原始设定和其他复现中比较保守的组合。对抗损失的 1e-3 不是随便取的:把权重调到 1e-2 以上,内容损失压不住 GAN 的自由度,输出会出现反复震荡的纹理;调到 1e-4 以下,对抗损失等于没接,生成器退化成纯人工网络,锐度会明显下降。
3. SRGAN-master 工程骨架:文件划分、生成器结构与训练主循环
3.1 别照抄压缩包目录,按五个边界重建
见过不少 SRGAN-master 风格的工程,最常见的形态是一个 train.py 塞下模型定义、数据加载和训练循环,跑通没问题,一旦要换噪声模型就到处找参数。我一般拿到这类工程后会先按五块重排:模型定义、损失函数、数据管线、训练循环、推理入口。目录不需要花哨,边界清楚最重要。
srgan_denoise/ ├── model.py # 生成器、判别器、残差块 ├── loss.py # 感知损失、对抗损失、权重组合 ├── dataset.py # HR 下采样 + 噪声注入的配对数据 ├── train.py # 预训练与对抗训练的调度 ├── infer.py # 批量推理、分块拼接、残差噪声估计 └── config.yaml # 学习率、权重、噪声参数全部集中dataset.py 只负责输出两个东西:干净的 HR 目标,以及由它退化得到的带噪 LR 输入。退化过程必须对训练和验证用同一套随机种子,否则验证曲线会毛糙得没法判断。loss.py 里不出现模型定义,model.py 里不做数据处理。
3.2 生成器用残差块堆叠,判别器输出 logit
超分生成器的主干是一串残差块,后面接亚像素卷积完成上采样。显存没到 16G 的情况下,残差块数量我一般从 16 起步,超过 24 训练周期明显变长,收益也有限。
import torch import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, n_feats=64): super().__init__() self.conv1 = nn.Conv2d(n_feats, n_feats, 3, 1, 1) self.bn1 = nn.BatchNorm2d(n_feats) self.relu = nn.PReLU(num_parameters=1, init=0.25) self.conv2 = nn.Conv2d(n_feats, n_feats, 3, 1, 1) self.bn2 = nn.BatchNorm2d(n_feats) def forward(self, x): out = self.bn2(self.conv2(self.relu(self.bn1(self.conv1(x))))) return self.relu(x + out) def make_generator(upscale=2, n_blocks=16): body = [nn.Conv2d(3, 64, 9, 1, 4), nn.PReLU()] body += [ResidualBlock(64) for _ in range(n_blocks)] body += [nn.Conv2d(64, 64, 3, 1, 1), nn.BatchNorm2d(64)] body += [nn.Conv2d(64, 64 * upscale ** 2, 3, 1, 1), nn.PixelShuffle(upscale)] body += [nn.Conv2d(64, 3, 9, 1, 4), nn.Tanh()] return nn.Sequential(*body)PixelShuffle 要求输入通道等于输出通道乘以放大倍数的平方,所以二维卷积输出通道写成64 * upscale ** 2,拆分后它会重组出 64 通道的高分辨率特征图。最后一个卷积输出 3 通道,Tanh 把结果压到 [-1, 1],对应数据集的归一化范围。忘了这个约束,判别器会一直无法收敛。
判别器我习惯输出一个 logit,而不是经过 sigmoid 的 0 到 1 概率,配合 BCEWithLogitsLoss,数值上比单独跑 sigmoid 再算 BCE 更可控,推理时再用 sigmoid 取概率。
class Discriminator(nn.Module): def __init__(self, in_ch=3): super().__init__() self.features = nn.Sequential( nn.Conv2d(in_ch, 64, 3, 1, 1), nn.LeakyReLU(0.2, inplace=True), self._block(64, 64, stride=2), self._block(64, 128, stride=1), self._block(128, 128, stride=2), self._block(128, 256, stride=1), self._block(256, 256, stride=2), ) self.head = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(256, 1), ) @staticmethod def _block(in_c, out_c, stride): return nn.Sequential( nn.Conv2d(in_c, out_c, 3, stride, 1), nn.BatchNorm2d(out_c), nn.LeakyReLU(0.2, inplace=True), ) def forward(self, x): return self.head(self.features(x))3.3 训练循环里的更新顺序和 detach
对抗训练顺序有两种容易出错的位置:更新判别器时生成器输出要 detach,否则生成器的梯度会穿过判别器反向传播;更新生成器时又不能 detach,要让对抗梯度流回生成器。下面的循环是大多数 SRGAN 实现的标准写法。
# generator/discriminator 已在前面实例化并送到 GPU bce = nn.BCEWithLogitsLoss() for lr_img, hr_img in train_loader: lr_img, hr_img = lr_img.cuda(), hr_img.cuda() # 一、训练判别器,让真实图判真、生成图判假 fake = generator(lr_img) real_pred = discriminator(hr_img) fake_pred = discriminator(fake.detach()) d_loss = 0.5 * (bce(real_pred, torch.ones_like(real_pred)) + bce(fake_pred, torch.zeros_like(fake_pred))) opt_d.zero_grad() d_loss.backward() opt_d.step() # 二、训练生成器,内容损失保证结构,对抗损失补高频 fake = generator(lr_img) g_adv = -torch.mean(discriminator(fake)) g_perc = perceptual_loss(fake, hr_img) g_loss = g_perc + 1e-3 * g_adv opt_g.zero_grad() g_loss.backward() opt_g.step()d_loss取 0.5 平均是为了让数值范围和 batch 大小解耦,真正要单独调的是d_step这类更新频率。对抗损失写成负均值而不是 BCE,目的是避免 D 已经给出接近 0 的输出时梯度消失。如果你在训练日志里看到判别器对生成图的 logit 输出长期低于 -8,多半是判别器一骑绝尘了,可以先让判别器每两个批次更新一次。
4. 把 SRGAN 调成去噪模型:四个训练调度细节
4.1 先做生成器预训练,再打开对抗损失
从零直接跑完整 SRGAN,最容易看到的现象是:前几百个批次输出全是偏色的噪点,判别器率先学会区分,生成器却什么都没掌握。常见做法是先只用 L1 损失单独训练生成器 50 到 100 个 epoch,让它先学会“大概长得对”,再引入感知损失和对抗损失。切到对抗阶段时学习率通常要回退一半,因为判别器刚起步,它的梯度在数值上很大,保持预训练的学习率容易把生成器一步推坏。
for epoch in range(total_epochs): for lr_img, hr_img in train_loader: if epoch < 80: loss = nn.functional.l1_loss(generator(lr_img), hr_img) else: # 切到对抗阶段:学习率回退一半,损失换成感知+对抗 fake = generator(lr_img) loss = (perceptual_loss(fake, hr_img) + 1e-3 * (-torch.mean(discriminator(fake)))) opt_g.zero_grad() loss.backward() opt_g.step()epoch=80 这个值不是固定规则。预训练阶段的 L1 loss 曲线大概在什么时候进入平台,就在平台之后切入对抗,一般比固定 epoch 可靠。如果 L1 预训练只跑了 20 个 epoch 就切对抗,生成器连基本轮廓都没成型,后续反复震荡的时间往往更长。
4.2 噪声注入不能只加高斯白噪声
超分训练数据是 HR 下采样成 LR,这一步本身只是模糊退化。如果训练时不做噪声注入,生成器只学了上采样,遇到带有传感器噪声的真实图片时,它会把这些噪声当成纹理锐化。去噪项目的数据管线应该把退化写成一个可组合的流程:下采样、加噪声、压缩,按顺序执行。
| 噪声类型 | 模拟方法 | 对应现实场景 |
|---|---|---|
| 高斯白噪 | hr + sigma * torch.randn_like(hr),sigma 取 5/255~25/255 | 传感器读取噪声 |
| 泊松噪声 | 在光子计数域用np.random.poisson(x)采样,再归一化 | 低照度成像 |
| JPEG 压缩 | 保存为 q=60~85 再读回 | 传输和存档 |
| 混合退化 | 高斯叠加 JPEG | 手机夜景直出 |
提示:噪声强度要按应用现场标定。我一般会在现场拍几十张暗场帧,用它们的标准差作为训练时高斯噪声的 sigma,而不是从论文里抄一个固定值。现场噪声是固定模式和随机噪声混合的,纯高斯训练出来的模型在条纹型固定噪声上会失效。
对于训练速度,dataset.py 里离线把退化结果存成预处理后的数组,比在线跑 JPEG 编码快很多。缺点是换了噪声参数要重新生成一次数据集,但换来的是训练过程可复现。
4.3 生成器和判别器学习率分开设置
生成器和判别器共用一套学习率是另一种常见问题。SRGAN 的对抗项权重很小,如果判别器学得太快,它把生成器所有输出都判为假,对抗梯度的方向变成随机噪声;如果判别器学得太慢,生成器会反复在一两个伪纹理模式上循环,D 给不出有区分度的惩罚。我的起步值:生成器 2e-4,判别器 1e-4,优化器都选 Adam,beta=(0.9, 0.999)。
opt_g = torch.optim.Adam(generator.parameters(), lr=2e-4) opt_d = torch.optim.Adam(discriminator.parameters(), lr=1e-4) # 判别器训练节奏:每 N 步更新一次 if step % d_step == 0: d_loss.backward() opt_d.step() opt_d.zero_grad()d_step 怎么判断:如果判别器对生成图输出的 logit 长期低于 -8,表示 D 太强势,把 d_step 从 1 改成 2;如果对抗项在生成器总损失里占比一直下降,再把 d_step 调回 1。另一个常被忽略的参数是残差块的 PReLU 初始化,init=0.25 相对保守,从源码里默认 0.0 改到 0.25 后,前几百个 batch 的梯度幅值明显更收敛。
4.4 用验证集 LPIPS 做早停,而不是只盯 PSNR
SRGAN 训练过程中很容易出现 PSNR 向上、人眼观感向下的阶段。原因是对抗项开始生成“看着真实但不属于原图”的高频细节,逐像素误差反而小了,感知质量却在下降。只保存最后一个 checkpoint 的做法会把这个阶段的问题保留下来。我习惯每 5 个 epoch 在固定验证集上算一次 LPIPS 和 PSNR,以 LPIPS 为主决定保存点。
best_lpips = 1e9 if lpips_val < best_lpips: best_lpips = lpips_val torch.save(generator.state_dict(), "srgan_denoise_best.pt")保存条件不要写成“LPIPS 或 PSNR 任一上升”,那样会退化成保存所有 checkpoint。LPIPS 用 AlexNet 或 VGG 的特征距离都可以,数据少时用 AlexNet 版本更不容易被噪声带偏。训练集本身如果有噪声,LPIPS 的参考图必须是干净的 HR,否则它会把噪声也算进“相似”。
5. 落地推理:批量处理、分块拼接和帧间闪烁
5.1 用命令行参数把推理和训练分开
训练和推理不要共用同一个入口。infer.py 只做四件事:读 checkpoint、读输入目录、分块推理、写输出目录。模型结构和训练时保持一致,checkpoint 里只存 state_dict,不存整个模型,这样更换 GPU 和 PyTorch 版本时不会因为类定义位置不同而加载失败。
python infer.py \ --input ./noisy_rgb \ --output ./restored \ --checkpoint srgan_denoise_best.pt \ --upscale 2 \ --tile 256 \ --overlap 32 \ --device cuda:0tile 和 overlap 是超分推理的两个关键参数。tile 决定单个 patch 的边长,按显存选;overlap 是相邻 patch 的重叠宽度,一般取 tile 的八分之一到四分之一。overlap 太小,拼接处会有明显亮暗交替;overlap 太大,计算量成倍增加。
5.2 大图分块推理与重叠区融合
超过 4K 的航拍、病理切片直接整图进生成器,BatchNorm 在推理时用的是训练集的统计量,patch 之间没有上下文,拼接缝几乎必现。除 overlap 外,我会再给每个 patch 乘一个余弦窗,让重叠区域按权重渐变叠加,而不是直接覆盖。
import torch import numpy as np def tile_inference(model, img, tile=256, overlap=32): h, w = img.shape[-2:] out = torch.zeros_like(img) weight = torch.zeros_like(img) step = tile - overlap for y in range(0, h - tile + 1, step): for x in range(0, w - tile + 1, step): patch = img[:, :, y:y + tile, x:x + tile] with torch.no_grad(): restored = model(patch) # 余弦窗降低 patch 边缘权重,重叠区按权重相加 win = torch.hamming_window(tile, periodic=False).to(img.device) win = torch.outer(win, win) out[:, :, y:y + tile, x:x + tile] += restored * win weight[:, :, y:y + tile, x:x + tile] += win return out / (weight + 1e-8)hamming_window生成一维窗,再用outer变成二维;如果不做窗,只在 overlap 区域平均,拼缝附近仍然会有轻微亮暗跳变。尾行不足 tile 的部分会被循环漏掉,处理时要在左边补零或把最后一块右对齐到图像边缘,我常用后一种,避免边缘区域被整体缩小。
5.3 逐帧超分的闪烁问题
视频逐帧跑同一张 SRGAN,单帧看效果没问题,放到时间线上就出现纹理抖动。原因是 GAN 每一帧都在独立生成高频成分,两帧之间的纹理细节不连续,视觉上表现为高速闪烁。如果推理管线需要输出视频,我一般加一个时间维度的后处理:当前帧和前后两帧先做运动补偿对齐,再对三帧取中值。这个方案不改模型,只增加推理成本。
更彻底的做法是把连续三帧叠成 9 通道输入生成器,让模型在训练时就学会利用时间信息,但这需要重新训练,且运动大时的鬼影问题要单独处理。没有重训时间预算时,帧间中值后处理就已经能消除大部分高频闪烁,代价是画面稍微偏柔。
5.4 和小波阈值去噪的分工
SRGAN 擅长保留纹理,但它不会像小波阈值那样精确地区分噪声和信号。传统小波阈值去噪在速度上有绝对优势,CPU 上处理 1080p 也就几十毫秒,SRGAN 要上 GPU;反过来小波软阈值处理大纹理边缘时容易发虚。一个折中的落地方式是:SRGAN 负责超分和主体去噪,剩余轻微噪声再用 3 层 sym8 小波软阈值收尾。
| 方法 | 速度 | 边缘保持 | 噪声模型假设 | 主要风险 |
|---|---|---|---|---|
| 小波阈值去噪 | 快,多平台可跑 | 较弱 | 高斯白噪声为主 | 纹理细节一起抹掉 |
| SRGAN 去噪 | 慢,需要 GPU | 强 | 数据驱动,可覆盖混合噪声 | 伪纹理、训练难收敛、纹理假细节 |
使用 SRGAN 输出后再跑小波阈值,要注意分解层数不超过 3 层,阈值用 Donoho 公式的 σ√(2 log N) 估算,并通过 λ 乘子把强度控制在 0.2 到 0.4 之间。阈值过大时,SRGAN 重建出来的主体纹理会被二次平滑,损失的主要是真细节而不是噪声。
6. 一个替代肉眼的验证方法:平坦区噪声残差估计
GAN 生成的伪纹理很难用肉眼在单帧里发现,也难在训练日志里量化。最实用的一个参考图无关指标是平坦区残余噪声:图像里亮度梯度很小的区域本来应该接近纯色,模型如果还在制造噪声,这些区域的局部标准差就是直接的证据。
原理不复杂。真实世界的平坦区域(天空、墙面、光滑金属面)在传感器响应上会有轻微噪声,但去噪模型的正确输出应该是近似平坦的。计算每个小窗口的梯度幅值,选出梯度最小的那些窗口作为平坦区,然后统计这些窗口的灰度标准差。标准差高于阈值说明噪声仍被保留;低于阈值且验证集 LPIPS 没有同步变好,就得怀疑模型把纹理也抹平了。
import numpy as np def flat_region_noise(model_output, flat_threshold=0.01): img = model_output.squeeze() if img.ndim == 3: gray = img.mean(dim=0).cpu().numpy() else: gray = img.cpu().numpy() gy, gx = np.gradient(gray) edge_mag = np.hypot(gx, gy) flat_mask = edge_mag < flat_threshold flat_count = flat_mask.mean() residual_std = gray[flat_mask].std() return residual_std, flat_countflat_threshold 要按输出图像的范围调整:如果输入输出归一化在 [-1, 1],取 0.01;如果是 0 到 255 的 8 位图,要放大到 2.0 附近。对归一化到 [-1, 1] 的模型,残差标准差在 0.004 到 0.008 之间通常表示噪声压得不错;超过 0.02 说明噪声仍在;低于 0.002 且验证集 LPIPS 没有同步变好,就得怀疑模型把纹理也抹平了。这个代码放进上一章 infer.py 的--eval-residual参数里,批量推理完成后自动打印每张图的平坦区噪声,比盯着 PSNR 曲线判断要直接得多。
本文还有配套的精品资源,点击获取