RCAN超分模型PyTorch实现:原理、训练与部署全解析
2026/9/12 10:16:55 网站建设 项目流程

简介: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.pyrcan.py里看到四个关键类:CALayerRCABResidualGroup(或直接写RIR)、RCAN。RCAN 的关键创新不是堆层数,而是引入了通道注意力(Channel Attention)机制。超分任务中,不同特征通道对重建的贡献不一样,高频通道负责纹理,低频通道负责结构。通道注意力通过全局平均池化提取每个通道的统计量,再经过卷积和 Sigmoid 激活生成 0 到 1 之间的权重,把这些权重乘回原特征图,实现对通道重要度的重标定。

RCAB(Residual Channel Attention Block)是 RCAN 的基本构建单元。一个 RCAB 包含两层卷积、一个 ReLU 激活和一个 CALayer,整个块使用残差连接。残差连接的存在让梯度可以直接从网络深处流回浅层,而通道注意力让梯度在通道维度上有了选择性。

import torch import torch.nn as nn class CALayer(nn.Module): def __init__(self, channel, reduction=16): super(CALayer, self).__init__() self.avg_pool = nn.AdaptiveAvgPool2d(1) self.conv_du = nn.Sequential( nn.Conv2d(channel, channel // reduction, 1, padding=0, bias=True), nn.ReLU(inplace=True), nn.Conv2d(channel // reduction, channel, 1, padding=0, bias=True), 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_feats=64, kernel_size=3, reduction=16): super(RCAB, self).__init__() self.body = nn.Sequential( nn.Conv2d(n_feats, n_feats, kernel_size, padding=kernel_size // 2, bias=True), nn.ReLU(inplace=True), nn.Conv2d(n_feats, n_feats, kernel_size, padding=kernel_size // 2, bias=True), CALayer(n_feats, reduction) ) def forward(self, x): return x + self.body(x)

reduction参数是通道压缩率,默认取 16。比如输入是 64 通道,中间先压缩成 4 通道,再扩张回 64 通道,实际上形成了一个 bottleneck 结构。这个参数直接影响注意力的表达能力和参数量,reduction越小,注意力模块参数量越大,拟合能力越强,但也更容易过拟合。RCAB 的残差连接把注意力模块的输出和输入逐元素相加,这样的设计让网络在训练初期退化为普通卷积堆叠,训练更稳定。

RIR(Residual 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_resgroups=10, n_resblocks=20, n_feats=64, scale=4, reduction=16): super(RCAN, self).__init__() kernel_size = 3 self.scale = scale self.head = nn.Conv2d(3, n_feats, kernel_size, padding=kernel_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, padding=kernel_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 out

n_resgroupsn_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(单图推理)、datadataset.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 python=3.10 -y conda activate rcan # 优先走 PyTorch 官网提供的安装命令,按 CUDA 版本选择对应安装命令 pip install torch torchvision --index-url https://download.pytorch.org/whl/cpu pip install numpy opencv-python

如果机器上没有 GPU,CPU 版 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_resgroups=10, n_resblocks=20, n_feats=64, scale=4).to(device) # 加载预训练权重,strict=False 允许部分权重缺失 state_dict = torch.load('checkpoints/RCAN_BIX4.pt', map_location=device) model.load_state_dict(state_dict, strict=True) 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_dictstrict=True会在权重键名不匹配时直接报错,如果你的压缩包里没有RCAN_BIX4.pt这个文件名,先打印state_dict.keys()核对键名,再修改加载逻辑。

4. 训练 RCAN 模型:数据集、超参数与 loss 设计的完整策略

4.1 训练数据准备与 patch 采样逻辑

从零训练 RCAN 对数据和机器都是有门槛的。最标准的训练集是 DIV2K,包含 800 张高分辨率训练图,每张图像尺寸在 2K 级别。完整训练一个 scale=4 的 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 是 BD(Blur Downscale),即先模糊再下采样,训练出的模型对不同模糊核更鲁棒。

import torch.utils.data as data import random class TrainDataset(data.Dataset): def __init__(self, hr_paths, scale=4, patch_size=192): 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] # 生成 LR:bicubic 下采样 lr_patch = cv2.resize(hr_patch, (self.patch_size // self.scale, self.patch_size // self.scale), interpolation=cv2.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_tensor

patch 大小推荐 192×192,对应 LR patch 在 scale=4 时为 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(), lr=1e-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'] = lr

batch 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_value=255.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_location='cpu') 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_location='cuda:0'map_location='cpu'控制加载设备,默认情况下会把权重加载到保存时所在的设备,如果你的机器没有对应设备,会报 CUDA error。所以先统一map_location='cpu'再手动移到 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_version=11 )

注意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 接到实时视频超分流程时值得做的一步优化。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询