EDSR超分辨率实战:从原理到PyTorch推理全流程
1. 为什么超分辨率值得你花时间折腾第一次接触超分辨率重建是在一个老照片修复的需求上。手头有一批2005年左右用卡片机拍的数码照片分辨率只有640×480放到现在的大屏上惨不忍睹。当时试过PS的锐化、Topaz的放大效果都差强人意——要么糊要么假。后来接触到EDSR这个模型才真正理解什么叫“用深度学习把丢失的像素找回来”。EDSR全称Enhanced Deep Residual Networks是2017年CVPR上NTIRE超分辨率挑战赛的冠军方案。它的核心贡献说起来很简单把当时流行的ResNet残差块里的Batch Normalization层全部去掉同时把通道数从64扩展到256模型规模翻了十几倍效果直接碾压了之前的SRResNet。这个思路在当时其实挺反直觉的——大家都在往网络里加东西它反而做减法。但正是这个减法让EDSR在PSNR指标上比SRResNet高了0.5dB以上视觉上的提升更明显。这篇文章适合谁看如果你手头有PyTorch环境想跑一个完整的超分辨率项目从下载预训练模型到生成高清图像走通全流程那这篇就是为你写的。如果你只是想了解超分辨率的基本原理文章里的原理拆解部分也能帮你建立直观认知。我默认你有基本的Python和PyTorch使用经验但即使没有跟着步骤走也能跑通。整个流程我会拆成四块先讲清楚EDSR的设计思路和为什么这么设计再讲环境准备和模型下载的具体操作然后是完整的推理代码和参数解析最后是我在实际操作中踩过的坑和排查方法。每一步我都会解释“为什么这么做”而不是只给命令让你复制粘贴。2. EDSR模型架构拆解与设计逻辑2.1 去掉BN层到底解决了什么问题要理解EDSR为什么去掉BN层得先知道BN在超分辨率任务里干了什么“坏事”。Batch Normalization在分类任务里是标配它通过归一化每层的输入分布来加速训练、防止梯度消失。但在超分辨率任务里网络要学的是从低分辨率到高分辨率的映射关系这个映射对数值的绝对大小很敏感。举个例子假设某个像素在低分辨率图里是128对应的高分辨率目标值是200。BN层会把128归一化到均值为0、方差为1的分布上这个过程中像素之间的相对大小关系被压缩了。虽然理论上后续层可以恢复但实际训练中这种归一化会引入不稳定性尤其是在batch size较小的时候BN的统计量估计不准反而拖累效果。EDSR的作者做了对比实验同样的网络结构有BN的版本PSNR是32.15dB去掉BN后直接跳到32.65dB。这0.5dB在超分辨率领域是很大的差距相当于从“能看”到“清晰”的跨越。而且去掉BN后模型参数量减少了约15%推理速度反而更快。注意去掉BN并不意味着训练不需要归一化。EDSR在训练时对输入图像做了简单的均值减法把像素值从[0,255]映射到[-127.5,127.5]左右这个操作在推理时同样要做否则输出会偏色。2.2 残差缩放与多尺度训练策略EDSR的另一个关键设计是残差缩放Residual Scaling。每个残差块的输出在加回输入之前会乘以一个0.1的系数。这个操作看起来不起眼但它是训练深层网络的关键。没有这个缩放当网络堆到32个残差块时梯度会爆炸训练根本跑不起来。为什么是0.1而不是0.5或0.01作者在论文里做了消融实验0.1时训练最稳定收敛后的PSNR最高。0.5时梯度仍然偏大需要更小的学习率0.01时残差分支的贡献被过度压制模型退化成浅层网络。这个0.1是实验调出来的经验值不是拍脑袋定的。多尺度训练是EDSR的另一个亮点。训练时不是固定放大2倍或4倍而是随机从[1.0, 4.0]里采样一个缩放因子把高分辨率图下采样到对应尺寸作为输入。这样训出来的模型可以处理任意放大倍数而且不同尺度之间互相正则化效果比单独训一个4倍模型还好。不过官方发布的预训练模型还是分尺度的因为固定尺度的模型在对应任务上表现更专精。2.3 模型规模与硬件需求的权衡EDSR有多个版本最常用的是EDSR-Baseline16个残差块64通道和EDSR32个残差块256通道。Baseline版本参数量约1.5M推理一张1080p图在GTX 1060上大概0.8秒完整版参数量43M同样的图要3秒以上。如果你只是做个人项目Baseline版本完全够用效果比传统插值好太多。完整版EDSR的43M参数是什么概念对比一下ResNet-50是25MVGG-16是138M。EDSR比ResNet-50还大但它的计算量主要集中在卷积上没有全连接层所以显存占用相对可控。4倍放大的EDSR在推理时需要约4GB显存处理512×512的输入块如果显存不够可以分块处理再拼接。3. 环境准备与模型下载实操3.1 PyTorch环境搭建的避坑指南PyTorch的安装方式直接决定了后续会不会遇到奇怪的兼容性问题。我的建议是不要用pip直接装用conda创建独立环境。原因很简单conda能同时管理Python版本和CUDA版本pip装出来的PyTorch经常和系统CUDA对不上跑起来报“CUDA driver version is insufficient”这种错。具体操作如下conda create -n edsr python3.9 conda activate edsr conda install pytorch torchvision torchaudio pytorch-cuda11.8 -c pytorch -c nvidia这里选Python 3.9是因为它在PyTorch各版本里兼容性最好3.10以上有些老版本的torchvision会出问题。CUDA 11.8是目前最稳的版本12.x虽然新但有些显卡驱动还没跟上。装完之后验证一下import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))如果cuda.is_available()返回False先检查显卡驱动版本。在命令行跑nvidia-smi看右上角的CUDA Version。如果显示的是12.2但你装的是11.8的PyTorch理论上向下兼容没问题但保险起见还是装对应版本。提示如果你用的是Windows建议用WSL2而不是原生Windows。WSL2的CUDA支持已经成熟而且文件系统性能比Windows原生好很多读大模型文件时差距明显。3.2 预训练模型下载的三种途径EDSR的官方预训练模型托管在GitHub上但直接下载经常断流。我试过三种方式按推荐程度排序第一种是用git clone把整个仓库拉下来模型文件在experiment/目录下。这种方式最稳因为git支持断点续传。命令是git clone https://github.com/sanghyun-son/EDSR-PyTorch.git cd EDSR-PyTorch仓库大概200MB模型文件在experiment/edsr_baseline_x4/下面文件名是model_best.pt。第二种是用wget或curl直接下载单个模型文件。官方仓库的release页面有直链但速度不稳定。可以加-c参数支持断点续传wget -c https://github.com/sanghyun-son/EDSR-PyTorch/releases/download/1.0/EDSR_x4.pt第三种是用Hugging Face的镜像。Hugging Face上有人上传了转换好的EDSR模型格式是safetensors加载更方便。搜索“EDSR super resolution”就能找到下载速度比GitHub快不少。模型文件下载后放在项目根目录的models/文件夹下后面代码里指定路径就行。3.3 项目目录结构规划一个清晰的目录结构能省掉很多路径报错的时间。我习惯这样组织edsr_project/ ├── models/ │ └── EDSR_x4.pt ├── inputs/ │ └── test.jpg ├── outputs/ │ └── test_x4.png ├── src/ │ ├── model.py │ ├── inference.py │ └── utils.py └── requirements.txtmodel.py放EDSR的网络定义inference.py是推理脚本utils.py放图像预处理和后处理的函数。这样拆分的好处是如果你想换其他超分辨率模型比如RCAN或SwinIR只需要改model.py推理逻辑不用动。4. 推理代码逐行解析与参数调优4.1 EDSR网络定义的极简实现官方仓库的代码有上千行包含了训练、测试、各种数据增强。但推理只需要网络定义和前向传播我把它精简到了80行左右import torch import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, channels, res_scale0.1): super().__init__() self.res_scale res_scale self.conv1 nn.Conv2d(channels, channels, 3, padding1) self.conv2 nn.Conv2d(channels, channels, 3, padding1) self.relu nn.ReLU(inplaceTrue) def forward(self, x): residual self.conv1(x) residual self.relu(residual) residual self.conv2(residual) return x residual * self.res_scale class EDSR(nn.Module): def __init__(self, scale4, channels64, num_blocks16): super().__init__() self.head nn.Conv2d(3, channels, 3, padding1) self.body nn.Sequential(*[ResidualBlock(channels) for _ in range(num_blocks)]) self.body_conv nn.Conv2d(channels, channels, 3, padding1) self.upsample nn.Sequential( nn.Conv2d(channels, channels * scale * scale, 3, padding1), nn.PixelShuffle(scale) ) self.tail nn.Conv2d(channels, 3, 3, padding1) def forward(self, x): x self.head(x) residual x x self.body(x) x self.body_conv(x) x x residual x self.upsample(x) x self.tail(x) return x这里有几个细节值得说。PixelShuffle是上采样的关键它把通道维度的数据重新排列到空间维度。比如输入是[B, 64*16, H, W]经过PixelShuffle(4)后变成[B, 64, 4H, 4W]。这种方式比转置卷积快而且没有棋盘格伪影。res_scale0.1就是前面说的残差缩放。num_blocks16对应Baseline版本改成32就是完整版。channels64是Baseline的配置完整版是256。4.2 图像预处理与后处理的数值细节超分辨率推理最容易出错的地方不是网络本身而是预处理和后处理的数值范围。EDSR训练时输入是RGB三通道像素值归一化到[0,1]然后减去0.5再除以0.5映射到[-1,1]。推理时必须做同样的操作否则输出会偏暗或偏亮。import numpy as np from PIL import Image def preprocess(image_path): img Image.open(image_path).convert(RGB) img np.array(img).astype(np.float32) / 255.0 img (img - 0.5) / 0.5 img np.transpose(img, (2, 0, 1)) img np.expand_dims(img, 0) return torch.from_numpy(img) def postprocess(tensor): img tensor.squeeze(0).cpu().numpy() img np.transpose(img, (1, 2, 0)) img (img * 0.5 0.5) * 255.0 img np.clip(img, 0, 255).astype(np.uint8) return Image.fromarray(img)np.clip这一步不能省。网络输出可能超出[0,255]范围不裁剪的话PIL会报错或者产生奇怪的色块。astype(np.uint8)也要显式调用否则保存出来的图是浮点格式体积大而且有些看图软件不认。4.3 分块推理解决显存不足如果你要处理4K甚至8K的图直接整图推理会爆显存。解决办法是分块处理每块之间有重叠最后拼接时取重叠区域的平均值。重叠大小建议是放大倍数的2倍比如4倍放大就用8像素重叠。def tile_inference(model, img_tensor, scale4, tile_size256, overlap8): b, c, h, w img_tensor.shape output torch.zeros(b, c, h*scale, w*scale, deviceimg_tensor.device) weight torch.zeros_like(output) for i in range(0, h, tile_size - overlap): for j in range(0, w, tile_size - overlap): i_end min(i tile_size, h) j_end min(j tile_size, w) i_start max(0, i_end - tile_size) j_start max(0, j_end - tile_size) tile img_tensor[:, :, i_start:i_end, j_start:j_end] with torch.no_grad(): out_tile model(tile) output[:, :, i_start*scale:i_end*scale, j_start*scale:j_end*scale] out_tile weight[:, :, i_start*scale:i_end*scale, j_start*scale:j_end*scale] 1 return output / weight这个实现里有个小技巧i_start max(0, i_end - tile_size)保证了最后一块不会越界同时所有块的大小都是tile_size避免了边缘块尺寸不一致导致的拼接错位。5. 常见问题排查与性能优化实录5.1 模型加载报错的三种典型情况第一种是KeyError: model。官方发布的模型文件里state_dict是嵌套在model键下面的但有些第三方转换的模型直接就是state_dict。解决办法是加载后检查一下checkpoint torch.load(models/EDSR_x4.pt, map_locationcpu) if model in checkpoint: state_dict checkpoint[model] else: state_dict checkpoint model.load_state_dict(state_dict)第二种是size mismatch。这通常是因为你用的网络配置和模型文件不匹配。比如模型是Baseline16块64通道但你实例化的是完整版32块256通道。检查num_blocks和channels参数是否和模型文件名对应。第三种是CUDA out of memory。除了分块推理还可以用半精度推理model model.half() img_tensor img_tensor.half()半精度能把显存占用减半速度提升30%左右画质损失几乎看不出来。但要注意有些老显卡比如GTX 9系对半精度支持不好可能反而更慢。5.2 输出图像偏色或发灰的排查思路偏色问题90%出在预处理和后处理不匹配。检查清单如下现象可能原因解决方法整体偏暗预处理没做[-1,1]归一化加上(img-0.5)/0.5整体偏亮后处理没做反归一化加上img*0.50.5颜色发灰输入是BGR但模型按RGB训练用convert(RGB)边缘有色带分块重叠不够增大overlap到scale*2局部过曝没做clip加np.clip(img,0,255)我遇到过一次特别诡异的情况输出图整体偏绿。排查了半天发现是Image.open读进来的图带了alpha通道np.array之后是4通道但网络只接受3通道。convert(RGB)能解决这个问题但如果你忘了加网络会把alpha通道当颜色通道处理输出就偏色了。5.3 推理速度优化的四个实用技巧第一个是torch.no_grad()。这个上下文管理器会关闭梯度计算显存占用减少约40%速度提升20%以上。推理时一定要加训练时才需要梯度。第二个是model.eval()。这会把BN层和Dropout层切换到推理模式。虽然EDSR没有BN但如果你用的是其他模型忘了加这个会导致输出不稳定。第三个是channels_last内存格式。在支持Tensor Core的显卡上RTX系列把模型和输入转成channels_last能利用Tensor Core加速model model.to(memory_formattorch.channels_last) img_tensor img_tensor.to(memory_formattorch.channels_last)实测在RTX 3080上4倍放大的EDSR从3.2秒降到2.1秒提升明显。第四个是torch.compilePyTorch 2.0。这个能把模型编译成优化后的计算图首次运行慢但后续推理快30%左右model torch.compile(model)不过torch.compile对EDSR这种纯卷积网络效果一般对Transformer类模型提升更大。而且编译过程可能报错建议先跑通再尝试。5.4 批量处理的工程化建议如果你要处理大量图片不要一张一张跑那样GPU利用率很低。正确的做法是攒一个batch比如8张图一起推理。但要注意不同尺寸的图不能直接拼batch需要先padding到相同尺寸推理完再裁掉多余部分。def batch_inference(model, image_paths, scale4, batch_size8): results [] for i in range(0, len(image_paths), batch_size): batch_paths image_paths[i:ibatch_size] tensors [preprocess(p) for p in batch_paths] max_h max(t.shape[2] for t in tensors) max_w max(t.shape[3] for t in tensors) padded [] for t in tensors: pad_h max_h - t.shape[2] pad_w max_w - t.shape[3] padded.append(torch.nn.functional.pad(t, (0, pad_w, 0, pad_h), modereflect)) batch torch.cat(padded, dim0) with torch.no_grad(): outputs model(batch) for j, t in enumerate(tensors): h, w t.shape[2], t.shape[3] out outputs[j:j1, :, :h*scale, :w*scale] results.append(postprocess(out)) return resultsmodereflect比modeconstant效果好因为边缘像素是镜像反射的不会在边界产生突变推理完裁掉padding区域后边界更自然。6. 从单张推理到批量生产的扩展思路跑通单张推理只是起点。实际项目中你可能会遇到视频超分辨率、实时超分辨率、或者和其他模型串联的需求。这里分享几个我实践过的扩展方向。视频超分辨率最简单的做法是逐帧处理但这样会有闪烁问题因为相邻帧的推理结果不一致。解决办法是用光流做帧间对齐或者用循环网络把前一帧的特征传过来。EDSR本身不支持这些但你可以把EDSR作为基础模型在外面套一层时序模块。实时超分辨率需要模型轻量化。EDSR-Baseline在RTX 3060上处理720p到4K大概0.5秒一帧离实时30fps还差得远。可以考虑用MobileNetV2的倒残差块替换EDSR的残差块参数量降到0.3M速度提升10倍画质损失在可接受范围内。和其他模型串联的典型场景是“超分辨率人脸修复”。先用EDSR把整图放大再用GFPGAN或CodeFormer修复人脸区域。这样比单独用人脸修复模型效果好因为人脸修复模型通常只处理512×512的输入整图放大后人脸区域的分辨率刚好合适。提示串联模型时要注意数值范围的一致性。EDSR输出是[0,255]的uint8但GFPGAN期望输入是[0,1]的float32。中间加一层转换否则人脸修复会失败。最后说一个我踩过的坑模型文件不要放在中文路径下。PyTorch的torch.load在某些版本对中文路径支持不好会报FileNotFoundError。项目路径全用英文省心。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →