拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

SwinIR自定义训练测试代码:从数据准备到PSNR计算全流程

SwinIR自定义训练测试代码:从数据准备到PSNR计算全流程 简介这份资源是面向图像恢复方向学习者与开发者的SwinIR自定义训练与测试代码实现基于Swin Transformer架构可完成图像超分辨率与图像去噪等任务。代码在官方源码基础上重新梳理了训练与测试流程并补充了关键注释逻辑完整易懂适合希望深入理解SwinIR内部机制、动手复现实验的读者。压缩包共17个文件约28KB以5个py脚本为核心涵盖网络结构、数据加载、损失函数与工具函数等模块另含4个pyc缓存、4个xml与1个iml等IDE配置、2个m脚本及1个说明文本便于快速搭建工程。图像去噪任务修改数据路径后即可直接运行超分任务则需取消数据集加载类中的patchsize操作作者在说明中已给出提示。目前已有5374人学习下载读者可据此掌握完整的训练测试链路、参数配置与排错思路并在此基础上开展自己的图像恢复实验。1. SwinIR 自定义训练测试代码从开箱到跑通自己的数据如果你手头有一批低质图像想用 SwinIR 做恢复却卡在“官方代码能跑 demo换成自己的数据就翻车”这一步那这份自定义训练测试代码就是为你准备的。它把 SwinIR 的训练、验证、测试三条链路拆开重写逻辑完整、注释到位不是那种只留一个main.py让你猜的仓库。SwinIR 本身是把 Swin Transformer 引入图像恢复的经典结构擅长超分、去噪、去压缩伪影但官方实现偏研究向数据加载和配置耦合较深。这份代码的价值在于你能清楚看到一张低质图从Dataset到DataLoader、从forward到loss.backward()、再到test阶段拼接输出的全过程适合想真正吃透 SwinIR 训练细节、而不是只调一次推理的从业者。2. 环境与数据准备把 SwinIR 跑起来的前置条件2.1 依赖版本与目录结构SwinIR 依赖 PyTorch、timm、einops 等库版本不匹配是新手第一个跟头。我一般会先固定一套能跑通的组合再谈调参。下面这份requirements.txt是我在 3090 和 4090 上都验证过的CUDA 11.8 环境# requirements.txt torch2.0.1cu118 torchvision0.15.2cu118 timm0.9.7 einops0.7.0 numpy1.24.3 opencv-python4.8.1.78 pillow10.0.1 tqdm4.66.1安装时注意torch和torchvision要一起装别先装 torch 再单独升级 torchvision否则容易出现undefined symbol这类玄学报错。目录结构建议按下面组织训练脚本里用相对路径引用换机器时只改一个--data_root就能迁移SwinIR_custom/ ├── data/ │ ├── train/ │ │ ├── HR/ # 高质图训练目标 │ │ └── LR/ # 低质图网络输入 │ └── val/ │ ├── HR/ │ └── LR/ ├── models/ │ └── swinir.py ├── options/ │ └── train_swinir.json ├── train.py ├── test.py └── utils/ ├── dataset.py └── metric.pyHR 和 LR 的文件名必须一一对应比如HR/0001.png对应LR/0001.png。常见做法是用脚本批量生成 LR超分任务用 bicubic 下采样去噪任务加高斯噪声去压缩伪影则用 JPEG 重新编码。别手动改文件名几百张之后一定会错位。2.2 数据集类的关键参数dataset.py里最容易写错的是配对逻辑和归一化。下面这段是我常用的实现支持超分和去噪两种模式# utils/dataset.py import os import cv2 import torch import numpy as np from torch.utils.data import Dataset class PairedDataset(Dataset): def __init__(self, hr_dir, lr_dir, patch_size64, scale4, modesr): self.hr_dir hr_dir self.lr_dir lr_dir self.patch_size patch_size self.scale scale self.mode mode # 只保留两边都存在的文件名避免训练中途报 FileNotFoundError hr_names set(os.listdir(hr_dir)) lr_names set(os.listdir(lr_dir)) self.names sorted(list(hr_names lr_names)) assert len(self.names) 0, HR 和 LR 目录没有同名文件检查配对 def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] hr cv2.imread(os.path.join(self.hr_dir, name), cv2.IMREAD_COLOR) lr cv2.imread(os.path.join(self.lr_dir, name), cv2.IMREAD_COLOR) hr cv2.cvtColor(hr, cv2.COLOR_BGR2RGB) lr cv2.cvtColor(lr, cv2.COLOR_BGR2RGB) # 训练时随机裁剪保证 HR 和 LR 空间对应 if self.mode sr: h, w, _ lr.shape if h self.patch_size or w self.patch_size: # 小图直接 resize避免裁剪越界 lr cv2.resize(lr, (self.patch_size, self.patch_size)) hr cv2.resize(hr, (self.patch_size * self.scale, self.patch_size * self.scale)) else: top np.random.randint(0, h - self.patch_size 1) left np.random.randint(0, w - self.patch_size 1) lr lr[top:top self.patch_size, left:left self.patch_size] hr hr[top * self.scale:(top self.patch_size) * self.scale, left * self.scale:(left self.patch_size) * self.scale] # 转 tensor 并归一化到 [0,1]SwinIR 内部按这个范围处理 lr torch.from_numpy(lr.transpose(2, 0, 1)).float() / 255.0 hr torch.from_numpy(hr.transpose(2, 0, 1)).float() / 255.0 return {lr: lr, hr: hr, name: name}逻辑说明hr_names lr_names取交集是血泪经验曾经因为 LR 少了一张图训练到第 300 个 iteration 才崩排查半天。patch_size指 LR 上的裁剪尺寸HR 对应裁剪patch_size * scale这样超分任务的空间对齐不会错。归一化用/255.0而不是 ImageNet 均值方差是因为 SwinIR 原论文就是在 [0,1] 上训练的换标准化会改变输入分布收敛变慢。modedenoise时不做裁剪直接整图送入噪声在生成 LR 时已经加好。提示如果显存吃紧把patch_size从 64 降到 48batch size 从 8 降到 4先保证能跑起来再逐步加。3. 训练脚本拆解损失、优化器与日志3.1 模型初始化与损失选择SwinIR 的模型定义在models/swinir.py这份代码保留了原结构的核心浅层特征提取、深层 Swin Transformer 块、重建模块。初始化时几个参数直接决定显存和效果# train.py 片段 import torch from models.swinir import SwinIR from utils.dataset import PairedDataset from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) model SwinIR( upscale4, in_chans3, img_size64, window_size8, img_range1.0, depths[6, 6, 6, 6, 6, 6], embed_dim180, num_heads[6, 6, 6, 6, 6, 6], mlp_ratio2, upsamplerpixelshuffle, resi_connection1conv ).to(device) # 超分用 L1去噪可换 Charbonnier压缩伪影任务 L1 也稳 criterion torch.nn.L1Loss() optimizer torch.optim.Adam(model.parameters(), lr2e-4, betas(0.9, 0.99)) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max250000)参数说明window_size8是 Swin Transformer 的注意力窗口显存不够就降到 6 或 4但太小会削弱长程建模。depths和embed_dim是模型容量180 通道、6 层是 SwinIR 的轻量配置想追指标可以上 240 通道但显存翻倍。upsamplerpixelshuffle对应超分去噪任务改成pixelshuffledirect或直接去掉上采样。优化器用 Adam 而不是 AdamW是因为原论文配置如此betas 第二项 0.99 比默认 0.999 收敛更快。学习率 2e-4 配余弦退火是超分任务里比较稳的组合。3.2 训练循环与验证节奏训练循环里最容易被忽略的是梯度裁剪和验证频率。下面这段是核心骨架# train.py 片段 train_set PairedDataset(data/train/HR, data/train/LR, patch_size64, scale4, modesr) train_loader DataLoader(train_set, batch_size8, shuffleTrue, num_workers4, drop_lastTrue) for epoch in range(1, 201): model.train() for step, batch in enumerate(train_loader): lr batch[lr].to(device) hr batch[hr].to(device) optimizer.zero_grad() pred model(lr) loss criterion(pred, hr) loss.backward() # 梯度裁剪防止 Swin 块偶发爆炸 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() if step % 50 0: print(fepoch {epoch} step {step} loss {loss.item():.5f} flr {scheduler.get_last_lr()[0]:.2e}) # 每 5 个 epoch 存一次权重别只存最后一个 if epoch % 5 0: torch.save(model.state_dict(), fweights/swinir_epoch{epoch}.pth)逻辑说明drop_lastTrue避免最后一个 batch 只有 1 张图导致 BatchNorm 统计不稳虽然 SwinIR 主要用 LayerNorm但保持习惯。梯度裁剪max_norm1.0是踩坑后加的Swin 的窗口注意力在个别样本上会产生大梯度不裁剪偶尔会 loss 变 NaN。验证不要每个 epoch 都做太耗时我一般每 5 个 epoch 跑一次test.py的验证模式看 PSNR 是否单调上升。权重按 epoch 存而不是只存 best是因为后期可能过拟合需要回退到中间版本。注意num_workers在 Windows 上设 0否则 DataLoader 可能卡死Linux 上设 4 到 8看 CPU 核数。4. 测试与推理把 PSNR 算对把图拼对4.1 测试脚本的滑动窗口推理SwinIR 测试时如果直接整图送入显存会爆尤其是 2K 以上的图。常见做法是滑动窗口加重叠拼接# test.py 片段 import torch import cv2 import numpy as np from models.swinir import SwinIR def inference(model, img_lr, window64, overlap16, scale4): model.eval() _, _, h, w img_lr.shape output torch.zeros(1, 3, h * scale, w * scale).to(img_lr.device) weight torch.zeros_like(output) stride window - overlap for y in range(0, h, stride): for x in range(0, w, stride): y1 min(y, h - window) x1 min(x, w - window) patch img_lr[:, :, y1:y1 window, x1:x1 window] with torch.no_grad(): pred model(patch) output[:, :, y1*scale:(y1window)*scale, x1*scale:(x1window)*scale] pred weight[:, :, y1*scale:(y1window)*scale, x1*scale:(x1window)*scale] 1 return output / weight.clamp(min1)逻辑说明window64是 LR 上的窗口overlap16是重叠像素重叠是为了消除拼接缝。y1 min(y, h - window)保证最后一块不越界。权重累加后归一化比直接覆盖边缘更平滑。如果图比窗口还小直接整图推理别走循环。这个写法比官方 demo 的test.py更直观你能看到每一块怎么拼回去。4.2 PSNR 与 SSIM 的正确计算方式算指标时最常见的翻车是通道顺序和数值范围不一致。下面这段是标准做法# utils/metric.py import torch import numpy as np from skimage.metrics import peak_signal_noise_ratio, structural_similarity def calculate_psnr_ssim(pred, gt): # pred, gt: tensor [1,3,H,W], range [0,1] pred pred.squeeze(0).permute(1, 2, 0).cpu().numpy() gt gt.squeeze(0).permute(1, 2, 0).cpu().numpy() pred np.clip(pred, 0, 1) gt np.clip(gt, 0, 1) psnr peak_signal_noise_ratio(gt, pred, data_range1.0) ssim structural_similarity(gt, pred, channel_axis2, data_range1.0) return psnr, ssim参数说明data_range1.0对应 [0,1] 归一化如果你前面用了 [0,255]这里要改成 255否则 PSNR 会差出 48dB 的常数偏移。channel_axis2是 skimage 新版的写法老版本用multichannelTrue。np.clip防止模型输出超出 [0,1] 导致 PSNR 计算异常。SSIM 的win_size默认 7小图要手动调小否则报错。提示测试集上算指标前先确认 LR 和 HR 的配对关系没被 resize 破坏否则 PSNR 再高也是假的。5. 避坑与排查SwinIR 训练中最容易翻车的五件事5.1 现象loss 从第一个 iteration 就是 NaN原因学习率过大或输入数据里有全黑/全白图导致梯度爆炸。SwinIR 对输入范围敏感如果误用了 [0,255] 未归一化的图第一层卷积直接溢出。解决检查dataset.py里是否除了 255.0把学习率从 2e-4 降到 1e-4 试一个 epoch加torch.nn.utils.clip_grad_norm_。我一般还会在训练前用脚本扫一遍数据把像素均值小于 1 或大于 254 的图挑出来。5.2 现象训练 loss 正常下降但测试 PSNR 只有 20dB 出头原因训练和测试的归一化方式不一致或者测试时忘了model.eval()Dropout 和 LayerNorm 行为不同。解决统一用 [0,1] 归一化测试脚本开头加model.eval()和torch.no_grad()检查测试时 HR 是否被错误 resize。曾经有个项目训练用 RGB、测试用 BGRPSNR 直接掉 10dB血泪教训。5.3 现象显存溢出batch size 降到 1 还是 OOM原因window_size和patch_size组合过大或者测试时整图推理没切块。解决训练时把patch_size降到 48、window_size降到 6测试时用第 4 章的滑动窗口推理。另外检查num_workers是否设得过高DataLoader 的 pinned memory 也会占显存。5.4 现象验证 PSNR 波动大时高时低原因验证集太小或者验证时随机裁剪导致每次评估的 patch 不同。解决验证集至少 20 张图验证时不做随机裁剪用整图或固定中心裁剪。我一般会单独写一个val.py固定随机种子保证每次评估的输入一致。5.5 现象训练到后期 loss 突然上升PSNR 下降原因过拟合或者余弦退火的学习率降到太低后模型在局部最优附近震荡。解决保留每 5 个 epoch 的权重回退到 PSNR 最高的那个加早停策略连续 10 次验证 PSNR 不升就停数据增强可以加随机翻转和旋转但注意 HR 和 LR 要同步变换。6. 进阶技巧用配置文件管理实验与多尺度训练6.1 把超参抽到 JSON 配置硬编码超参的实验没法复现。我习惯把关键参数抽到options/train_swinir.json{ scale: 4, patch_size: 64, batch_size: 8, lr: 2e-4, epochs: 200, window_size: 8, embed_dim: 180, depths: [6, 6, 6, 6, 6, 6], data_root: data/train }训练脚本用argparse读配置路径再覆盖个别参数。这样换数据集只改data_root换模型容量只改embed_dim实验记录一目了然。我一般还会在权重文件名里带上关键参数比如swinir_w8_e180_ps64.pth回头找模型不用猜。6.2 多尺度训练提升泛化固定patch_size训出来的模型对某些尺寸的图效果会差一截。常见做法是每隔几个 epoch 随机换一次patch_size比如在 [48, 64, 96] 里抽# train.py 片段 import random if epoch % 10 0: new_patch random.choice([48, 64, 96]) train_set.patch_size new_patch train_loader DataLoader(train_set, batch_size8, shuffleTrue, num_workers4, drop_lastTrue) print(fepoch {epoch} switch patch_size to {new_patch})逻辑说明patch_size变化后要重建 DataLoader因为 Dataset 的裁剪逻辑依赖这个参数。多尺度训练能让模型适应不同分辨率的输入测试时 PSNR 通常能涨 0.1 到 0.2dB。但别太频繁每 10 个 epoch 换一次够了换太勤模型来不及收敛。6.3 验证方法固定种子跑三次取平均SwinIR 训练有随机性单次结果不能说明问题。我的习惯是固定torch.manual_seed(42)、np.random.seed(42)、random.seed(42)跑三次不同初始化的训练取 PSNR 均值。如果三次方差超过 0.3dB说明数据或超参不稳得回头查。验证时用同一组测试图别每次换图否则指标没有可比性。从那以后我每次开新实验都强制先跑一遍 5 个 epoch 的小规模训练确认 loss 下降、PSNR 上升、显存不爆再上完整训练。这套流程帮我省下了至少几十小时的无效 GPU 时间。希望帮到你。本文还有配套的精品资源点击获取
返回列表