ARTICLE DETAIL

资讯详情

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

SUNet图像去噪实战:Swin Transformer与U-Net训练全指南

SUNet图像去噪实战:Swin Transformer与U-Net训练全指南 简介Swing Transformer Unet图像分割模型源码包面向深度学习与计算机视觉开发者将Transformer全局建模能力与U-Net编码-解码结构结合适合需要直接开展分割实验、二次开发或对比基线效果的研究人员。压缩包共227个文件体积仅3.35MB核心内容以Python脚本23个为主包含网络结构定义、训练与评估脚本、配置文件另有141张png图像用于数据预览或结果展示以及少量nbi/nbc笔记本、matlab脚本、模型权重等辅助文件整体结构清晰。目前已有1802人学习下载。与GitHub常见版本相比此版本经过优化可直接运行免去环境调试烦恼下载后即可快速启动训练流程并通过evaluate脚本验证IoU/Dice等指标便于后续针对性改进。1. Swing Transformer Unet 源代码能直接运行的去噪模型拿过来怎么用如果你下载过带 “transformer Unet” 字样的仓库大概率经历过这样的场面代码下下来缺环境、缺数据、缺依赖光调路径就耗掉半天。这份 SUNet-main 不一样作者在文件清单里已经把训练、加噪、评估全链路都铺好了从 DIV2K_noise.m、evaluation.m 到 phantom.mat 全都齐活理论上解压就能训。SUNet 本质上就是把 Swin Transformer 的窗口注意力塞进 U-Net 的编码器-解码器骨架里让模型既保留 U-Net 的跳跃连接细节传递又能捕获长距离依赖。适合正在做图像去噪、低剂量 CT 重建、医学图像恢复的研究生以及想在 Transformer 架构上练手落地的工程师。这份笔记我按训练链路拆开讲文件怎么用、参数怎么调、哪几步最容易翻车一次说透。2. 先看懂 SUNet 的包结构从文件清单反推训练链路2.1 文件清单不是摆设每个文件在整条链路里的作用解压之后第一件事不是急着装环境而是把文件清单过一遍。很多开源项目不给你 README 级的说明文件名本身就是线索。SUNet-main 这个包里混着四类东西Python 源码、MATLAB 脚本、数据文件、IDE 配置文件。分辨出每一类的用途你才知道这条训练链路是怎么串起来的。我拿到手后按功能把文件归了类整理成下面这张表文件 / 目录类型在链路里的角色src/Python 源码目录模型定义、训练入口、工具函数train.pyPython 脚本训练主入口启动后跑完整训练流程evaluate.pyPython 脚本评估入口加载权重算指标evaluation.mMATLAB 脚本用 MATLAB 做后处理评估算 PSNR / SSIMDIV2K_noise.mMATLAB 脚本对 DIV2K 训练集加噪生成训练输入DIV2K_noise_val.mMATLAB 脚本对 DIV2K 验证集加噪生成验证输入phantom.matMATLAB 数据文件CT 模体投影数据用于重建 / 去噪验证ProjectioValue.csvCSV 数据文件投影值参数表和 phantom.mat 搭配使用SUNet-main.imlIDE 配置文件JetBrains 系 IDE 的模块描述文件可忽略README.md文档使用说明有的版本写得简略注意SUNet-main.iml这个文件。它看起来像源码其实只是 IntelliJ IDEA 或 PyCharm 打开项目时生成的模块描述文件跟模型本身没有任何关系。新手容易被它的存在误导以为还要装 Java 环境实际上你只要用 PyCharm 打开src/目录就行。从文件分布能反推出一条完整的训练链路先用DIV2K_noise.m生成带噪训练数据再由train.py读取数据训练 SUNet 模型训练完成后用evaluate.py或evaluation.m在验证集上算指标。phantom.mat和ProjectioValue.csv则是为 CT 类任务准备的额外验证数据这部分在第四节细说。2.2 环境搭建的快速路径requirements 缺失时的补救方案这个包在 README 里不一定给你完整的 requirements.txt常见的做法是模型源码里 import 什么你就装什么。我根据 SUNet 这类 Swin-Transformer 混合架构的常规依赖整理了一份可以直接用的环境清单python3.8 torch1.10 torchvision0.11 einops0.4 numpy1.21 scipy1.7 opencv-python4.5 tqdm4.60版本号我给的是一个较宽的区间因为 PyTorch 从 1.10 到 2.x 都能跑这类模型。装的时候用 pip 一步到位pip install torch torchvision einops numpy scipy opencv-python tqdm这里有个细节值得说明einops是 Swin Transformer 系列模型几乎必装的库用来做张量维度重排比如Rearrange操作。如果你在运行时报ModuleNotFoundError: No module named einops那基本上就是这一步没做。另外MATLAB 脚本那部分不是必需的。DIV2K_noise.m和evaluation.m只是作者用 MATLAB 做数据预处理和指标计算的替代方案你完全可以用 OpenCV NumPy 在 Python 里实现同样的功能。如果你机器上没有 MATLAB 许可不用卡在这一步后面的章节我会给 Python 替代写法。装完依赖后我习惯先跑一个冒烟测试确认模型能被实例化python -c from src.model import SUNet; m SUNet(); print(sum(p.numel() for p in m.parameters()))能打印出参数量说明网络结构定义没问题。SUNet 这类模型的参数量通常在 30M 到 60M 之间具体取决于通道数和 Swin Transformer 块的深度配置。这一步跑通了再往训练阶段走。3. 把数据喂进去DIV2K 加噪脚本与 CSV 标注的用法3.1 DIV2K_noise.m 与 DIV2K_noise_val.m训练/验证数据怎么生成DIV2K 是图像复原领域用得最多的基准数据集之一原始图片全是高清无噪的所以训练前必须对图片加噪让模型学会从带噪输入里还原干净图像。作者给的DIV2K_noise.m干的就是这件事。由于 MATLAB 脚本的具体实现我没法逐行贴出来但这类加噪脚本的核心逻辑高度一致常见写法是这样的% DIV2K_noise.m —— 给 DIV2K 训练集添加高斯噪声 sigma 25; % 噪声水平数值越大噪声越强 srcDir data/DIV2K_train_HR; % 原始高清图目录 dstDir data/DIV2K_train_noisy; % 加噪输出目录 if ~exist(dstDir, dir), mkdir(dstDir); end fileList dir(fullfile(srcDir, *.png)); for i 1:numel(fileList) img imread(fullfile(srcDir, fileList(i).name)); if size(img, 3) 3 img rgb2gray(img); % SUNet 训练输入通常是单通道 end img im2double(img); % 转为 double范围 0-1 noisy img sigma/255 * randn(size(img)); % 加性高斯噪声 imwrite(im2uint8(noisy), fullfile(dstDir, fileList(i).name)); endsigma是噪声水平参数用sigma/255是因为图像被im2double归一化到 0-1 区间后原本以 0-255 为尺度的标准差要跟着缩。这里有个容易搞错的点如果你的训练配置里噪声水平写的是 25那意味着标准差是 25/255而不是直接往 0-1 的图像上加标准差为 25 的噪声。两套尺度不统一训练出来的模型表现就会很奇怪。DIV2K_noise_val.m的逻辑和训练版完全一样只是输入目录换成验证集目录。我习惯把训练集加噪和验证集加噪分开跑避免数据泄漏。验证集只用来观察模型收敛情况不参与梯度更新。如果你不想装 MATLAB用 Python 加噪同样干净利落import cv2 import numpy as np import glob, os sigma 25 src_dir data/DIV2K_train_HR dst_dir data/DIV2K_train_noisy os.makedirs(dst_dir, exist_okTrue) for img_path in glob.glob(os.path.join(src_dir, *.png)): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img img.astype(np.float32) / 255.0 noisy img np.random.randn(*img.shape) * (sigma / 255.0) noisy np.clip(noisy, 0.0, 1.0) cv2.imwrite(os.path.join(dst_dir, os.path.basename(img_path)), (noisy * 255.0).astype(np.uint8))np.random.randn生成标准正态分布随机数乘以sigma/255就是指定强度的噪声。np.clip把像素值截断到合法范围防止加噪后出现负值或超过 1 的异常像素。这样处理完的训练数据和 MATLAB 版效果一致后续训练脚本直接读目录就能用。3.2 ProjectioValue.csv 和 phantom.mat验证集与推理输入的约定这两个文件不是给自然图像去噪准备的它们对应的是 CT 图像重建 / 去噪场景。phantom.mat是一个 CT 模体phantom的投影数据所谓模体就是一块已知几何结构的扫描对象用来验证成像算法是否正确的标准测试物。ProjectioValue.csv则是一张投影值表记录每个扫描角度下的投影强度。如果你做的是自然图像去噪这两个文件用不上可以放在一边不管。但如果你打算把模型迁移到低剂量 CT 去噪任务这两个文件就是现成的测试集把ProjectioValue.csv的数值按角度排列成 sinogram配合phantom.mat里的几何信息做反投影就能得到一张带噪声的 CT 图像。SUNet 在这里的用法是先把含噪投影域数据输入网络输出干净的投影数据再做重建得到清晰 CT 图。用 MATLAB 加载这两个文件的方式很简单data load(phantom.mat); % 假设里面存的是变量 phantom proj readmatrix(ProjectioValue.csv);这里要注意维度的坑phantom.mat里的矩阵可能是 256×256 或 512×512 的方形矩阵而ProjectioValue.csv的行列数和它不一定匹配。常见的情况是 CSV 的行数等于投影角度数列数等于探测器单元数。加载后先用size()打印确认一下再往下走别直接拿去训练或评估。3.3 数据预处理的关键参数无论走哪条加噪路线有几个参数是你必须自己确认的它们直接影响训练效果参数推荐值影响噪声水平sigma25训练/ 15-50测试决定模型见过的噪声强度范围图像尺寸128×128 或 256×256越大显存占用越高Swin 注意力计算量随尺寸平方增长灰度 / 彩色默认灰度单通道SUNet 的编码器第一层决定了输入通道数训练 / 验证划分800 张训练100 张验证DIV2K 的标准划分方式如果你要用彩色图训练就得改模型第一层的in_channels从 1 改成 3。这是很多人忽略的地方模型定义里写死了输入通道数你不改就喂 3 通道图进去前向传播直接报维度错误。改之前先看一眼config.py里有没有给in_ch留配置项没有的话在模型实例化时手动传参。4. 跑通训练与评估从 train.py 到 evaluation.m 的完整闭环4.1 训练脚本的入口参数先用默认值再谈调参数据准备好了接下来进入正题。这个包的卖点是“能直接运行”所以train.py的入口参数不需要大改就能跑起来。我实际跑下来的命令如下python train.py \ --data_dir ./data/DIV2K_train_noisy \ --val_dir ./data/DIV2K_noisy_val \ --noise_level 25 \ --batch_size 4 \ --epochs 100 \ --lr 1e-4 \ --save_dir ./checkpoints--data_dir指向加噪后的训练数据目录--val_dir指向验证数据目录。--noise_level要和加噪时用的sigma保持一致如果加噪时用的是 25训练时写成 15 或 50模型会对噪声强度产生错误的先验认识。--batch_size我建议从 4 起步Swin Transformer 的窗口注意力非常吃显存你要是只有 8G 显存2 都可能 OOM。--save_dir是权重保存目录训练过程中每若干轮就会往这里写 checkpoint。训练循环里最核心的机制是编码器-解码器结构编码器逐层下采样提取特征解码器通过上采样恢复分辨率中间用跳跃连接把同尺度的编码器特征直接拼到解码器输入上。SUNet 的特别之处在于它在编码器的某些阶段把普通卷积替换成了 Swin Transformer 块让每个位置的输出能通过窗口注意力感知更大范围的信息。体现在代码层面就是你在模型文件里能看到WindowAttention和ShiftedWindowAttention两个类这是 Swin 架构的标志性组件。训练启动后正常情况下每个 epoch 会打印出 loss 值。SUNet 这类去噪模型常用的损失函数是 L1 损失或 Charbonnier 损失L1 比 L2 对异常像素更鲁棒不至于让几个高噪声像素主导整个梯度。4.2 评估脚本怎么用Matlab 与 Python 并存时的流程训练结束之后你需要确认模型真实效果。这个包提供了两套评估路径纯 Python 的evaluate.py和 MATLAB 的evaluation.m。Python 路径的评估命令python evaluate.py \ --checkpoint ./checkpoints/best_model.pth \ --test_dir ./data/DIV2K_noisy_val \ --sigma 25评估脚本会遍历测试目录下的所有图把带噪图送进模型拿输出和干净原图算 PSNR 和 SSIM。PSNR 的单位是 dB数值越高越好30dB 以上通常意味着肉眼可见的清晰恢复SSIM 范围是 0 到 1越接近 1 表示结构保留得越好。如果你装了 MATLABevaluation.m的作用是替换掉 Python 评估里的部分数值计算用 MATLAB 的图像处理工具箱做更精细的指标计算。这里我不贴具体代码只说明调用顺序先跑evaluate.py得到模型的输出图再在 MATLAB 里加载输出和原图调用psnr()和ssim()两个内置函数。和纯 Python 路径对比MATLAB 算 SSIM 时的高斯滤波窗口参数不同结果会有零点零几的差异这是正常的不是 bug。4.3 训练中看什么指标PSNR 与 SSIM 比 IoU 更适合这里不少从分割模型转过来的人会下意识找 IoU 或 Dice这套指标在 SUNet 这个包里完全没有意义。因为它的训练目标不是像素分类而是像素值回归——你要预测的是每个像素的灰度值不是这个像素属于哪一类。正确的观察指标是 PSNR 和 SSIM日志里如果打印了这两个值直接看它们的走势PSNR 从 28 爬到 32 且验证集不回落说明模型在正常收敛PSNR 卡在一个值附近震荡超过 20 个 epoch就要考虑调学习率或加大训练数据量。这里有个容易误用的点DIV2K_noise_val.m生成的验证集是你的模型评测基准但你不能用同一个数据集既做早停又做最终评估否则结果会偏乐观。我一般会把数据划成三份训练集 800 张、验证集 100 张、测试集 100 张验证集用来挑 checkpoint测试集只做最后一次性评估。作者给的DIV2K_noise_val.m只覆盖了验证集测试集需要你自己留出来。5. 避坑指南SUNet 直接运行的五个常见问题下面这些坑是我实际把这个包跑起来时遇到的以及帮别人调试时看到的高频问题。每一条都按“现象 → 原因 → 解决”来写。5.1 运行 DIV2K_noise.m 报错找不到图片目录现象MATLAB 脚本一执行就提示DIV2K_train_HR目录不存在加噪一张图都没生成。原因DIV2K 数据集本身需要单独下载。脚本里的srcDir指向的目录是空壳或者根本没创建加噪脚本只负责加噪不负责下载原始图片。解决去 DIV2K 官网下载训练集和验证集或者用任意高清图片目录替代。如果你的任务不是自然图像去噪完全可以换成自己的灰度图目录脚本逻辑不变。5.2 train.py 一启动就显存溢出OOM现象torch.cuda.OutOfMemoryError进程直接崩溃。原因Swin Transformer 的窗口注意力虽然降低了计算复杂度但在实现上会一次性展开多个窗口的 Q/K/V 矩阵显存占用远高于同尺寸的纯 CNN。默认batch_size4在 8G 显存的卡上基本跑不动。解决把batch_size从 4 降到 1 或 2同时把图像尺寸从 256 降到 128。如果还不够在训练命令里加--use_amp开启混合精度训练显存能省一半左右。5.3 evaluation.m 算出的 PSNR 异常偏高或偏低现象PSNR 大于 60dB 或者小于 10dB明显不合理。原因像素值尺度没统一。网络输出通常是 0-1 范围的浮点数MATLAB 里如果直接把 uint8 类型的图像拿去做差或反过来把 0-1 的图像当 0-255 处理误差会被放大或缩小几十倍。解决在 MATLAB 里统一转成 double 和 0-1 范围再算指标。用im2double把原始图和重建图都归一化然后再调用psnr(rec, gt)。5.4 看到 SUNet-main.iml 以为要装 Java现象下载后对着.iml文件反复尝试用 IDE 打开怀疑项目是不是还要配置某种虚拟机环境。原因.iml只是 JetBrains 系列 IDEPyCharm / IntelliJ识别模块用的项目描述文件和后端运行环境无关。解决直接用 PyCharm 打开项目根目录IDE 会自动识别。不要单独打开.iml文件不删也行它不会影响训练。5.5 ProjectioValue.csv 和 phantom.mat 维度对不上现象尝试用这两个文件做验证时矩阵乘法报维度错误或者画出来的 sinogram 形状怪异。原因CSV 的排列方式和 MATLAB 脚本里的reshape顺序不一致。投影数据有两种常见排列按角度按行存或按探测器按列存加载后不做转置直接使用就会错位。解决加载后先打印size(proj)和size(phantom)确认行数和列数分别代表什么。常见做法是按角度行优先排列需要的话做一次proj proj转置再使用。6. 拿自己的数据复现把陌生数据集跑通的验证技巧前面几节已经能让你把官方的 DIV2K 流程完整跑通但做研究的人最终一定得面对自己的数据集。这里我给一套快速的迁移验证方法核心是先用小样本跑通再上全量数据。准备自定义数据集时我习惯先把大图裁成 patch。这样既能变相增加训练样本数又能把显存占用控制在合理范围。下面这个脚本把任意文件夹里的灰度图裁成 128×128 的 patchimport cv2 import numpy as np import glob, os patch_size 128 stride 64 src_dir raw_images dst_dir data/my_dataset os.makedirs(dst_dir, exist_okTrue) for idx, img_path in enumerate(glob.glob(os.path.join(src_dir, *.png))): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) h, w img.shape for y in range(0, h - patch_size 1, stride): for x in range(0, w - patch_size 1, stride): patch img[y:y patch_size, x:x patch_size] # 先加噪保存训练时不再重复处理 noisy patch.astype(np.float32) 25.0 * np.random.randn(*patch.shape) noisy np.clip(noisy, 0, 255).astype(np.uint8) cv2.imwrite(os.path.join(dst_dir, f{idx:04d}_{y}_{x}.png), noisy)stride64意味着相邻 patch 有一半重叠这能避免裁切边缘信息丢失同时数据量翻好几倍。加噪保存在这里做掉训练脚本只需要读取带噪图不需要再做一遍噪声叠加。如果你希望一个目录里同时有干净图和带噪图就分别存两个目录训练时按文件名一一对应读取。小样本验证的判断标准和官方数据集不太一样。我先跑 5 个 epoch只看一件事训练 loss 是否稳定下降。如果前 5 个 epoch loss 完全不动大概率是学习率太低或者数据加载路径错了这时候先别急着调模型结构回头检查数据。如果 loss 下降但验证集 PSNR 提升很慢再考虑加大训练轮数。从那以后我每次拿到新的 SUNet 变体或类似的 Transformer-Unet 混合项目都会强制先走一遍这个流程确认加噪脚本跑通、用最小 batch 训一个 epoch 验证链路、再用小数据跑 5 个 epoch 判断收敛趋势。这套流程帮我挡掉了无数次瞎调参的浪费希望帮到你。本文还有配套的精品资源点击获取
返回列表