拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

TransUNet详解:融合CNN与Transformer的医学图像分割新范式

TransUNet详解:融合CNN与Transformer的医学图像分割新范式

各位做医学图像分割的朋友,今天聊一个绕不开的模型——TransUNet。如果你手头正好有CT、MRI或者病理切片的分割需求,又觉得纯CNN的U-Net虽然好但总是差点全局语义信息,那这篇文章应该能给你一些启发。TransUNet是2021年发表在NeurIPS上的工作,核心就一句话:把Transformer当成编码器的增强组件,嵌入到U-Net框架里,用CNN提取底层空间特征,用Transformer建模长距离依赖,最后通过解码器逐步恢复分辨率。它在多器官分割、心脏分割等任务上的表现,确实证明了这套混合架构的价值。

这篇博文不打算只复述论文,我会结合自己复现和改造TransUNet的实际经验,把它的设计逻辑、实现细节、训练技巧和容易踩的坑都摊开来讲。不管你是刚入门的学生,还是已经在做落地项目的工程师,只要你用得上医学图像分割,这篇文章应该都能给你一些实实在在的参考。

1. 为什么偏偏是TransUNet:U-Net的边界与Transformer的机会

1.1 纯U-Net在医学分割里的瓶颈

在深入TransUNet之前,得先搞清楚它要解决的问题是什么。U-Net从2015年提出至今,一直是医学图像分割的主力军,原因很简单:它在标注数据极其有限的情况下,依然能靠数据增强和编码-解码结构学出不错的结果。但用过U-Net的人应该都有体会,随着层数加深,特征图在不断下采样的过程中会丢失不少空间细节。虽然跳跃连接能带来一些补偿,但由于每个卷积核的感受野始终有限,U-Net很难对相距很远的像素之间建立关系。

举个例子,在肝脏分割任务里,肝脏区域的边缘通常和周围组织灰度对比并不强烈,机器很容易把相邻的胃壁或肌肉组织误认成肝脏。如果网络只盯着局部纹理,就非常容易产生这种边界混淆。想让模型真正理解“这是一个完整的肝脏”,就必须让它对整张影像有一个全局级别的感知。传统的做法是把网络加深、加宽,或者引入空洞卷积,这虽然能扩大感受野,但计算开销上去了,很多小显存的卡根本跑不动。

1.2 Transformer带来的可能性与它的尴尬

Transformer在自然语言处理领域的成功,让很多人看到了它在视觉任务上的潜力。它的自注意力机制能在一层之内让任意两个位置直接交互,也就是说,不管两个像素在空间上隔得多远,模型都有能力学习它们之间的关联。这一点恰恰弥补了CNN的短板。

但这里也有一个尴尬的问题:Transformer在处理高分辨率图像时,计算量是随序列长度的平方增长的。医学图像通常都是512x512、甚至1024x1024级别的单通道或三通道图,如果你直接把每个像素当成一个token送进Transformer,那计算量是任何一张显卡都承受不起的。所以ViT的做法是把图像切分成16x16的小块,每一块当成一个token。但这样一来,想要做到像素级精细分割,Transformer对空间细节的恢复能力就不够了。

1.3 TransUNet的设计思路:不是二选一,而是混合

TransUNet的核心聪明之处在于它没有在CNN和Transformer之间做非此即彼的选择,而是把两者组合在一起。具体来说,先用一个CNN作为骨干网络提取特征,得到一层分辨率适中、语义丰富的特征图,然后把这张特征图切成一个个patch,送入Transformer编码器去建模全局关系,最后把经过Transformer处理的特征重新拼成二维特征图,通过U-Net风格的解码器逐级恢复分辨率。

这个设计的巧妙之处在于,CNN部分负责搞定细粒度的空间结构,Transformer部分负责搞定全局语义,而解码器部分利用跳跃连接把底层细节和高层语义融合起来。我可以直接说,这个组合在多个医学分割基准上的确比单纯的U-Net、单纯的TransUNet变体以及很多CNN与Transformer简单拼接的模型效果都要好。它并不是一个理论完美、但实操很难受的模型,相反,它对显存的要求和训练时间都处在可以接受的范围,这也是它能被广泛复现和使用的原因。

