ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

基于Pytorch的EDSR图像超分辨率实战:原理、配置与训练

基于Pytorch的EDSR图像超分辨率实战:原理、配置与训练 简介面向深度学习初学者的EDSR图像超分辨率重建PyTorch代码包适合论文复现、课程设计与工程实践。基于AnacondaPycharm开发PyTorch 1.9.1CUDA 11.1代码经过修改测试可直接运行。压缩包共27个文件以5个Python脚本、5个Matlab生成的.mat测试图、3个数据制作.m脚本、1个h5数据集和1个pth权重文件为主覆盖网络定义、训练、评估与可视化完整流程。edsr.py中包含残差块与亚像素卷积上采样模块适合对照论文逐段阅读dataset.py负责将h5数据转为DataLoader输入便于理解数据流。配套Set5测试集采用.mat格式PSNR计算更规范model_edsr.pth是训练得到的最优权重可跳过训练直接验证效果。data目录中的m脚本可自行制作h5训练数据便于扩展到其他数据集。已有1593人学习下载。按main_edsr.py、eval.py、test.py的顺序可依次完成训练、计算Set5平均PSNR和单张图像超分对比适合想快速复现EDSR并理解重建细节的入门者。 刚接触图像超分辨率的时候最劝退我的不是算法本身有多难懂而是跑通一个能用的复现项目要踩的坑实在太多。环境装好了权重对不上权重对上了数据格式又不对整个晚上全耗在报错和版本匹配上。所以当我拿到这套基于Pytorch复现的EDSR代码时第一反应是终于有人把新手最头疼的那些事情提前处理掉了。代码结构清晰、无bug连DIV2K数据的h5文件和训练好的最优PSNR权重都一并打包好真正做到了开箱即用。这篇文章我结合自己的复现经验和实际测试过程把EDSR的核心原理、环境配置、h5数据集制作、训练和测试细节全部梳理了一遍尤其把容易出错的地方单独摘出来讲希望能帮你少走弯路。1. 项目定位与整体设计思路1.1 EDSR是什么为什么选它EDSR全称是Enhanced Deep Residual Networks for Single Image Super-Resolution出自2017年CVPR的一篇论文作者用一套增强的残差网络结构在超分辨率比赛NTIRE2017上拿了冠军。在它之前的SRCNN、VDSR还停留在比较浅的网络结构而EDSR通过加深网络和残差学习把图像重建质量拉高了一个台阶。从实用的角度看EDSR特别适合作为超分辨率入门的第一个模型。它的结构不复杂核心就是一堆残差块堆叠再配合尾部上采样模块做尺寸放大。相比后来那些带注意力机制、Transformer结构的模型EDSR理论简单、实现容易、训练稳定而且效果放到今天依然能打。我身边不少朋友做超分方向最初都是从复现EDSR开始的。这套项目代码基于Pytorch实现Pytorch本身就是学术界用得最多的深度学习框架之一动态图机制让调试变得非常方便。项目里已经帮你把h5格式的DIV2K训练数据集做好了还有一个在验证集上PSNR效果最好的模型权重文件你拿到之后可以直接测试单张图片不需要从零开始训练也能看到超分效果。1.2 项目结构和文件分工我拿到这套代码之后先把目录结构过了一遍整体布局很清晰。核心文件包括model.pyEDSR网络结构定义、dataset.pyh5数据集读取、train.py训练入口、test.py测试与验证还有一个utils模块主要负责PSNR计算、图像保存这类辅助功能。这种分模块的设计对新手特别友好你想理解哪一块就直接打开对应文件阅读不用在一大坨代码里翻来翻去。比如你只想搞清楚EDSR的上采样是怎么实现的就看model.py里的Upsampler部分就够了。项目里h5数据和权重文件都是单独存放的训练代码里通过路径参数引用不会把数据文件混在代码里这点做得很规范后续你想换数据集或者换模型权重都只需要改配置路径。2. 环境配置与Pytorch部署细节2.1 版本匹配是整个项目跑通的前提很多新手复现项目经常卡在环境安装这一步。EDSR这套代码对Pytorch版本要求不算苛刻一般Pytorch 1.8以上都能正常跑。如果你用的是GPU版本那么CUDA、cuDNN、Pytorch三者的版本必须匹配否则会报各种奇怪的底层错误。以我目前的环境为例用的是Python 3.8配合Pytorch 1.12.1加上CUDA 11.3跑这套代码没有任何问题。如果你是全新安装建议直接用Anaconda创建独立环境conda create -n edsr python3.8 conda activate edsr pip install torch1.12.1 torchvision0.13.1 --extra-index-url https://download.pytorch.org/whl/cu113 pip install h5py numpy pillow tqdm如果你是CPU的环境那把Pytorch换成CPU版本就行了训练速度会慢一些但代码本身的逻辑是完全一样的。我建议新手先在自己的电脑上把测试流程跑通哪怕用CPU推理一张小图也比折腾半天环境最后跑不起来强得多。2.2 依赖库逐个说明除了Pytorch本体这套项目还依赖几个库h5py用于读取h5格式的数据集numpy做数组运算Pillow处理图像tqdm用来显示训练进度条。这几个库都是Python生态最常见的基础组件安装不会遇到什么坑。有一点需要提醒你h5py和你读取h5文件的方式有关。项目代码里用的是h5py读取不是pandas所以安装的时候确认装的是h5py而不是h5py文件解析相关的其他库。如果读取数据时报错“Unable to synchronously open file”多半是h5py版本和h5文件创建版本之间不兼容重新安装较新版本的h5py一般能解决。3. h5数据集制作原理与全流程3.1 为什么把训练数据封装成h5格式图像超分辨率训练通常用的是DIV2K数据集包含800张高清训练图和100张验证图。如果直接在训练时从原始图片文件读取每次都要做图像解码、随机裁剪、双三次下采样这些操作IO开销非常大训练速度会被明显拖慢。把数据提前处理成h5格式就是为了解决这个问题。h5是HDF5层次化数据格式的一种可以一次性把多个数组存到一个文件里读取时支持按key快速访问速度远比逐张读图片快。训练时只需要从h5里取出已经裁剪和归一化好的低分辨率和高分辨率图像对直接转成Tensor丢给模型就行省掉了大量重复的图像预处理操作。3.2 制作h5数据的核心脚本逻辑虽然这份项目已经提供了制作好的h5文件我还是建议新手自己动手跑一遍制作脚本这能帮你彻底理解训练数据的生成过程。核心逻辑分成四步读取高清原图、用双三次插值生成低分辨率图、随机裁剪成固定大小的patch、把低分和高分图像对存入h5文件。import h5py import numpy as np from PIL import Image import os hr_dir DIV2K_train_HR lr_scale 4 patch_size 192 h5_path div2k_train_x4.h5 with h5py.File(h5_path, w) as f: for idx, img_name in enumerate(sorted(os.listdir(hr_dir))): hr Image.open(os.path.join(hr_dir, img_name)).convert(RGB) lr hr.resize((hr.width // lr_scale, hr.height // lr_scale), Image.BICUBIC) hr_array np.array(hr).astype(np.float32) / 255.0 lr_array np.array(lr).astype(np.float32) / 255.0 # 随机裁剪中心区域保证尺寸一致 h, w lr_array.shape[:2] x np.random.randint(0, w - patch_size // lr_scale 1) y np.random.randint(0, h - patch_size // lr_scale 1) lr_patch lr_array[y:y patch_size // lr_scale, x:x patch_size // lr_scale] hr_patch hr_array[y * lr_scale:(y patch_size // lr_scale) * lr_scale, x * lr_scale:(x patch_size // lr_scale) * lr_scale] f.create_dataset(f{idx:04d}_lr, datalr_patch) f.create_dataset(f{idx:04d}_hr, datahr_patch)这段脚本有几个细节需要注意。lr_scale表示下采样倍率EDSR常见的有x2、x3、x4不同倍率需要单独制作对应的h5文件。归一化用除以255把像素值缩放到[0,1]区间这能加快模型收敛。裁剪时lr和hr要严格对齐lr的坐标乘以缩放倍率就是hr的坐标这里最容易出错一旦偏移模型学到的映射关系就是错的。3.3 训练时如何正确读取h5数据制作好h5文件之后dataset.py里的读取逻辑用h5py按key访问就行。训练时每次取出一个patch对转成Pytorch的Tensor格式并且要确认维度顺序是(C, H, W)而h5里存的是(H, W, C)需要用permute或transpose转换。import h5py import torch from torch.utils.data import Dataset class H5Dataset(Dataset): def __init__(self, h5_path): self.h5 h5py.File(h5_path, r) self.keys list(self.h5.keys()) # 只保留lr或hr按需读取这里简单演示 self.lr_keys [k for k in self.keys if k.endswith(_lr)] def __len__(self): return len(self.lr_keys) def __getitem__(self, idx): lr self.h5[self.lr_keys[idx]][:] hr_key self.lr_keys[idx].replace(_lr, _hr) hr self.h5[hr_key][:] lr_t torch.from_numpy(lr).permute(2, 0, 1).float() hr_t torch.from_numpy(hr).permute(2, 0, 1).float() return lr_t, hr_t实际训练代码里往往还会做随机翻转、旋转等数据增强增强的目的是增大数据多样性降低过拟合风险。你如果只看h5读取的逻辑核心就是记住h5文件里的数据是numpy数组转Tensor时注意维度和类型就行。4. EDSR模型核心结构与代码解读4.1 残差块为什么去掉了BatchNormEDSR模型的主体结构可以拆成三部分浅层特征提取、残差块堆叠、上采样重建。浅层特征提取就是一个3x3卷积把输入图像从RGB三通道映射到64通道的特征空间。中间是若干残差块每个残差块由两个3x3卷积加ReLU激活组成并且有一条跳跃连接把输入直接加到输出上。这里最关键的设计是去掉了BatchNorm层。原版SRResNet用了BN但EDSR的作者发现在图像超分任务里BN层会破坏图像原有的统计信息让训练变得更不稳定而且BN层本身消耗大量显存。去掉BN之后残差块变成了纯粹的卷积加激活显存占用更小训练也更加稳定。这也解释了为什么EDSR能用更深的结构去提升性能。4.2 上采样模块和PixelShuffle低分辨率特征经过残差块堆叠后尺寸还是和输入LR图一样大。要把特征图放大到HR尺寸EDSR用的是亚像素卷积也就是Pytorch里的PixelShuffle操作。它的原理是把通道维度的信息重新排列到空间维度上比如你想放大4倍就把一个形状为(N, Cx16, H, W)的特征图重排成(N, C, 4H, 4W)。class Upsampler(nn.Sequential): def __init__(self, scale, n_feats): m [] if (scale (scale - 1)) 0: for _ in range(int(math.log2(scale))): m.append(nn.Conv2d(n_feats, 4 * n_feats, 3, padding1)) m.append(nn.PixelShuffle(2)) elif scale 3: m.append(nn.Conv2d(n_feats, 9 * n_feats, 3, padding1)) m.append(nn.PixelShuffle(3)) super(Upsampler, self).__init__(*m)用PixelShuffle做上采样比直接resize图像再卷积的好处是上采样所需的参数是网络自己学出来的重建细节更丰富。从代码里也能看到对于2的整数倍缩放采用log2次数的PixelShuffle(2)实现方式非常简洁。4.3 损失函数与PSNR计算原理EDSR训练用的损失函数是L1损失也就是预测图像和真实高分辨率图像之间的绝对误差均值。L1损失相比L2损失的梯度更加平稳在超分任务上训练出来的模型PSNR往往更高。如果论文里用的是L1你复现的时候就不要换成L2否则效果会有下降。评估超分效果最常用的指标就是PSNRPeak Signal-to-Noise Ratio峰值信噪比。它的计算基于MSE公式是PSNR 10 * log10(MAX^2 / MSE)MAX是图像像素的最大值8bit图像就是255。PSNR越高说明重建图像和原图越接近通常EDSR在Set5数据集x4倍率下PSNR能到28.9以上在Urban100上接近26.5。import numpy as np def calculate_psnr(img1, img2, max_val255.0): mse np.mean((img1.astype(np.float64) - img2.astype(np.float64)) ** 2) if mse 0: return float(inf) return 10 * np.log10(max_val ** 2 / mse)PSNR不是万能的它和人的主观感知并不完全一致但在超分这个领域它始终是使用最广泛的量化指标。这套项目给你提供的权重文件就是基于验证集PSNR挑选的最优模型训练过程中每个epoch结束都会在验证集上算一次PSNR保留最好的那个。5. 训练与测试实操5.1 训练参数的设定逻辑训练EDSR时batch size、学习率、patch size这几个参数需要搭配合理。原论文用的是batch size 16、patch size 48x48、初始学习率1e-4每训练200个epoch学习率乘以0.1。实际复现时受限于GPU显存很多人的batch size会调小一些比如8或者4同时相应加大patch size到96或者192保证每轮迭代看到的像素总量差不多。python train.py --scale 4 --batch_size 16 --patch_size 96 --lr 1e-4 --n_resblocks 32 --n_feats 64 --num_epochs 300训练过程中需要注意学习率的调整时机。如果训练到后面PSNR提升非常缓慢就是学习率偏大了需要手动降低或者按照预设的schedule衰减。我自己训练时一般是前150个epoch用1e-4后面降到1e-5再跑100个epoch效果比固定学习率要稳不少。5.2 加载最优权重进行测试项目里下载好的模型权重文件一般是以epoch和PSNR命名的比如epoch_215_psnr_29.01.pth。测试时只需指定权重路径和测试图片路径程序会先读图、转Tensor、做归一化、送进模型、得到超分结果最后保存图片并打印PSNR。device torch.device(cuda if torch.cuda.is_available() else cpu) model EDSR(scale4, n_resblocks32, n_feats64) state_dict torch.load(best_psnr_x4.pth, map_locationdevice) model.load_state_dict(state_dict) model.to(device) model.eval()加载权重这里有个新手高频报错点权重文件如果是用DataParallel多卡训练保存的键名会多出“module.”前缀。直接load会报unexpected key错误解决办法有两个要么在load之前把键名的“module.”去掉要么保存的时候就直接保存model.module.state_dict()。这套项目的权重是单卡保存的正常load没问题。5.3 单张图片自定义超分演示为了快速验证模型效果我通常会写一个简单的单图测试脚本输入任意一张图片输出放大后的超分图像。以x4为例如果输入是100x100输出就是400x400。from PIL import Image import torchvision.transforms as transforms def sr_image(model, img_path, scale4): img Image.open(img_path).convert(RGB) lr img.resize((img.width // scale * scale, img.height // scale * scale), Image.BICUBIC) lr_t transforms.ToTensor()(lr).unsqueeze(0).to(device) with torch.no_grad(): sr_t model(lr_t) sr sr_t.squeeze(0).cpu().clamp(0, 1) sr transforms.ToPILImage()(sr) return sr output sr_image(model, test.png, scale4) output.save(test_sr.png)跑下来最大的体会是EDSR对清晰的边缘和纹理恢复效果非常明显原图比较模糊的文字区域超分之后锐利程度肉眼可见地提升。如果你发现输出图像有色彩异常多半是归一化和反归一化的范围没对齐检查一下输入和输出是否都在[0,1]区间。6. 常见问题与排查技巧实录6.1 维度对不齐和device不一致我见过最多的问题就是新手自己改数据加载部分时维度顺序弄错。Pytorch要求图像Tensor维度是(N, C, H, W)但用h5py读出来是(H, W, C)如果没做permute直接丢给模型马上就会报RuntimeError。另外模型放到了GPU上输入Tensor还在CPU上也会报expected device cuda但got device cpu。解决这类问题最好的习惯是添加断言比如assert lr_t.shape (1, 3, H, W)这样报错时能快速定位问题所在而不是从头到尾排查。device不一致就统一用tensor.to(device)处理。6.2 训练不收敛或PSNR始终很低训练时如果loss一直在降但验证集PSNR上不去先检查数据集是否对齐。低分辨率图和高分辨率图的对应关系必须严格一致如果LR是原图resize到1/4再resize回来HR是原图那么这个映射本身是合理的。但如果LR和HR来自不同图片或者裁剪位置偏了模型是不可能学到正确映射的。还有一点容易被忽略验证集计算PSNR之前要把输出图像clamp到[0,1]范围否则像素值超出区间会导致PSNR偏低。另外计算PSNR的图片尺寸必须一致如果HR是真值原图SR是放大后的图两者尺寸一定要对齐最好的做法是HR也提前裁剪成和SR一致的尺寸。6.3 显存溢出和推理速度慢训练时显存溢出一般发生在batch size过大或者patch size过大。EDSR虽然有几十个残差块但去掉BN之后显存占用已经优化很多了如果还是溢出优先调低batch size同时适当提高梯度累积步数保持总batch size不变。推理速度慢通常是因为没有关闭梯度计算。测试阶段一定要记得加torch.no_grad()否则Pytorch会为每个中间变量保存梯度信息显存和计算时间都会成倍增加。GPU比CPU快十几倍如果有条件尽量用GPU推理。问题现象可能原因解决方案load权重报unexpected key多卡模块前缀去掉键名中的module.输出图像颜色偏色归一化范围不一致确保输入输出都在[0,1]区间PSNR忽高忽低不稳定验证集裁剪未对齐用中心裁剪固定区域计算训练loss为nan学习率过大或数据含NaN调低l并检查数据集我自己在实际操作中最大的体会是这套项目把新手最容易卡住的三个环节都提前打通了环境配置方案明确h5数据集可以直接读预训练权重够好不需要从头训就能看到效果。但我也建议你千万别只停留在跑通代码这一步一定自己动手重写一遍dataset.py和model.py理解每一个模块的作用。图像超分辨率是一个实践性很强的方向只有把代码吃透了后面遇到新模型、新数据集才不会寸步难行。最后再分享一个小细节训练过程中每隔几个epoch把当前模型在验证图上的输出保存下来肉眼观察重建效果这比只看PSNR数值更能帮助你判断模型是否真的学到了纹理细节。毕竟PSNR高了不代表图像一定好看超分辨率的最终目的还是让人眼看着舒服这个习惯我从跑EDSR开始一直保持到现在。本文还有配套的精品资源点击获取
返回列表