PyTorch实战ESRGAN:图像超分从原理到训练推理全流程解析

📅 2026/8/27 6:04:43
PyTorch实战ESRGAN:图像超分从原理到训练推理全流程解析
简介图像超分辨率是计算机视觉中一项极具实用价值的基础技术旨在从低分辨率图像中恢复高频细节与真实纹理。随着生成对抗网络GAN的引入基于深度学习的超分模型在视觉质量上取得了显著突破其中ESRGAN作为GAN-based超分的代表性工作通过改进残差密集连接块RRDB、相对判别器与感知损失设计有效解决了早期模型纹理失真与训练不稳定的问题。本文从图像超分的基本概念出发系统梳理生成对抗网络在图像重建中的应用原理并深入解析ESRGAN在PyTorch框架下的工程实现细节涵盖网络架构设计、损失函数组合、训练策略调优、推理部署以及常见问题排查。无论是深度学习算法工程师还是希望提升图像质量的开发者均可在实际项目中利用PyTorch构建高效的图像超分流程获得清晰自然的重建效果。 图像超分辨率这事玩过一阵子的人都知道看着简单真想做出能用的效果坑多到能写一本书。ESRGAN这个名字搞过超分的基本都绕不开它在2018年提出之后很长一段时间都是各种超分比赛和应用的标配baseline。我这次用PyTorch把ESRGAN完整跑通了一遍从原理到训练再到推理中间踩了不少雷也攒了一些挺实用的经验整理出来分享给正在折腾或者准备折腾图像超分的朋友。这个项目适合谁一类是想入门图像超分辨率、想搞懂GAN-based超分原理的算法工程师和学生另一类是手头有低分辨率图像、想把它变清晰但不想只靠现成在线工具的普通开发者。看完这篇文章你不仅能搞懂ESRGAN的核心思路还能直接拿到一份能跑的PyTorch训练和推理流程。1. 项目整体设计与思路拆解1.1 为什么选ESRGAN而不是其它超分模型图像超分这个方向方案很多。传统的插值算法双三次、 Lanczos速度快但细节是糊的SRCNN、VDSR这类早期深度学习模型能恢复一些纹理但高频细节还是欠缺。后来GAN被引入超分领域SRGAN在2016年提出通过感知损失加对抗训练让生成图像在视觉上更锐利但SRGAN有个明显问题——生成的纹理经常是假纹理不真实而且训练很不稳定。ESRGAN是SRGAN的直接改进版核心变化有三个用RRDBResidual in Residual Dense Block替换原来的残差块引入相对判别器Relativistic GAN以及用激活前的特征计算感知损失。这三个改动直接解决了SRGAN的痛点尤其是RRDB的结构设计让网络容量足够大能学到更丰富的纹理表达。我在选型时也对比过其它方案Real-ESRGAN虽然更贴近真实场景退化但它的训练复杂度高很多而且对硬件资源要求更苛刻SwinIR是Transformer结构的超分模型效果确实好但推理速度慢训练起来也更吃显存。相比之下ESRGAN的性价比很高结构清晰、训练相对稳定、推理速度也够用是个很适合作为第一站完整跑通的模型。1.2 PyTorch做ESRGAN的优势在哪选PyTorch而不是TensorFlow主要看中三点。第一动态计算图让调试友好太多GAN这种生成器和判别器交替训练的模型用TensorFlow的静态图写起来是真的痛苦第二PyTorch生态里超分相关的预训练权重和实现资料最全基础代码参考的价值非常大第三HuggingFace、Ultralytics这些社区的事实标准都是PyTorch后续要接部署或者二次开发生态优势是实打实的。1.3 整体架构规划我设计的项目结构很清晰训练、推理、数据、模型、配置分开方便日常操作和后续扩展esrgan-project/ ├── configs/ # 训练和推理的配置文件yaml/json ├── data/ # 训练数据存放位置 │ ├── train_hr/ # 高分辨率训练图像 │ ├── train_lr/ # 对应的低分辨率图像可实时生成 │ └── val_hr/ # 验证集 ├── models/ # 网络结构定义 │ ├── rrdb.py # RRDB基础块 │ ├── generator.py # 生成器网络 │ └── discriminator.py# 判别器网络 ├── losses/ # 损失函数定义 ├── utils/ # 工具函数数据加载、指标计算等 ├── train.py # 训练主脚本 ├── inference.py # 推理脚本 └── test.py # 测试评估脚本这套结构跟很多开源项目保持一致如果你之前跑过BasicSR这类框架会感觉很熟悉。做项目不要一上来就把结构搞得特别花哨保持简单直观是第一位后面有需要再拆。2. 核心原理拆解ESRGAN到底改了什么2.1 RRDB残差中的密集连接RRDB是ESRGAN生成器的核心构建块。简单理解它是在残差结构里面嵌套了密集连接模块。作者的设计思路是网络的深度和容量直接决定了特征表达能力但简单的堆层数会导致训练困难RRDB通过残差和密集连接让超深网络生成器有23个RRDB块也能稳定训练。一个RRDB块内部包含三个Dense Block每个Dense Block内有4层卷积每层卷积的输入会拼接上之前所有层的输出。这种密集连接让每一层都能直接看到前面所有层提取的特征梯度传递路径短信息流动充分。RRDB还使用了残差缩放residual scaling一般设置为0.2也就是残差路径的输出乘以0.2再和主线相加这个细节对稳定训练很重要。# RRDB中一个Dense Block的PyTorch实现简化版 class DenseBlock(nn.Module): def __init__(self, num_features64, num_growth32): super().__init__() self.conv1 nn.Conv2d(num_features, num_growth, 3, 1, 1) self.conv2 nn.Conv2d(num_features num_growth, num_growth, 3, 1, 1) self.conv3 nn.Conv2d(num_features 2 * num_growth, num_growth, 3, 1, 1) self.conv4 nn.Conv2d(num_features 3 * num_growth, num_growth, 3, 1, 1) self.lrelu nn.LeakyReLU(0.2, inplaceTrue) def forward(self, x): x1 self.lrelu(self.conv1(x)) x2 self.lrelu(self.conv2(torch.cat((x, x1), dim1))) x3 self.lrelu(self.conv3(torch.cat((x, x1, x2), dim1))) x4 self.conv4(torch.cat((x, x1, x2, x3), dim1)) return torch.cat((x, x1, x2, x3, x4), dim1)每经过一个Dense Block通道数都会增长。从初始的64通道经过4层后变成644×32192通道然后再通过一个1×1卷积把通道压缩回64保持整体特征维度一致。这种设计让RRDB块的信息容量远大于普通残差块代价是参数量和计算量增大一个包含23个RRDB块的生成器参数量在16M左右。2.2 相对判别器让对抗训练更有意义SRGAN用的标准判别器判断“输入图像是真的HR还是假的生成的SR”这是一个绝对判断。ESRGAN改成了相对判别器它关注的不再是“这张图单纯真或假”而是“生成的图像是否比真实图像更真”。这种设计的直观理解是绝对判别器只要发现生成图像有瑕疵就能判断为假但有些瑕疵人眼根本注意不到不值得让生成器花大力气去修而相对判别器衡量的是真假图像之间的差异能让生成器把精力放在“比真图更像真图”的方向上——实际上就是让生成图像在真假概率上与真实图像拉开差距这更符合图像超分的任务本质。从数学上看标准GAN的判别器输出σ(C(x))相对判别器的输出是σ(C(x_real) - E[C(x_fake)])。生成器的损失相应变成# 相对判别器相关的对抗损失简化伪代码 def ragan_loss(discriminator, fake, real): # real和fake是判别器输出的logits real_logits discriminator(real) fake_logits discriminator(fake) # 判别器损失让真实图像比生成图像更真 d_real_loss torch.mean(nn.functional.binary_cross_entropy_with_logits( real_logits - torch.mean(fake_logits), torch.ones_like(real_logits))) d_fake_loss torch.mean(nn.functional.binary_cross_entropy_with_logits( fake_logits - torch.mean(real_logits), torch.zeros_like(fake_logits))) # 生成器损失让生成图像尽量接近真实图像的判别概率 g_loss torch.mean(nn.functional.binary_cross_entropy_with_logits( fake_logits - torch.mean(real_logits), torch.ones_like(fake_logits))) return d_real_loss d_fake_loss, g_loss实际训练中这个改动带来的提升是肉眼可见的。用标准GAN训练ESRGAN生成的图像经常出现局部过曝或纹路不自然的区域换成相对判别器后整体观感自然了很多。2.3 感知损失在激活前还是激活后感知损失是让生成图像在VGG特征空间里接近真实图像而不是在像素空间里接近。ESRGAN有个小改动在VGG网络的激活层之前计算特征距离而不是用激活后的特征。原因是ReLU激活层会把负值截断为0损失函数梯度无法通过激活层的负区域反向传播导致信息丢失。激活前特征保留更完整的纹理和结构信息让梯度更新更平滑。实现上也很简单就是把预先定义的ReLU激活层替换成恒等映射# 使用激活前特征计算感知损失 class PerceptualLoss(nn.Module): def __init__(self, vgg_model): super().__init__() # 取vgg19的features模块并去掉激活层 self.features vgg_model.features[:35] for i in [1, 3, 6, 8, 11, 13, 15, 17, 20, 22, 24, 26, 29, 31, 33]: self.features[i] nn.Identity() # 用恒等映射替换ReLU def forward(self, pred, target): pred_feat self.features(pred) target_feat self.features(target) return nn.functional.l1_loss(pred_feat, target_feat)2.4 损失函数的组合与权重分配ESRGAN的生成器总损失是三个损失的加权和这是整套训练的平衡木损失分量权重作用Perceptual Loss (L1)1.0保持感知结构一致Adversarial Loss0.005让纹理更加真实自然Pixel Loss (L1)0.01保持颜色和亮度准确实际训练中我感受最明显的是对抗损失的权重千万不能给大否则训练很容易震荡图像会变得“油腻”或者出现奇怪的伪影。Perceptual Loss权重最大因为它主导了整体训练方向Pixel Loss控制在0.01左右既保证像素级约束又不过度压制生成器的纹理生成能力。3. 环境搭建与工程准备3.1 PyTorch环境配置实操PyTorch的安装说难不难说简单也容易踩坑尤其是GPU版本和CUDA版本之间的匹配。我的建议是先确定自己的CUDA版本再选择对应的PyTorch版本。查看CUDA版本用命令nvidia-smi注意这里的版本是驱动支持的CUDA版本不是说你机器上必须装了这个版本的CUDA toolkit。PyTorch只会用到驱动不一定需要安装对应版本的CUDA toolkit这是个常见的误区。我自己用的是PyTorch 2.6.x搭配CUDA 12.1的组合安装命令很直接# 用conda创建独立环境避免污染系统Python conda create -n esrgan python3.10 conda activate esrgan # 安装PyTorch用官方源或者国内镜像源 pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121 # 安装其他依赖 pip install numpy opencv-python pillow tqdm tensorboard pyyaml关于PyTorch 2.6.x有一个必须注意的细节从2.6开始torch.load的weights_only参数默认值改成了True。这意味着如果你用旧版本PyTorch训练的模型权重直接在新版本里torch.load会报错需要显式设置weights_onlyFalse。这个坑我踩过一次排查了半天才反应过来是版本兼容问题。3.2 数据处理训练集怎么准备才靠谱ESRGAN训练需要成对的HR高分辨率和LR低分辨率图像。数据集通常使用DIV2K但如果只是先跑通流程用自己的图片也能训练。关键细节是LR图像的生成方式一般用4倍双三次下采样。import cv2 hr_img cv2.imread(hr_path) # 读取HR图像 h, w hr_img.shape[:2] lr_img cv2.resize(hr_img, (w // 4, h // 4), interpolationcv2.INTER_CUBIC) lr_img cv2.resize(lr_img, (w, h), interpolationcv2.INTER_CUBIC)这两步resize的逻辑是先把HR下采样到1/4尺寸再放大回原尺寸得到的就是一张有模糊和细节损失的低质量图。训练时把HR和对应的LR按patch小块输入网络一般用128×128的HR patch和它对应的LR patch。数据增强我用的是随机旋转90度、水平翻转和垂直翻转每次迭代随机组合。这种简单增强能让有限的数据集发挥更大价值实测下来对模型泛化能力有不错的提升。4. 实操过程与核心环节实现4.1 训练流程的完整代码实现训练主循环其实并不复杂核心是GAN的交替更新思路。下面是我整理过的一个可运行版本的核心部分import torch from torch.utils.data import DataLoader from torch.optim import Adam from tqdm import tqdm def train_one_epoch(generator, discriminator, dataloader, opt_g, opt_d, criterion_pixel, criterion_perceptual, device, lambda_adv5e-3, lambda_pixel1e-2): generator.train() discriminator.train() pbar tqdm(dataloader, descTraining) for lr_imgs, hr_imgs in pbar: lr_imgs lr_imgs.to(device) hr_imgs hr_imgs.to(device) # 1. 更新判别器 opt_d.zero_grad() with torch.no_grad(): sr_imgs generator(lr_imgs) real_pred discriminator(hr_imgs) fake_pred discriminator(sr_imgs) d_loss -torch.mean(torch.log(torch.sigmoid(real_pred - torch.mean(fake_pred)) 1e-12) torch.log(1 - torch.sigmoid(fake_pred - torch.mean(real_pred)) 1e-12)) d_loss.backward() opt_d.step() # 2. 更新生成器 opt_g.zero_grad() sr_imgs generator(lr_imgs) fake_pred discriminator(sr_imgs) real_pred discriminator(hr_imgs) adv_loss -torch.mean(torch.log(torch.sigmoid(fake_pred - torch.mean(real_pred)) 1e-12)) pixel_loss criterion_pixel(sr_imgs, hr_imgs) perceptual_loss criterion_perceptual(sr_imgs, hr_imgs) g_loss perceptual_loss lambda_adv * adv_loss lambda_pixel * pixel_loss g_loss.backward() opt_g.step() # 记录到tensorboard或控制台 pbar.set_postfix({ d_loss: f{d_loss.item():.4f}, g_loss: f{g_loss.item():.4f} })训练过程中有几个细节很关键。生成器更新时需要把判别器的梯度冻结也就是opt_g.step()之前不做opt_d.zero_grad()之外的判别器操作。更规范的做法是更新生成器时直接传入discriminator但不更新它的参数因为opt_g里本来就只有生成器的参数。4.2 训练策略与参数调节心得初始学习率我设置为1e-4使用Adam优化器betas(0.9, 0.99)。每20万次迭代学习率衰减一半。当然对于小规模训练比如自己的一两百张图直接跑几千个iteration看效果学习率保持在5e-5到2e-4之间都是可接受的。训练轮数上我的经验是想得到能看的效果至少要让生成器“过一遍”所有训练数据30次以上30个epoch。如果发现生成的图像细节模糊、整体偏平滑说明感知损失权重小了或者模型还没收敛如果出现明显噪点和伪影那就减少对抗损失权重或者调低学习率。梯度裁剪值得做。我在训练中给生成器和判别器都加了梯度裁剪max_norm设为1.0这一招对防止训练崩坏很有用尤其是GAN这种对抗训练容易震荡的场景。4.3 推理与模型部署训练完成后推理阶段只需要加载生成器权重判别器直接丢掉def inference(generator, lr_img_path, output_path, upscale_factor4, devicecuda): lr_img cv2.imread(lr_img_path) lr_tensor torch.from_numpy(lr_img).permute(2, 0, 1).float().unsqueeze(0) / 255.0 lr_tensor lr_tensor.to(device) with torch.no_grad(): sr_tensor generator(lr_tensor) sr_img sr_tensor.squeeze(0).permute(1, 2, 0).cpu().numpy() * 255.0 sr_img np.clip(sr_img, 0, 255).astype(np.uint8) cv2.imwrite(output_path, sr_img)推理时有两个容易忽略的点。一是输入图像要归一化到0-1不要直接喂0-255的像素值否则训练和推理尺度不一致效果会差很多二是输出要clip到0-255再转uint8GAN生成器偶尔会输出超出范围的像素值。4.4 评估指标怎么算超分任务最常用的指标是PSNR和SSIM。PSNR是从像素误差角度评估数值越高越好SSIM从结构相似性角度评估越接近1越好。但这两个指标都有局限——PSNR高不代表视觉效果好GAN类的超分模型PSNR通常不如纯L1训练的模型但视觉观感可能更好。所以评估时我一般会搭配着看不能只看单个指标。import numpy as np from skimage.metrics import structural_similarity as ssim def psnr(img1, img2): img1 img1.astype(np.float64) / 255.0 img2 img2.astype(np.float64) / 255.0 mse np.mean((img1 - img2) ** 2) if mse 0: return float(inf) return 20 * np.log10(1.0 / np.sqrt(mse)) psnr_value psnr(sr_img, hr_img) ssim_value ssim(sr_img, hr_img, channel_axis2)5. 常见问题与排查技巧实录5.1 训练时Loss变成NaN这是初学者最常遇到的问题。可能原因有几个学习率过大、数据里存在全黑或全白图像、网络初始化不稳定。我遇到过一次是因为数据集中有张全黑的图像它的梯度计算出来数值非常大直接导致loss爆炸。排查思路先用小batch size比如2跑几个iteration如果没出现NaN再逐步加大检查训练数据剔除全黑、全白、异常损坏的图像给优化器添加梯度裁剪把max_norm设为1.0能拦截大部分梯度异常的情况。5.2 PyTorch 2.6加载旧模型报错报错信息类似这样RuntimeError: Weights only load failed. This file can still be loaded, to do so you can passweights_onlyFalsetotorch.load.这是PyTorch 2.6把weights_only默认值改成True导致的。解决办法是在加载模型时显式声明state_dict torch.load(weight_path, map_locationcpu, weights_onlyFalse)更彻底的做法是加载旧权重后重新保存一次让新的权重文件直接兼容。这一步可以避免后续每一次加载都要加参数。5.3 生成图像颜色偏灰或偏暗这种问题多数是图像预处理不一致造成的。训练时如果用了归一化比如把像素从0-255归一化到0-1推理时也必须做同样的归一化。反之也一样。另一个原因可能是数据增强时用了亮度、对比度调整但推理时没有对应调整。颜色偏暗的另外一个常见原因是生成器输出的数值经过了sigmoid或者tanh激活函数但你忘了做对应的反变换。ESRGAN的生成器默认不使用这些激活函数但如果你改过结构就要检查这一层。5.4 显存不够怎么办训练时batch size设大、patch尺寸设置大都会挤压显存。遇到显存不够我的建议是优先降低batch size保patch size。128×128的patch对模型学习纹理很关键batch size降到4甚至2影响相对小一些。如果还是不够可以用混合精度训练# PyTorch自带AMP自动混合精度实现 from torch.cuda.amp import autocast, GradScaler scaler GradScaler() with autocast(): output model(input) loss criterion(output, target) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度在40系显卡上不仅能省显存还因为Tensor Core的加持提升了训练速度实测大概能快20%-30%。5.5 训练多久才能看到效果这个问题最容易被问但答案最不好给。我用单张RTX 4090在DIV2K训练集约800张图上训练大概4个小时左右就能看到比较像样的效果。如果只是拿一百来张自己的图片跑通流程1到2个小时就能看到结果。关键不是时间而是定期保存模型检查点并做验证。我每500个iteration就在验证集上跑一次推理把生成图像存下来肉眼对比看效果变化。这种看板式的记录方式比单纯盯loss曲线靠谱得多。6. 项目复盘与可扩展方向把ESRGAN完整跑通之后我最大的感受是这套方案的成熟度和可迁移性。生成器和判别器的设计思路可以迁移到很多其他图像生成任务比如去噪、去模糊、修复只需要改数据和损失函数就行。RRDB作为特征提取主干放在其他超分模型比如Real-ESRGAN里依然是核心组件。后续如果想让效果再上一层楼有几个扩展方向可以考虑引入真实退化模型做数据合成走Real-ESRGAN的路线增加退化消除模块或高频特征增强模块用蒸馏的方式把16M的模型压缩到几M适配移动端部署。我最推荐先跑通论文里的原始配置再按自己的数据特点微调不要一上来就堆模块改结构否则出了问题很难定位是哪里引入的。最后分享一个我自己的小习惯每次训练前固定随机种子PyTorch里torch.manual_seed、torch.cuda.manual_seed、CUDNN的deterministic模式这样同样的配置可以复现结果对比实验才公平。图像超分这个方向很多细节决定了最终效果多跑几次、多对比几组参数慢慢就能找到手感。本文还有配套的精品资源点击获取