尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

RCAN超分模型PyTorch实现:原理、训练与部署全解析

RCAN超分模型PyTorch实现:原理、训练与部署全解析 简介RCAN是图像超分辨率领域经典的深度学习模型由尹正等人于2018年提出核心在于将残差学习与通道注意力机制结合以增强特征表达能力并提升重建细节。这份资源是RCAN的PyTorch实现适合希望复现实验、训练自定义数据或深入研究超分辨率算法的开发者和科研人员。压缩包共16个文件约1.96MB涵盖模型定义model.py、数据集处理dataset.py、训练与测试主脚本main.py、通用工具utils.py以及示例入口example.py并且附有README说明文档和用于效果对比的PNG/BMP图像便于快速熟悉工程结构。目前已有813人学习下载。解压配置好PyTorch环境后即可运行通过自带图像对比低分辨率、双三次插值以及RCAN重建结果同时可对照源码理解残差注意力组RAG和通道注意力层的实现细节为后续网络改进与应用扩展提供参考。1. RCAN 是什么一份 .rar 里装的超分模型值多少拿到一个名为“RCAN-pytorch.rar”的压缩包大概率是超分辨率super-resolution方向的一份经典 PyTorch 实现。RCAN 全称 Very Deep Residual Channel Attention Networks for Image Super-Resolution是 2018 年提出的基于残差通道注意力的超分模型。它解决的问题很具体在图像超分任务中网络加深之后如何让梯度流动顺畅同时让模型真正学会“关注”对重建最有价值的通道和区域。你在搜索引擎里输入“RCAN 代码”“RCAN pytorch”找到的仓库、网盘、课程附件基本就是同一个东西一份包含模型定义、训练脚本、测参数和几个 scale 预训练权重的代码包。适合的人群也很明确刚入门超分方向的学生、要在自己数据集上微调超分模型的算法工程师以及想把注意力机制嵌入到其他图像恢复任务里的人。这份代码的价值不在于“能跑通”而在于它把通道注意力、残差嵌套结构、长跳连这些思想用极清晰的 PyTorch 写了出来值得拆开逐行读。2. RCAN 的原理与 PyTorch 代码拆解RIR 与通道注意力的实现细节2.1 RCAN 的核心组合RCAB、通道注意力与 RIR把 RCAN 的模型文件打开通常能在model.py或rcan.py里看到四个关键类CALayer、RCAB、ResidualGroup或直接写RIR、RCAN。RCAN 的关键创新不是堆层数而是引入了通道注意力Channel Attention机制。超分任务中不同特征通道对重建的贡献不一样高频通道负责纹理低频通道负责结构。通道注意力通过全局平均池化提取每个通道的统计量再经过卷积和 Sigmoid 激活生成 0 到 1 之间的权重把这些权重乘回原特征图实现对通道重要度的重标定。RCABResidual Channel Attention Block是 RCAN 的基本构建单元。一个 RCAB 包含两层卷积、一个 ReLU 激活和一个 CALayer整个块使用残差连接。残差连接的存在让梯度可以直接从网络深处流回浅层而通道注意力让梯度在通道维度上有了选择性。import torch import torch.nn as nn class CALayer(nn.Module): def __init__(self, channel, reduction16): super(CALayer, self).__init__() self.avg_pool nn.AdaptiveAvgPool2d(1) self.conv_du nn.Sequential( nn.Conv2d(channel, channel // reduction, 1, padding0, biasTrue), nn.ReLU(inplaceTrue), nn.Conv2d(channel // reduction, channel, 1, padding0, biasTrue), nn.Sigmoid() ) def forward(self, x): y self.avg_pool(x) y self.conv_du(y) return x * y class RCAB(nn.Module): def __init__(self, n_feats64, kernel_size3, reduction16): super(RCAB, self).__init__() self.body nn.Sequential( nn.Conv2d(n_feats, n_feats, kernel_size, paddingkernel_size // 2, biasTrue), nn.ReLU(inplaceTrue), nn.Conv2d(n_feats, n_feats, kernel_size, paddingkernel_size // 2, biasTrue), CALayer(n_feats, reduction) ) def forward(self, x): return x self.body(x)reduction参数是通道压缩率默认取 16。比如输入是 64 通道中间先压缩成 4 通道再扩张回 64 通道实际上形成了一个 bottleneck 结构。这个参数直接影响注意力的表达能力和参数量reduction越小注意力模块参数量越大拟合能力越强但也更容易过拟合。RCAB 的残差连接把注意力模块的输出和输入逐元素相加这样的设计让网络在训练初期退化为普通卷积堆叠训练更稳定。RIRResidual in Residual是 RCAN 的顶层结构多个 ResidualGroup 串联每个 Group 内部又是多个 RCAB 的串联每个 Group 外层再套一层长跳连。从整体看整个网络可以看作一个巨大的残差块输入到输出之间有一条直接的恒等映射通路。这样做的好处是显而易见的网络深度可以跑到 400 层以上而不会出现明显的梯度消失。2.2 从 RCAB 到整个 RCAN 的前向过程完整 RCAN 的前向流程分为三部分浅层特征提取、深层特征映射和重建。先用一个 3×3 卷积把输入的低分辨率图像映射到特征空间然后送入 RIR 深度网络做特征变换最后通过 PixelShuffle 实现上采样。class RCAN(nn.Module): def __init__(self, n_resgroups10, n_resblocks20, n_feats64, scale4, reduction16): super(RCAN, self).__init__() kernel_size 3 self.scale scale self.head nn.Conv2d(3, n_feats, kernel_size, paddingkernel_size // 2) body [] for _ in range(n_resgroups): group [] for _ in range(n_resblocks): group.append(RCAB(n_feats, kernel_size, reduction)) body.append(nn.Sequential(*group)) self.body nn.Sequential(*body) self.tail nn.Sequential( nn.Conv2d(n_feats, n_feats * (scale * scale), kernel_size, paddingkernel_size // 2), nn.PixelShuffle(scale) ) def forward(self, x): res self.head(x) out self.body(res) out out res out self.tail(out) return outn_resgroups和n_resblocks分别控制分组的数量和每组内的残差块数量官方默认配置是 10 组、每块 20 层总共约 200 个 RCAB。这个配置是论文里测试过的性能基线直接使用 ReLU 和 3×3 卷积没有使用批归一化因为在超分任务里单图输入没有 batch 统计意义上的归一化需求而批归一化反而会引入额外的计算开销。PixelShuffle 是这里的核心上采样操作它把c * r * r个通道重新排列成c个通道、宽高各放大r倍的图像。相比直接使用转置卷积PixelShuffle 没有可学习的插值参数棋盘伪影更少训练也更稳定。3. PyTorch 环境搭建与 RCAN 最小推理从 .rar 到 demo 出图3.1 解压 .rar 后先看什么代码结构识别的顺序拿到压缩包先不要急着跑训练脚本第一步应该把文件列表展开看一遍。一般 RCAN 的 PyTorch 实现包含以下几个文件model.py模型定义、option.py超参数配置、train.py训练入口、demo.py单图推理、data或dataset.py数据加载、checkpoints预训练模型目录。这些文件的命名在不同仓库里略有差异但结构基本一致。先看option.py里的参数定义注意几个关键项--scale表示超分倍率常见值是 2、3、4、8--n_resgroups和--n_resblocks表示模型深度配置直接影响显存占用--data_train和--data_test分别是训练和测试数据集路径--save_results决定是否在测试时保存输出图片。如果demo.py存在那说明压缩包附带了最简单的推理脚本不需要完整的数据集环境就能跑通。# 解压后先做两件事确认 Python 版本、确认是否有预训练权重 unzip rcan-pytorch.rar -d rcan-pytorch cd rcan-pytorch ls -la checkpoints/3.2 PyTorch 环境搭配CPU 版本也能跑最小推理RCAN 的推理代码没有复杂的第三方依赖核心只需要 PyTorch、NumPy 和图像处理库。环境搭建最快捷的方式是用 Anaconda 创建一个独立环境避免和日常开发环境产生依赖冲突。conda create -n rcan python3.10 -y conda activate rcan # 优先走 PyTorch 官网提供的安装命令按 CUDA 版本选择对应安装命令 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu pip install numpy opencv-python如果机器上没有 GPUCPU 版 PyTorch 跑一张 128×128 的低分辨率图像推理时间是足够的大约 2 到 5 秒。pytorch 版本选择上RCAN 的代码大多基于 PyTorch 1.x 编写直接用 PyTorch 2.x 也能运行只是要注意旧代码里的torch.nn.functional.upsample_bilinear如果存在则可能需要替换成torch.nn.functional.interpolate。这点在依赖pytorch 基础框架的版本升级时格外容易踩坑。3.3 demo.py 最小推理与参数含义如果没有 demo 脚本自己写一个推理脚本只需要四十行左右。核心流程是读图、转 Tensor、归一化、模型前向、还原像素范围、保存图片。下面这个脚本可以直接复制使用。# demo_infer.py import torch import cv2 import numpy as np from model import RCAN def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) model RCAN(n_resgroups10, n_resblocks20, n_feats64, scale4).to(device) # 加载预训练权重strictFalse 允许部分权重缺失 state_dict torch.load(checkpoints/RCAN_BIX4.pt, map_locationdevice) model.load_state_dict(state_dict, strictTrue) model.eval() img cv2.imread(input.jpg) # 读入 BGR 图像 img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) lr_tensor torch.from_numpy(img.transpose(2, 0, 1)).float().div_(255.0) lr_tensor lr_tensor.unsqueeze(0).to(device) # [1, 3, H, W] with torch.no_grad(): sr_tensor model(lr_tensor) sr_img sr_tensor.squeeze(0).clamp_(0.0, 1.0).mul_(255.0) sr_img sr_img.byte().cpu().numpy().transpose(1, 2, 0) sr_img cv2.cvtColor(sr_img, cv2.COLOR_RGB2BGR) cv2.imwrite(output.png, sr_img) print(f输入尺寸: {img.shape}, 输出尺寸: {sr_img.shape}) if __name__ __main__: main()这个脚本的输入输出没有做边界 padding 处理如果输入图像尺寸不是 scale 的整数倍输出尺寸会向下取整偶发像素错位。更严谨的写法是在前向之前把 H 和 W 向上对齐到 scale 的整数倍填充方式任选裁掉多余的边界即可。注意model.load_state_dict的strictTrue会在权重键名不匹配时直接报错如果你的压缩包里没有RCAN_BIX4.pt这个文件名先打印state_dict.keys()核对键名再修改加载逻辑。4. 训练 RCAN 模型数据集、超参数与 loss 设计的完整策略4.1 训练数据准备与 patch 采样逻辑从零训练 RCAN 对数据和机器都是有门槛的。最标准的训练集是 DIV2K包含 800 张高分辨率训练图每张图像尺寸在 2K 级别。完整训练一个 scale4 的 RCAN 模型单卡 V100 大约需要 3 到 5 天具体取决于 patch 大小和 batch size。如果你没有 DIV2K使用 DIV2K 的子集或者自己收集 200 张高清图片也能训出效果尚可的模型只是泛化能力会弱一些。RCAN 的数据加载逻辑通常分两条线一是用torch.utils.data.Dataset把 HR 图像裁成固定大小的 patch运行时随机裁剪二是实时生成对应的 LR 图像先对 HR patch 做高斯模糊加下采样再把 LR patch 送入网络。注意RCAN 训练时使用的 LR 是由 HR 经过 bicubic 插值下采样得到的称为 Bicubic degradation简称 BI。还有另一种 degradation 是 BDBlur Downscale即先模糊再下采样训练出的模型对不同模糊核更鲁棒。import torch.utils.data as data import random class TrainDataset(data.Dataset): def __init__(self, hr_paths, scale4, patch_size192): self.hr_paths hr_paths self.scale scale self.patch_size patch_size def __getitem__(self, idx): hr cv2.imread(self.hr_paths[idx]) hr cv2.cvtColor(hr, cv2.COLOR_BGR2RGB) ih, iw, _ hr.shape x random.randint(0, iw - self.patch_size) y random.randint(0, ih - self.patch_size) hr_patch hr[y:y self.patch_size, x:x self.patch_size] # 生成 LRbicubic 下采样 lr_patch cv2.resize(hr_patch, (self.patch_size // self.scale, self.patch_size // self.scale), interpolationcv2.INTER_CUBIC) lr torch.from_numpy(lr_patch.transpose(2, 0, 1)).float().div_(255.0) hr_tensor torch.from_numpy(hr_patch.transpose(2, 0, 1)).float().div_(255.0) return lr, hr_tensorpatch 大小推荐 192×192对应 LR patch 在 scale4 时为 48×48。patch 太大会导致单个 batch 显存飙升patch 太小则会导致感受野不足影响重建质量。数据加载的.div_(255.0)把像素归一化到 0 到 1 区间这是 PyTorch 图像训练的常见做法RCAN 论文也采用同样的归一化方式。4.2 损失函数、优化器与学习率调度RCAN 论文里使用的是 L1 损失函数也就是 Mean Absolute Error。相比 L2 损失L1 损失在超分任务中通常能带来更高的 PSNR 和更好的感知质量。损失函数定义可以通过torch.nn.L1Loss一行代码实现。优化器使用 Adam初始学习率1e-4权重衰减默认不设置或设成极小值。RCAN 在训练过程中使用的学习率策略是里程碑衰减在第 200 个 epoch 时把学习率降到初始值的十分之一在第 300 个 epoch 时再降一次。完整训练周期通常设为 500 个 epoch。如果数据集规模较小可以把里程碑提前比如 100 和 200 轮各降一次。criterion torch.nn.L1Loss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) def adjust_lr(epoch): lr 1e-4 if epoch 200: lr 1e-5 if epoch 300: lr 1e-6 for param_group in optimizer.param_groups: param_group[lr] lrbatch size 的设置有很强的硬件约束。RCAN 完整模型参数量约 16M输入 48×48 LR patch 时FP16 混合精度下 11GB 显存可以跑 batch size 16。如果在 8GB 显存的卡上训练建议 batch size 降到 8同时把n_resgroups降到 5 个左右来换取速度。训练 loss 曲线的下降在 L1 loss 下看起来会比较平缓前 50 个 epoch 可能只能从 0.08 降到 0.06不要因为这个数值幅度小就觉得模型没在学。epoch lr train_loss val_psnr(x4) 50 1.00e-4 0.0423 26.81 100 1.00e-4 0.0361 27.64 200 1.00e-4 0.0312 28.20 250 1.00e-5 0.0294 28.52 300 1.00e-5 0.0281 28.71 350 1.00e-6 0.0276 28.88上表是一份模拟的 500 轮训练日志用于说明 loss 和学习率变化的对应关系。真实训练时验证集 PSNR 的波动在 0.1 到 0.2 dB 之间是正常的milestone降学习率后 PSNR 会有一个明显跳升这是常见的超分训练信号可以据此判断调度策略是否生效。4.3 验证阶段的 PSNR / SSIM 评估逻辑验证评估要遵循一个标准流程将 LR 输入送入模型得到 SR 输出然后与 HR 图像比较。计算 PSNR 之前要把 SR 裁剪到和 HR 完全一致的尺寸一般做法是去掉边界几个像素因为卷积的 padding 会导致边界像素重建质量较差。import math def calc_psnr(img1, img2, max_value255.0): mse torch.mean((img1 - img2) ** 2) if mse 0: return float(inf) return 10 * math.log10(max_value ** 2 / mse.item())calc_psnr的输入是 0 到 255 范围的浮点 Tensor。注意不要在 0 到 1 范围和 0 到 255 范围混用否则数值差 20 dB 左右。SSIM 建议直接用skimage.metrics.structural_similarity或者pytorch_msssim库自己写 SSIM 很容易在边界处理和方差计算上出问题。5. 参数调试与异常排查训练和推理阶段最常遇到的 6 个问题5.1 加载预训练权重时报错 key 不匹配RCAN 的预训练权重通常来自原作者的官方训练脚本不同仓库的键名可能不同最常见的问题是module.前缀。用torch.nn.DataParallel训练保存的权重会在所有 key 前面加一个module.而直接加载时模型没有这个前缀。解决办法是加载后去掉前缀。state_dict torch.load(RCAN_BIX4.pt, map_locationcpu) new_state_dict {} for k, v in state_dict.items(): name k[7:] if k.startswith(module.) else k new_state_dict[name] v model.load_state_dict(new_state_dict)PyTorch 1.10 之后的版本还允许在torch.load里直接用map_locationcuda:0或map_locationcpu控制加载设备默认情况下会把权重加载到保存时所在的设备如果你的机器没有对应设备会报 CUDA error。所以先统一map_locationcpu再手动移到 GPU是最稳妥的写法。5.2 显存不足OOM的排查路径训练 RCAN 时 OOM 的高发点有三处输入 patch 太大、batch size 太大、梯度图累积。最容易忽略的是验证阶段也会占显存因为验证时同样要前向传播而 RCAN 前向过程会保存中间激活值用于计算图。解决方法是在验证代码块里显式使用with torch.no_grad():这个上下文管理器会关闭自动求导的梯度计算图节省的显存相当可观。如果 batch size 调整后仍 OOM使用梯度累计来模拟更大的 batch。梯度累计的原理是把多个 mini-batch 的梯度累加后再更新一次参数数学上等价于更大的 batch size只是 BN 层的统计会有细微差异。RCAN 本身不使用 BN所以可以放心累计。accumulation_steps 4 optimizer.zero_grad() for i, (lr, hr) in enumerate(train_loader): lr, hr lr.to(device), hr.to(device) sr model(lr) loss criterion(sr, hr) loss loss / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()把loss / accumulation_steps之后再做反向传播等效于把每个 batch 的梯度按比例缩小后再累加这样梯度值不会因为累计步数变多而膨胀。5.3 训练 loss 不下降的 4 个常见原因RCAN 训练 50 个 epoch 后 loss 几乎不动最直接的原因就是学习率过大或过小。学习率1e-4是 RCAN 的经典配置去掉权重衰减之后Adam 的前期训练会非常稳定。如果换了数据集且图像整体偏暗或偏亮需要检查输入是否做了归一化常见错误是直接用 PIL 读图得到 0 到 255 的整数数组输入网络导致梯度数值爆炸。排查方法打印模型输出张量的标准差如果 SR 输出值范围超过 0 到 1 或集中在某个奇异区间多半是输入归一化的问题。第二个原因是 patch 采样逻辑导致数据分布偏差。比如随机裁剪时 HR 图的边缘区域占比过高或者模糊下采样时使用了错误插值方法。RCAN 的数据管线统一使用 bicubic 下采样ONNX 推理验证时也保持一致否则训练和测试的 degradation 不一致模型会在两个 domain 间摇摆。第三个原因是 CPU 和 GPU 数据加载速度不匹配导致 GPU 利用率锯齿状波动这种情况训练流程本身没问题但显卡空转时间长有效训练轮次变少。第四个最容易忽视的问题是学习率调度器写了但没生效。使用adjust_lr(epoch)这种手动调整方式时要确保每个 epoch 结束时确实调用了这个函数并在几个关键 epoch 点上打印当前学习率验证。常见的错误是在epoch 200的判定里写成了epoch % 200导致学习率每 200 轮反复跳跃。5.4 推理结果偏暗或出现伪影如果加载预训练权重后输出的 SR 图整体偏暗先检查 PixelShuffle 后有没有做 255 的像素值缩放。再检查输入图像是否被转成 BGR模型是在 RGB 上训练的用 OpenCV 的imread读入后如果不转换直接送入网络通道错位会导致颜色失真严重。伪影问题主要集中在棋盘格效应这通常出现在使用转置卷积的实现版本里。官方 RCAN 实现用的是 PixelShuffle不会出现这个现象。如果你的代码版本里出现伪影可以尝试把上采样换成 PixelShuffle并在最后加一个 3×3 卷积层做平滑。6. 把 RCAN 迁移到任意尺寸输入ONNX 导出与动态形状验证RCAN 的模型结构本身是全卷积网络理论上支持任意尺寸输入但实际部署到服务端或 FPGA 时需要考虑到固定形状的性能优化。把 PyTorch 模型导出为 ONNX 是一种常见的部署路径RCAN 导出 ONNX 时最需要注意的就是dynamic_axes参数它决定是否允许输入的宽高维度动态变化。model.eval() x torch.randn(1, 3, 48, 48).to(device) torch.onnx.export( model, x, rcan_x4.onnx, input_names[lr_input], output_names[sr_output], dynamic_axes{ lr_input: {0: batch, 2: height, 3: width}, sr_output: {0: batch, 2: height, 3: width} }, opset_version11 )注意dynamic_axes里的第 0 维同时配置了 batch 维度这在 RCAN 这种没有 batch 归一化的网络上是安全的。如果模型里有 BN动态 batch 会要求 BN 在导出时处于 eval 模式否则每次推理的 batch 统计量都会被重新计算结果不稳定。导出后可以使用onnxruntime来验证输出与 PyTorch 原模型是否一致。pip install onnxruntime onnx python -c import onnx; m onnx.load(rcan_x4.onnx); onnx.checker.check_model(m)验证时通过前后两次输入不同尺寸来确认动态形状是否生效。正确的预期是输入 48×48 输出 192×192输入 64×48 输出 256×192两者都能成功推理且 PSNR 差异小于 0.01 dB。如果只有固定尺寸能跑检查opset_version是否过低ONNX 算子集中 PixelShuffle 的台形推理在低版本 opset 中支持不完整。额外一个验证技巧是尝试导出 FP16 版本使用model.half()并传入半精度输入能在保持几乎相同 PSNR 的前提下把显存和模型体积各减半这也是把 RCAN 接到实时视频超分流程时值得做的一步优化。本文还有配套的精品资源点击获取
返回列表