简介:面向图像超分重建任务,这是一份基于深度学习的SRGAN算法Python实现,适合具备一定深度学习基础、希望复现超分模型或开展算法研究的开发者。压缩包共199个文件、约297.44MB,包含105个jpg与19个bmp图像样本、8个Python源码文件、2个pth权重文件、5个xml与4个json配置,另有avi/mp4演示视频与TensorBoard训练日志,覆盖从训练数据、模型代码到结果验证的完整链路。代码已调试通过,并配有全中文注释,图像样本数据和权重文件可直接用于复现实验,帮助读者快速理解生成对抗网络在图像超分中的实现细节。已有4457人学习下载,适合课程设计、毕业设计及科研入门使用。 第一次把SRGAN训出来的时候,我盯着输出图看了很久。放大后的效果和老式超分算法那种“硬锐化”完全不同,它更像是把缺失的高频细节“编”了出来,皮肤毛孔、布料纹理、建筑边缘,明明低分辨率图里完全没有这些信息,它却给你补得挺像回事。这就是生成对抗网络在图像超分重建领域的典型代表——SRGAN,全称Super-Resolution Generative Adversarial Network,2017年发表在CVPR,从它开始,超分领域正式从“PSNR竞赛”转向“感知质量竞赛”。即使现在ESRGAN、Real-ESRGAN、SwinIR这些新模型层出不穷,SRGAN依然是理解超分原理、入门GAN训练的最佳起点,而且开源资源非常完整,模型代码、训练脚本、数据集准备工具都能找到成套实现。这篇文章就围绕“SRGAN图像超分重建算法的Python实现”展开,把我从选数据集到调损失,再到稳定训练的完整经验写出来,适合正在接触超分、或者想自己跑一个GAN图像重建项目的同学参考。
1. 为什么超分重建绕不开SRGAN:先跳出PSNR陷阱
1.1 像素空间MSE的局限
很多第一次做超分的同学上来就用MSE(均方误差)作为损失函数训练一个深度网络,训练完以后看着PSNR指标确实涨了,但把图放大看,边缘像被水泡过一样发虚。原因不复杂:MSE是逐像素求差的均值,模型在像素空间里找到的最优解,是对所有可能的高频细节取平均。取平均的后果就是细节被抹平,整体看着平滑,人眼却觉得不锐利。
我举个例子。一张脸部的低分辨率图,放大后眼睫毛到底朝哪个方向弯?MSE训练的网络不知道,它只能输出一个“概率上最像”的毛茸茸区域,于是就成了模糊的一团。而SRGAN的思路很直接:既然像素空间的最优不代表人眼感知的最优,那就换一个更能表达“看起来真”的空间来做约束。
1.2 感知损失与对抗损失的配合
SRGAN论文里最关键的一步,是把内容损失从像素空间挪到了VGG19的特征空间。具体做法是把生成的高分辨率图和真实高分辨率图分别送进VGG19,取conv5_4层输出的特征图,再算两边的MSE。这样模型不再逐像素对齐,而是对齐结构、轮廓、语义层面的相似性,给生成器留下了灵活生成纹理的空间。
同时加上生成对抗网络的对抗损失,判别器负责判断输入图像是真实高清图还是生成图,生成器则努力让生成的SR图骗过判别器。两者交替博弈的结果是:生成器既要在感知特征上贴近原图,又要生成足以以假乱真的高频纹理。这个“内容保真+对抗逼真”的设计,就是SRGAN解决模糊感的核心思路,后面ESRGAN、Real-ESRGAN这些变体,基本都是在这个框架上做的修改。
2. 生成器与判别器的PyTorch逐模块实现
2.1 生成器:残差块堆叠是主干
SRGAN生成器的最初输入是一个3通道低分辨率图像,先用一个9x9卷积把通道数扩到64,然后过16个残差块。每个残差块内部是“卷积-BN-PReLU-卷积-BN”,最后把输入直接加回输出,也就是恒等映射。论文作者在实验中发现16个残差块的性价比最好,再加深收益有限,训练成本和显存占用却涨得明显。
残差块的好处是让梯度在深层网络中传递更顺畅,这在我自己复现时体会很深:预训练阶段如果不用残差结构,生成器的loss收敛速度肉眼可见地慢,而且容易在中途震荡。用上残差后,前几个epoch就能看到输出从噪声逐渐变成像样的图像。
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(channels) 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 + identity2.2 上采样:PixelShuffle比转置卷积稳
生成器里最容易被替换也最值得注意的模块是上采样部分。SRGAN论文用的是子像素卷积,在PyTorch里对应nn.PixelShuffle,把低分辨率特征图通过卷积生成r^2倍数量的通道,再重排成分辨率放大r倍的特征图。
用PixelShuffle而不是反卷积(转置卷积),主要原因是反卷积核在重叠区域会产生不均匀的响应,容易导致输出图像出现棋盘格伪影。子像素卷积在多数实验里更稳,参数也更少,因此后续的ESRGAN等模型也都延续了这个设计。
4倍超分对应两个2倍上采样阶段,代码可以这么写:
class GeneratorTail(nn.Module): def __init__(self, channels=64, scale=4): super().__init__() upsample = [] for _ in range(2): # 4x = 2次2x upsample.append(nn.Conv2d(channels, channels * 4, 3, 1, 1)) upsample.append(nn.PixelShuffle(2)) upsample.append(nn.PReLU(channels)) self.upsample = nn.Sequential(*upsample) self.out = nn.Conv2d(channels, 3, 9, 1, 4) def forward(self, x): return self.out(self.upsample(x))这里需要额外注意一个细节:nn.PixelShuffle(2)要求输入通道数是64*4=256,重排后变成64个通道、宽高各乘2。如果你自定义上采样倍数,一定要算清通道数和重排的对应关系,否则会直接报维度错误。
2.3 判别器:不追求花哨
判别器结构相对简单,输入是3通道图像,连续8个卷积块逐步把分辨率减半、通道数翻倍,最后经过全连接层输出一个0到1之间的实数,表示图像为真实高清图的概率。原版判别器在第一个卷积块不加BatchNorm,这是很多复现版本容易忽略的细节。后面几个卷积块在卷积之后接BN再接LeakyReLU(负斜率0.2),下采样用步长为2的卷积完成,没有额外加池化层。
class Discriminator(nn.Module): def __init__(self, in_channels=3): super().__init__() def block(in_c, out_c, stride=1, bn=True): layers = [nn.Conv2d(in_c, out_c, 3, stride, 1)] if bn: layers.append(nn.BatchNorm2d(out_c)) layers.append(nn.LeakyReLU(0.2, inplace=True)) return nn.Sequential(*layers) self.features = nn.Sequential( block(in_channels, 64, bn=False), block(64, 64, 2), block(64, 128), block(128, 128, 2), block(128, 256), block(256, 256, 2), block(256, 512), block(512, 512, 2), ) self.classifier = 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.classifier(self.features(x))注意判别器最后一层没有接Sigmoid,因为我习惯在损失函数里用BCEWithLogitsLoss而不是BCELoss,数值稳定性更好。这也是复现时的一个小优化点。
3. 数据集与降质细节:不是随便找一批图就能训
3.1 高分辨率图源怎么选
SRGAN训练用DIV2K数据集就够了,800张训练图加100张验证图,每张都是接近2K分辨率的高清图。别看只有800张,超分任务里模型学的是从低分辨率到高分辨率的映射规则,而不是图像语义,所以数据量要求没有分类任务那么夸张。
验证和测试阶段常用Set5、Set14、BSD100、Urban100这4个经典基准集。Set5只有5张图,主要是图个方便,肉眼对比效果时用它最省事;Urban100包含大量城市建筑纹理,对超分模型的细节还原能力考验最严,论文里Urban100通常最能拉开模型差距。
3.2 bicubic下采样:antialias参数的坑
低分辨率图像不是随便找的,标准做法是从HR图用双三次插值(bicubic)下采样得到。这里有个非常容易踩的坑:PIL的Image.resize默认antialias=True,但某些版本或某些框架的resize语义不同,直接导致生成的LR和论文里的降质方式不一致,训练出来的模型在别人数据集上效果变差。
我用的稳定做法是:
from PIL import Image hr = Image.open("image.png").convert("RGB") lr = hr.resize((hr.width // 4, hr.height // 4), Image.BICUBIC)这里的hr是原始高清图,lr是真正输入网络的低分辨率图。有些实现会把lr先resize回HR尺寸再做比较,但SRGAN原版是把SR输出直接和HR比较,网络输入就是小尺寸图,所以实际使用时要保持一致。
3.3 训练时切patch与数据增强
显存有限的话,一般不把整张2K图直接送进网络,而是先随机裁剪成固定尺寸的patch。我常用HR patch尺寸384x384,对应LR patch就是96x96,批次大小设16,在单张24G显存的卡上可以稳定训练。如果显存只有16G,建议把patch降到256x256或者批次降到8,再配合梯度累积。
数据增强方面,随机水平翻转、随机90度旋转就够了。像光照变化、色彩抖动这类增强在超分任务里意义不大,反而可能影响模型学习真实的色调映射关系。
4. 损失函数组合与训练节奏:两阶段策略是稳定关键
4.1 预训练阶段
我强烈建议不要上来就直接用GAN损失训练,几乎必崩。标准做法是先只用MSE(或L1)损失训练生成器,把模型先训练到一个“能看”的状态。
预训练阶段的损失就是nn.MSELoss,优化器用Adam,初始学习率2e-4,批大小和patch尺寸按第3章说的来。一般训练20到30个epoch,看验证集PSNR不再明显提高就停。这个阶段的目的是让生成器学会基本的上采样映射,给后续GAN微调一个合理的起点。
4.2 GAN微调阶段:三个损失如何平衡
微调阶段生成器的总损失由三部分组成:内容损失的变体(VGG特征空间的MSE)、对抗损失、以及可选的一个小权重像素损失。SRGAN原文中对抗损失权重为1e-3,感知损失使用VGG54(VGG19的conv5_4层)特征空间的MSE。
def vgg_loss(pred, target, vgg_extractor): pred_feat = vgg_extractor(pred) target_feat = vgg_extractor(target) return nn.functional.mse_loss(pred_feat, target_feat) criterion_adv = nn.BCEWithLogitsLoss() g_loss = vgg_loss(sr, hr) + 1e-3 * criterion_adv(disc(sr), real_labels)判别器损失就是标准的二分类BCE,真实图标签为1,生成图标签为0。
4.3 判别器更新节奏与标签平滑
实际训练中,生成器和判别器强度经常失衡,判别器太强时生成器的梯度会消失,损失不再下降。我处理这个问题有几点心得:
第一,判别器只更新一次,生成器也更新一次,不做多余的额外更新。保持1:1节奏即可,除非发现判别器loss长期接近0,才临时把它换成每2步更新一次。
第二,判别器的标签采用平滑处理,真实图的标签用0.9而不是1,生成图的标签用0.1而不是0,这样给判别器留一点“犯错空间”,有效减少训练震荡。
第三,GAN微调阶段的学习率从2e-4开始,每10个epoch衰减为原来的0.5,观察生成器loss的下降趋势,稳定后可以提前停止。
5. 训练中踩过的坑:伪影、色彩偏移与显存溢出
5.1 棋盘格伪影
早期版本我用转置卷积做上采样,训练出来的图在细密纹理区域明显有格子状条纹,这种就是典型的棋盘格伪影。根源是转置卷积的卷积核重叠区域分布不均匀,网络学到的权重放大了这种不均匀。换成PixelShuffle之后,这个现象基本消失。
如果你已经用了PixelShuffle还是出现格子纹,那就要检查是不是最后输出卷积层的核大小和padding不匹配,导致边界信息处理不当。
5.2 色彩整体偏绿或偏紫
输出图色彩不对,问题往往不在损失函数,而在数据预处理环节。最常见的是训练时图像归一化用了ImageNet的mean/std,但推理时忘记对输出做对应的反归一化;或者在保存图片时浮点数没有clip到[0,255]就直接转uint8,溢出部分会变成奇怪的色斑。
我的最佳实践是:训练和推理共用同一套预处理函数,保存图片时显式做torch.clamp(out, 0, 1),再乘255转uint8。
5.3 显存溢出
显存不足时,优先不是换更大显存的卡,而是调小patch尺寸、调小batch size、打开梯度累积。我遇到过24G显存跑patch尺寸384都溢出的情况,后来定位到是测试阶段把整张2K图送进网络导致的,推理时只要用滑动窗口切块处理,就不存在这个问题。
5.4 PSNR下降了,但看着更真实了
GAN微调开始后,验证集PSNR数字大概率会掉一点,很多人到这里就慌了,以为模型训坏了。这不是BUG,恰恰说明对抗损失开始发挥作用。PSNR衡量的是逐像素一致程度,它天然偏向平滑结果;而人眼更看重的边缘锐利度和纹理真实感,在PSNR里几乎没有体现。
| 指标 | 衡量内容 | GAN微调后的表现 |
|---|---|---|
| PSNR | 逐像素失真 | 可能微降 |
| SSIM | 结构相似度 | 整体稳定,小幅波动 |
| LPIPS | 感知距离 | 更接近真实图时明显下降 |
所以我训练超分模型时一般同时看三个指标:PSNR看基础失真,SSIM看结构相似性,LPIPS看感知距离,后者越低代表人的主观感受越接近真实图。SRGAN这类生成式超分模型,追求的就是LPIPS更好而PSNR略降的取舍。
6. 完整资源说明与后续扩展
6.1 一次说清“完整资源”里该有什么
经常有同学下载到号称“完整资源”的SRGAN代码包,打开一看只有一个模型文件和一个train.py,数据集链接还是失效的。真正能独立跑通的完整资源至少要包含这些部分:模型定义文件、数据集准备与加载脚本、训练脚本、验证与测试脚本、已训练好的预训练权重、README里写清环境依赖和目录结构。
我整理的项目里把数据准备单独拆成了一个脚本,第一次运行时自动检测数据集目录,不存在就提示下载地址。整个目录结构大致是:
srgan/ ├── models.py # 生成器和判别器定义 ├── dataset.py # 数据集加载与降质逻辑 ├── losses.py # 感知损失与对抗损失 ├── train.py # 预训练 + GAN微调入口 ├── test.py # 单图和文件夹推理 ├── evaluate.py # PSNR/SSIM/LPIPS计算 └── weights/ # 预训练权重存放目录这样无论是从零训练还是直接拿权重推理,都能快速上手。
6.2 环境搭建与快速启动
代码基于Python 3.8+、PyTorch 1.10+,CUDA 11.3或更高版本。第一次跑实验时,我建议直接用conda创建虚拟环境:
conda create -n srgan python=3.8 conda activate srgan pip install torch torchvision numpy pillow tqdm tensorboard实测在Windows和Linux下都能稳定运行。用VSCode调试的时候,注意让解释器指向conda环境,避免pip装到了全局环境里。启动训练很简单,预训练阶段直接python train.py --phase pretrain,GAN微调阶段改成python train.py --phase gan,中间会定时保存checkpoint和TensorBoard日志,方便观察loss曲线。
6.3 下一步往哪走
SRGAN本身是理解超分原理的好起点,但生产环境里一般会升级成ESRGAN或Real-ESRGAN,后者在生成器里去掉了BN层、改用残差密集块,并且用更复杂的降质模型模拟真实退化,对老照片、压缩截图的修复效果明显更好。如果你手头有“模糊+噪声+压缩”混合退化的真实图片,直接上Real-ESRGAN会更省事。
如果你对GAN训练感兴趣,SRGAN也是一个非常适合练手的最小完整框架,把生成器、判别器、感知损失三个组件的设计吃透,后面再去看Diffusion模型做超分或者视频超分,理解成本会低很多。
最后说一点个人体会。SRGAN复现的难点从来不在网络结构本身,而在训练策略和数据细节。我第一次训练时吃亏就吃在预处理和损失权重上,模型结构完全一样,结果一个版本输出发灰,一个版本训练不稳定。所以如果你照着这篇文章搭完代码发现效果不理想,优先检查两件事:低分辨率图的生成方式是否和论文一致,GAN微调阶段的损失权重是否真的在1e-3的量级。这两个地方对了,训练就成功了一半。
本文还有配套的精品资源,点击获取