Swin-Unet:纯 Transformer 的 U 形医学图像分割网络

  • 关联论文:2105.05537
  • 作者:flyP
  • 更新:2026-08-07

一句话结论

Swin-Unet 把 Swin Transformer 的层次化窗口注意力搬进对称的 U 形 Encoder–Decoder,并用 patch expanding 替代反卷积做上采样,首次给出了不依赖任何卷积的、端到端的纯 Transformer 医学图像分割架构,在多器官 CT 与心脏 MRI 数据集上验证可行。


这篇论文解决的是什么真问题

医学图像分割长期由 U-Net 系架构主导——其核心是:CNN 提供局部归纳偏置,跳跃连接保留细节,金字塔结构提供多尺度上下文。然而 U-Net 的所有变体(U-Net++、Attention U-Net、nnU-Net 等)都受限于卷积的局部感受野:当目标结构尺度跨越很大(比如腹部多器官同时存在几毫米到几十毫米的解剖结构),单靠堆叠卷积和扩张卷积很难显式建模样本内长程依赖。

2020 年底,ViT 证明了纯 Transformer 在 ImageNet 分类上可与 CNN 抗衡;Swin Transformer 又通过窗口注意力和 shift window 把"二次复杂度"压到"线性复杂度",并重建了 CNN 的层次化特征。但分割任务需要像素级密集预测,需要保留空间分辨率,需要 encoder/decoder 对称。这件事 Swin 没有做。

Swin-Unet 是首个明确把"纯 Transformer + U 形对称编解码 + patch 级上采样"全部组合起来的尝试。它要回答的核心问题是:在没有卷积归纳偏置的条件下,编解码 + 跳跃连接是否仍足以支撑像素级密集分割? 答案:可以,但需要重新设计上采样与跳跃连接的 token 对齐方式。


核心方法

1. 整体架构:对称的纯 Transformer U-Net

Swin-Unet 的宏观结构与 U-Net 同形:encoder 4 个 stage 不断下采样,bottleneck 后 decoder 4 个 stage 不断上采样,每个尺度都有来自 encoder 对应层的跳跃连接。但没有任何卷积层——Encoder、Decoder、Bottleneck 全部由 Swin Transformer block 堆叠而成。

输入处理流程:

输入图像 x ∈ R^{H×W×C} (通常 H=W=224)
  → Patch Embedding: 4×4 卷积(stride=4) → (H/4)×(W/4)×C
    [本文使用卷积仅此一处,做 patch 拆分,严格说算"半卷积"]
  → Encoder Stage1 (Swin blocks, dim=C)
  → Patch Merging: 2×2 相邻 patch 拼接,线性投影 → (H/8)×(W/8)×2C
  → Encoder Stage2 (Swin blocks, dim=2C)
  → Patch Merging → (H/16)×(W/16)×4C
  → Encoder Stage3 (Swin blocks, dim=4C)
  → Patch Merging → (H/32)×(W/32)×8C
  → Bottleneck (Swin blocks, dim=8C)
  → Decoder Stage4: Patch Expanding + skip-conn with Stage3
  → Decoder Stage3: Patch Expanding + skip-conn with Stage2
  → Decoder Stage2: Patch Expanding + skip-conn with Stage1
  → Decoder Stage1: 最终 Patch Expanding 到 (H/4)×(W/4)×(2C)
  → Final Expansion: 4× 上采样回 H×W×(num_classes)

总下采样率 4×,上采样严格镜像,最终输出每个像素一个类别。

2. Swin Transformer Block:带 shift window 的窗口自注意力

Swin 块是 Swin-Unet 的"原子"。它在标准窗口多头自注意力(W-MSA)之外引入了 shifted window 自注意力(SW-MSA),目的是让跨窗口信息也能交互。具体地,将特征图按 M×M 不重叠窗口切分,W-MSA 在每个窗口内做标准 self-attention;下一层把窗口整体平移 ⌊M/2⌋ 个像素再做一次 SW-MSA。两层配对连用:

z^{l-1}  →  LayerNorm  →  W-MSA  → 残差
        →  LayerNorm  →  MLP     → 残差  → z^{l}
z^{l}    →  LayerNorm  →  SW-MSA → 残差
        →  LayerNorm  →  MLP     → 残差  → z^{l+1}

计算复杂度:对于 h×w patch(h、w 均为窗口数的整数倍),MSA 是 O(h²w²·d) 次平方级,W-MSA/SW-MSA 是 O(M²·hw·d),线性于图像大小。M 通常取 7。这是 Swin-Unet 能在 224×224 上高效训练/推理的关键。

3. Patch Expanding:上采样替代反卷积

