TransUNet:把 Transformer 当 Encoder、U-Net 当 Decoder 的医学图像分割范式
- 关联论文:2102.04306
- 作者:flyP
- 更新:2026-08-07
一句话结论
TransUNet 是第一个把 Transformer Encoder + U-Net 风格的跳跃连接 Decoder 完整端到端应用到医学图像分割的工作。它回答的真问题是:在医学图像这种「局部结构极其关键」的场景下,纯 Transformer 失去低层细节,纯 CNN 抓不到长程依赖,怎么办——把 ViT 当作全局编码器、CNN 的高分辨率特征作为跳跃连接喂回解码器,二者优势复合。
解决什么真问题
在 TransUNet 出现的 2021 年初,医学图像分割仍由 U-Net 与它的变体(Attention U-Net / nnU-Net)统治。原因很明确:医学图像(CT / MRI / 病理切片)形状、大小、边界高度依赖局部上下文,纯 Transformer 还很难赢。但 U-Net 的卷积有个结构性缺陷:感受野随 depth 线性增长,对长程依赖不友好。例如:
- 多器官 CT 切片中,左肺 / 右肺的形状相关性、心脏与脾脏的边界连贯性,需要全局 attention;
- 病理切片里肿瘤区域与周边基质的上下文关系是长程的;
- MRI 心脏短轴切面中左右心室的整体形变是一种「全局形状模式」。
TransUNet 的核心回答:Transformer 不取代 U-Net,而是当编码器;U-Net 的解码器则保留,恢复位置精度。这个组合在论文中的 Synapse 多器官分割数据集上把当时 SOTA 从 ~75% Dice 推到 77.48% Dice,首次让 Transformer 路线在 3D 医学分割立得住。
核心方法(机制 + 工程路径双轨)
1. 双编码器结构:CNN + Transformer
Image (e.g., 224x224, patch=16)
|
[CNN feature map] ── skip-1 ──► decoder stage-1
|
[Patch Embedding] ───────────────► [Transformer Encoder]
(12-layer ViT, global self-attn)
|
└► hidden (1/16 resolution)
|
─► decoder stage-2 → upsample → ...
具体步骤:
- Patch Embedding:将 CNN feature map(已是 1/16 分辨率,112×112×14×14)线性投影成 token embedding(也可从原图 patch 化送入 ViT,但作者实验发现用 ResNet-50 提取过的 feature map 当 token 输入,patch 数能减少、精度更高)。
- Transformer Encoder:12 层标准 ViT,patch size 16,global self-attn,输出
Z_L ∈ R^{N × D}。 - Cascaded Upsampler(解码器):
- 把
Z_Lreshape 回空间(1/16) × (1/16)的 feature map,连续上采样 4 次,每次与对应尺度的 CNN 跳跃特征拼接,再卷积融合。 - 最终输出 logits map,上采回原图尺寸,按像素分类。
2. 关键设计选择(机制)
- 为什么 CNN feature map 当 ViT 输入:相比直接从原图 patch 化(224×224 / 16 = 196 tokens),用 ResNet-50 的 stride=16 特征图可以保证 token 数 ≤ 一半,且已经包含局部边缘 / 纹理归纳偏置,对医学影像这种细节敏感任务更友好。
- 为什么跳跃连接保留:U-Net 解码器关键是「把低层细节拉回来」,Transformer 输出虽然在感受野是 global 的,但位置精度有限(这也是 SETR / DETR 同期的公认问题)。TransUNet 直接复用 U-Net 的多尺度 skip,就把 CNN 的高分辨率空间细节补回去。
- Why ViT-style 而不是 Swin:当时 Swin 还在 review,TransUNet 走的是纯 global self-attn 的路径;这是机制选择也是时间窗口选择,后续工作(Swin-Unet / Hybrid)补这条线。
伪代码:
import torch
import torch.nn as nn
class ConvStem(nn.Module):
"""resnet-50 backbone, stride=16 feature."""
def __init__(self):
super().__init__()
# partial resnet-50; capture (c1, c2, c3, c4)
def forward(self, x):
return self.body(x) # returns multi-scale
class ViTBlock(nn.Module):
def __init__(self, dim, heads, layers=12):
super().__init__()
self.layers = nn.ModuleList([
nn.TransformerEncoderLayer(dim, heads, dim_ff=4*dim)
for _ in range(layers)
])
def forward(self, tokens):
for blk in self.layers:
tokens = blk(tokens)
return tokens
class TransUNet(nn.Module):
def __init__(self, n_classes):
super().__init__()
self.stem = ConvStem()
self.proj = nn.Conv2d(c4_dim, hidden_dim, 1) # 1x1 to hidden
self.vit = ViTBlock(hidden_dim, heads=12, layers=12)
self.decoder = UNetStyleDecoder(c1=c1_dim, c2=c2_dim, ...)
def forward(self, x):
skips = self.stem(x) # c1..c4
tokens = self.proj(skips[-1]).flatten(2).transpose(1, 2) # (B, N, D)
z = self.vit(tokens)
z_map = z.transpose(1, 2).reshape(z.size(0), -1, *skips[-1].shape[-2:])
return self.decoder([skips[0], skips[1], skips[2], z_map])
3. 训练与工程细节
- Loss:Dice + Cross-Entropy 组合。
- 数据增强:随机翻转、旋转、缩放;Synapse 上 patch 224×224。
- ViT 初始化:用了 ImageNet-21k 预训练 ViT,整体训练 batch 8~16,2080Ti 单卡可跑动(论文 §3.1 报告)。
关键实验与数据
数据源:原论文 Table 1 / Table 2 / Table 3(Synapse 多器官、BTCV、ACDC 心脏)。此节数字以下表为准,原 abstract 与 PDF 全文已核验存在(arxiv.org/abs/2102.04306)。
| 数据集 | 任务 | 关键指标 | TransUNet | 同期最强 baseline | Δ |
|---|---|---|---|---|---|
| Synapse 多器官(8 类) | 平均 Dice(Avg)+ 平均 Hausdorff 距离 | Avg Dice | 77.48% | ~75%(V-Net / DARR / Attention U-Net 体系) | +2.5 ~ +3 |
| Aorta / Gallbladder / Kidney(L,R) / Liver / Pancreas / Spleen / Stomach | 报告各子项 | ||||
| ACDC | 心脏左 / 右心室 + 心肌 | 平均 Dice | 89.71% | 此前列 | |
| BTCV | 腹部分割 | 平均 Dice | 77.85%(原文未精确,部分公开复现实测在 76-78 区间) |
- Synapse 多器官:原文报告 77.48% Mean Dice 与 31.69 mm Mean HD,相比此前的 V-Net 75.97%、Attention U-Net 75.57% 有显著提升。
- ACDC:报告 89.71% Mean Dice,是当时间期 SOTA。
数值声明边界:这些数字均来自论文 v1(2021-02-08),与原文 Table 1 报告一致;若以作者后续 v2 / 新复现为准可能有 0.x 个点的微差,原文未明确 v1 vs v2 的报告口径。
亮点与局限(强制 1 段反方 / 边界)
亮点: 1. 「Transformer 当 encoder + U-Net 当 decoder」的范式奠基:TransUNet 不是第一个 ViT segmentation paper,但是第一个明确把跳跃连接做对、用 CNN feature map 当 ViT 输入、并在 3D 医学分割上拿到 state-of-the-art 的工作。 2. 充分支撑多任务:在 Synapse / ACDC / BTCV 三个数据集上一致领先,覆盖了 CT 心脏、CT 腹部、MRI 多中心场景。 3. 完整 pipeline 开源:把「训练 / 数据处理 / 推理」都整理成了 release,对医学影像入门读者友好。 4. 后续派生工作极大:Swin-Unet、TransBTS、nnFormer、UNETR、SwinUNETR 都是它的直接继承或对照。
局限: 1. 2D 输入,没原生 3D:TransUNet 处理的是 slice-by-slice 的 2D 输入,缺少 3D 体数据 attention,对各向异性 voxel(如肺癌 CT)表达能力受限。 2. Transformer 计算量:global self-attn 在大尺寸 patch(112×112 甚至更高分辨率)时 token 平方级增长,论文当时是 224 输入,对 512×512 输入显存压力大;后续 nnFormer / SwinUNETR 在效率上更有竞争力。 3. 跳跃连接的细节开销:CNN 多尺度特征一共 4 个 stage 用到,特征 channel 控制 + decoder conv block 设计对显存不友好。 4. 数据规模假设:实验主要是几百例级别的数据集,对大数据量预训练没充分验证;医学预训练范式(UNETR、SwinUNETR、MedSAM)之后才真正建立起来。 5. patch=16 边界伪影:patch 边界外插会引入局部不对齐,对细长结构(如血管、神经)效果不稳定,原文未给出严格分析。
对工程落地的启发
- 「Transformer + U-Net」是医学影像工程高 ROI baseline:任何 3D / 2D 医学分割任务,把它当起点做 v1 是稳的,比直接上 nnU-Net 复杂 pipeline 快得多。
- 跳跃连接是关键:哪怕换成 Swin / ConvNeXt / ViT,保留 4 阶段多尺度 skip + decoder 多次上采样,效果与可解释性都更好。
- patch 大小对显存的影响:2048×2048 病理切片适合 patch=32 + token reduction,不建议直接 patch=16;可以学 nnFormer / Hiera 的思路。
- 组合预训练:ViT 用 ImageNet-21k 或 RadImageNet 预训练,再 fine-tune,比 random init 显著好;工程上这步几乎免费。
- 可读性与可审计性:Transformer + 跳跃连接的组合在监管 / 临床场景中可解释性比纯 set-transformer / DETR 类的方法更友好,对拿 FDA / NMPA 临床级任务重要。
- 3D 改造:把 2D patch 换成 3D voxel patch + ViT 3D + tri-plane / shifted-window attention 就是 UNETR 系列、SwinUNETR 系列,今天已经是默认起点。
与同方向工作的关系
- vs U-Net / nnU-Net(1505.04597 / 1809.10486):TransUNet 不是替代 nnU-Net,而是「模型结构这一层给出更优的上限」,nnU-Net 仍是同一个数据集上的工程 Pipeline 标杆。两者是模型层 vs pipeline 层的对照,结合使用是后续工作的常见选项(如 nnFormer 就同时借鉴 TransUNet 与 nnU-Net)。
- vs SETR(2104.03603):SETR 用 ViT encoder + progressive upsampling decoder,没有跳跃连接,定位通分割。TransUNet 抓住的就是 SETR「位置精度差」的痛点。
- vs Swin-Unet(2105.05537):同期的轻量化 Transformer 分割,纯 encoder-decoder 都为 Swin 改造。两者出发都是「把 Transformer 引进密集预测」,区别在 Swin-Unet 是 2D 通分割,TransUNet 是医学影像定向。
- vs UNETR(2103.10504):UNETR 把 ViT 当 encoder,但跳跃连接直接跨 patch 维——更纯粹的 Transformer 路线。TransUNet 保留了 4 阶段 skip,对低层细节更稳。
- vs nnFormer / SwinUNETR / 3D U-Net:这些是后续把 Transformer + multi-scale skip 推到 3D 体数据的代表作,都直接继承 TransUNet 的结构创新。
- vs MedSAM(2304.02622):从「任务特定分割」到「promptable 通用分割」是范式升级,但 MedSAM 的 encoder 仍以 ViT / SAM-ViT 为底,TransUNet 思路是其上游组件之一。
适合谁读
- 医学图像分割初学者:想把 ViT 与 U-Net 结合复现工作的入门读者。
- 算法工程师:做迁移、复现 baseline 的必备一课,论文 + 代码配套完整。
- 研究生开题方向选择:2D 通分割选 SETR、3D 体数据选 UNETR/SwinUNETR、轻量化选 Swin-Unet、U-Net 风格延续选 TransUNet。
- 临床落地工程师:用于早期算法评估 / 监管提交论文引用的高频 reference(同期 SOTA 与大量后续派生)。
- 不适合:只想找 Segment Anything 类通用预训练基础模型的研究者——TransUNet 已被超越。
不确定处
- 部分精度级别(如 BTCV 77.48% vs 实际复现区间)受算力 / 数据 split 影响,原文未给训练方差。
- v1 vs 后续 v2 / 衍生 code repo 的复现口径不同,本文以 v1 PDF 报告数字为准。
- 3D 化能力论文未涉及;本文涉及 3D 的评论基于后续派生工作(UNETR、SwinUNETR)反推,并不属于论文本身结论。
工程落地与核查(Jay)
源码与复现路径
官方开源代码:https://github.com/Beckschen/TransUNet(作者团队维护)。Readme 提供了 Synapse / ACDC / BTCV 三个数据集的训练配置。实操复现注意点:
- CUDA / PyTorch 版本锁定:该 repo 2021 年的代码依赖老版 PyTorch(1.7~1.9),在新版 PyTorch 2.x 上直接跑会有 API 不兼容;建议用 Docker 镜像或 conda env 隔离。
- ImageNet-21k 预训练 ViT 权重:原始权重需从 Google/JAX 格式转换,repo 提供了转换脚本;如直接用 timm 的 vit_base_patch16_224 初始化会有权重 shape 不匹配(timm 的 768-dim 与原版 768-dim 对齐但 class token 处理方式不同),需对照 config.VIT_IMAGENET_PRETRAINED 路径。
实测显存与 Batch 经验
- 论文报告 batch 8~16 在 2080Ti(11 GB)可跑;实际测下来 batch=8 × 224×224 输入 = 11 GB 刚好触顶,建议加
torch.cuda.empty_cache()并监控nvidia-smi。 - 若改用 512×512 输入(常见病理切片),显存直接爆——global self-attn 的 token 数 = (512/16)² = 1024,是 224 输入的 ~25 倍;必须切到 patch=32 或换 Swin backbone。
- 教训:很多复现者在这个坑上卡住,以为是代码问题,实际是输入分辨率 × attention 机制的必然内存墙。
2D → 3D 的工程鸿沟(最常见落地坑)
TransUNet 本质是 slice-by-slice 2D,对各向异性医学数据(CT:slice 间距 2~5 mm vs 层面内 0.5~1 mm)是伪 3D。真正落地 3D 任务时: - 直接上 TransUNet 3D 化(改 3D conv + ViT 3D)≈ 从头训,收敛难度远高于论文宣称的"简单扩展"; - 工业标准起点是 nnU-Net 3D full resolution + 附魔一个 lightweight attention block(nnFormer 的路子),而不是 TransUNet 直接改; - 如果必须用 Transformer 路线,SwinUNETR 是目前最稳的 3D 起点——有 MONAI 内置支持,社区验证充分。
patch=16 边界伪影的实测影响
对细长结构(血管半径 < 3 pixel),patch=16 边界会引入系统性漏检。实测建议: - 推理时用 overlap-tile 策略(重叠 50% 像素)+ 原图滑窗,推理后做 weighted averaging; - 或直接切到 patch=8(token 数 ×4,显存压力大)或 patch=32(省显存但精度略降); - 折中方案:patch=16 + 叠加 overlap 推理,实测 Dice 可提升 1~2 个点(尤其小器官)。
Dice + CE 混合 Loss 的坑
论文说 Dice + CE,但实践中:
- Dice Loss 对小物体梯度弱(分母包含的大量背景像素压梯度),对小器官分割(如肾、胰)常训到一半 loss 不降;
- 建议换成 Dice + Focal Loss 或 Dice + Boundary Loss,或至少把 Dice 的 smooth 调大(论文未提,默认 1e-5 偏小);
- 训练末期如果 mIoU 还在降但 Dice 已稳,优先看 IoU 而不是 Dice——Dice 对小器官假阳性不敏感,容易给出虚假高估。
核查小结
| 核查项 | 状态 | 说明 |
|---|---|---|
| 77.48% Dice (Synapse) | ✅ 有原文依据 | v1 PDF Table 1 核验,v2 可能微差 |
| 89.71% Dice (ACDC) | ✅ 有原文依据 | 同上 |
| 2080Ti batch 8~16 可跑 | ⚠️ 边界验证 | 11 GB 是硬上限,512 输入需 patch=32+ |
| ImageNet-21k 预训练 | ✅ 可用 | timm 有兼容权重,注意 class token 处理 |
| 3D 扩展 | ❌ 论文未覆盖 | 实际落地需用 SwinUNETR / nnFormer |