SwinIR图像复原模型:从零构建完整训练与测试流程

发布时间:2026/9/4 3:56:13
SwinIR图像复原模型:从零构建完整训练与测试流程 简介本资源是一套基于SwinIR模型的自定义图像恢复训练与测试代码实现面向计算机视觉方向的研究者与深度学习初学者聚焦图像超分辨率重建与去噪等典型低级视觉任务。代码逻辑完整、注释清晰支持开箱即用图像去噪任务仅需调整数据路径即可直接运行超分任务则需在数据加载类中取消patchsize相关操作便于理解底层数据处理机制。压缩包共17个文件含5个核心Python源码如main.py、net.py、data.py、4个编译缓存pyc、4个IDE配置XML及辅助MATLAB脚本psnr.m、compute_psnr.m等总大小仅28KB轻量易部署。已有5370人学习下载配套README与目录结构含train/test数据子目录、result输出路径及.idea工程配置设计合理显著降低复现门槛特别适合快速掌握SwinTransformer在图像恢复中的工程化落地流程。1. 项目缘起为什么需要一份“逻辑完整易懂”的SwinIR训练测试代码如果你在图像超分、去噪或者JPEG压缩伪影去除这些领域折腾过大概率听说过SwinIR这个名字。作为一个基于Swin Transformer架构的图像复原模型它在多个公开基准测试上都刷出了相当亮眼的成绩。但当你兴冲冲地打开它的官方仓库准备用自己的数据集大干一场时很可能会被泼一盆冷水官方代码库功能强大但结构复杂各种配置文件和分散的脚本对于想要快速上手、理解核心流程的开发者来说学习曲线有点陡峭。特别是当你只想专注于“如何准备数据”、“如何定义训练循环”、“如何评估模型”这几个核心环节时官方代码里大量的工程化封装和分布式训练支持反而成了一种干扰。这就是我动手整理这份“逻辑完整易懂”的自定义训练测试代码的初衷。市面上能找到的很多教程要么是直接跑官方脚本对内部逻辑一笔带过要么就是过于简化缺失了数据加载、损失计算、验证集评估等关键环节导致代码跑起来却无法真正应用到自己的项目里。我的目标是剥离那些复杂的工程外壳呈现一个最核心、最骨架化的SwinIR训练测试流程。这份代码不会追求极致的性能或最前沿的改进而是确保每一个步骤都清晰可见每一行代码都有明确的意图让使用者能真正理解从数据到模型再到结果的完整闭环。无论你是想学习Transformer在底层视觉任务中的应用还是急需一个能快速适配自己数据集的超分模型基线这份代码都希望能提供一个扎实的起点。2. 核心逻辑拆解SwinIR训练与测试的五个关键模块要构建一个完整的训练测试流程我们可以将其分解为五个相对独立又相互衔接的模块。理解每个模块的职责和它们之间的数据流是掌握整个项目逻辑的关键。2.1 数据准备与加载模块任何机器学习项目的根基都是数据。对于图像复原任务我们的数据通常是以“配对”形式存在的低质量图像LR和对应的高质量图像HR。数据加载模块的核心任务就是以高效、灵活的方式将这些图像对读入内存并转换为模型可处理的张量。首先我们需要定义一个Dataset类。这个类会接收一个包含图像路径的列表或一个指明目录的字符串。在__getitem__方法中我们需要完成以下几件事读取图像使用PIL或OpenCV同时读取LR和HR图像。配对验证确保LR和HR图像在内容上是对齐的。通常通过文件名匹配如“001_lr.png”和“001_hr.png”来实现。基础预处理包括将图像数据转换为[0, 1]范围的浮点张量以及将HWC格式转换为CHW格式PyTorch标准。可选的数据增强对于训练集我们可以引入随机裁剪、水平/垂直翻转、旋转等操作来增加数据多样性防止过拟合。需要注意的是必须对LR和HR图像施加完全相同的几何变换否则配对关系就被破坏了。返回字典最终返回一个包含‘lr’和‘hr’键的字典方便后续步骤访问。这里有一个关键的细节图像归一化。很多教程会建议进行减均值除方差的标准归一化但对于图像复原任务更常见的做法是简单地将像素值从[0, 255]缩放到[0, 1]。因为复原任务的目标是像素级的精确重建使用基于数据集的统计量进行归一化可能会引入不必要的复杂度并且在推理时还需要反归一化。保持简单的[0, 1]范围能让损失计算和可视化更加直观。2.2 模型构建模块SwinIR模型本身结构复杂但幸运的是我们可以直接引用其官方实现中的核心类。我们的目标不是重写Swin Transformer Block而是正确地组装和使用它。SwinIR主要包含三个部分浅层特征提取一个简单的卷积层将输入的3通道LR图像映射到指定特征维度的深层特征。深层特征提取这是模型的核心由多个Swin Transformer层组成的残差组构成。每个组内部包含多个Swin Transformer块用于进行长距离依赖建模和特征增强。组与组之间通过残差连接和卷积层进行特征融合与下采样/上采样取决于具体配置。图像重建模块一个上采样子网络负责将学习到的深层特征映射回高分辨率的HR图像空间。对于超分任务这通常是一个像素洗牌层加上卷积层。在我们的简化代码中我们会直接导入官方的SwinIR模型类。我们需要关注的是模型初始化参数upscale: 上采样倍数如234。in_chans: 输入通道数RGB图像为3。img_size: 训练时输入LR图像的块大小patch size。这决定了Transformer处理的基本单元大小。window_size: Swin Transformer中局部窗口的大小通常设置为8。depths和num_heads: 这两个列表决定了每个Swin Transformer Stage的层数和注意力头数是模型容量和计算量的关键参数。对于轻量级模型可以使用[6, 6, 6, 6]和[6, 6, 6, 6]对于更强大的模型可能会使用[6, 6, 6, 6, 6, 6]和[6, 6, 6, 6, 6, 6]。注意直接使用官方大型模型的配置可能会导致参数量巨大对显存要求很高。在自定义训练时尤其是数据集不大时适当减少depths和embed_dim特征维度是控制模型规模、防止过拟合的有效手段。这引出了一个常见误区“全参训练”与“微调”对显存的要求差异。全参训练是指随机初始化模型所有权重并进行训练显存占用主要取决于模型前向传播的激活值和反向传播的梯度。而微调通常是在预训练模型的基础上用较小学习率更新部分或全部参数。虽然学习率小但只要模型结构不变前向传播的激活值显存占用是一样的因此微调并不会显著降低显存需求。显存瓶颈主要在模型本身的大小和输入数据尺寸。2.3 训练循环模块这是整个流程的“发动机”。一个典型的训练循环Epoch包含以下步骤模式切换model.train()启用Dropout、BatchNorm的训练模式。遍历数据加载器对于每一个batch的数据。数据迁移将batch数据从CPU转移到GPU。前向传播output model(lr_img)。损失计算loss criterion(output, hr_img)。图像复原任务最常用的损失函数是L1损失MAE和L2损失MSE。L1损失对异常值不那么敏感通常能产生视觉上更清晰的边缘L2损失则更强调像素级的平均误差。实践中L1损失更为常用和稳定。有时也会结合感知损失如VGG特征损失或对抗损失GAN来提升视觉质量但在基础版本中我们优先保证逻辑清晰使用L1损失。反向传播loss.backward()计算梯度。参数更新optimizer.step()利用梯度更新模型权重。梯度清零optimizer.zero_grad()为下一个batch做准备。日志记录定期打印当前epoch、batch的损失值有时也计算并打印PSNR/SSIM等指标虽然这会在验证阶段更正式地计算。优化器选择Adam或AdamW是Transformer模型训练的首选。AdamW因其更好的权重衰减处理方式通常能带来更优的泛化性能。学习率初始值一般设置在1e-4到5e-4之间。学习率调度使用余弦退火Cosine Annealing或带热重启的余弦退火Cosine Annealing with Warm Restarts是非常有效的策略。它能在训练初期快速下降后期精细调整有助于模型收敛到更好的局部最优点。2.4 验证与评估模块模型在训练集上损失下降并不代表它在未见过的数据上表现好。因此我们需要一个独立的验证集在每个epoch结束后评估模型的泛化能力。验证循环与训练循环类似但有三个关键区别模式切换model.eval()关闭Dropout固定BatchNorm的统计量。禁用梯度计算使用torch.no_grad()上下文管理器大幅减少内存消耗并加速计算。计算量化指标除了损失函数我们更关心人类视觉或下游任务相关的指标。最常用的两个是PSNR峰值信噪比基于均方误差MSE计算数值越大代表重建图像与原始图像误差越小质量越高。计算简单但与主观视觉质量的相关性有时不高。SSIM结构相似性指数从亮度、对比度、结构三个方面比较图像其值越接近1表示两图像越相似。SSIM通常比PSNR更能反映人眼感知到的质量差异。在验证循环中我们会累加整个验证集上的损失并计算平均PSNR和SSIM。这些指标是判断模型性能、决定是否保存检查点checkpoint以及是否早停early stopping的核心依据。2.5 测试与推理模块训练完成后我们需要用训练好的模型对全新的测试集图像进行推理并保存结果。这个模块相对独立。其核心步骤包括加载训练好的模型权重model.load_state_dict(torch.load(‘best_model.pth’))。预处理输入图像读取测试LR图像进行与训练时相同的归一化如/255.0和格式转换HWC - CHW。这里可能涉及一个重要的步骤填充。由于SwinIR模型可能对输入尺寸有要求与窗口大小整除有时需要对图像进行边缘填充以满足要求并在输出后裁剪掉填充部分。执行推理同样在model.eval()和torch.no_grad()模式下进行前向传播。后处理输出将模型输出的[0, 1]范围张量转换回[0, 255]范围的整数并转换回HWC格式。保存结果图像使用OpenCV或PIL保存最终的超分结果。对于超分任务一个常见的测试技巧是自集成将输入图像及其翻转版本水平、垂直、对角线分别输入模型然后将得到的输出再反翻转回来取平均这通常能稳定提升最终输出的质量是一种简单有效的测试时增强。3. 从零搭建代码逐行详解与避坑指南接下来我们将把上述逻辑转化为具体的PyTorch代码。我会假设一个经典的图像超分辨率任务上采样倍数为4倍。3.1 环境配置与依赖安装首先确保你的环境已安装PyTorch。SwinIR官方实现依赖于timm库和一些基础工具。pip install torch torchvision pip install timm pip install opencv-python pip install Pillow pip install scikit-image # 用于计算PSNR和SSIM pip install matplotlib # 用于可视化3.2 数据加载器实现我们创建一个简单的配对数据集。假设你的数据存放在一个目录下LR和HR图像分别放在‘train_lr’和‘train_hr’子文件夹中且文件名一一对应。import os from PIL import Image import torch from torch.utils.data import Dataset, DataLoader import torchvision.transforms as transforms class PairedImageDataset(Dataset): def __init__(self, lr_dir, hr_dir, patch_size64, scale4, is_trainTrue): Args: lr_dir: 低分辨率图像目录 hr_dir: 高分辨率图像目录 patch_size: 训练时从HR图像随机裁剪的块大小 scale: 超分倍数 is_train: 是否为训练模式决定是否进行数据增强 self.lr_dir lr_dir self.hr_dir hr_dir self.patch_size patch_size self.scale scale self.is_train is_train # 获取配对的文件名列表假设文件名相同 self.lr_filenames sorted([f for f in os.listdir(lr_dir) if f.endswith((.png, .jpg, .bmp))]) self.hr_filenames sorted([f for f in os.listdir(hr_dir) if f.endswith((.png, .jpg, .bmp))]) # 简易检查确保文件列表长度一致 assert len(self.lr_filenames) len(self.hr_filenames), LR和HR图像数量不匹配 # 基础转换ToTensor会将PIL图像转换为[C, H, W]且范围[0,1] self.to_tensor transforms.ToTensor() if is_train: # 训练时的数据增强 self.transform transforms.Compose([ transforms.RandomCrop(patch_size * scale), # 对HR图像随机裁剪 transforms.RandomHorizontalFlip(p0.5), transforms.RandomVerticalFlip(p0.5), ]) else: # 验证/测试时通常不进行随机裁剪而是处理整图或中心裁剪 self.transform None def __len__(self): return len(self.lr_filenames) def __getitem__(self, idx): lr_path os.path.join(self.lr_dir, self.lr_filenames[idx]) hr_path os.path.join(self.hr_dir, self.hr_filenames[idx]) lr_img Image.open(lr_path).convert(RGB) hr_img Image.open(hr_path).convert(RGB) if self.transform: # 对HR图像进行几何变换 hr_img self.transform(hr_img) # 对LR图像进行完全相同的变换需要先调整到对应大小 # 注意RandomCrop的参数是针对HR的LR的crop大小需要除以scale # 这里更安全的做法是先对HR进行变换然后根据变换参数同步处理LR。 # 为简化我们采用一种常见做法先对HR进行变换然后下采样得到LR模拟退化过程。 # 但我们的数据是已经配对的LR所以更合理的做法是定义一个联合变换。 # 由于代码复杂度这里暂不实现精确的联合随机裁剪假设数据已对齐。 # 一个实用的简化在训练时我们只对HR做中心裁剪并同步裁剪LR。 pass # 简化处理假设数据已对齐裁剪 # 转换为Tensor lr_tensor self.to_tensor(lr_img) hr_tensor self.to_tensor(hr_img) return {lr: lr_tensor, hr: hr_tensor} # 创建数据加载器 train_dataset PairedImageDataset(lr_dir./data/train_lr, hr_dir./data/train_hr, is_trainTrue) val_dataset PairedImageDataset(lr_dir./data/val_lr, hr_dir./data/val_hr, is_trainFalse) train_loader DataLoader(train_dataset, batch_size16, shuffleTrue, num_workers4, pin_memoryTrue) val_loader DataLoader(val_dataset, batch_size1, shuffleFalse, num_workers4, pin_memoryTrue) # 验证时batch_size常设为1避坑指南数据配对与裁剪是第一个大坑。确保你的LR图像是HR图像经过正确的下采样如双三次插值得到的而不是简单的缩放。在实现随机裁剪时必须保证LR和HR的裁剪区域在内容上严格对应。上面的简化代码省略了精确的联合随机裁剪在实际应用中你需要自己实现一个RandomPairedCrop变换或者使用已经对齐好的小块数据集如DIV2K的patch数据集。3.3 模型初始化与工具函数我们将使用timm中的SwinIR实现。你需要从官方仓库找到models目录下的相关文件如swinir.py并复制到你的项目里或者直接安装包含SwinIR的扩展包。import torch.nn as nn from swinir import SwinIR # 假设你已经将SwinIR模型定义放在swinir.py中 def create_swinir_model(upscale4, training_patch_size64): 创建一个SwinIR模型实例。 这里选择一个相对轻量的配置以适应大多数显卡。 model SwinIR(upscaleupscale, in_chans3, img_sizetraining_patch_size, window_size8, img_range1., # 输入图像范围[0, 1] depths[6, 6, 6, 6], embed_dim60, # 特征维度原论文大模型为180这里减小以降低显存 num_heads[6, 6, 6, 6], mlp_ratio2, upsamplerpixelshuffle, resi_connection1conv) return model # 初始化模型、损失函数、优化器 device torch.device(cuda if torch.cuda.is_available() else cpu) model create_swinir_model(upscale4, training_patch_size48).to(device) # 训练patch size可小于HR patch size criterion nn.L1Loss() # 使用L1损失 optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max200) # 假设训练200个epoch注意img_size参数是训练时输入模型的LR图像块大小。如果你的HR patch是192x192上采样4倍那么对应的LR patch就是48x48。因此在创建模型时img_size应设置为48而不是192。这是一个容易混淆的点。3.4 训练循环实现下面是一个完整的训练epoch函数包含了损失计算、反向传播、日志记录。def train_one_epoch(model, dataloader, optimizer, criterion, device, epoch): model.train() total_loss 0.0 for batch_idx, batch in enumerate(dataloader): lr_imgs batch[lr].to(device) hr_imgs batch[hr].to(device) # 前向传播 sr_imgs model(lr_imgs) loss criterion(sr_imgs, hr_imgs) # 反向传播与优化 optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() # 每50个batch打印一次日志 if batch_idx % 50 0: current_lr optimizer.param_groups[0][lr] print(fEpoch: {epoch} [{batch_idx}/{len(dataloader)}] Loss: {loss.item():.6f} LR: {current_lr:.7f}) avg_loss total_loss / len(dataloader) return avg_loss3.5 验证循环与指标计算验证循环需要计算PSNR和SSIM。我们使用skimage.metrics中的函数注意它们需要numpy数组格式且在[0, 1]或[0, 255]范围内。from skimage.metrics import peak_signal_noise_ratio as psnr from skimage.metrics import structural_similarity as ssim import numpy as np def validate(model, dataloader, criterion, device): model.eval() total_loss 0.0 total_psnr 0.0 total_ssim 0.0 count 0 with torch.no_grad(): for batch in dataloader: lr_imgs batch[lr].to(device) hr_imgs batch[hr].to(device) sr_imgs model(lr_imgs) loss criterion(sr_imgs, hr_imgs) total_loss loss.item() # 将张量转换为numpy数组用于计算指标 # 假设数据范围是[0,1] sr_np sr_imgs.squeeze(0).cpu().numpy().transpose(1, 2, 0) # [C,H,W] - [H,W,C] hr_np hr_imgs.squeeze(0).cpu().numpy().transpose(1, 2, 0) # 裁剪到有效区域如果模型输出有填充 # 这里假设输出尺寸与HR匹配无需裁剪 # 计算PSNR和SSIM # 注意skimage的psnr和ssim默认输入范围是[0, 1]。如果我们的数据是[0,1]则data_range1 current_psnr psnr(hr_np, sr_np, data_range1) # SSIM计算可能需要多通道分别计算再平均或指定multichannelTrue current_ssim ssim(hr_np, sr_np, data_range1, multichannelTrue, channel_axis-1) total_psnr current_psnr total_ssim current_ssim count 1 avg_loss total_loss / count avg_psnr total_psnr / count avg_ssim total_ssim / count return avg_loss, avg_psnr, avg_ssim重要提示skimage的ssim函数在较新版本中multichannel参数已改为channel_axis。确保你的代码与库版本匹配。另外PSNR/SSIM的计算非常耗时尤其是在验证集较大时。在实际项目中可以考虑每N个epoch计算一次或者在训练时不计算仅在测试时详细评估。3.6 主训练流程与模型保存将以上所有部分串联起来形成主训练函数。def main_train(): num_epochs 200 best_psnr 0.0 for epoch in range(1, num_epochs 1): # 训练一个epoch train_loss train_one_epoch(model, train_loader, optimizer, criterion, device, epoch) # 更新学习率 scheduler.step() # 每5个epoch验证一次 if epoch % 5 0: val_loss, val_psnr, val_ssim validate(model, val_loader, criterion, device) print(f[Validation] Epoch: {epoch}, Loss: {val_loss:.6f}, PSNR: {val_psnr:.4f}, SSIM: {val_ssim:.4f}) # 根据PSNR保存最佳模型 if val_psnr best_psnr: best_psnr val_psnr torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), scheduler_state_dict: scheduler.state_dict(), best_psnr: best_psnr, }, best_swinir_model.pth) print(fBest model saved at epoch {epoch} with PSNR {val_psnr:.4f}) # 保存最近的检查点 if epoch % 20 0: torch.save({ epoch: epoch, model_state_dict: model.state_dict(), optimizer_state_dict: optimizer.state_dict(), scheduler_state_dict: scheduler.state_dict(), }, fcheckpoint_epoch_{epoch}.pth) if __name__ __main__: main_train()4. 测试推理与结果可视化训练完成后我们加载最佳模型进行单张图像测试。def inference_single_image(model_path, lr_image_path, output_path, scale4): 对单张低分辨率图像进行超分重建并保存。 # 加载模型 checkpoint torch.load(model_path, map_locationcpu) model create_swinir_model(upscalescale, training_patch_size48) # 注意训练patch size model.load_state_dict(checkpoint[model_state_dict]) model.eval() model.to(device) # 读取并预处理LR图像 lr_img Image.open(lr_image_path).convert(RGB) lr_tensor transforms.ToTensor()(lr_img).unsqueeze(0).to(device) # [1, C, H, W] # 推理 with torch.no_grad(): sr_tensor model(lr_tensor) # 后处理张量 - PIL图像 sr_tensor sr_tensor.squeeze(0).cpu() # [C, H, W] sr_tensor torch.clamp(sr_tensor, 0, 1) # 确保范围在[0,1] sr_img transforms.ToPILImage()(sr_tensor) # 保存结果 sr_img.save(output_path) print(fSuper-resolved image saved to {output_path}) # 可选可视化对比 fig, axes plt.subplots(1, 2, figsize(10, 5)) axes[0].imshow(lr_img) axes[0].set_title(Low-Resolution Input) axes[0].axis(off) axes[1].imshow(sr_img) axes[1].set_title(SwinIR Output (4x)) axes[1].axis(off) plt.show() # 使用示例 inference_single_image(best_swinir_model.pth, test_lr.png, test_sr.png)4.1 处理大尺寸图像与边界伪影在实际应用中你可能会遇到尺寸很大的测试图像直接输入模型可能导致显存溢出。常见的策略是重叠分块处理将大图分割成有重叠的小块分别输入模型再将输出的小块拼接起来并采用加权融合的方式平滑重叠区域以避免边界处出现明显的接缝。此外SwinIR等基于窗口的Transformer模型可能对输入尺寸有整除要求如窗口大小的倍数。如果输入尺寸不满足简单的做法是对输入图像进行填充padding在模型输出后再裁剪掉填充的部分。timm中的SwinIR实现通常已经内部处理了这一点但了解这个原理有助于你调试输出尺寸不匹配的问题。这份代码从数据加载到训练、验证、测试形成了一个完整的闭环。它省略了分布式训练、复杂的日志系统、高级的数据增强等生产级功能但核心逻辑是完整且自洽的。你可以以此为基础根据实际需求添加数据预处理、更复杂的损失函数、学习率热身、指数移动平均等高级技巧。最重要的是通过阅读和运行这段代码你能清晰地把握SwinIR模型训练与测试的每一个关键步骤这才是“逻辑完整易懂”的价值所在。本文还有配套的精品资源点击获取