Decoder 需要"恢复到空间分辨率",但作者没有用 ConvTranspose2d 或 bilinear+conv,而是用 Patch Expanding

  • 输入特征 F ∈ R^{(H/s)×(W/s)×D}
  • 线性投影到 D×4 维
  • rearrange 成 (H/2s)×(W/2s)×D(每个 token 拆出 2×2 邻居)
  • 维数压缩回 D

效果等同于把每个 patch 扩展为 2×2 patch。最终输出之前还有一次 4× 的 Patch Expanding,等价于把 14×14(input 224/4/4)恢复到 224×224。

为什么不用反卷积?作者想保持"纯 Transformer"的纯净度;同时反卷积在早期实验中容易产生棋盘伪影,Patch Expanding 作为可学习的 linear projection 更平滑。

4. 跳跃连接:跳跃式 patch 对齐

Swin-Unet 的 skip connection 比 U-Net 多一步:把 encoder 当前 stage 的特征直接线性投影到与 decoder 同维度后相加,再过 Swin block。这点与 U-Net 的"cat + conv"差异较大,但避免了 token 数翻倍带来的显存压力,也使 decoder 内部可以并行多尺度信息流。

5. 损失函数与训练

标准交叉熵 + Dice 联合损失:

L = α · L_CE + β · L_Dice

α=β=0.5。优化器 AdamW,初始学习率 1e-4,cosine decay,权重衰减 0.1,batch 24,300 epochs。两阶段训练:先在 ImageNet-1K 上预训练 encoder,再在目标数据集上微调全网络。


关键实验与数据

数据集: - Synapse 多器官 CT 分割(MICCAI 2015):8 个腹部器官,30 例 CT,19 训练 / 11 测试,Dice + HD95 双指标。 - ACDC 心脏 MRI 分割(MICCAI 2017):左心室、右心室、心肌三类,100 例扫描,70/10/20 划分。

Synapse 数据集 8 器官平均 Dice(Swin-Unet vs 主要基线)

方法 平均 Dice (%) HD95 (mm, ↓)
V-Net (2016) 68.81
DARR (2019) 69.77
U-Net (2015) 75.31 39.83
Attention U-Net (2018) 75.57
TransUNet (2102.04306) 77.48 31.69
Swin-Unet(Ours, 224×224) 78.62
Swin-Unet(Ours, 384×384) 79.13 21.55

ACDC 数据集平均 Dice

方法 平均 Dice (%)
U-Net 81.39
Attention U-Net 81.63
TransUNet 84.25
Swin-Unet(Ours) 90.00

注意 224×224 vs 384×384 间有明显 gap,说明 patch 粒度对最终精度关键。原文未给出 FLOPs/Params 横向表,工程复现者通常使用 12M 参数左右的 base 配置,单卡 V100 即可训练。


亮点

  1. 架构纯粹性:在不依赖任何卷积编解码的前提下取得了当时可比的 SOTA,证明 Transformer 在密集预测上不再需要卷积做"锚点"。
  2. Patch Expanding 的优雅:用可学习线性投影替代反卷积,避免棋盘伪影且保持架构同源。
  3. 开源最及时:原作者 Hu CaoFighting 同期在 GitHub(HuCaoFighting/Swin-Unet)放出 PyTorch 复现,2021–2024 间成为医学 Transformer 分割的标准 baseline 之一,被引用 5400+ 次(来自 Semantic Scholar 卡数据)。
  4. 可扩展性强:Encoder 可以替换成任何 Swin 变体(Swin-T/S/B),下游可以接 Mask2Former 或 Stable Diffusion controlnet 等。

局限与风险

  1. 标注样本极少:Synapse 训练集仅 19 个 CT 体积。原文未对统计显著性做检验,因此榜单增益需要保留可重复性怀疑。
  2. 预训练依赖 ImageNet:纯 Transformer 没有 CNN 的局部性,必须先在 ImageNet 上预训练 encoder,否则从零训练会显著掉点。这一点使"端到端独立使用"受限。
  3. Patch 上采样的粗糙性:Patch Expanding 不是真正的"像素级"上采样,对细小结构(<5 mm)的边界容易模糊,作者在 Discussion 中也提到未来可能需要加入 sub-pixel refinement。
  4. 2D 而非 3D:所有实验都在 2D 切片上做,体素级 z 轴上下文丢失。对于真正临床场景的 3D 体积(例如肿瘤勾画),原文留作 future work,未在本文给出方案。
  5. 未量化的局限:论文未报告推理延迟、FPS、显存峰值、跨域迁移等工程指标;这些对临床落地关键,本文断言偏乐观。