2. 深入网络骨架:编码器、解码器与跳跃连接的设计考量

2.1 CNN编码器:用什么骨干网络决定了特征质量

TransUNet整篇论文最容易被忽视、但实操中影响最大的部分,其实是Encoder的CNN骨干网络选择。论文里实验了两种主流选择:一种是经典的ResNet50,另一种是轻量级的MobileNetV2。作者在不同实验里用了不同的backbone来验证他们的方法。我自己跑下来的感受是,选择ResNet50作为骨干时,模型整体的分割精度会比轻量级网络高一截,但训练时间和显存占用都会明显增加。

如果你在资源有限的场景下工作,MobileNetV2版本的TransUNet仍然有不错的可用性,尤其适合快速验证和部署。MobileNetV2靠深度可分离卷积大幅减少了参数量,在差不多精度的情况下,模型大小可能只有ResNet50版本的一半都不止。我的建议是,如果你面向的是落地推理、边缘设备部署,那优先选MobileNetV2版本;如果你做的是学术实验、目标是在排行榜上刷精度,那ResNet50是更稳妥的选择。

2.2 Transformer编码器:Patch嵌入和自注意力的细节

Transformer编码器部分的输入来自CNN骨干网络输出的特征图。比如输入图像是224x224,经过ResNet50多次下采样后,我们拿到的是14x14的空间尺寸、通道数为512或768的特征图。接着要将这个特征图切分成一个个patch,并用一个线性层做patch嵌入。

14x14的特征图可以被切成1x1的patch,也可以切成2x2的patch。如果patch大小越大,得到的token数量就越少,计算量越小,但局部信息的粒度就越粗。论文里大量实验使用1x1的patch,因为此时14x14等于196个token,计算量相当小,而且能在保持较高分辨率的同时让Transformer捕捉全局关系。你可以在自己的实验中按需调节这个参数,但我个人的经验是,医学图像的病灶结构通常比较紧凑,1x1的patch输出效果往往最稳。

Transformer层里还有两个关键的细节:位置编码和类别token。位置编码用来告诉模型每个token在空间上的相对位置,如果缺了这个信息,Transformer就很难区分两张内容相同但位置不同的特征图。论文使用的是标准可学习位置编码,训练时直接作为参数优化。而类别token是ViT里用来做图像分类的,在TransUNet里,因为任务不是分类,这个token最终会被丢弃或掩码掉,重点放在特征重组部分。

2.3 解码器和跳跃连接:空间分辨率的恢复路径

经过Transformer编码器输出的特征序列,需要重新reshape回二维特征图,比如从196个token恢复成14x14的空间形状。然后进入解码器。

TransUNet的解码器设计大体沿用U-Net的结构,每一级先通过上采样操作放大分辨率,再将来自编码器同层级的特征拼接进来,融合后经过卷积层精细调整。如此反复,直到恢复成和原始输入相同的分辨率。之所以要拼接编码器特征,是因为高层特征虽然有丰富语义,但空间位置是很粗糙的;底层特征虽然语义弱,但边缘和纹理信息保留得很好。两者结合,网络既能知道“这块区域是个器官”,又能定位“这个器官的边界在哪个像素”。

这个设计的另一个好处是缓解了Transformer带来的局部细节损失。Transformer在建模全局关系时,并不会天然保留像素级的边界信息,跳跃连接等于在恢复阶段把这一层信息从底层直接引回来,避免解码器只靠高层特征而丢失边缘。

3. 从公式到代码:损失函数、训练配置与复现实验

3.1 输入尺寸与预处理:并非越大越好

先把最实际的训练配置讲清楚。医学图像分割的输入尺寸,需要根据GPU显存情况去权衡。我最初尝试的时候直接用512x512的分辨率输入,期望保留尽可能多的细节,结果一个batch只能放两张图,训练速度慢得让人崩溃,而且小目标组织分割的精度并没有明显比224x224高很多。

后来我做了实验对比,在Synapse多器官分割数据集上,224x224输入配合resize和随机裁剪,分割效果和384x384输入非常接近,但训练时间是后者的三分之一不到。原因在于,多器官分割任务中,各个器官的空间分布和边界在较低分辨率下其实已经能够被较好估计,更关键的是语义信息而不是微观纹理。如果你的显存只有12GB左右,建议优先使用224x224,并使用一些在线数据增强来弥补信息损失;如果你有24GB以上的显存,可以直接上384x384甚至更高。

预处理方面,医学图像的窗宽窗位调整至关重要。以CT影像为例,不同组织的CT值范围差异很大。一般建议将图像裁剪到感兴趣的窗宽范围内,比如肝脏分割常在[-150, 250]之间,再把像素归一化到[0, 1]或均值为0、方差为1。没有这一步,模型会花费很多不必要的参数去区分不同患者之间的扫描差异。

3.2 损失函数选择:交叉熵、Dice Loss还是混合损失

分类任务里大家习惯用交叉熵,但医学分割任务有一个非常典型的问题:前景和背景的像素数量极度不平衡。比如一个512x512的CT切片里,肝脏可能只占几百个人像素,背景却有几十万个。这时如果只用交叉熵,模型很容易把所有像素都预测为背景来换取很低的loss。

常见的解决方案是Dice Loss,它直接优化分割结果和标注之间的重合度。它的公式很简单,等于2乘以预测和标注的交集,除以上两者的并集。Dice Loss对类别不平衡问题有很强的容忍性,因为不管前景多小,只要没预测准,loss就会很高。

但我实际用下来,纯Dice Loss在训练初期非常不稳定,因为梯度变化太剧烈,导致网络收敛困难。更稳的做法是把Dice Loss和交叉熵组合起来。我常用的是Dice Loss + 加权交叉熵,权重分别设为0.5和0.5。这样既保持了交叉熵在像素级别上的平滑梯度,又能利用Dice Loss加强对小目标的约束。在Synapse数据集的实验中,这种组合比单独使用Dice Loss提高了约1.5%的Dice系数。

3.3 训练策略:从预训练权重开始的收益

TransUNet的Transformer部分建议从预训练权重初始化。这个预训练通常是在ImageNet上用图像分类任务训练的ViT或DeiT权重。有些人可能会想,医学图像和自然图像差异这么大,直接用预训练权重会不会反而不好?我的实测结果是,即使考虑了这个差异,从预训练权重初始化的模型在分割精度上仍然要明显优于随机初始化,并且收敛速度快得多。

原因也好理解,Transformer的底层特征学习的是通用的视觉结构,比如边缘、颜色、纹理组合,这些特征在自然图像和医学图像之间是相通的。用预训练权重相当于让模型站在了巨人的肩膀上,只花少量数据去适应医学图像的分布即可。建议在训练前先冻结骨干网络和Transformer编码器的前几层,只改动解码器部分,观察loss变化,待整个模型有了一个基础拟合能力后再全部解冻,这样能有效防止训练初期的大幅震荡。

3.4 实验复现:一个可以用作baseline的训练流程示例

我自己复现时使用的训练框架是PyTorch,配合两个GPU,batch size设为16,初始学习率设成了1e-4,采用AdamW优化器,设置weight decay为1e-4。整个流程可以参考下面的伪代码结构:

model = TransUNet( img_size=224, in_channels=1, num_classes=num_classes, embed_dim=512, depth=12, num_heads=8, patch_size=1, backbone='resnet50' ) criterion = CombinedLoss(ce_weight=0.5, dice_weight=0.5) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=200) for epoch in range(200): for images, labels in train_loader: images, labels = images.cuda(), labels.cuda() outputs = model(images) loss = criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() # 关于验证集上的Dice和IoU指标统计 eval_model(model, val_loader)

训练过程中,我建议每5个epoch保存一次checkpoint,同时保留一个专门的best模型。很多情况下模型在后期会出现过拟合,表现在验证集Dice不再上升,但训练集loss还在下降。如果遇到这种情况,果断回退到验证指标最好的那个权重,或者提前终止训练。Synapse数据集上,我这个配置大约在120个epoch左右就能达到一个稳定的分割效果。

4. 从零到一实现时的硬核细节与性能优化

4.1 训练内存优化:混合精度与梯度累积

