☰
Swin-UNet从源码到实战:Swin Transformer与UNet医学图像分割指南
2026/10/2 8:32:13 网站建设 项目流程

简介:一个融合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_size=224, # 输入分辨率,注意后续 window 整除约束 patch_size=4, # patch embedding 的卷积核大小和步长 in_chans=3, # 输入通道,灰度图改为 1 之后要处理预训练权重 num_classes=5, # 按自己数据集的类别数改,和输出 head 对齐 embed_dim=96, # 第一阶段通道数 depths=[2, 2, 2, 2], # 每阶段 Swin Block 数量 num_heads=[3, 6, 12, 24], window_size=7, # 窗口大小,强烈建议保持 7 mlp_ratio=4., qkv_bias=True, drop_rate=0., drop_path_rate=0.1, ape=False, # 是否使用绝对位置编码 patch_norm=True ) 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。下面这份依赖清单是常见源码里可复现性最好的配置:

python>=3.8 torch>=1.10.0 torchvision>=0.11.0 timm==0.6.12 numpy>=1.21 einops tqdm SimpleITK

安装时建议新建干净的 conda 环境,避免把别的项目的 torch 版本带进来:

conda create -n swin_unet python=3.9 -y conda activate swin_unet pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 --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_size=0.15, random_state=42) 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_ok=True) os.makedirs(out_lbl, exist_ok=True) 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, interpolation=cv2.INTER_LINEAR) # 标签缩放必须用最近邻,避免插值产生不存在的类别 lbl = cv2.resize(lbl, IMG_SIZE, interpolation=cv2.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-2,Swin 的 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_location="cpu") proj_weight = checkpoint["model"]["patch_embed.proj.weight"] # shape 为 [embed_dim, 3, 4, 4] checkpoint["model"]["patch_embed.proj.weight"] = proj_weight.mean( dim=1, keepdim=True ) 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 得 56,56 能被 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, value=0 ) # 推理得到 pred,shape 为 (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 timm==0.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_counts=True) 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.py:Dice 与 CrossEntropy 组合损失 import torch import torch.nn as nn import torch.nn.functional as F class DiceCrossEntropyLoss(nn.Module): def __init__(self, weight_ce=0.4, weight_dice=0.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, dim=1) 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 的预训练权重带来的迁移优势大概率能压过训练时间成本。希望这篇对你有帮助。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询