ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

SRGAN超分辨率重建:从对抗训练原理到PyTorch工程落地

SRGAN超分辨率重建:从对抗训练原理到PyTorch工程落地 简介图像超分辨率重建是计算机视觉中一项基础而实用的技术旨在从低分辨率输入中恢复出清晰、细腻的高分辨率图像。传统插值算法与早期卷积网络虽然能提升分辨率却难以生成真实细腻的纹理细节。生成对抗网络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, channels64): super().__init__() self.conv1 nn.Conv2d(channels, channels, 3, 1, 1) self.bn1 nn.BatchNorm2d(channels) self.relu nn.ReLU(inplaceTrue) 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_blocks16): super().__init__() self.conv1 nn.Conv2d(3, 64, 3, 1, 1) self.relu nn.ReLU(inplaceTrue) 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(inplaceTrue) ) self.up2 nn.Sequential( nn.Conv2d(64, 256, 3, 1, 1), nn.PixelShuffle(2), nn.ReLU(inplaceTrue) ) 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_channels3): super().__init__() self.features nn.Sequential( nn.Conv2d(in_channels, 64, 3, 1, 1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, 64, 3, 2, 1), nn.BatchNorm2d(64), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, 128, 3, 1, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(128, 128, 3, 2, 1), nn.BatchNorm2d(128), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(128, 256, 3, 1, 1), nn.BatchNorm2d(256), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(256, 256, 3, 2, 1), nn.BatchNorm2d(256), nn.LeakyReLU(0.2, inplaceTrue), ) 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 也常见影响不大二是最后的分类层不要接 SigmoidBCEWithLogitsLoss 内部会处理数值稳定性直接输出 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(pretrainedTrue).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_advTrue): 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_advFalse之后改成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 r2。还不行就在判别器里也加一层高斯模糊预处理强迫生成器输出干净的高频信号。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×96batch 从 16 降到 8。也可以用torch.cuda.amp混合精度显存节省约 40%速度还更快。5.5 加载预训练权重维度不匹配现象加载 VGG19 特征提取层时报 size mismatch集中在第一层卷积。原因你的训练输入是单通道灰度图但 VGG19 的预训练权重是针对 3 通道 RGB 的或者你用了vgg19(pretrainedFalse)然后手动 load权重结构对不上。解决输入图统一转成 3 通道感知损失网络用pretrainedTrue并冻结参数再在加载后把第一层卷积改成nn.Conv2d(3, 64, 3, 1, 1)复制 RGB 通道权重做初始化或者直接用weights_onlyTrue加载官方 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 图足够实时预览场景用。最后说一句训练习惯每次跑实验都把损失曲线、验证集指标、生成的样例图存一个带时间戳的目录特别是对抗训练阶段十次里有一两次会出现前期正常中期崩坏的情况没有历史记录很难定位是哪一步权重变化导致的。这个习惯帮我少走了很多弯路希望帮到你。本文还有配套的精品资源点击获取
返回列表