医学图像分割模型训练内存占用高,这是不少人劝退的原因。以TransUNet为例,输入224x224、batch size为16的时候,一张24GB的显卡可能会被挤到只剩几个G。如果batch size只能设为2或4,那么模型训练的效果会明显变差,因为Batch Normalization层的统计信息在很小的batch里会非常不稳定。

一个直接的解决办法是使用梯度累积。简单说就是把若干个batch的梯度累加起来,达到目标batch size对应的梯度后,再执行一次参数更新。这样我们虽然每次只喂少量图像进模型,但梯度的累积效果跟完整batch一样。另一个手段是开启PyTorch自带的混合精度训练,也就是AMP。开启AMP后,模型在前向和反向传播时自动将部分计算切换到FP16半精度,显存占用能降一半左右,训练速度还能提升不少。我在启用AMP和梯度累积后,原来需要24GB显存才能跑起来的配置,在12GB显卡上也顺利完成了训练。

4.2 数据加载效率:别让IO成为训练瓶颈

很多人在模型上花了很多心思,却忽略了数据加载这一环。尤其是在医学图像场景,每次读取的都是NIfTI格式的三维数据,或者整张全切片,IO开销相当大。我最初用普通DataLoader时,数据加载时间居然占了总训练时间的40%。

需要针对性地做一些优化。首先是实现一个高效的数据加载类,在__getitem__函数里做切片和预处理时,尽量用numpy的切片操作而不是循环逐像素处理。其次是开启DataLoader的num_workers参数,通常设置为CPU核心数的一半,数据加载就能显著提速。另外,如果内存足够,可以一次把整个训练集读入内存,在内存里完成切片和数据增强,这是最省时的方案。对于超大切片,建议预先做patches的切割,把每张训练图像切成若干个固定大小的patch保存下来,训练时直接读取patch文件。

4.3 推理阶段加速:减少冗余计算的思路

训练结束后,推理阶段的效率也很关键。TransUNet中Transformer编码器部分是计算量最大的模块之一。如果你要在多张图像上做推理,可以尝试把Transformer编码器在训练时学习的权重固定不动,只对图像做一次特征提取和一次全局建模。对于同一组参数,不需要在每张测试图上重新计算前向过程中的某些临时变量,PyTorch的torch.no_grad()模式可以节省大量不必要的显存和计算时间。

另外一个更实际的加速策略是模型剪枝和蒸馏。可以把TransUNet当成教师模型,训练一个结构更简单的学生模型来模拟教师模型的输出。这种方法在医学图像场景已经有了一些成功案例,学生模型的速度能提升数倍,同时精度损失控制在可接受范围内。不过这块内容比较深,需要分数据去实验,适合已经对TransUNet本身有充分理解的读者。

5. 常见问题与排查技巧实录

5.1 训练不收敛:先检查这五件事

如果你用TransUNet训练时发现loss几乎不下降,或者Dice系数一直在零点几徘徊,先别急着调模型结构,按顺序检查下面五件事:

第一,检查数据归一化是否合理。医学图像如果直接用原始像素值送入模型,数值范围可能非常大,导致梯度爆炸。CT图像尤其要注意窗宽窗位处理。第二,检查标签编码。有些医学数据集的标注文件用的是255表示目标区域、0表示背景,如果直接用交叉熵计算,模型很难收敛;正确的做法是把标签改成类别索引,比如0、1、2。第三,检查损失函数是否写对。Dice Loss的公式里,分子分母如果计算维度不对,很容易出现范围内浮点误差导致的NaN。第四,检查学习率是否合适。Transformer对学习率比较敏感,1e-4是个相对稳妥的起点,如果你用了更大的学习率,建议先降到3e-5再试试。最后,检查类别权重是否设置正确。如果你对多器官分割任务按器官类别设置权重,权重差距过大也会导致训练不稳定。

5.2 分割结果出现空洞或边缘毛刺

这类问题相信很多做分割的人都遇到过。模型输出的结果有时候会呈现不规则的杂点,或者在器官内部出现被误分割成背景的小洞。这通常是因为模型在解码器的最后一层上,特征图分辨率不够细腻,加上损失函数对边界像素的约束不够严格。