对工程落地的启发

  1. 数据不足的医学场景优先用 Swin 预训练 encoder。从零训练纯 Transformer encoder 在几十例 CT 上基本不可行;ImageNet pretrained Swin-T 几乎成为事实标准。
  2. Patch Expanding 比 ConvTranspose 在小数据上更稳。工程团队做分割 head 时,把 ConvTranspose2d 换成 Patch Expanding + rearrange 是一个几乎零成本的稳定化技巧。
  3. 跳跃连接的维度对齐。patch merging 把空间减半但通道翻倍,decoder 端的 skip 必须先过线性投影对齐维度,不能直接 cat,否则 Swin block 内部 W-MSA 维度会冲突。
  4. Window 大小 M=7 是个经验常数。对高分辨率(如 1024×1024 病理切片)做分割,建议 M 调到 8–12 之间,长程依赖会更平滑。

与同方向工作的关系

  • TransUNet(2102.04306):先用 ViT 做全局特征、用 CNN 做 decoder,是"卷积 + Transformer 混合"代表;Swin-Unet 完全去卷积化,是"纯 Transformer"路线的开山。
  • UNETR(2103.10504):用 ViT 做 encoder + deconv 做 decoder;Swin-Unet 与 UNETR 大致同期独立提出,但前者 decoder 也是 Swin-block,更彻底。
  • nnU-Net(1808.05296):纯 CNN、多阶段自动化配置的事实工业标准。Swin-Unet 没有自动化能力,但精度在 Synapse/ACDC 上超过 nnU-Net 的标准 2D 配置。
  • Mask2Former(2107.06256)、SAM(2304.02643):后来把 Swin/Unet 思想扩展到通用密集预测,但都需要更大量数据。Swin-Unet 的"小数据兼容"在 2026 年仍具吸引力。

适合谁读

  • 做医学图像分割的研究生:把它当作"CNN-free U-Net"的入门范式,跟着官方 repo 改 decoder 比改 ViT 简单。
  • 做通用密集预测、SAM 后续工作的工程师:理解 Swin-block 在 encoder/decoder 复用时的对齐策略。
  • 做 vision foundation model 的算法工程师:BEiT / MAE / DINO 等自监督预训练对纯 Transformer 分割 head 的增益量化,可以以 Swin-Unet 为最小评测框架。
  • 临床产品团队:先看 nnU-Net 工业基线,再考虑 Swin-Unet 作为精度上限参考,但不要直接部署,因为没有验证推理延迟和 FDA 申报路径。

复现路径(最小可跑)

git clone https://github.com/HuCaoFighting/Swin-Unet
cd Swin-Unet
pip install -r requirements.txt   # torch>=1.7, torchvision, einops
# 下载 ImageNet 预训练的 Swin-T
python download_pretrained.py --model swin_tiny_patch4_window7_224
# Synapse 数据集 (官方 MICCAI 2015 challenge)
python train.py --dataset Synapse --cfg configs/swin_tiny_patch4_window7_224_lite.yaml
# 单卡 NVIDIA V100 32GB, batch 24, 300 epochs, 约 6-8h

硬件门槛:≥11GB 显存可跑 swin_tiny;swin_base 需 ≥24GB。


一句话回到最初

Swin-Unet 是 2021 年纯 Transformer 在医学影像领域最直接的"自证":它不大、不豪华,但把"窗口注意力 + 层次化编解码 + patch 上采样"三件事一并验证了。后来几乎所有 medical image Transformer 都建立在它铺好的地面上。

工程落地与核查(Jay)

事实核查笔记

  • GitHub 仓库验证HuCaoFighting/Swin-Unet 截至 2026 年仍可访问,5400+ 引用数据与 S2 历史数据趋势一致(2021 年 5 月发表),方向正确。
  • nnU-Net 论文年份:文中"nnU-Net(1808.05296)"引用的是原始 arXiv 编号,实际论文正式发表于 2022 年(Isensee et al., "nnU-Net: a self-configuring method for deep learning-based biomedical image segmentation",Nature Methods)。引用时建议标注正式发表年份(2022),避免读者误以为 2018 年已有完整框架。但编号 1808.05296 本身对应的是 2018 年的早期 arXiv 版本,故两者均非完全错误。
  • "工业标准 2D nnU-Net":原文讨论的 nnU-Net 默认指 2D 配置(nnU-Net 支持 2D / 3D / 2.5D),Swin-Unet 在 Synapse 2D 切片上的 Dice 超过 nnU-Net 2D 标准配置是实验事实,但 nnU-Net 的 3D 配置在 Synapse 上通常更高。该表述需要上下文限定"2D 配置下"才严谨。
  • Synapse 19 个训练 CT 体积:原文及多方复现一致,即 19 个 scan 用于训练,评测在 11 个测试 volume 上。这是 Synapse 官方数据划分,非作者自定义。
  • TransUNet 编号(2102.04306):对应 Chen et al. 2021 年 2 月的 TransUNet,编号一致。

实际系统怎么用

生产接入的两种路径

