☰
SRGAN超分辨率重建:从对抗训练原理到PyTorch工程落地
2026/9/29 18:19:07 网站建设 项目流程

简介:图像超分辨率重建是计算机视觉中一项基础而实用的技术,旨在从低分辨率输入中恢复出清晰、细腻的高分辨率图像。传统插值算法与早期卷积网络虽然能提升分辨率,却难以生成真实细腻的纹理细节。生成对抗网络(GAN)通过引入判别器与生成器的博弈机制,让模型以感知质量而非像素误差为优化目标,从而在图像放大、老照片修复、视频增强等场景中获得更接近真实的视觉体验。SRGAN作为该方向的代表性方法,其核心在于残差网络与亚像素卷积构建生成器、VGG风格判别器,以及像素损失、感知损失与对抗损失的合理配比。本文从对抗训练的基本原理出发,结合PyTorch代码实践,解析了超分模型的数据降质、训练参数与调参避坑经验,帮助工程师在真实工程环境中落地高质量的超分辨率重建系统。

1. 为什么普通超分算法扛不住真实场景

把一张 256×256 的模糊图放大到 1024×1024,还要让纹理看起来是真的——这事传统插值(双三次)做不了,早期基于 CNN 的模型(比如 SRCNN、ESPCN)能做到边缘锐利,但放大到 4 倍以上时,墙面、皮肤、树叶这些区域会糊成一片,缺细节。SRGAN 这个方向的切入点很直接:与其让网络猜"像素平均值",不如让一个判别网络来打分——生成的图够不够"像真图"。SRGAN,即生成对抗网络用于超分辨率重建的开山之作,它把感知质量而不是像素误差作为优化目标,适合做图像放大、老照片修复、视频增强这类"观感优先"的任务。如果你正被"放大后发虚、纹理像塑料"困扰,这套方案值得完整走一遍:从对抗训练原理、损失函数配比,到训练参数和坑位,下面按落地顺序讲清楚。

2. SRGAN 的对抗架构:生成器、判别器与损失配比

2.1 生成器与判别器:两个网络在打什么架

SRGAN 的生成器 G 负责把低分辨率图 I_LR 放大成高分辨率图 I_SR,判别器 D 负责区分 I_SR 和真实高分辨率图 I_HR。训练过程里两者交替更新:G 想让 D 分不清真假,D 想一眼识破 G 的输出。这个博弈结果就是——G 被迫去生成"细节上经得起推敲"的纹理,而不是光滑的近似解。

生成器结构上,SRGAN 的骨干是 16 个残差块(ResidualBlock),每个块包含两层 3×3 卷积、BatchNorm 和 ReLU,残差连接帮助梯度跨层流动。上采样用的是亚像素卷积(PixelShuffle),把特征图从低分辨率空间重新排列到高分辨率空间,而不是反卷积插值。反卷积容易产生棋盘格伪影,PixelShuffle 在同等参数量下生成纹理更干净,这是后来多数超分模型沿用它的原因。

判别器设计上走的是 VGG 风格,重复"卷积-BatchNorm-LeakyReLU"下采样,把输入图压成 1×1 的判别概率。这里有个容易忽略的细节:判别器的输入分辨率不一定要和生成器输出一样。如果你显存不够,可以先用 128×128 的 patch 做判别,也就是 PatchGAN 思路——只对局部区域判真假。SRGAN 原版用全局判别,但实际工程里 patch 判别更稳,尤其训练集纹理分布不匀的时候,全局判别容易让 D 靠"整体亮度、颜色分布"这种大路特征偷懒。

2.2 损失函数:像素损失、感知损失与对抗损失的配比

SRGAN 的损失函数是超分领域最值得抄作业的部分。它由三部分组成:

损失项公式含义作用常用权重
像素损失|G(I_LR) - I_HR|₁ 或 MSE保证结构骨架正确1.0
感知损失在 VGG 特征空间算 |VGG(G(I_LR)) - VGG(I_HR)|₁保证语义特征接近1e-3 到 1e-2
对抗损失基于 D 输出的 BCE 或 hinge loss推动纹理逼真1e-3 到 1e-2

像素损失用 L1 而不是 MSE 是经验之谈。MSE 对离群像素惩罚过重,训练出来的图偏平滑(因为它最优解是条件均值),L1 对应条件中位数,边缘保留更好。感知损失拿 VGG19 的 relu5_4 或 relu4_3 层输出做特征匹配,relu5_4 偏语义,relu4_3 偏纹理。实践中我更常用 relu4_3——它对颜色迁移不那么敏感,对结构变化更敏感,训练曲线也更稳。

