ARTICLE DETAIL

资讯详情

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

用PyTorch实现SRGAN:从残差块到感知损失的超分实战

用PyTorch实现SRGAN:从残差块到感知损失的超分实战 简介基于 CVPR 2017 年论文《使用生成对抗网络实现照片级真实感单图像超分辨率》的 SRGAN PyTorch 实现面向计算机视觉研究者、算法工程师与深度学习入门者适用于图像与视频超分任务的原理论证、效果复现和二次开发。压缩包共 30 个文件、约 16.33MB包含 8 个 Python 脚本、15 张 PNG 示例与基准测试图片以及 README 说明和 LICENSE 开源许可脚本覆盖生成器与判别器结构、损失函数、SSIM 评估、训练与测试流程图片则用于展示不同放大倍率下的重建结果。已有 1029 人学习下载。目录划分清晰数据、训练结果、基准结果、统计信息等模块一目了然便于按需取用。使用者可直接运行训练脚本完成模型训练借助单图测试脚本评估任意图像也能通过基准测试脚本和视频脚本对比 PSNR、SSIM 等指标配套的损失与 SSIM 模块为质量评价和后续优化提供了可扩展框架是理解生成对抗网络在低层视觉任务中应用的合适参考。1. SRGAN 是什么图像超分辨率重建的对抗思路把一张 96×96 的模糊小图放大成 384×384 的高清图还能补出毛发、皮肤纹路这些“本不存在”的细节——这是 SRGAN 在 2017 年给超分辨率重建领域带来的核心改变。它用生成对抗网络让输出不再只追求像素差最小而是骗过判别器让结果在人眼感知上更接近真实高清图。对做图像增强、视频修复、遥感影像处理或老照片翻新的工程师来说SRGAN 是理解“感知质量优先于 PSNR”这一思路的起点。下面直接进入结构实现和数据流程给出可复现的 PyTorch 代码。2. SRGAN 的结构生成器与判别器在 PyTorch 里的落地2.1 生成器16 个残差块加亚像素卷积SRGAN 生成器主体是 16 个残差块每个残差块的内部结构是“卷积-批归一化-PReLU-卷积-批归一化”再与输入相加。选择残差结构的原因是深层网络在超分任务上容易出现梯度消失残差连接让梯度能直接回传到浅层。PReLU 与 ReLU 的区别在于负半轴有可学习的斜率参数实验里它对生成图像的色彩过渡更平稳。import torch import torch.nn as nn 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.prelu nn.PReLU() self.conv2 nn.Conv2d(channels, channels, 3, 1, 1) self.bn2 nn.BatchNorm2d(channels) def forward(self, x): identity x out self.prelu(self.bn1(self.conv1(x))) out self.bn2(self.conv2(out)) return out identity这个残差块的 kernel_size3、padding1保证特征图空间尺寸不变。这里有一个容易漏掉的细节第二个 BN 后面没有激活函数这是为了不让非线性破坏残差路径上的恒等映射。有些实现会在第二个 BN 后加 ReLU效果差异不大但按原论文的写法第二个 BN 直接输出。上采样部分用 PixelShuffle亚像素卷积替代早期超分网络常用的转置卷积。转置卷积会在放大后的特征图边缘产生棋盘伪影PixelShuffle 则通过重新排列通道来放大空间尺寸棋盘效应明显更轻。以 4 倍放大为例需要做两次 ×2 的 PixelShuffle每次先把通道数变为原来的 4 倍再重排为空间上 2 倍大小的单通道特征。class Generator(nn.Module): def __init__(self, num_res_blocks16, base64, upscale4): super().__init__() self.entry nn.Sequential( nn.Conv2d(3, base, 9, 1, 4), nn.PReLU() ) blocks [] for _ in range(num_res_blocks): blocks.append(ResidualBlock(base)) self.body nn.Sequential(*blocks) self.skip nn.Sequential( nn.Conv2d(base, base, 3, 1, 1), nn.BatchNorm2d(base) ) convs, shuffles [], [] for _ in range(upscale // 2): convs.append(nn.Conv2d(base, base * 4, 3, 1, 1)) shuffles.append(nn.PixelShuffle(2)) self.upsample nn.ModuleList() for conv, shuffle in zip(convs, shuffles): self.upsample.extend([conv, shuffle, nn.PReLU()]) self.tail nn.Conv2d(base, 3, 9, 1, 4) def forward(self, x): entry self.entry(x) # [B, 64, H, W] body self.body(entry) # 经过16个残差块 skip self.skip(body) out entry skip # 长跳连接 for layer in self.upsample: out layer(out) return self.tail(out)entry 和 skip 相加的时机值得注意网络用长跳连接把浅层特征直接接到深层相当于跨层连接保证放大过程不丢失低频信息。上采样循环里每次卷积把通道数扩到 base×4再经过 PixelShuffle 把最后两个维度各扩到 2 倍、通道数回到 base这正是“通道换空间”的核心操作。生成器各模块的作用可以从下表看更清楚模块层输出通道作用entryConv9×9 PReLU64大卷积核捕获较大感受野body16 × ResidualBlock64深层特征提取skipConv3×3 BN64长跳连接的转换层upsampleConv3×3 PixelShuffle64两次 ×2 放大tailConv9×93输出 RGB 图像残差块里的 BN 在 batch 较小时波动很大。显存只放得下 batch4 左右时可以考虑去掉生成器里的 BN这是 SRGAN 训练中最常见的改动之一后文参数部分会专门说。2.2 判别器VGG 风格堆叠与 LSGAN 输出设计判别器的作用是区分真实高清图和生成器输出的超分图。SRGAN 判别器沿用了 VGG 网络风格一组 3×3 卷积stride2 时降低特征图分辨率配合 LeakyReLU(0.2) 避免负半轴梯度死亡。class Discriminator(nn.Module): def __init__(self, in_channels3): super().__init__() def conv_block(i, o, stride1, bnTrue): layers [nn.Conv2d(i, o, 3, stride, 1)] if bn: layers.append(nn.BatchNorm2d(o)) layers.append(nn.LeakyReLU(0.2, inplaceTrue)) return nn.Sequential(*layers) self.features nn.Sequential( conv_block(in_channels, 64, 1, bnFalse), conv_block(64, 64, 2), conv_block(64, 128, 1), conv_block(128, 128, 2), conv_block(128, 256, 1), conv_block(256, 256, 2), conv_block(256, 512, 1), conv_block(512, 512, 2), ) self.out nn.Sequential( nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(512, 1024), nn.LeakyReLU(0.2, inplaceTrue), nn.Linear(1024, 1), ) def forward(self, x): return self.out(self.features(x))判别器最后一个全连接输出的是 logits而不是经过 Sigmoid 的概率值。这是论文里 LSGAN 的实现方式损失函数直接用最小二乘项而不是二分类交叉熵。LSGAN 的好处是训练更稳定生成器梯度不会因为判别器过于自信而消失。使用 LSGAN 时判别器损失变成 (D(x_real) - 1)² D(x_fake)²生成器损失为 (D(x_fake) - 1)²。换成代码就是取 logits 直接做 MSE。如果仍想用 BCE则需要在判别器最后加 Sigmoid。两种写法在效果上差别不大但 LSGAN 不容易出现判别器 loss 快速归零的问题。2.3 上采样倍数与通道配置超分辨率重建通常指 2×、3×、4× 三种倍数SRGAN 默认 4×。我的做法是给 Generator 传 upscale 参数通过循环次数控制上采样层数量而不是为不同倍数维护三个模型文件。2× 只需要一次 PixelShuffle4× 需要两次。3× 是一个麻烦的边界情况PixelShuffle 只做整数倍通道重排不能用两次整数倍放大的组合得到 3。常见做法是先上采样到 4× 再中心裁剪到目标尺寸这会多计算约 30% 的像素性能敏感场景可以用转置卷积核为 3、stride3 的替代方案。选型前先确认需求是否支持 4×多数真实场景里 4× 已经够用。显存不足时不要改 base_channels而应该降低 HR patch 尺寸因为显存占用随 patch 尺寸平方增长。16 残差块的生成器在 128×128 输入下的显存开销大约是 96×96 输入的 1.8 倍这个增长主要在 body 阶段。3. 从数据到训练循环在 PyTorch 里跑通 SRGAN3.1 环境与数据准备先把 PyTorch 环境搭起来。下面是最小可运行的 GPU 环境配置假设已装好 NVIDIA 驱动。用 conda 创建独立环境可以避免污染系统 Python。conda create -n srgan python3.10 -y conda activate srgan pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 pip install numpy opencv-python pillow tqdm scikit-image lpipsCUDA 12.1 PyTorch 2.x 是当前常见组合。对 SRGAN 来说8GB 显存可以跑 batch16 的 96×96 输入16GB 显存可以尝试 128×128。如果从官网下载 torch 速度慢可以把 pip 的 index-url 换成国内镜像但要注意镜像源只加速 torch 本体CUDA 依赖库仍会从官方源拉取。数据集方面SRGAN 原论文使用 DIV2K 的 800 张训练图完整下载几个 GB。如果只是复现流程可以用 COCO val 集或从 OpenImages 抽几百张图。关键是训练图要足够大至少 256×256才能裁出带高频细节的 patch。参数推荐值说明HR patch96×96训练时随机裁剪的高清块scale4放大倍数决定 LR 尺寸batch_size168GB 显存可接受Adam lr1e-4生成器与判别器初始学习率3.2 数据集类与在线退化策略超分训练需要的样本对是低分辨率图高分辨率图。SRGAN 的数据生成方式是先从高清图里随机裁一块 96×96 的高分辨率 patch再用双三次插值缩到 24×24成为低分辨率输入。这样 LR 和 HR 严格对齐便于计算逐像素损失。from torch.utils.data import Dataset import cv2 import numpy as np class SRDataset(Dataset): def __init__(self, image_paths, hr_crop96, scale4): self.paths image_paths self.hr_crop hr_crop self.scale scale def __len__(self): return len(self.paths) def __getitem__(self, idx): img cv2.imread(self.paths[idx]) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) h, w img.shape[:2] ih iw min(self.hr_crop, h, w) top np.random.randint(0, h - ih 1) left np.random.randint(0, w - iw 1) hr img[top:topih, left:leftiw] lr_size ih // self.scale lr cv2.resize(hr, (lr_size, lr_size), interpolationcv2.INTER_CUBIC) return lr, hr返回的 lr 和 hr 都是 H×W×3 的 numpy 数组。训练循环里再转成 PyTorch Tensor归一化到 [-1, 1]。生成器尾部用 tanh输出范围正好是 [-1, 1]这两者必须匹配否则早期损失会反复震荡。cv2.resize 默认的插值方式可能改变图像颜色范围显式指定 INTER_CUBIC 是为了保证和训练时的退化模型一致。如果只做双三次下采样模型对真实低分辨率图像泛化会稍差更贴近实际的退化流程是先后三次下采样再叠加 JPEG 压缩或模糊核。初学阶段先用纯双三次即可。3.3 训练循环判别器与生成器交替更新SRGAN 的每一步迭代是把真实 HR 和生成器输出的 SR 喂给判别器先计算判别器损失并回传 D 的梯度再固定判别器计算生成器的总损失并回传 G 的梯度。两个网络的优化器交替 step。device torch.device(cuda) G Generator(upscale4).to(device) D Discriminator().to(device) opt_G torch.optim.Adam(G.parameters(), lr1e-4, betas(0.9, 0.999)) opt_D torch.optim.Adam(D.parameters(), lr1e-4, betas(0.9, 0.999)) mse_loss nn.MSELoss() ce_loss nn.BCEWithLogitsLoss() for epoch in range(epochs): for lr_img, hr_img in dataloader: lr_img lr_img.to(device) / 127.5 - 1.0 hr_img hr_img.to(device) / 127.5 - 1.0 batch lr_img.size(0) # 更新判别器真实图为真超分图为假 d_real D(hr_img) d_fake D(G(lr_img).detach()) d_loss ce_loss(d_real, torch.ones(batch, 1, devicedevice)) \ ce_loss(d_fake, torch.zeros(batch, 1, devicedevice)) opt_D.zero_grad() d_loss.backward() opt_D.step() # 更新生成器超分图骗过判别器 sr_img G(lr_img) g_adv ce_loss(D(sr_img), torch.ones(batch, 1, devicedevice)) g_content mse_loss(sr_img, hr_img) * 0.1 # perceptual_loss 类定义见 4.1先占位 g_loss g_content 0.001 * g_adv perceptual_loss(sr_img, hr_img) opt_G.zero_grad() g_loss.backward() opt_G.step()判别器损失由两部分组成真实图判定为真的损失加上超分图判定为假的损失。生成器损失里对抗项的目标是让超分图被判定为真。注意 d_fake 的输入带了 detach()如果不 detach判别器的反向传播会把梯度传到生成器导致一次 backward 同时更新两个网络判别器这一步就白训了。我一般建议前几个 epoch 用较小的生成器学习率或降低对抗损失权重避免生成器过早被 GAN 带偏。对抗项权重的调节是 SRGAN 训练里操作最多的地方具体数值放在下一章。4. 损失函数组合与关键超参数SRGAN 收敛的核心4.1 感知损失用 VGG19 特征代替像素比较像素级 MSE 只会让输出尽量逼近所有候选结果的平均值这正是超分结果偏模糊的原因。感知损失的思路是把 SR 和 HR 都送入预训练 VGG19在某一层取出特征图比较两者特征图的 MSE。这样损失不再逐像素对比而是在语义内容层面比较。import torchvision class PerceptualLoss(nn.Module): def __init__(self, devicecuda, layer35): super().__init__() vgg torchvision.models.vgg19(pretrainedTrue).features.to(device) self.features nn.Sequential(*list(vgg.children())[:layer 1]) for p in self.features.parameters(): p.requires_grad False self.features.eval() def forward(self, sr, hr): sr_feat self.features(sr) hr_feat self.features(hr) return nn.functional.mse_loss(sr_feat, hr_feat)layer35 对应 conv5_4。从 relu1_2 到 relu5_4感受野逐渐变大高层特征更关注整体结构低层特征更关注边缘纹理。SRGAN 论文采用 relu5_4这是“全局一致性”的折中。如果发现生成图高频细节不足可以改为 relu2_2或同时取多层特征加权求和这属于 ESRGAN 提出的感知损失改进方向。感知损失里最容易踩的坑是 VGG 特征提取器处于 train 模式。BatchNorm 层在 train 模式会更新 running stats导致预训练特征分布漂移感知损失的数值随训练逐渐失真。构造函数里用了 .eval()训练过程不要调用 .train()。4.2 三种损失如何配比SRGAN 生成器损失由三部分构成。像素空间损失用 MSE让输出在低层接近 HR感知损失用 VGG 特征 MSE约束语义内容对抗损失让输出骗过判别器补出高频纹理。三者的权重直接决定了训练倾向。权重方案适用阶段观察到的现象W_adv0.001W_mse1.0复现论文效果PSNR 中等纹理自然W_adv0.01W_mse0.1追求细节丰富LPIPS 更低可能出现伪影W_adv0W_mse1.0预训练阶段PSNR 高图像偏软我一般使用两阶段训练第一阶段不接判别器只用像素损失加感知损失训练 30 个 epoch让生成器先达到一个稳定的超分水平第二阶段把对抗损失加上判别器和生成器交替训练 20 个 epoch。这比一开始就对抗训练更可控也解决了“生成器在 GAN 早期就被带偏”的问题。对抗损失权重的调整有一个判断技巧如果生成器输出的纹理多但出现彩色噪点说明对抗项过强如果图像很平滑但判别器 loss 已经接近 0.5说明对抗项失效需要回调学习率而不是继续加大权重。4.3 学习率、批大小与训练步数SRGAN 训练里常见的是 Adam 优化器初始学习率 1e-4前 40 个 epoch 固定之后每 10 个 epoch 衰减为原来的 0.1。batch size 在 96×96 输入下取 16显存不够时优先减 batch 而不是降分辨率过小的 patch 会让 BN 统计不稳定。判别器和生成器的学习率可以不同。常见做法是 D 保持 1e-4G 用 2e-4 甚至 5e-4加速生成器逼近真实分布。但如果 D 更新太快G 的对抗梯度会失去稳定语义表现为 g_loss 几个小时不降反而震荡。这时把 D 的学习率降到 5e-5往往比调 G 更有效。如果训练一段时间后判别器 loss 变成 0说明 D 完全碾压了 G。此时降低 D 的步长或给 D 的输入加少量高斯噪声标准差 0.01。加噪声能让 D 的决策边界平滑一些给 G 留出学习空间。还要注意训练集 patch 的多样性。如果训练图大多是大片天空或墙面判别器学到的特征会偏向平坦区域生成器在纹理复杂的区域表现差。建议按图像的梯度方差筛选裁剪位置确保每个 batch 里同时包含边缘、纹理与平滑区域。5. 验证 SRGAN 效果PSNR、SSIM、LPIPS 与下采样一致性测试只看 PSNR 和 SSIM 不足以反映生成纹理的视觉质量训练 SRGAN 时我至少会同时记录 LPIPS。LPIPS 计算两个图像在预训练网络特征空间的加权距离分数越低说明感知越接近。它更匹配“人眼看着像不像”这个目标。from skimage.metrics import peak_signal_noise_ratio, structural_similarity import lpips import torch lpips_fn lpips.LPIPS(netalex) def evaluate(hr, sr): # hr/sr 为 0-255 的 RGB 图像shape: H x W x 3 psnr peak_signal_noise_ratio(hr, sr, data_range255) try: ssim structural_similarity(hr, sr, channel_axis2, data_range255) except TypeError: ssim structural_similarity(hr, sr, multichannelTrue, data_range255) hr_t torch.from_numpy(hr).permute(2, 0, 1).float().unsqueeze(0) / 127.5 - 1.0 sr_t torch.from_numpy(sr).permute(2, 0, 1).float().unsqueeze(0) / 127.5 - 1.0 l lpips_fn(hr_t, sr_t).item() return psnr, ssim, lPSNR 与 SSIM 都要求两张图尺寸一致且 SR 必须与 HR 严格对齐。如果模型输出尺寸与 HR 差几个像素指标会整体失真先做中心裁剪再计算。scikit-image 旧版本用 multichannelTrue新版改成了 channel_axis代码里用 try 兼容了两者。定量指标之外我会加一个下采样一致性测试选一张不在训练集里的清晰图先做高斯模糊和 JPEG 压缩模拟退化再双三次缩小到目标 LR送入模型。把输出与原始 HR 对比重点看边缘有没有振铃、平坦区域有没有色斑。正常模型的 PSNR 比训练集指标低 2dB 以内如果降幅超过 2dB回查退化模型是否与训练分布一致或者 patch 多样性是否不足。验证时建议同时跑多张包含文字和人脸的图这两类内容对伪影最敏感。本文还有配套的精品资源点击获取
返回列表