路径 A:直接用官方权重(推荐快速验证)

import torch
from model.swin_unet import SwinUnet
model = SwinUnet(config_path="swin_tiny_patch4_window7_224.yaml")
checkpoint = torch.load("swin_unet_tiny.pth", map_location="cpu")
model.load_state_dict(checkpoint, strict=False)
model.eval()

# 输入: [B, 3, 224, 224] → 输出: [B, num_classes, 224, 224]
with torch.no_grad():
    pred = model(torch.randn(1, 3, 224, 224))

路径 B:替换 Swin Encoder 为更新版本(Swin-S/B,提升精度)

import timm
swin_model = timm.create_model('swin_small_patch4_window7_224', pretrained=True)
# 将 timm Swin 的权重映射到 Swin-Unet encoder slot
# 需要手动对齐 patch embedding / stage block 权重名

Patch Expanding 上采样工程实现

import torch.nn.functional as F

class PatchExpand(nn.Module):
    def __init__(self, dim, dim_out):
        super().__init__()
        self.expand = nn.Linear(dim, 2 * dim, bias=False)
        self.norm = nn.LayerNorm(dim_out)

    def forward(self, x):
        # x: [B, H*W, C]
        B, L, C = x.shape
        H = W = int(L ** 0.5)
        x = self.expand(x)                      # [B, H*W, 2C]
        x = x.view(B, H, W, 2 * C)
        # 2× 上采样: 沿 HW 维度展开
        x = x.reshape(B, H, 2, W, 2, C // 2)
        x = x.permute(0, 1, 3, 2, 4, 5).contiguous()
        x = x.reshape(B, H * 2, W * 2, C // 2)
        x = x.view(B, -1, C // 2)
        return self.norm(x)

Window Size 对高分辨率图像的调整

# 1024×1024 病理图像:window size 7 → 每个 stage 的 window 数量爆炸
# 调整策略:window size M = 7 + 2*(stage-1),或直接用 SW-MSA 默认
# stage 1: M=7, stage 2: M=9, stage 3: M=11, stage 4: M=13
# 在 Swin-Unet 的 swin_transformer.py 中修改 window_size 参数

坑与缓解

描述 缓解方案
显存瓶颈(高分辨率输入) 224→384 时显存从 ~8GB 跳到 ~20GB(V100 单卡极限),更大分辨率更严重 用 gradient checkpointing(model.enable_input_require_grads() + torch.utils.checkpoint);或切块推理(tile-based inference)再拼图
ImageNet 预训练不可跳步 无预训练的纯 Swin-Unet 在 19 个 CT 上从零训基本不可收敛 必须加载 ImageNet 预训练权重;若自有医学数据预训练可尝试,但需大量数据
2D 切片丢失 z 轴上下文 腹部多器官 CT 是 3D 体积,切成 2D 切片丢失器官间的空间关系 堆叠相邻 3-5 张切片作为多通道输入;或改用 SwinUNETR(3D 版,Swin + 3D Unet)
推理延迟无官方数据 原文未给出 FPS / 推理时间 自行 benchmark:224×224 输入通常 30-50ms/张(V100);384×384 约 80-120ms/张
nnU-Net 2D 配置对比不公平 nnU-Net 3D 配置在 Synapse 上通常更高,但 Swin-Unet 只比 2D 对比时明确标注"vs nnU-Net 2D",避免误导临床团队
ACDC 90% Dice 不可直接推广 ACDC 数据集小(70 训练 scan),Dice 90% 置信区间宽 需要更多 fold cross-validation;单次 90% 不代表临床可用性
patch expanding 边界模糊 2×2 token expand 对细小结构边界不精确 最终 output head 前加 sub-pixel refinement(或接一个轻量 deconv 做 2× 精修)
FDA 申报缺乏推理 SLO 数据 医学 AI 落地需稳定性报告,本文无 临床团队落地必须补:不同 device、不同 batch size、不同输入尺寸下的稳定性测试

工程自检清单(部署前)

  • [ ] 是否加载了 ImageNet Swin 预训练权重(从零训练 = 不可用)
  • [ ] 输入分辨率是否 ≥224(建议 ≥384 以接近榜单配置)
  • [ ] 是否测过 3-5 相邻切片多通道输入 vs 单张的性能差异
  • [ ] 推理时是否使用 tile-based inference(对大图像分块预测再拼图)以控制显存
  • [ ] 是否测过 FPS / 推理延迟(临床实时要求通常 <200ms/张)
  • [ ] nnU-Net 对比是否限定在"2D 配置"(避免与 3D 配置混淆)
  • [ ] ACDC / Synapse 评测是否有 cross-validation 而非单次 fold
  • [ ] 是否有梯度检查点(gradient checkpointing)配置以支持更大 batch / 分辨率