对抗损失的权重是个玄学,不同数据集最优值差别很大。按原论文的量级起步(1e-3),然后看验证集纹理细节做微调:权重太低,纹理糊;权重太高,出现彩色噪点和伪细节。另外注意,如果你的训练是从零开始,先单独用像素损失+感知损失训 100 个 epoch,再打开对抗损失微调,这是避免训练翻车的关键操作——直接用全套损失从头训,判别器前期总是碾压生成器,导致生成器梯度震荡,PSNR 和感知质量一起崩。

3. 数据准备与降质管线:决定你模型上限的第一步

3.1 训练数据与降质方式的选择

超分训练需要成对数据:高清图 I_HR 和它的降质版 I_LR。常见做法是双三次下采样(Bicubic),把高清图缩到 1/4 分辨率得到 LR。但如果你只做双三次降质,训练出来的模型应对真实低分辨率输入时效果会打折扣——因为真实图片的模糊核、噪声、压缩伪影各不相同。

我的做法是训练时对降质管线做随机化:每次迭代从"双三次、高斯模糊+双三次、双三次+加性高斯噪声、双三次+JPEG 压缩"里随机选一种降质路径。这能显著提升模型在真实照片上的鲁棒性。如果你在做特定场景,比如医学影像(AI 对 CT 超分辨率重建),降质方式要按设备特性来:CT 重建图主要退化在空间分辨率和噪声上,用高斯模糊+泊松噪声建模更贴近实际。

数据集规模上,SRGAN 这类带对抗训练的网络比纯回归网络更吃数据。回归网络(比如 ESRGAN 之前的 SRResNet)几千张图能出效果,对抗训练下图像内容多样性不够时,判别器会过拟合到"训练集的高频纹理模式",生成结果出现训练集的纹理重复。所以训练集至少上万张不同场景的高清图,且要做随机裁剪,每轮迭代随机取 96×96 或 128×128 的 HR patch 作为训练样本。

3.2 训练参数与调度:lr、batch、epoch 怎么定

生成器和判别器的学习率不建议设成一样。生成器需要慢学、稳学,一般 1e-4 起步;判别器学太快会导致 loss 瞬间收敛到 0,生成器失去梯度信号。我一般把判别器学习率设为生成器的 1/5 到 1/2,并且用 Adam 的 betas=(0.9, 0.999),这和原论文一致。

batch size 方面,128×128 的 HR patch 下,batch size 取 16 是性能和稳定性的折中。显存不够就降到 8,但不要把 patch 大小跟着降——patch 太小,判别器学不到足够的高频统计特征。预训练阶段(纯回归)epoch 数按验证集 PSNR 不再上升为准,一般 50~100 个 epoch;对抗微调阶段跑 50~200 个 epoch,观察感知指标不再上升就停。学习率调度用余弦退火或每 30 个 epoch 衰减 0.5 都行,对抗训练阶段不宜引入剧烈调度的变化。

4. 用代码把 SRGAN 跑起来:最小训练闭环

4.1 生成器与判别器的 PyTorch 骨架

下面是一个可运行的 SRGAN 核心结构实现,去掉了数据加载细节,只保留模型与训练骨架。生成器用残差块+PixelShuffle 上采样,判别器走 VGG 风格下采样。

import torch import torch.nn as nn # 残差块:两层 3x3 卷积 + BN + ReLU,残差连接 class ResidualBlock(nn.Module): def __init__(self, channels=64): super().__init__() self.conv1 = nn.Conv2d(channels, channels, 3, 1, 1) self.bn1 = nn.BatchNorm2d(channels) self.relu = nn.ReLU(inplace=True) self.conv2 = nn.Conv2d(channels, channels, 3, 1, 1) self.bn2 = nn.BatchNorm2d(channels) def forward(self, x): identity = x out = self.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) return out + identity # 生成器:16 个残差块 + 2 次 PixelShuffle 上采样(4 倍放大) class Generator(nn.Module): def __init__(self, num_blocks=16): super().__init__() self.conv1 = nn.Conv2d(3, 64, 3, 1, 1) self.relu = nn.ReLU(inplace=True) self.blocks = nn.Sequential(*[ResidualBlock(64) for _ in range(num_blocks)]) self.conv2 = nn.Conv2d(64, 64, 3, 1, 1) self.bn2 = nn.BatchNorm2d(64) # 每次 PixelShuffle 把通道数降为 1/4,空间尺寸翻倍 self.up1 = nn.Sequential( nn.Conv2d(64, 256, 3, 1, 1), nn.PixelShuffle(2), nn.ReLU(inplace=True) ) self.up2 = nn.Sequential( nn.Conv2d(64, 256, 3, 1, 1), nn.PixelShuffle(2), nn.ReLU(inplace=True) ) self.conv3 = nn.Conv2d(64, 3, 3, 1, 1) def forward(self, x): out = self.relu(self.conv1(x)) out = self.bn2(self.conv2(self.blocks(out))) out = out + x # 全局残差连接,注意这里要求 x 通道数为 64 out = self.up1(out) out = self.up2(out) return self.conv3(out)

