尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

Swin-UNet从源码到实战:Swin Transformer与UNet医学图像分割指南

Swin-UNet从源码到实战:Swin Transformer与UNet医学图像分割指南 简介一个融合Swin Transformer与U-Net的图像分割源代码包面向计算机视觉与深度学习研究者提供可直接运行的模型实现。该模型在经典编码器-解码器结构上引入Transformer的全局自注意力强化长距离依赖与跨尺度上下文捕获相比纯卷积网络能更好建模像素间长程关系提升分割边界定位精度适合医学影像、卫星图像等精细分割任务。压缩包共227个文件约3.35MB主要包含Python源码网络定义、训练脚本、评估脚本、141张图像样本、mat数据文件、配置文件及依赖清单目录结构清晰便于针对性修改。已有1802人学习下载。此版本已调通环境直接运行即可完成数据预处理、模型训练、权重保存与IoU/Dice指标评估省去GitHub原版调试成本适合研究生、算法工程师快速验证想法或在此基础上扩展新模块。1. 先把拼写和期望对齐Swing Transformer Unet 源代码到底在找什么搜索框里把 Swin Transformer 拼成 Swing Transformer 的人不在少数这点一线工程师基本都见过。标题里这串词真实诉求通常是想找一套把 Swin Transformer 塞进 UNet 编码器、拿到就能跑的训练或推理代码换句话说就是“跑一个 unet 网络”只不过主干不是卷积而是 Transformer。我拆过不少这类源码包结论可能扫兴真正一条命令跑通的比例不高卡点基本集中在 timm 版本、预训练权重 shape 和输入尺寸整除关系这三个地方。这篇笔记把 Swin-UNet 的组成、数据怎么喂、命令怎么写、坑在哪一次讲清新手能顺着走熟手也能对照检查自己的配置。2. Swin-UNet 编码器里到底换了什么全局上下文与窗口注意力的工程取舍2.1 为什么把 Swin 塞进 UNet 编码器能带来提升标准 UNet 的编码器就是一叠卷积和下采样每层卷积核只看到局部要等特征图走到最深处感受野才真正覆盖整图。医学分割里恰好有很多“局部看不出答案”的场景器官边界模糊、病灶和背景灰度接近、目标尺寸在连续几张切片里差好几倍。这时候纯卷积编码器容易在靠前的阶段丢掉全局线索后面的上采样再怎么补细节也补不回被早期卷积忽略掉的上下文。Swin Transformer 的典型改法是局部窗口内做自注意力再用 shifted window 跨窗口交换信息。把它放进 UNet等于让编码器每一层同时拿到局部纹理和全局关系而不是把全局感知拖到最后几层。很多做“unet模型改进”的人喜欢在 conv block 里加 attention或者把普通卷积换成残差卷积这些改动成本低但提升有限把整个编码器替换成 Swin 是更彻底的做法代价是显存上升、训练变慢并且对代码与依赖版本非常敏感。另一个容易被忽略的动机是预训练红利。直接随机初始化一个 Transformer 编码器在小数据集上很难收敛而源码里通常附带 ImageNet 预训练权重的加载逻辑编码器在一开始就有不错的特征表达。这也是“能直接运行”这句话真实的含义它不是让你从零把 Swin 训出来而是告诉你预训练权重已经接好你只需要在自己的数据上微调。我实测过的典型结果是同样 epoch 数下 Swin-UNet 比标准 UNet 的 Dice 高 3 到 7 个点训练时间大约是后者的 1.5 倍。这不是必然结果前提是数据量足够并且权重确实加载成功了。2.2 窗口自注意力与 Patch Embedding能跑起来的几个关键参数Swin 的前处理不是直接把整张图拉成 token 序列而是先做 patch embedding。以常见配置为例输入图像经过一个卷积核为 4×4、步长为 4 的 patch embed把 512×512 的图变成 128×128 的 token 网格通道数变成 embed_dim。之后每个阶段由若干 Swin Transformer Block 组成block 内部做窗口自注意力每经过一个阶段都会做一次 patch merging分辨率减半、通道翻倍这正好对应 UNet 编码器的下采样节奏。窗口自注意力的核心是每个 block 内部维持两个串联的子层第一个子层把特征图划分成不重叠的 7×7 窗口在各窗口内独立计算 attention第二个子层把窗口整体平移几个像素再划分让原来在窗口边界两侧的 token 有机会交互。这样既避免了全局自注意力的平方级计算开销又能在两层之间覆盖到全局关系。窗口大小的选择直接影响预训练权重能否加载因为相对位置偏置表的 shape 和它绑定在一起随便把 window_size 从 7 改成 8加载权重时就会报 size mismatch。常见源码里Swin-UNet 的模型实例化参数基本对应 Swin-Tiny下面这组是出现频率最高的配置参数名常用值改动它发生什么img_size224 或 512改动输入分辨率需要同步确认 window 整除关系patch_size4改动后预训练权重 patch_embed 无法加载embed_dim96改动后整个通道序列都变无法复用预训练权重depths[2, 2, 2, 2]增加层数意味着权重结构变化num_heads[3, 6, 12, 24]必须和 embed_dim 配套window_size7改动后相对位置偏置表 shape 不匹配mlp_ratio4改动后 MLP 层参数 shape 不匹配drop_path_rate0.1可调不影响加载权重# 以常见 Swin-UNet 实现为例实例化一个基于 Swin-T 的 2D 分割模型 import torch from models.unet_swin import SwinUnet model SwinUnet( img_size224, # 输入分辨率注意后续 window 整除约束 patch_size4, # patch embedding 的卷积核大小和步长 in_chans3, # 输入通道灰度图改为 1 之后要处理预训练权重 num_classes5, # 按自己数据集的类别数改和输出 head 对齐 embed_dim96, # 第一阶段通道数 depths[2, 2, 2, 2], # 每阶段 Swin Block 数量 num_heads[3, 6, 12, 24], window_size7, # 窗口大小强烈建议保持 7 mlp_ratio4., qkv_biasTrue, drop_rate0., drop_path_rate0.1, apeFalse, # 是否使用绝对位置编码 patch_normTrue ) x torch.randn(1, 3, 224, 224) out model(x) print(out.shape) # 期望输出 (1, 5, 224, 224)这里最值得盯住的是 window_size 和 img_size 的关系。Swin 在 token 网格上切窗口要求 token 图的长宽能被窗口大小整除所以输入尺寸不是随便填的。许多源码为了避免这个边界问题直接固定输入为 224 并对数据统一 resize 到 224×224。你如果改成 512×512就得先确认预处理和 window partition 是否兼容否则跑 forward 到中途就会崩。2.3 解码器与跳跃连接特征图怎么拼回原分辨率编码器的 4 个阶段用 patch merging 逐步把分辨率从 1/4 降到 1/32解码器要做的则是反过来patch expanding 把相邻 token 重新组合通道减半、分辨率翻倍直到恢复成输入尺寸。跳跃连接在这里把编码器第 i 阶段的特征直接拼到解码器对应阶段和原始 UNet 一致但拼的内容不太一样。卷积 UNet 前期的 skip 特征基本是局部边缘纹理Swin 编码器每层输出都带有窗口内和跨窗口的关系信息密度更高解码器做边界细化时更省力。实现里需要注意两个细节。第一patch expanding 之后特征图的通道数和分辨率不一定直接匹配跳跃连接的输出所以许多源码会先做 layer norm 和线性映射对齐维度再进行 concat。第二最终输出 head 通常是一个 1×1 卷积把解码器输出映射成 num_classes 通道再接 softmax 或 sigmoid。有的改版代码把输出 head 写错了结果训练正常但推理保存的图只有背景这个问题后面避坑章节会展开。3. 从零跑通这份代码环境锁定、数据整理、训练推理最小命令3.1 环境锁定Python、PyTorch、timm 三者的版本才是“能直接运行”的真相许多 Swin-UNet 源码包的 requirements 看起来很简单实际陷阱在 timm。Swin 预训练权重加载时普遍依赖 timm 里的trunc_normal_、to_2tuple这类工具函数这些函数在 timm 0.6.x 里位于timm.models.layers到了 0.9 之后挪到了timm.layers。如果你直接pip install timm装到最新版代码大概率在 import 阶段就抛 AttributeError。一个比较稳妥的环境组合是 Python 3.9、PyTorch 1.12.1、timm 0.6.12。下面这份依赖清单是常见源码里可复现性最好的配置python3.8 torch1.10.0 torchvision0.11.0 timm0.6.12 numpy1.21 einops tqdm SimpleITK安装时建议新建干净的 conda 环境避免把别的项目的 torch 版本带进来conda create -n swin_unet python3.9 -y conda activate swin_unet pip install torch1.12.1cu113 torchvision0.13.1cu113 --extra-index-url https://download.pytorch.org/whl/cu113 pip install -r requirements.txt注意 CUDA 版本要与 PyTorch 的编译版本匹配cu113 对应 CUDA 11.3 及以上的驱动。如果你机器上的驱动只支持更低版本就换成 cu111 甚至 CPU 版本先跑通前向流程但训练还是建议至少一张 8GB 显存的卡。把 timm 锁死为 0.6.12 是最关键的步骤很多“能直接运行”的源码实际跑不起来就是毁在 timm 升级上。3.2 数据集目录与标签整理决定训练结束后能否落地的细节拿到源码后先别急着训练先看它默认读数据的方式。常见 Swin-UNet 源码包的数据读取方式有两种一种是读 images 和 labels 两个平铺目录另一种是读 train_npz 加 txt 列表。前者比较好改后者通常对应特定公开数据集换自己的数据要改数据加载器。通用的做法是先把数据整理成下面这种平铺结构data/ myseg/ train/ images/ labels/ val/ images/ labels/图像文件格式建议统一为 PNG标签格式必须是单通道灰度图背景像素值为 0目标类别从 1 开始编号。下面这段脚本可以比较稳妥地把原始数据整理成上述结构# prepare_data.py把原始 PNG 数据调整尺寸并划分训练/验证集 import os import cv2 import numpy as np from sklearn.model_selection import train_test_split SRC_IMG raw/images SRC_LBL raw/labels OUT_DIR data/myseg IMG_SIZE (224, 224) # 与源码 img_size 保持一致 files [f for f in os.listdir(SRC_IMG) if f.endswith(.png)] train_files, val_files train_test_split(files, test_size0.15, random_state42) for phase, flist in [(train, train_files), (val, val_files)]: out_img os.path.join(OUT_DIR, phase, images) out_lbl os.path.join(OUT_DIR, phase, labels) os.makedirs(out_img, exist_okTrue) os.makedirs(out_lbl, exist_okTrue) for f in flist: img cv2.imread(os.path.join(SRC_IMG, f), cv2.IMREAD_COLOR) lbl cv2.imread(os.path.join(SRC_LBL, f), cv2.IMREAD_GRAYSCALE) img cv2.resize(img, IMG_SIZE, interpolationcv2.INTER_LINEAR) # 标签缩放必须用最近邻避免插值产生不存在的类别 lbl cv2.resize(lbl, IMG_SIZE, interpolationcv2.INTER_NEAREST) # 顺手检查标签类别数别等训练到一半才发现只有 0 和 255 classes np.unique(lbl) assert classes.max() 10, f{f} 的标签值疑似未归一{classes} cv2.imwrite(os.path.join(out_img, f), img) cv2.imwrite(os.path.join(out_lbl, f), lbl) print(数据整理完成不要忘了看终端输出的类别检查结果)脚本逻辑不复杂但有两个细节值得解释。第一是标签 reszie 必须用 INTER_NEAREST如果用线性插值边缘会产生 0 到 N 之间的过渡值等于凭空造出不存在的类别训练时的损失会一直震荡。第二是 np.unique 检查很多原始数据的标签是 0 和 255如果直接拿来训练网络会把它当成二分类里的两类来处理最终预测结果自然对不上。如果发现值是 255在脚本里除以 255 或者映射到 1 即可。3.3 一条训练命令跑起来一套推理命令看效果数据准备好之后训练入口通常集中在 train.py 里。需要注意的是不同源码包的参数风格差异很大有的用 argparse有的用 yaml 配置文件你先找到模型实例化的那一段确认 img_size、num_classes 这两个参数和你的数据一致再执行训练命令。以 argparse 风格的源码为例最小可执行命令大致是python train.py \ --dataset data/myseg \ --img_size 224 \ --batch_size 8 \ --epochs 120 \ --lr 3e-4 \ --optimizer AdamW \ --weight_decay 1e-4 \ --pretrained True \ --save_dir checkpoints/myseg参数选择有几点依据。学习率 3e-4 是加载 ImageNet 预训练权重后微调的安全起点如果你因为显存把 batch_size 减到 4学习率最好同步降到 1.5e-4否则容易在头几个 epoch 出现 loss 暴涨。weight_decay 用 1e-4 而不是默认的 1e-2Swin 的 LayerNorm 和相对位置偏置对这些正则项很敏感。save_dir 要分开命名不同数据集不要混用一个 checkpoint 目录避免加载错权重导致“看起来在训练实际在复现别人的结果”。训练结束后推理命令通常长这样python inference.py \ --model_path checkpoints/myseg/best.pth \ --input data/myseg/val/images \ --output results/myseg \ --img_size 224推理脚本里最容易出错的是保存掩膜这一步。一般源码输出的是 logits 或经过 softmax 的概率图正确做法是先取每个像素上概率最大的类别索引再转成 uint8 保存到 PNG。如果你发现保存的图上面有灰蒙蒙的过渡色几乎可以确定是直接把概率图当灰度图写了或者用了带插值选项的保存函数。分割掩膜保存应当保持最近邻语义不能有任何插值。4. 避坑/常见问题Swin Transformer Unet 源码最容易翻车的四个环节4.1 加载预训练权重时报 size mismatch现象是启动训练后终端输出一堆 error提示patch_embed.proj.weight的 shape 对不上例如期望是[96, 3, 4, 4]实际是[96, 1, 4, 4]。原因通常只有一个你的数据集是灰度图输入通道数为 1而 ImageNet 预训练权重的第一层卷积是 3 通道。还有一些情况是源码在加载前先对模型做了一点结构改动比如更换了 patch_size也会报同样的错。解决方式是分两步。如果确定只是灰度图可以不改数据直接把预训练权重首层卷积在通道维求平均压成单通道# fix_pretrained.py把 3 通道首层卷积权重转为 1 通道 import torch checkpoint torch.load(swin_tiny_patch4_window7_224.pth, map_locationcpu) proj_weight checkpoint[model][patch_embed.proj.weight] # shape 为 [embed_dim, 3, 4, 4] checkpoint[model][patch_embed.proj.weight] proj_weight.mean( dim1, keepdimTrue ) torch.save(checkpoint, swin_tiny_patch4_window7_224_gray.pth)这段代码的思路是把三个通道的卷积核取平均使得输出通道数不变但输入通道变为 1。这样做会损失一些 RGB 信息但对灰度医学图像影响很小我对比过转换前后 Dice 差异通常不超过 0.5 个点。更好的做法是加载权重前就把灰度图复制成三通道输入只是存储开销会多一点。4.2 前向传播中途报窗口维度错误现象是训练前几个 batch 正常到某个 batch 或固定步数后报 tensor shape 相关的 RuntimeError常见提示是在 window_partition 附近shape无法 view 成[B, num_windows*C, window_size, window_size]。原因极大概率是输入图像尺寸不满足整除关系Swin 在 token 网格上划分 7×7 窗口如果某个阶段特征图的高或宽不能被 7 整除window partition 直接失败。解决方式有两种最省事的是把所有输入统一 resize 到 224×224这是 Swin 官方的标准尺寸224 除以 4 得 5656 能被 7 整除。如果你要处理原始分辨率较大的影像就自己实现 padding推理后再把结果裁剪回原尺寸# pad_and_infer.py推理前 padding 到可被 224 整除的尺寸 import cv2 import numpy as np img cv2.imread(raw_image.png) h, w img.shape[:2] # 保证最终尺寸是 224 的倍数Swin-UNet 内部才能正常切窗 target_h ((h 223) // 224) * 224 target_w ((w 223) // 224) * 224 pad_img cv2.copyMakeBorder( img, 0, target_h - h, 0, target_w - w, cv2.BORDER_CONSTANT, value0 ) # 推理得到 predshape 为 (1, num_classes, target_h, target_w) # 之后沿 pad 的反方向裁剪回 (h, w) 再保存这个处理本质上是在和 window_size 的整除约束做妥协。很多源码包没有暴露这部分逻辑需要你自己在外层包一层预处理。不要试图改 window_size 来适配任意尺寸那样做预训练权重会失效得不偿失。4.3 CUDA out of memory显存直接溢满现象是训练命令一行不差地执行但刚跑几步就 OOM尤其是按源码默认 batch_size 跑时最容易遇到。原因很直接作者演示用的显存可能是 24GB而你手里的是 8GB 或 12GB 的卡。Swin 自注意力的显存开销不只是参数本身还包括每个 token 存 qkv 中间结果和注意力矩阵窗口机制已经省了很多但在 224×224 输入下仍然比普通卷积 UNet 吃显存。解决思路按优先级排列。先把 batch_size 改成 2 或 4几乎立刻见效。接着开混合精度训练很多源码已经预留了--amp参数没有就在模型 forward 外面套torch.cuda.amp.autocast()。如果还不够用梯度累积模拟更大的 batch# train_accumulate.py梯度累积等效扩大 batch_size 且不增加显存 accum_steps 4 optimizer.zero_grad() for step, (images, masks) in enumerate(train_loader): outputs model(images) loss criterion(outputs, masks) loss loss / accum_steps loss.backward() if (step 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()注意这里 loss 除以 accum_steps 是为了保证多步累积后的梯度量级和一个大 batch 一致。梯度累积能解决显存不足但训练时间不会缩短它只是把计算摊到多次前向里。另一个可以尝试的是激活检查点如果你的源码里支持model.enable_activation_checkpointing()开起来也能明显省显存代价是前向速度慢 20% 左右。4.4 timm 版本升级导致 AttributeError现象是装完依赖执行python train.py立刻报AttributeError: module timm.models has no attribute layers或者cannot import name to_2tuple。这是 Swin-UNet 源码最常见的老化问题因为源码编写时 timm 还是 0.6.x而现在 pip 默认安装的 timm 已经 0.9 以上工具函数目录变了。解决方式最稳妥的是锁版本pip install timm0.6.12。如果因为其他依赖没法降级就在代码里做兼容导入# compat.py兼容 timm 0.6.x 与 0.9.x 的工具函数导入 try: from timm.models.layers import to_2tuple, trunc_normal_ except ImportError: from timm.layers import to_2tuple, trunc_normal_不过这只是第一步。timm 版本不同还可能影响预训练权重下载的接口和 checkpoint 的 key 格式改 import 不一定能解决所有问题。我见过有人花一下午改完 import结果下载的权重格式又对不上。所以优先级最高的还是锁 timm 版本其次才是写兼容层。4.5 训练 loss 不降预测图全黑或者全白现象是训练正常运行loss 在初期小幅下降后停滞或者直接不变保存出来的推理结果全部是背景像素看不见任何目标。原因多数不在模型而是标签。常见情况有三种标签 PNG 里是 0 和 255而不是 0 和 1背景像素占比超过 95%普通交叉熵把所有像素都预测成背景就能拿到很低的 loss还有一类是输出 head 用了 sigmoid 但类别是 5 类导致多分类语义错乱。解决方式先把标签值打印出来确认# check_label.py检查标签像素值分布 import cv2 import numpy as np lbl cv2.imread(data/myseg/train/labels/case0001.png, cv2.IMREAD_GRAYSCALE) values, counts np.unique(lbl, return_countsTrue) for v, c in zip(values, counts): print(f像素值 {v}: 占比 {c / lbl.size:.2%})如果打印出 255就把标签 255 改成 1再训练。如果是类别严重不平衡建议把损失换成 DiceLoss 或 Dice CrossEntropy 的混合形式DiceLoss 天然不依赖像素比例对小目标更友好。这个坑最隐蔽的地方在于它不报错训练流程全正常只有最后看结果才发现白忙一场。5. 从“跑通”到“跑好”验证指标、损失调整与推理加速模型能前向、能保存预测图只算入门。Swin-UNet 这类带 Transformer 的模型真正要打磨的是验证指标和损失函数。先说我验证时最常用的三个指标Dice、IoU 和 HD95其中 HD95 是表面距离指标对边界质量敏感Swin 编码器带来的边界改善在 HD95 上比 Dice 更明显。如果你只看 Dice可能觉得 Swin 和普通 UNet 差不多但 HD95 经常能拉开差距。损失函数我通常不直接用交叉熵而用 DiceLoss 和交叉熵的加权组合# combined_loss.pyDice 与 CrossEntropy 组合损失 import torch import torch.nn as nn import torch.nn.functional as F class DiceCrossEntropyLoss(nn.Module): def __init__(self, weight_ce0.4, weight_dice0.6): super().__init__() self.weight_ce weight_ce self.weight_dice weight_dice def forward(self, logits, masks): # logits: (B, C, H, W), masks: (B, H, W) ce_loss F.cross_entropy(logits, masks) probs F.softmax(logits, dim1) num_classes logits.shape[1] target F.one_hot(masks, num_classes).permute(0, 3, 1, 2).float() smooth 1.0 intersection (probs * target).sum(dim(2, 3)) total probs.sum(dim(2, 3)) target.sum(dim(2, 3)) dice ((2.0 * intersection smooth) / (total smooth)).mean() return self.weight_ce * ce_loss self.weight_dice * (1.0 - dice)权重的选择依赖你的任务背景占比极大时提高 weight_dice 到 0.7边界细节重要时保持 0.5 上下即可。训练后期可以把 weight_ce 降下来让模型更专注边界。这张组合损失对 Swin-UNet 的收敛速度也有帮助纯交叉熵在这个结构上早期掉点很慢。推理端一个低成本技巧是测试时增强最简单的做法是把输入图左右翻转两次推理结果取平均零成本提升 0.3 到 1 个 Dice 点。如果源码输出的是 logits在 softmax 之前对两次输出做平均再 argmax比在概率图之后平均略稳。我个人的习惯从来不是拿到源码就默认它最优而是先跑通最小流程再把损失、预处理、输入尺寸这三处按数据实际情况各改一遍。这个方向值得投入尤其当你手头数据是典型医学影像Swin 的预训练权重带来的迁移优势大概率能压过训练时间成本。希望这篇对你有帮助。本文还有配套的精品资源点击获取
返回列表