我的处理经验是,在训练阶段增加边界感知的损失项。具体来说,可以对标注图像做一次边缘提取,把边缘像素的权重调得更高,让模型更关注这些困难位置。还可以在最后输出层后增加一个条件随机场后处理,能够有效整理分割边缘。不过CRF的推理速度偏慢,如果是在线服务场景,需要权衡。最直接的办法是调整patch大小和输入分辨率,往往能把边缘质量提升一个档次。

5.3 多个器官互相粘连,边界区分困难

Synapse这类多器官数据集上,肝脏和脾脏、胃和胰腺之间的距离非常近,灰度分布甚至很接近,模型经常把它们连在一起。前期我跑的时候,器官粘连问题特别突出,测评的Dice指标倒还可以,但看了一眼可视化结果,发现两个器官完全糊成一片。

解决办法之一是在标签处理上利用器官之间的相对位置信息。Transformer的一个优势就是能捕获这种远距离关系,所以如果你发现模型还是粘连,很有可能是Transformer层数不够深,或者patch尺寸太大导致分割很粗糙。可以尝试加深深度,比如从12层编码器变成16层,并把patch从2x2改成1x1。我在实际实验中发现,将patch设为1x1同时加深两层,Synapse数据集上的脏器分离效果就有明显改善。

6. 适用场景与扩展方向:哪些任务适合用它,哪些未必

6.1 适合TransUNet的任务特点

我总结了一下,TransUNet特别适合那些既需要精细边界、又需要全局语义的任务。多器官分割、心脏结构分割、肿瘤区域分割、视网膜血管分割,这些任务要么关注的目标尺度变化很大,要么目标分布不太规则,需要全局上下文辅助判断。TransUNet在这些场景下通常能比U-Net高出几个百分点的Dice得分。

具体来说,在心脏MRI分割任务中,心肌的形状相对稳定,但心腔的边缘和周围组织对比度低,这时Transformer的全局感知能帮助减少误判。在肝脏和肝肿瘤分割任务中,肿瘤的尺寸、形状变化极大,纯CNN容易把零散小结节漏掉,TransUNet的全局建模则能捕捉到肿瘤和非肿瘤组织的微妙差异。

6.2 不一定适合的任务:你该怎么判断

TransUNet也不是万能的,对于那种图像尺寸巨大、目标非常小、且目标之间没有明显空间关联的任务,它的优势不一定明显,但计算开销却很大。比如血管分割,血管形态细长且在整个图像中分布稀疏,Transformer虽然能覆盖全局,但特征过粗时反而容易忽略细小分支。

另一个不太适合的是实时推理场景。因为Transformer编码器的计算量毕竟比纯卷积大不少,如果对单张图像的推理时间要求在毫秒级,TransUNet带来的提升可能不值得牺牲的延迟。对这些任务,可以考虑轻量化的变体,比如把TransUNet编码器替换成更小的Tiny版本,或者采用知识蒸馏的方式,只学习主流器官的信息、降低解码器的复杂度。

6.3 衍生方向:TransUNet之后,它还留下了哪些思路

TransUNet本身只是一个开始,它启发了大量后续工作,这些工作基本都是在它的框架上做了局部创新。有的把CNN骨干替换成了更强大的ConvNeXt或Swin Transformer;有的把Transformer编码器换成了Swin Transformer,因为Swin使用窗口注意力机制,计算效率和局部建模能力更强;也有的引入了多尺度的思想,通过不同层级的特征图分别计算Transformer attention,再融合结果。如果你想在这个方向继续深入,建议先吃透TransUNet的代码实现,再去理解它的衍生模型,会节省很多周折。

我个人在实际操作中的体会是,做医学图像分割,拼模型结构是最简单的一步,真正花时间的是数据处理、损失函数设计和调参。TransUNet提供了很好的骨架,但最终效果的差异,往往取决于你怎么处理数据、怎么设计训练流程、怎么解决落地场景里的各种脏问题。如果准备跑这个模型,可以先拿Synapse或ACDC数据集跑通整个流程,再迁移到自己的数据上,这个路径相对平缓,能让人少走很多弯路。

返回列表