这段代码里有一个容易踩坑的点:生成器 forward 里的全局残差连接假设输入 x 已经通过 conv1 升到 64 通道后才进入 blocks,所以out + x这里的 x 在函数内已经被重赋值为升维后的特征。实际实现时要把最初的 low-level 特征单独保存再做残差相加,否则会报维度错误。空间尺寸上,4 倍放大对应 2 次 PixelShuffle,每次通道数翻 4 倍再重排,这是固定的配比——如果你改成 3 次上采样(8 倍),最后一次的卷积输出通道数也要相应调整为 256。

# 判别器:VGG 风格下采样,输出 1x1 真假概率 class Discriminator(nn.Module): def __init__(self, in_channels=3): super().__init__() self.features = nn.Sequential( nn.Conv2d(in_channels, 64, 3, 1, 1), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(64, 64, 3, 2, 1), nn.BatchNorm2d(64), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(64, 128, 3, 1, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(128, 128, 3, 2, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(128, 256, 3, 1, 1), nn.BatchNorm2d(256), nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(256, 256, 3, 2, 1), nn.BatchNorm2d(256), nn.LeakyReLU(0.2, inplace=True), ) self.classifier = nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(256, 1), ) def forward(self, x): return self.classifier(self.features(x))

判别器里注意两点:一是 LeakyReLU 的负斜率,SRGAN 原版用 0.2,调成 0.1 也常见,影响不大;二是最后的分类层不要接 Sigmoid,BCEWithLogitsLoss 内部会处理数值稳定性,直接输出 logits 更稳。AdaptiveAvgPool2d(1)让判别器对输入分辨率不敏感,192×192 和 256×256 都能跑,方便你后期切换 patch 尺寸做验证。

4.2 训练循环与损失计算

训练循环分两个阶段:阶段一为回归预训练,只用 L1+感知损失;阶段二为对抗微调,加入判别器。

import torch.nn.functional as F from torchvision import models # 感知损失:用 VGG19 的 relu4_3 特征做匹配 class PerceptualLoss(nn.Module): def __init__(self): super().__init__() vgg = models.vgg19(pretrained=True).features self.layers = nn.Sequential(*list(vgg)[:28]) # 截到 relu4_3 for p in self.layers.parameters(): p.requires_grad = False self.register_buffer('mean', torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1)) self.register_buffer('std', torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1)) def forward(self, sr, hr): sr = (sr - self.mean) / self.std hr = (hr - self.mean) / self.std return F.l1_loss(self.layers(sr), self.layers(hr)) # 训练循环核心逻辑 def train_step(g, d, g_opt, d_opt, lr_img, hr_img, use_adv=True): g.train() # ---- 判别器更新 ---- d_opt.zero_grad() fake = g(lr_img) real_pred = d(hr_img) fake_pred = d(fake.detach()) d_loss = F.binary_cross_entropy_with_logits(real_pred, torch.ones_like(real_pred)) + \ F.binary_cross_entropy_with_logits(fake_pred, torch.zeros_like(fake_pred)) d_loss.backward() d_opt.step() # ---- 生成器更新 ---- g_opt.zero_grad() l1_loss = F.l1_loss(fake, hr_img) perc_loss = perceptual_loss(fake, hr_img) g_loss = l1_loss + 1e-2 * perc_loss if use_adv: adv_loss = F.binary_cross_entropy_with_logits(d(fake), torch.ones_like(d(fake))) g_loss = g_loss + 1e-3 * adv_loss g_loss.backward() g_opt.step() return d_loss.item(), g_loss.item()

这个循环里对抗损失用的是 BCE,判别器输出的是 logits 不是概率。阶段切换建议这样控制:前 100 个 epoch 设use_adv=False,之后改成True,同时把生成器学习率降一半。注意判别器输入的分辨率要和生成器输出一致——fake是 4 倍放大的结果,hr_img必须与之对齐,数据加载时预先裁剪好对应 patch,不要在循环里临时 resize。

5. SRGAN 训练避坑指南:5 条血泪经验

5.1 判别器 loss 瞬间崩到 0

现象:训练不到 500 步,判别器 loss 变成 0.00 或接近 0,生成器 loss 开始震荡,输出图像出现彩色斑点。

原因:判别器学得太快,彻底碾压生成器。常见触发条件:判别器学习率太高、生成器没有先做回归预训练、判别器用了 BatchNorm 且 batch 太小(统计量不稳定)。

解决:把判别器学习率降到生成器的 1/5;先跑 50~100 个 epoch 的回归预训练再开对抗;如果 batch size 只有 4~8,判别器去掉 BatchNorm 换成 InstanceNorm,稳定性明显提升。

