医学图像分割一直是深度学习落地的重要方向,从器官轮廓勾画到病灶区域提取,都对模型的精度要求极高。传统的U-Net依靠卷积神经网络(CNN)强大的局部特征提取能力长期占据主导地位,但CNN固有的局部感受野限制了它对长距离依赖的建模能力。TransUNet提出了一种将Transformer与CNN串联混合的方案,在编码器前半段使用CNN提取细粒度局部特征,后半段引入Transformer捕获全局上下文,再通过U型解码器逐步恢复空间分辨率,最终在多器官分割和病灶分割任务上取得了优于纯CNN方案的表现。本文将系统拆解TransUNet的架构设计与实现细节。
一、为什么需要CNN与Transformer的混合架构
要理解TransUNet的设计动机,首先要明白两种骨干网络各自的短板。CNN依靠卷积核在局部区域内滑动计算,天然具有平移不变性和局部归纳偏置,这使得它在小数据集上也能训练出不错的模型。但卷积操作的感受野是有限的,即使经过多层堆叠和下采样,深层特征图上每个位置能"看到"的原图区域仍然受限,对于器官之间的空间关系、大范围解剖结构这类全局信息,CNN的表达能力明显不足。
另一方面,纯Transformer架构如ViT虽然凭借自注意力机制可以建模任意位置之间的依赖关系,全局感受野不存在任何限制,但它有三个明显问题。第一,自注意力的计算复杂度随图像块数量呈平方增长,直接对高分辨率医学图像做全局注意力开销巨大;第二,Transformer缺乏局部归纳偏置,在医学影像这种数据量通常不大的场景下容易过拟合;第三,纯ViT通常需要把图像切分成较大的patch(例如16x16),会损失大量细粒度空间信息,而医学分割恰恰对边界精度极为敏感。
TransUNet的核心思路就是扬长避短:用CNN先提取高分辨率的局部特征并降低特征图尺寸,减少送入Transformer的token数量;再由Transformer在低分辨率特征上建模全局依赖;最后通过带跳跃连接的CNN解码器逐步恢复分辨率,把局部细节补回来。这种设计同时解决了CNN全局建模不足和Transformer局部信息丢失两个问题。
二、TransUNet的整体架构设计
TransUNet整体上是一个U型结构,由混合编码器、Transformer模块和级联上采样解码器三部分组成。编码器基于ResNet-50的混合设计:前几个阶段使用卷积层逐级下采样,假设输入是512x512的图像,经过卷积编码后得到下采样16倍的特征图(32x32大小),这个特征图就是送入Transformer的输入。
Transformer部分的输入处理与ViT类似。将32x32的特征图按1x1的patch展平,得到1024个token,每个token的维度等于通道数。接着拼接一个可学习的类别token,加上位置编码后送入多个标准的多头自注意力层。由于token数量已经从原始图像的262144个像素块压缩到1024个,自注意力的计算量大幅降低,这就是先CNN后Transformer串联的好处。位置编码使用经典的正弦余弦函数,帮助模型保留token之间的空间顺序信息。
解码器采用级联上采样(Cascaded Upsampler,简称CUP)的设计,通过多层转置卷积或双线性插值上采样,逐步把Transformer输出的低分辨率特征恢复到原始分辨率。关键的跳跃连接从CNN编码器的中间层引出,将浅层的高分辨率细节特征与解码器特征在通道维度拼接,这样最终分割图既能保持全局语义的正确性,又能还原器官边界等精细结构。实验表明,跳跃连接对TransUNet的边界精度提升非常显著,去掉跳跃连接后Dice系数会明显下降。
三、基于PyTorch的核心实现思路
下面用简化后的PyTorch代码展示TransUNet的关键组件实现,包括CNN编码器、Transformer编码器和解码器的上采样逻辑。实际工程中可以直接使用timm库提供的预训练ResNet和ViT权重加速收敛。
import torch
import torch.nn as nn
# CNN编码器:基于ResNet的前半部分,输出下采样16倍的特征图
class HybridEncoder(nn.Module):
def __init__(self, in_channels=3, feature_dim=768):
super().__init__()
self.conv1 = nn.Sequential(
nn.Conv2d(in_channels, 64, kernel_size=7, stride=2, padding=3),
nn.BatchNorm2d(64), nn.ReLU(inplace=True)
)
# 后续可替换为ResNet的layer1和layer2
self.layer1 = nn.Sequential(
nn.Conv2d(64, 256, kernel_size=3, padding=1),
nn.BatchNorm2d(256), nn.ReLU(inplace=True)
)
self.layer2 = nn.Sequential(
nn.Conv2d(256, feature_dim, kernel_size=3, stride=2, padding=1),
nn.BatchNorm2d(feature_dim), nn.ReLU(inplace=True)
)
# 1x1卷积将特征图按patch展平成token序列
self.patch_size = 1
def forward(self, x):
x = self.conv1(x) # 下采样2倍
skip = self.layer1(x) # 保留跳跃连接特征
x = self.layer2(x) # 下采样,得到Transformer输入
B, C, H, W = x.shape
# 展平为token序列: (B, N, C)
tokens = x.flatten(2).transpose(1, 2)
return tokens, skip, (H, W)
# Transformer编码器:标准多头自注意力堆叠
class TransformerBlock(nn.Module):
def __init__(self, dim=768, heads=12, num_layers=12):
super().__init__()
layer = nn.TransformerEncoderLayer(
d_model=dim, nhead=heads,
dim_feedforward=dim * 4, dropout=0.1,
batch_first=True
)
self.encoder = nn.TransformerEncoder(layer, num_layers=num_layers)
def forward(self, tokens):
return self.encoder(tokens)
# 解码器:级联上采样 + 跳跃连接融合
class CUPDecoder(nn.Module):
def __init__(self, dim=768, num_classes=9):
super().__init__()
self.upsample1 = self._up_block(dim, 256)
self.upsample2 = self._up_block(256, 128)
self.upsample3 = self._up_block(128, 64)
self.upsample4 = self._up_block(64, 16)
self.head = nn.Conv2d(16, num_classes, kernel_size=1)
def _up_block(self, in_c, out_c):
return nn.Sequential(
nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True),
nn.Conv2d(in_c, out_c, kernel_size=3, padding=1),
nn.BatchNorm2d(out_c), nn.ReLU(inplace=True)
)
def forward(self, tokens, skip, hw):
H, W = hw
x = tokens.transpose(1, 2).reshape(-1, tokens.shape[-1], H, W)
x = self.upsample1(x)
x = self.upsample2(x)
# 拼接浅层细节特征,还原边界信息
x = torch.cat([x, skip], dim=1)
x = self.upsample3(x) if x.shape[1] == self.upsample3[0][1].in_channels else x
x = self.upsample4(x)
return self.head(x)
# 完整的TransUNet组装
class TransUNet(nn.Module):
def __init__(self, num_classes=9):
super().__init__()
self.encoder = HybridEncoder()
self.transformer = TransformerBlock()
self.decoder = CUPDecoder(num_classes=num_classes)
def forward(self, x):
tokens, skip, hw = self.encoder(x)
tokens = self.transformer(tokens)
return self.decoder(tokens, skip, hw)
net = TransUNet()
out = net(torch.randn(2, 3, 224, 224))
print(out.shape) # 输出与输入同尺寸的分割图
上述代码中需要注意几个细节。首先是patch划分方式,原始TransUNet把下采样16倍后的特征图每个像素当作一个token,也可以改成2x2的patch进一步减少token数量。其次是位置编码,如果用可学习的位置编码需要固定输入尺寸,而插值式正弦编码可以支持可变分辨率输入,医学场景中数据尺寸不统一时更灵活。最后是跳跃连接的拼接位置,建议把最浅层、分辨率最高的特征接到最后一级上采样之后,这样边界细节损失最小。
四、性能表现、局限性与适用场景
在公开的Synapse多器官CT数据集上,TransUNet的平均Dice系数达到77.48%,相比基线V-Net的65.04%和Attention U-Net的69.72%有显著提升,尤其是对小体积器官(如胆囊、胰腺、食管)的分割改善明显,这正体现了全局上下文建模的价值:小器官自身特征不明显,但与其他器官的相对位置关系可以作为有效线索。在自动心脏诊断挑战赛(ACDC)的数据集上,TransUNet同样表现稳定。
当然,TransUNet也有局限。参数量约为105M,远大于U-Net的约30M,推理速度和显存占用都是短板,部署到临床实时场景需要做剪枝或蒸馏。另外,Transformer部分通常依赖ImageNet或更大规模数据集上的预训练权重,从零训练小数据集效果会打折扣。后续的Swin-UNet引入了层级化窗口注意力,Swin UNet等改进版本在降低计算量的同时保持了精度,可以作为工程选型时的对比方案。
从适用场景来看,TransUNet适合对分割精度要求高、器官或病灶之间存在明显空间关系依赖的任务,例如腹部多器官分割、放疗靶区勾画等。如果任务本身局部特征就足够判别(比如对比度很高的细胞核分割),纯CNN方案可能已经够用,此时引入Transformer反而带来不必要的开销。选型时建议先跑通一个U-Net基线,再评估精度差距是否值得为混合架构付出额外的计算成本。
总结来说,TransUNet的成功不在于单独发明了什么新模块,而在于把CNN的局部归纳偏置与Transformer的全局建模能力做了合理的串联分工,并用U型结构和跳跃连接把两者的输出有效融合。这种混合设计思想后来被大量分割模型借鉴,理解它对掌握医学图像分割领域的模型演进脉络非常有帮助。
TransUNet医学图像分割Transformer修改时间:2026-08-31 08:04:39