拓冰建站拓冰建站
首页 / 资讯中心 / 正文

PyTorch复现SRCNN图像超分辨率:从数据管线到PSNR/SSIM评估

简介基于Pytorch的SRCNN图像超分辨率重建复现工程面向深度学习与底层视觉入门者覆盖x2、x3、x4三档放大倍率训练及推理全流程。压缩包共41个文件包含9个Python源码、6个H5格式数据集、3个PTH权重文件及BMP样例图等数据集涵盖91-image训练集与Set5标准测试集已预处理为h5格式可直接调用。代码涵盖模型构建、数据封装、指标计算与评估绘图权重为最优PSNR/SSIM模型可一键加载复现超分结果。配套教程链接提供了保姆级使用说明帮助理解训练细节与文件组织逻辑。包内还提供了单张图像测试、benchmark批量测试以及Loss/PSNR/SSIM曲线绘制脚本便于对比双三次插值与SRCNN在各放大倍数下的客观指标差异。从数据准备到模型评估均有对应脚本适合系统学习SRCNN的搭建思路。已有712人学习压缩包大小213.2MB整体结构清晰适合作为超分辨率入门的实践模板和二次开发基础。1. 图像超分辨率SRCNN的PyTorch复现卡在评估指标而不是网络本身把图像超分辨率SRCNN用PyTorch复现出来通常不会卡在网络结构上——三层卷积半天就能写完真正花时间的是三个环节bicubic降采样和论文不一致、PSNR和SSIM统计口径对不上、权重文件在换设备或换scale后加载失败。下面这套方案围绕带详细注释的PyTorch复现工程展开覆盖数据管线、模型定义、训练验证和科研绘图全流程并给出x2、x3、x4三档scale下按SSIM和PSNR筛选的最优权重文件用法。适合两类人第一次接触图像超分辨率重建、想借SRCNN跑通pytorch基础框架的初学者以及能训练但结果指标异常、需要核对评估细节的从业者。读完之后应该能复现出与论文同一量级的PSNR/SSIM并把训练曲线和对比图直接放进论文或报告。2. SRCNN原理与PyTorch数据管线bicubic降采样与YCbCr通道2.1 三步走的结构patch提取、非线性映射、重建SRCNNSuper-Resolution Convolutional Neural Network是图像超分辨率重建的奠基工作核心思路不是从低分辨率图直接生成高分辨率图而是先把低分辨率图用bicubic放大到目标尺寸再让网络学习“放大后的模糊图→清晰图”的映射。这个前提决定了整个PyTorch复现的数据流网络输入和输出shape完全一致训练时拿放大图当输入、原始高分辨率图当标签做的是逐像素回归。网络只有三个卷积层。第一层9×9卷积把输入切成重叠patch并提取特征输出64张特征图第二层1×1卷积做非线性映射把64维特征压到32维这一层是SRCNN参数量小的关键第三层5×5卷积把特征聚合成一张亮度图给出重建结果。层与层之间用ReLU第三层输出不加激活直接作为回归值。整个网络没有池化、没有BatchNorm输入输出分辨率不变padding按kernel算(9-1)//24(5-1)//221×1卷积padding为0。一个常见误解是SRCNN在学残差、输出会和输入相加。原始设计里没有全局残差连接网络直接输出完整重建图。残差学习属于VDSR那一代改进工作的范畴复现论文指标时别画蛇添足加了残差会改变收敛行为最终PSNR量级和论文对不上。2.2 modcrop与bicubic降采样先裁齐再缩放数据管线第一步是读图。训练数据常用T91、BSDS200或DIV2K读取后先把宽高裁成scale的整数倍否则降采样再放大回原尺寸时边界会多出或缺失像素训练标签错位。这个操作在SRCNN系列代码里约定叫modcrop。import cv2 import numpy as np def modcrop(img, scale): 把图像宽高裁剪为scale的整数倍避免降采样后尺寸无法对齐 h, w img.shape[:2] h h - h % scale w w - w % scale return img[:h, :w] def make_lr_hr_pair(img, scale): img modcrop(img, scale) h, w img.shape[:2] # 先缩小再放大回原尺寸得到与HR同shape的LR放大图 lr cv2.resize(img, (w // scale, h // scale), interpolationcv2.INTER_CUBIC) lr_up cv2.resize(lr, (w, h), interpolationcv2.INTER_CUBIC) ycrcb_hr cv2.cvtColor(img, cv2.COLOR_BGR2YCrCb) ycrcb_lr cv2.cvtColor(lr_up, cv2.COLOR_BGR2YCrCb) return ycrcb_hr[:, :, 0], ycrcb_lr[:, :, 0]关键在interpolationcv2.INTER_CUBIC。cv2的resize默认用INTER_LINEAR双线性插值换成它之后指标通常掉1dB以上因为训练输入和论文的实验设定不一致。先缩小再放大的两步不能省SRCNN的输入是“已经放大回目标尺寸的低分辨率图”不是原尺寸小图这一步在复现时最容易做错。返回时转成YCbCr颜色空间只取第0通道Y亮度Cb和Cr色度通道不参与训练。训练和评估都统一在Y通道这是复现论文数字的前提。2.3 训练patch采样33×33随机裁剪与数据增强全图直接进网络不是不行但SRCNN参数量小、每张图能提供的梯度有限常见做法是裁patch训练。固定patch_size33因为网络有效感受野是95-11333×33既覆盖感受野又留出上下文裁剪步长取14时patch之间有重叠数据利用率高。每个epoch随机裁一批不把所有patch一次性落盘内存和显存都省。数据管线参数取值作用插值方式cv2.INTER_CUBIC对齐论文的bicubic设定训练通道YCrCb的Y通道亮度通道对人眼更敏感评估口径统一patch_size33×33覆盖13×13感受野并留上下文裁剪stride14控制patch重叠与采样密度数据增强水平翻转90°旋转同等数据量扩充8倍下面是训练用的Dataset实现数值归一化到[0,1]同时做随机裁剪和翻转旋转。翻转和旋转组合起来等效8倍数据增强对T91这种只有91张图的小数据集这一项能让验证集PSNR提升0.2~0.4 dB。import torch from torch.utils.data import Dataset class SRCNNDataset(Dataset): def __init__(self, hr_paths, scale, patch_size33, trainTrue): self.hr_paths hr_paths self.scale scale self.patch_size patch_size self.train train def __len__(self): return len(self.hr_paths) def __getitem__(self, idx): img cv2.imread(self.hr_paths[idx]) hr_y, lr_y make_lr_hr_pair(img, self.scale) if self.train: # 随机裁33×33lr和hr在同位置裁因为两者shape一致 h, w hr_y.shape x np.random.randint(0, w - self.patch_size 1) y np.random.randint(0, h - self.patch_size 1) hr_p hr_y[y:yself.patch_size, x:xself.patch_size] lr_p lr_y[y:yself.patch_size, x:xself.patch_size] if np.random.rand() 0.5: # 水平翻转 hr_p hr_p[:, ::-1] lr_p lr_p[:, ::-1] if np.random.rand() 0.5: # 90°旋转 hr_p np.rot90(hr_p) lr_p np.rot90(lr_p) else: hr_p, lr_p hr_y, lr_y # float32 [0,1]归一化MSE损失在这个尺度下数值更稳 hr_t torch.from_numpy(hr_p.astype(np.float32) / 255.0).unsqueeze(0) lr_t torch.from_numpy(lr_p.astype(np.float32) / 255.0).unsqueeze(0) return lr_t, hr_t裁剪坐标用np.random.randint生成hr和lr在同一个坐标裁因为make_lr_hr_pair返回的两张图shape相同。验证阶段不做随机裁剪直接返回整图评估时在完整图上算PSNR和SSIM。最后用unsqueeze(0)补通道维得到(1, 33, 33)张量通道维在前是PyTorch卷积层的默认布局也符合NCHW约定。3. 三层卷积网络定义与训练配置损失函数、学习率和批量大小3.1 注释详细的SRCNN模型定义模型定义是整个工程里最短的部分但注释值得写细因为后面的科研绘图和权重筛选都依赖这个网络结构能被准确重建。环境准备上只要通过pip正常安装好的PyTorch就够不需要额外算力一块2GB显存的卡就能训练。import torch.nn as nn class SRCNN(nn.Module): SRCNN3层全卷积无池化、无BN、无残差输入输出shape一致 def __init__(self, num_channels1): super(SRCNN, self).__init__() # 9×9patch提取pad4保持分辨率不变输出64通道 self.conv1 nn.Conv2d(num_channels, 64, kernel_size9, padding4) # 1×1非线性映射把64维特征降到32维参数集中在这层 self.conv2 nn.Conv2d(64, 32, kernel_size1, padding0) # 5×5重建pad2输出1通道亮度图 self.conv3 nn.Conv2d(32, num_channels, kernel_size5, padding2) self.relu nn.ReLU(inplaceTrue) def forward(self, x): x self.relu(self.conv1(x)) x self.relu(self.conv2(x)) x self.conv3(x) # 第三层不经过ReLU输出可为负 return xpadding4和padding2分别对应(9-1)//2和(5-1)//2这三个padding值保证特征图尺寸逐层不变输入(1, H, W)输出也是(1, H, W)。第三层不加ReLU是因为回归目标可以是任意实数ReLU会把负的预测值截断强边缘处出现伪影肉眼可见的黑点就是这么来的。inplaceTrue省显存对小网络影响不大但写成习惯没坏处。3.2 损失函数与优化器MSE配Adam是复现默认组合损失函数选MSE而不是L1理由很直接PSNR由MSE换算得到最小化MSE等价于直接优化PSNR指标。L1损失是SRGAN、EDSR这类后续工作更常用的选择复现SRCNN时不要换换了之后训练曲线形态和最终PSNR都会偏移和论文对照就失去了意义。优化器有两条路线。原论文用带动量的SGD学习率1e-4PyTorch复现时更多人直接上Adam。Adam对学习率不敏感、前期收敛快在SRCNN这种小网络上通常几十个epoch就能看到平台期SGD更贴近论文原始行为但要配动量0.9和阶梯式降学习率。两条路线的最终PSNR差距在0.1 dB以内选哪个取决于你想更快看到结果还是严格复现论文配置。训练配置推荐值说明损失函数nn.MSELoss()与PSNR定义直接对应优化器Adamlr1e-4收敛快、对学习率不敏感或优化器SGDmomentum(0.9)lr1e-4更贴近原论文行为batch_size3216~64都行2GB显存可训练epoch50~100以验证集PSNR平台期为准学习率衰减StepLR(30, gamma0.5)每30轮衰减一半批量大小对SRCNN的影响没有分类网络那么敏感每个patch只有33×33batch32在2GB显存上就能跑这也是SRCNN适合当pytorch基础框架入门项目的直接原因。学习率不要拍脑袋调到1e-2MSE损失在[0,1]数据范围下梯度量级很小1e-4起步、等验证集PSNR连续5个epoch不涨再降是最省心的策略。关于固定随机种子复现SRCNN的另一个隐性要求是固定随机种子。训练前对Python、NumPy、PyTorch分别设seed并把cudnn的benchmark关掉否则两次训练虽然指标接近但“最优权重”的PSNR/SSIM会有零点几分贝的随机浮动。科研绘图时曲线对不上多半是没做这一步。import random import numpy as np import torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic True torch.backends.cudnn.benchmark False3.3 训练循环与checkpoint记录训练循环的写法决定“最优SSIM和PSNR的模型权重文件”能不能顺利产出。每轮epoch结束在验证集上算一次PSNR比历史最优值高就把模型存下来这是最朴素的早停加权重筛选逻辑。device torch.device(cuda if torch.cuda.is_available() else cpu) model SRCNN().to(device) criterion nn.MSELoss() optimizer torch.optim.Adam(model.parameters(), lr1e-4) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size30, gamma0.5) train_loader DataLoader(train_ds, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue) best_psnr 0.0 for epoch in range(100): model.train() total_loss 0.0 for lr_t, hr_t in train_loader: lr_t, hr_t lr_t.to(device), hr_t.to(device) optimizer.zero_grad() pred model(lr_t) loss criterion(pred, hr_t) loss.backward() optimizer.step() total_loss loss.item() * lr_t.size(0) scheduler.step() # 每轮结束验证一次验证逻辑见第4章 val_psnr, val_ssim evaluate(model, val_paths, scale) if val_psnr best_psnr: best_psnr val_psnr torch.save({ scale: scale, state_dict: model.state_dict(), val_psnr: val_psnr, val_ssim: val_ssim, epoch: epoch, }, fsrcnn_x{scale}_best.pth)checkpoint里同时存scale、state_dict和两个指标比只存权重多出一条关键信息加载时能直接读出这个文件是为哪一档scale训练出来的x2、x3、x4三个文件不会搞混。total_loss按样本量加权累加除以数据集长度就是当前epoch的平均MSE这个值留给第5章画训练曲线用。4. PSNR和SSIM计算精度坑Y通道评估与最优权重筛选4.1 shave边界、data_range与float64PSNR和SSIM公式本身简单但复现结果对不上论文十个里有八个出在计算口径。第一个口径是边界裁剪。卷积在图像边缘会有无效响应评估时把四周裁掉一圈裁多少以scale为准常见做法是shavescale即x2裁2像素、x3裁3像素、x4裁4像素。def calc_psnr(pred, gt, shave, data_range255.0): pred和gt都是uint8的Y通道图裁掉shave边界后按float64算 pred pred[shave:-shave, shave:-shave].astype(np.float64) gt gt[shave:-shave, shave:-shave].astype(np.float64) mse np.mean((pred - gt) ** 2) return 10.0 * np.log10(data_range * data_range / (mse 1e-10)) def calc_ssim(pred, gt, shave): from skimage.metrics import structural_similarity pred pred[shave:-shave, shave:-shave] gt gt[shave:-shave, shave:-shave] return structural_similarity(pred, gt, data_range255.0)第二个口径是通道。论文的SRCNN在Y通道训练和评估如果训练时把RGB三通道都送进网络、再在RGB图上算PSNR结果和论文参考值没有可比性通常还会偏低。第三个口径是data_range和数据精度。uint8图做差值计算前要先转float64避免整数溢出skimage的SSIM必须显式传data_range255.0不传的话按dtype推断float图默认1.0混用时指标直接失真。这三个口径任何一个出错标题里说的“最优PSNR和SSIM”就成了自说自话的数字。4.2 验证集与最优权重筛选逻辑验证集选Set5或Set14这两组图全图评估、公开结果多便于核对。评估函数把每张图转成Y通道、过模型、算指标最后取平均。def evaluate(model, img_paths, scale, device): model.eval() psnr_list, ssim_list [], [] shave scale # 边界裁剪以scale为准 for path in img_paths: img cv2.imread(path) hr_y, lr_y make_lr_hr_pair(img, scale) lr_t torch.from_numpy(lr_y.astype(np.float32) / 255.0) lr_t lr_t.unsqueeze(0).unsqueeze(0).to(device) with torch.no_grad(): sr model(lr_t).squeeze().cpu().numpy() sr_y np.clip(sr * 255.0, 0, 255) psnr_list.append(calc_psnr(sr_y, hr_y, shave)) ssim_list.append(calc_ssim(sr_y, hr_y, shave)) return float(np.mean(psnr_list)), float(np.mean(ssim_list))筛选逻辑以PSNR为主、SSIM为辅。存入checkpoint的val_ssim不是用来做早停的因为SSIM在高分区域区分度不如PSNR两个模型PSNR差0.3 dB、SSIM可能只差千分之几。把两个值同时存进权重文件论文审稿人问起“最优模型怎么选的”能直接给出当时的完整指标。输出先clip到[0, 255]再算指标不clip会把越界的预测值算进MSEPSNR会被异常像素拉低。4.3 公开参考值的核对区间论文在Set5上的公开结果大致是x2约36.6 dB、x3约32.7 dB、x4约30.4 dBSSIM对应在0.95、0.90、0.86附近。自己的复现受训练集、补丁采样和数据增强影响与这个值差0.5 dB以内都属正常差1 dB以上就要回头查数据管线而不是怀疑模型写错。评估项目常见错误正确做法PSNR偏低1dB输入是RGB或数据范围混用Y通道统一pred/gt同rangeSSIM异常偏高shave没裁或data_range传错裁掉scale圈边界data_range255训练正常但验证掉点验证时忘了model.eval()关Dropout和BN统计数字对不上论文用了FID/LPIPS等生成式指标SRCNN是MSE训练看PSNR/SSIMSRCNN这个阶段基本不看FID和LPIPS那两个指标是GAN类超分评测用的。MSE训练的模型拿FID对比没有区分度论文里也不报告别在复现工程里混用两套评估体系。5. 科研绘图与SRCNN效果对比图训练曲线、局部放大和结果表5.1 训练曲线loss和验证PSNR双y轴科研绘图的第一步是把训练过程的loss和验证集PSNR画到一张图里。两者量纲不同一个在0.01量级、一个在30以上直接共用y轴会压扁一边用双y轴是通用做法。import matplotlib matplotlib.use(Agg) import matplotlib.pyplot as plt def plot_training_curve(train_losses, val_psnrs, save_path): fig, ax1 plt.subplots(figsize(8, 5)) epochs range(1, len(train_losses) 1) # 左边y轴画MSE loss用蓝色 ax1.plot(epochs, train_losses, color#1f77b4, labelMSE Loss) ax1.set_xlabel(Epoch, fontsize12) ax1.set_ylabel(MSE Loss, color#1f77b4, fontsize12) ax1.tick_params(axisy, labelcolor#1f77b4) # 右边y轴画验证PSNR用红色 ax2 ax1.twinx() ax2.plot(epochs, val_psnrs, color#d62728, labelVal PSNR) ax2.set_ylabel(PSNR (dB), color#d62728, fontsize12) ax2.tick_params(axisy, labelcolor#d62728) fig.tight_layout() fig.savefig(save_path, dpi300, bbox_inchestight) plt.close(fig)ax1.plot画训练lossax1.twinx()创建共享x轴的第二个坐标系画验证PSNR两个曲线互不压缩。颜色用matplotlib默认配色里的蓝红黑白打印也能区分。savefig务必开dpi300和bbox_inchestight前者满足期刊分辨率要求后者避免坐标轴标签被裁掉。train_losses和val_psnrs就是第3章训练循环里每轮累加和evaluate函数返回的两个列表。图例与坐标轴范围双y轴图的图例容易叠在一起手动指定locupper rightx轴范围设成(0, epochs1)避免曲线贴边。PSNR的y轴起始值不要从0开始否则30和36 dB的区别在图上只是一条平线把ylim设为(min(val_psnrs)-1, max(val_psnrs)1)就能看出上升趋势这也是论文里常见的纵轴截断画法。5.2 效果对比图LR、Bicubic、SRCNN、HR四联图超分论文的核心图是视觉效果对比。一张图排四列低分辨率原图、bicubic放大图、SRCNN输出、高分辨率真值下面配PSNR/SSIM标注。def plot_comparison(lr_img, bicubic_img, sr_img, hr_img, scale, psnr_sr, ssim_sr, save_path): fig, axes plt.subplots(1, 4, figsize(16, 5)) titles [fLR (x{scale}), Bicubic, fSRCNN\nPSNR {psnr_sr:.2f} dB\n fSSIM {ssim_sr:.4f}, HR] for ax, img, title in zip(axes, [lr_img, bicubic_img, sr_img, hr_img], titles): ax.imshow(img, cmapgray, vmin0, vmax255) ax.set_title(title, fontsize12) ax.axis(off) fig.tight_layout() fig.savefig(save_path, dpi300, bbox_inchestight) plt.close(fig)imshow用cmapgray画Y通道灰度图vmin/vmax锁死在0和255不然matplotlib会自动拉伸对比度把噪声也“增强”出来视觉上误导读者。SRCNN和HR的PSNR/SSIM标注在标题里是审稿人最先看的位置。另外要补一张局部放大图取图像中纹理密集的区域比如建筑边缘或动物毛发裁剪后放大成子图超分重建有没有恢复出高频细节只有放大才能看出来。5.3 把逐图指标导出成表格除了曲线和对比图论文或报告还需要一张逐图指标表。用csv模块把每张验证图的文件名、PSNR、SSIM落盘后续粘进LaTeX表格或直接转Excel都方便。导出配置推荐值用途图片格式PDF或PNGPDF用于论文PNG用于网页分辨率dpi300满足期刊印刷要求字体默认或Times中文报告另设中文字体指标表CSV按图导出便于逐图核对与排版科研绘图的最基本红线是“图里出现的数据必须能由脚本复现”训练曲线对应train_losses列表对比图对应evaluate函数的输出指标表对应CSV里的行。图和数字脱钩这张图就失去了科研意义。6. x2、x3、x4权重文件的使用单图推理与排错清单6.1 权重加载与单图推理拿到x2、x3、x4三个权重文件后加载方式统一。checkpoint里存的是完整字典用map_locationcpu保证在无GPU机器上也能加载。def load_model(weight_path, device): ckpt torch.load(weight_path, map_locationcpu) model SRCNN() model.load_state_dict(ckpt[state_dict]) model.to(device).eval() print(fscale{ckpt[scale]}, val_psnr{ckpt[val_psnr]:.2f}, fval_ssim{ckpt[val_ssim]:.4f}) return model def infer_y(img_path, model, scale, device): img cv2.imread(img_path) hr_y, lr_y make_lr_hr_pair(img, scale) lr_t torch.from_numpy(lr_y.astype(np.float32) / 255.0) lr_t lr_t.unsqueeze(0).unsqueeze(0).to(device) with torch.no_grad(): sr model(lr_t).squeeze().cpu().numpy() return np.clip(sr * 255.0, 0, 255).astype(np.uint8)load_state_dict默认严格模式key对不上会直接报错。如果看到missing key或unexpected key先检查网络定义里卷积层的名字最常见的问题是自建类里层的命名和checkpoint不一致。6.2 三档scale的边界权重不能混用x2、x3、x4三个权重文件结构完全一样差在训练时的降采样倍数。用x4的权重去推x2的图输入图的分辨率关系就不对输出会出现明显的过度平滑或振铃反过来用x2权重推x4则细节严重不足。推理时输入图必须先按对应scale做一次bicubic放大再进网络Cb、Cr通道直接用bicubic放大结果只有Y通道走模型最后三通道合并回BGR再保存。提示换机器后权重文件报“size mismatch”几乎都是PyTorch版本差异或网络定义里num_channels改了核对第一行加载信息里的scale和val_psnr即可确认文件完整性。6.3 低指标的排错顺序拿到权重文件后指标偏低按下面顺序排查前三条覆盖了九成情况。现象原因处理PSNR比标注低2dBRGB全通道推理或range混用统一Y通道pred和gt同为[0,255]输出图像偏灰推理时忘了把[0,1]乘回255检查clip前的数据范围边缘有黑色条纹用了INTER_LINEAR做bicubic换成cv2.INTER_CUBIC结果有棋盘格伪影输入图没按scale先放大先bicubic放大到HR尺寸再过网络最低成本的验证方式是挑Set5里的baby或butterfly图跑一遍推理和公开参考值对比误差在0.5 dB内说明权重和管线都对误差超1 dB就按这张表从数据范围开始查。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门