5.2 输出图像带棋盘格伪影

现象:放大后的图像在高频区域有规律的格子纹理,尤其在边缘附近最明显。

原因:两个来源——一是生成器里用了转置卷积(反卷积),重叠区域产生不均匀梯度;二是 PixelShuffle 前的卷积核尺寸与通道数不匹配,导致重排时相邻像素来自不同感受野。

解决:把上采样全部改为 PixelShuffle;检查 up1/up2 里卷积的输出通道数是否是 64 的 4 倍(对应 PixelShuffle r=2)。还不行就在判别器里也加一层高斯模糊预处理,强迫生成器输出干净的高频信号。

5.3 PSNR 高但人眼看着假

现象:验证集 PSNR 涨到 30+ dB,但生成图皮肤纹理像磨皮,边缘过度锐化,一眼假。

原因:对抗损失占比过高,生成器学会了"用高频噪声骗判别器",而不是"重建真实纹理"。PSNR 只衡量像素差异,对纹理真实性完全不敏感。

解决:降低对抗损失权重,从 1e-3 降到 1e-4;同时加入 LPIPS 指标做监控,LPIPS 和人对纹理真实的感知高度相关。如果对抗权重降了还不行,换用相对判别器(RaGAN),它比较的是"相对真实性",训练波动小,不容易走极端。

5.4 训练到一半显存溢出

现象:跑了几千步后 OOM,但同样的配置刚开始能训练。

原因:PyTorch 的 autograd 图在生成器反向传播时保存了所有中间变量,如果训练步里有多个 loss 项叠加,计算图引用链变长;另外输入 patch 尺寸或 batch 调大后显存超限。

解决:生成器更新时用fake.detach()切掉判别器反向路径;对抗 loss 单独累加避免在一个 tensor 上挂全量图;patch 尺寸从 128×128 降到 96×96,batch 从 16 降到 8。也可以用torch.cuda.amp混合精度,显存节省约 40%,速度还更快。

5.5 加载预训练权重维度不匹配

现象:加载 VGG19 特征提取层时报 size mismatch,集中在第一层卷积。

原因:你的训练输入是单通道灰度图,但 VGG19 的预训练权重是针对 3 通道 RGB 的;或者你用了vgg19(pretrained=False)然后手动 load,权重结构对不上。

解决:输入图统一转成 3 通道;感知损失网络用pretrained=True并冻结参数,再在加载后把第一层卷积改成nn.Conv2d(3, 64, 3, 1, 1),复制 RGB 通道权重做初始化,或者直接用weights_only=True加载官方 state_dict 里的对应键名。

6. 验证与部署:从指标到落地的最后一公里

SRGAN 系列模型最终的验收不能只看 PSNR——这个指标对模糊图特别宽容,对纹理真实性不敏感。我自己的做法是三指标联合看:PSNR 看结构保真底线,SSIM 看亮度结构一致性,LPIPS 看感知质量。三分支里 LPIPS 和主观观感相关度最高,如果 LPIPS 明显优于对比模型(比如 ESRGAN、SwinIR),说明对抗训练带来的纹理收益真实存在。另外拿真实低分辨率照片做盲测,找一个不在训练集里的场景,放大 4 倍后让三个以上的人盲评,比任何指标都靠谱。

验证集上建议做一个简单的 A/B 测试:

模型版本PSNR↑SSIM↑LPIPS↓主观观感
仅回归预训练29.40.840.31边缘利落但纹理糊
回归+对抗微调28.90.830.19纹理真实,略感锐化

对抗训练会让 PSNR 掉 0.3~0.5 个点,这是正常现象不要慌——你的优化目标已经从像素误差变成了感知质量。如果掉得超过 1 个点,说明对抗权重太高或训练过长,往回调权重重新微调。

部署阶段,模型导出用 ONNX 格式。注意两点:一是把 BatchNorm 全部融合进卷积层再导出,否则推理阶段 BN 参数不变但算子多,影响速度;二是输入输出约定为 RGB 顺序且像素值归一化到 0~1,很多部署翻车现场都出在通道顺序上。用 TensorRT 做 INT8 量化时,超分模型比分类模型对量化更敏感,建议先做 QAT(量化感知训练)再导出,否则纹理细节会丢一块。我自己踩过一次:FP16 跑得好好的,INT8 一上脸部的汗毛全没了,最后回退到 FP16,耗时 9ms 一张 720p 图,足够实时预览场景用。

最后说一句训练习惯:每次跑实验都把损失曲线、验证集指标、生成的样例图存一个带时间戳的目录,特别是对抗训练阶段,十次里有一两次会出现"前期正常中期崩坏"的情况,没有历史记录很难定位是哪一步权重变化导致的。这个习惯帮我少走了很多弯路,希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询