简介:基于 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, channels=64): 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_size=3、padding=1,保证特征图空间尺寸不变。这里有一个容易漏掉的细节:第二个 BN 后面没有激活函数,这是为了不让非线性破坏残差路径上的恒等映射。有些实现会在第二个 BN 后加 ReLU,效果差异不大,但按原论文的写法,第二个 BN 直接输出。
上采样部分用 PixelShuffle(亚像素卷积)替代早期超分网络常用的转置卷积。转置卷积会在放大后的特征图边缘产生棋盘伪影,PixelShuffle 则通过重新排列通道来放大空间尺寸,棋盘效应明显更轻。以 4 倍放大为例,需要做两次 ×2 的 PixelShuffle,每次先把通道数变为原来的 4 倍,再重排为空间上 2 倍大小的单通道特征。
class Generator(nn.Module): def __init__(self, num_res_blocks=16, base=64, upscale=4): 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,这正是“通道换空间”的核心操作。
生成器各模块的作用可以从下表看更清楚:
| 模块 | 层 | 输出通道 | 作用 |
|---|---|---|---|
| entry | Conv9×9 + PReLU | 64 | 大卷积核捕获较大感受野 |
| body | 16 × ResidualBlock | 64 | 深层特征提取 |
| skip | Conv3×3 + BN | 64 | 长跳连接的转换层 |
| upsample | Conv3×3 + PixelShuffle | 64 | 两次 ×2 放大 |
| tail | Conv9×9 | 3 | 输出 RGB 图像 |
残差块里的 BN 在 batch 较小时波动很大。显存只放得下 batch=4 左右时,可以考虑去掉生成器里的 BN,这是 SRGAN 训练中最常见的改动之一,后文参数部分会专门说。
2.2 判别器:VGG 风格堆叠与 LSGAN 输出设计
判别器的作用是区分真实高清图和生成器输出的超分图。SRGAN 判别器沿用了 VGG 网络风格:一组 3×3 卷积,stride=2 时降低特征图分辨率,配合 LeakyReLU(0.2) 避免负半轴梯度死亡。
class Discriminator(nn.Module): def __init__(self, in_channels=3): super().__init__() def conv_block(i, o, stride=1, bn=True): layers = [nn.Conv2d(i, o, 3, stride, 1)] if bn: layers.append(nn.BatchNorm2d(o)) layers.append(nn.LeakyReLU(0.2, inplace=True)) return nn.Sequential(*layers) self.features = nn.Sequential( conv_block(in_channels, 64, 1, bn=False), 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, inplace=True), 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× 只需要一次 PixelShuffle,4× 需要两次。
3× 是一个麻烦的边界情况:PixelShuffle 只做整数倍通道重排,不能用两次整数倍放大的组合得到 3。常见做法是先上采样到 4× 再中心裁剪到目标尺寸,这会多计算约 30% 的像素;性能敏感场景可以用转置卷积核为 3、stride=3 的替代方案。选型前先确认需求是否支持 4×,多数真实场景里 4× 已经够用。
显存不足时不要改 base_channels,而应该降低 HR patch 尺寸,因为显存占用随 patch 尺寸平方增长。16 残差块的生成器在 128×128 输入下的显存开销大约是 96×96 输入的 1.8 倍,这个增长主要在 body 阶段。
3. 从数据到训练循环:在 PyTorch 里跑通 SRGAN
3.1 环境与数据准备
先把 PyTorch 环境搭起来。下面是最小可运行的 GPU 环境配置,假设已装好 NVIDIA 驱动。用 conda 创建独立环境可以避免污染系统 Python。
conda create -n srgan python=3.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 显存可以跑 batch=16 的 96×96 输入,16GB 显存可以尝试 128×128。如果从官网下载 torch 速度慢,可以把 pip 的 index-url 换成国内镜像,但要注意镜像源只加速 torch 本体,CUDA 依赖库仍会从官方源拉取。
数据集方面,SRGAN 原论文使用 DIV2K 的 800 张训练图,完整下载几个 GB。如果只是复现流程,可以用 COCO val 集或从 OpenImages 抽几百张图。关键是训练图要足够大,至少 256×256,才能裁出带高频细节的 patch。
| 参数 | 推荐值 | 说明 |
|---|---|---|
| HR patch | 96×96 | 训练时随机裁剪的高清块 |
| scale | 4 | 放大倍数,决定 LR 尺寸 |
| batch_size | 16 | 8GB 显存可接受 |
| Adam lr | 1e-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_crop=96, scale=4): 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:top+ih, left:left+iw] lr_size = ih // self.scale lr = cv2.resize(hr, (lr_size, lr_size), interpolation=cv2.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(upscale=4).to(device) D = Discriminator().to(device) opt_G = torch.optim.Adam(G.parameters(), lr=1e-4, betas=(0.9, 0.999)) opt_D = torch.optim.Adam(D.parameters(), lr=1e-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, device=device)) + \ ce_loss(d_fake, torch.zeros(batch, 1, device=device)) 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, device=device)) 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, device="cuda", layer=35): super().__init__() vgg = torchvision.models.vgg19(pretrained=True).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)layer=35 对应 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_adv=0.001,W_mse=1.0 | 复现论文效果 | PSNR 中等,纹理自然 |
| W_adv=0.01,W_mse=0.1 | 追求细节丰富 | LPIPS 更低,可能出现伪影 |
| W_adv=0,W_mse=1.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-4,G 用 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(net='alex') def evaluate(hr, sr): # hr/sr 为 0-255 的 RGB 图像,shape: H x W x 3 psnr = peak_signal_noise_ratio(hr, sr, data_range=255) try: ssim = structural_similarity(hr, sr, channel_axis=2, data_range=255) except TypeError: ssim = structural_similarity(hr, sr, multichannel=True, data_range=255) 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 旧版本用 multichannel=True,新版改成了 channel_axis,代码里用 try 兼容了两者。
定量指标之外,我会加一个下采样一致性测试:选一张不在训练集里的清晰图,先做高斯模糊和 JPEG 压缩模拟退化,再双三次缩小到目标 LR,送入模型。把输出与原始 HR 对比,重点看边缘有没有振铃、平坦区域有没有色斑。正常模型的 PSNR 比训练集指标低 2dB 以内,如果降幅超过 2dB,回查退化模型是否与训练分布一致,或者 patch 多样性是否不足。验证时建议同时跑多张包含文字和人脸的图,这两类内容对伪影最敏感。
本文还有配套的精品资源,点击获取