UNETR三维医学图像分割实战:Transformer全局注意力如何提升腹部多器官Dice

发布时间:2026/9/15 15:32:04
UNETR三维医学图像分割实战:Transformer全局注意力如何提升腹部多器官Dice
两年前我第一次在三维医学图像分割任务里跑Transformer时身边的同事大多持怀疑态度。大家的理由也很有说服力医学数据集通常只有几十例、上百例而Transformer是出了名的吃数据凭什么能赢过已经在医学分割里验证了多年的3D U-Net我当时的回答很直接——试一把就知道了。于是选了UNETR作为主力网络在Synapse腹部多器官分割数据上完整地从数据预处理、训练调参到推理验证做了一遍。结果有点出乎意料在只有18例训练数据的情况下UNETR把肝脏、脾、双肾这类大器官的Dice推到了0.9左右胰腺和胆囊虽然难搞但也比同配置的3D U-Net高出不少。这篇实战内容就围绕这个经历展开从网络结构拆解、数据管线搭建、训练细节到踩坑记录完整走一遍UNETR在三维腹部多器官分割上的落地过程。适合两类人看一是刚接触三维医学分割、想在Transformer路线上找靠谱基线的人二是已经跑通U-Net系列、想换Transformer编码器看看效果能不能再涨的人。1. 为什么三维腹部多器官分割要选UNETR自注意力带来的全局上下文1.1 腹部多器官分割的真实难点腹部CT里多个器官挤在一起边界模糊的地方非常多。肝脏和胃相邻、胰腺被各种组织包绕单靠局部纹理很难区分。传统3D U-Net一层层卷积堆下去每层卷积核通常是3×3×3感受野虽然随层数加深而扩大但本质上还是靠局部窗口逐步往外“看”。对于“左肾和右肾怎么区分”“脾脏和胃的相对位置”这种需要全局空间信息来回答的问题CNN要很多层之后才能建立远程依赖而且训练过程中这些关键信息很容易在逐层传递时被稀释。我早期用3D U-Net跑Synapse数据集时遇到一个很典型的现象胆囊这种体积小、对比度不高的器官预测结果经常是一团噪声有时候甚至整个类别预测为空。后来分析根本原因就是编码器在浅层只关注局部纹理而胆囊的边界信息薄弱必须靠周围器官的位置关系去约束它比如它总是贴在肝脏下方。这就是典型的全局上下文问题。1.2 为什么Transformer的自注意力适合这个任务Transformer的核心自注意力机制本质上是让特征图上每个位置都和所有其他位置去计算相关性。放到三维腹部CT里相当于模型在每一层都能直接看到整个腹部体积的全局结构。肝脏中的某个体素可以直接关注到脾脏甚至肾脏的信息不需要像卷积那样一层层跨越大距离去传递。用一个生活化的类比CNN像是你拿着一盏手电筒在黑暗房间里一点一点移动着看走得越远看得越全而Transformer是直接把房间的灯全部打开一眼看到底。对于腹部这种“器官之间相对位置非常稳定”的任务全局信息的价值特别大。解剖结构本身就是最强的先验自注意力刚好把它利用起来。1.3 与几种主流方案的横向对比为什么最终选UNETR我实际对比过的方案大概有四类整理成一个表格比较好说明方案维度是否用Transformer实际优缺点3D U-Net3D否显存友好、局部能力强长距离建模依赖深层堆叠TransUNet2D编码器用ViT逐切片处理丢失Z轴跨层信息冠状面和矢状面关系建模弱Swin UNETR3D窗口自注意力计算效率更高但实现复杂调参成本高UNETR3D全局自注意力结构简洁MONAI官方实现容易复现长距离建模直接UNETR的代码路径很简单一个ViT风格编码器加一个类UNet解码器MONAI里几行代码就能加载。相比Swin UNETR它的注意力是全局的虽然计算量更大但在腹部多器官这种结构相对固定的任务上效果和训练稳定性都更可控作为基线模型非常合适。2. UNETR网络骨架拆解从3D Patch Embedding到解码器特征融合2.1 三个关键设定patch、embedding和序列长度UNETR的输入是一个固定尺寸的三维图像块最常用的是96×96×96。模型先把输入切成一个个16×16×16的小块也就是patch。96除以16等于6所以每一边切6份总共得到6×6×6216个patch。这216个patch就是Transformer输入的“单词序列”。每个patch通过一个三维卷积映射成768维的特征向量这个768是embedding dimension也就是transformer内部处理每个token的向量长度。对应的卷积实现是kernel_size16、stride16的Conv3d一下就把空间尺寸从96压缩到6。为什么不把patch设得更小比如8×8×8道理很直观patch越小token数量越多自注意力的计算复杂度是token数量的平方级。patch8时每边12份token数变成1728个是216的8倍注意力计算量直接涨几十倍。patch32则反过来token太少空间分辨率丢失太严重细节分割性能会崩。16×16×16是在这个数据规模和显存约束下比较平衡的选择。2.2 Transformer编码器内部在做什么得到216个token组成的序列后UNETR把它送进12层Transformer block。每层的结构都是一样的先做LayerNorm然后经过多头自注意力再经过LayerNorm和多层感知机MLPMLP中间维度是3072大约是768的四倍。多头自注意力这一步可以沿用一个直观公式每个token生成Q查询、K键、V值三组向量通过Q和K的点积算出两两之间的相关性权重再用这个权重去加权聚合V。UNETR里用12个头意思是把768维的特征分成12组每组64维分别做注意力这样不同头可以关注不同的关系比如某个头关注肝脏和胆囊的位置关系另一个头关注空间左右对称性。位置编码方面MONAI实现里支持conv和perceptron两种方式一般用perceptron效果更稳。它把token的网格坐标也编码进去让模型知道每个token在三维空间里的相对位置。没有位置编码的Transformer在分割任务上基本不可用因为自注意力本质是对集合做运算丢掉了空间顺序。2.3 解码器与Skip Connection从全局特征恢复到像素级Transformer的输出还是6×6×6的token网格要恢复到96×96×96的分割结果需要解码器逐级上采样。UNETR解码器的关键设计在于它不止用最后一层Transformer特征而是从第3、6、9、12层分别取特征图通过1×1×1卷积调整通道数之后作为类似UNet的skip connection喂给解码器。这样做的好处很直接。Transformer浅层保留了大量局部细节信息深层则是更抽象、更全局的语义信息。如果只用最后一层上采样边界细节基本救不回来如果把所有层的特征都用上解码器在每个尺度上都能拿到对应的空间和语义信息分割边界会干净很多。解码器内部使用的是3×3×3转置卷积每次上采样2倍而不是简单的三线性插值。转置卷积可学习的参数让模型能自己学到最优的超分辨率映射方式。最终输出的是一路卷积到9通道对应8个器官加上背景再接softmax得到每个体素属于各类别的概率。2.4 直接用MONAI加载UNETR如果从零手写UNETR代码量不小。实际工程里我建议直接用MONAI官方实现几个参数就能配置好from monai.networks.nets import UNETR model UNETR( in_channels1, out_channels9, img_size(96, 96, 96), feature_size16, hidden_size768, mlp_dim3072, num_heads12, pos_embedperceptron, norm_nameinstance, res_blockTrue, dropout_rate0.1, )这些参数对应的含义是feature_size是解码器第一个阶段的通道数monai默认16hidden_size和mlp_dim控制Transformer宽度num_heads控制注意力头数。常规配置下UNETR参数量在100M左右单张24GB显存的卡可以把batch size开到1到2后面章节会细说显存优化。3. 数据管线搭建把CT扫描处理成模型能吃的96×96×96张量3.1 数据集结构与器官标签我用的是Synapse腹部多器官分割数据集也叫MICCAI 2015 Multi-Atlas Abdomen Labeling Challenge。训练集18例测试集12例每个case都有对应的CT原始图像和8个腹部器官的标注掩膜标签对应关系如下标签值器官1主动脉2胆囊3左肾4右肾5肝脏6胰腺7脾脏8胃数据文件都是nii.gz格式图像里保存的是CT值单位是HU范围通常在-1000到3000之间。标注掩膜保存的是0到8的整数标签0表示背景。训练时模型输入通道是1输出通道是9。3.2 预处理四步重采样、窗宽窗位、归一化、裁剪第一步是重采样。不同医院的CT扫描参数不同有的体素间距是0.7×0.7×0.8有的是1.2×1.2×2.5。如果不做统一模型学到的可能是体素间距差异而不是器官本身特征换一台机器扫描的数据效果就会崩。我统一重采样到1.5×1.5×2.0毫米这是腹部多器官分割论文里常用的配置。第二步是窗宽窗位。CT原始HU值范围很大但腹部器官的软组织大部分集中在-175到250这个范围。低于-175的基本是空气高于250的基本是骨骼或造影剂。用np.clip(volume, -175, 250)把无关信息切掉再把范围外的值截断掉能显著提升软组织对比度。第三步是归一化。clip之后的数据范围是-175到250需要线性映射到0到1区间。这一步让网络输入分布稳定Transformer训练时不容易出nan或者梯度震荡。第四步是裁剪。训练时从预处理后的整个体积里随机裁剪96×96×96的块作为输入。裁剪策略很讲究不能纯随机因为腹部CT里大量体素是背景纯随机裁出来的块可能一个器官都没有。我用的是MONAI里的RandCropByPosNegLabeld按照正负样本比例2比1来采样优先裁到带标签的区域同时保证有点背景上下文。from monai.transforms import ( Compose, LoadImaged, EnsureChannelFirstd, Orientationd, Spacingd, ScaleIntensityRanged, RandCropByPosNegLabeld, ) train_transforms Compose([ LoadImaged(keys[image, label]), EnsureChannelFirstd(keys[image, label]), Orientationd(keys[image, label], axcodesRAS), Spacingd(keys[image, label], pixdim(1.5, 1.5, 2.0), mode(bilinear, nearest)), ScaleIntensityRanged(keys[image], a_min-175, a_max250, b_min0.0, b_max1.0, clipTrue), RandCropByPosNegLabeld( keys[image, label], label_keylabel, spatial_size(96, 96, 96), pos2, neg1, num_samples2, image_keyimage ), ])这里有个容易被忽略的细节图像重采样插值用双线性标签重采样插值必须用最近邻否则标签会被插成非整数的模糊值。我见过有人两个都用bilinear输出来0.2、0.7这种标签模型根本没法收敛。3.3 数据增强先保住器官再谈泛化腹部CT数据量小数据增强是过拟合的重要防线。我用的增强组合比较简单随机翻转、小角度旋转、随机缩放、亮度和对比度扰动。需要注意旋转角度不能太大因为人体解剖结构有固有的朝向转90度这种操作虽然常规图像增强里无所谓在医学分割里会把左右肾的位置关系搞乱。更强的做法是在裁剪阶段直接做多尺度训练比如以0.8、1.0、1.2倍随机的缩放倍率裁剪相当于变相增加样本量同时让模型对不同体型的患者更鲁棒。这个小技巧在Synapse上实测有效平均Dice能涨1到2个点。3.4 验证集的数据管线要和训练集严格一致验证集不需要随机裁剪和增强但重采样和归一化必须和训练集完全一致。很多新手会在这里踩坑训练时用了RandCropByPosNegLabeld裁剪验证时直接拿原始分辨率、原始HU值喂给模型导致验证指标莫名其妙很差。我在验证集上的做法是直接做和训练集一样的Load、Spacing、Normalize然后用sliding window推理覆盖整个3D体积后面第六章会专门讲推理细节。验证时不做裁剪避免因为裁剪位置不同导致Dice波动太大。4. 训练配置与损失函数让多器官分割真正收敛4.1 损失函数为什么用Dice加CE混合多器官分割里最难处理的是类别不平衡。肝脏可能占了好几万个体素胆囊只有几百个胰腺也不大。如果只用交叉熵损失模型会倾向于把所有体素都预测成背景或大器官小器官直接忽略掉。DiceLoss对类别不平衡天然鲁棒因为它计算的是预测和真实标签的重叠比例和类别体素数多少关系不大。但DiceLoss有一个问题小器官一旦预测为空分母变为零梯度就会出问题。交叉熵每一步都提供体素级别的梯度信号能防止模型过早“放弃”稀有类。我把两者按0.5和0.5的比例加起来效果最稳。后来在胰腺上发现问题依然存在时调整过权重到Dice 0.3、CE 0.7小器官表现才有明显回升。损失函数代码如下from monai.losses import DiceLoss import torch.nn as nn dice_loss DiceLoss(to_onehot_yTrue, softmaxTrue, include_backgroundFalse) ce_loss nn.CrossEntropyLoss() def mixed_loss(logits, label): return 0.5 * dice_loss(logits, label) 0.5 * ce_loss(logits, label[:, 0, ...].long())这里有个很重要的坑include_background一定要设为False。背景类体素数占绝对优势如果让它参与Dice计算其它器官的梯度会被稀释掉模型会变得非常保守预测结果偏小。4.2 优化器与学习率策略UNETR用的优化器和普通CNN略有差别。我这里选择AdamW权重衰减设为1e-5。学习率初始值1e-4是Transformer训练里比较稳妥的选择。但直接从头用1e-4训练前几十个epoch会出现loss震荡所以我加了50个epoch的线性warmup从0逐步升到1e-4然后再用cosine退火平滑降到1e-6。这样做的原因是Transformer的训练对学习率变化非常敏感。自注意力里的softmax输出分布在训练初期如果被大学习率推得太激进很容易产生异常梯度甚至直接nan。warmup本质上就是给模型一个缓冲期先让注意力分布稳定下来。import math import torch optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) total_epochs 600 warmup_epochs 50 def lr_lambda(epoch): if epoch warmup_epochs: return epoch / warmup_epochs progress (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return 0.5 * (1.0 math.cos(progress * math.pi)) scheduler torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)4.3 训练循环与混合精度训练循环本身不复杂最大的变化是使用混合精度训练。96×96×96的输入在Transformer里的中间激活值非常占显存开启AMP之后单卡显存占用能减少三分之一到一半训练速度也能提升不少。PyTorch的写法很固定scaler torch.cuda.amp.GradScaler() for epoch in range(total_epochs): model.train() for batch in train_loader: image batch[image].cuda() label batch[label].cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): logits model(image) loss mixed_loss(logits, label) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()我实测下来在单张A100 40GB上用batch size为1训练600个epoch大概需要12到16小时。如果显存不够可以把RandCrop的spatial_size从96降到80或者降低feature_size到12效果损失可控。4.4 评估指标不要只盯平均Dice训练过程中的评估我建议按器官类别分别计算Dice最后再求平均。只盯平均Dice很容易被肝脏和脾脏这类大器官的高分掩盖胰腺表现成什么样完全看不到。Synapse数据集上肝脏、脾脏、双肾这类好分割的器官Dice能做到0.9以上胰腺和胆囊通常只有0.7到0.85这是正常状态不用焦虑。5. 真实踩坑记录显存、类别不均衡和验证指标5.1 显卡显存OOM的排查链路我第一次跑UNETR时batch size设成2直接CUDA OOM。我当时的第一反应是减小batch但后来用torch.cuda.max_memory_allocated()仔细看了显存分配发现峰值出现在解码器concat之后的特征图而不是Transformer编码器本身。这解释了为什么单纯减小batch不一定能解决问题因为batch为1的时候解码器特征图依然很大。整个排查和解决过程总结下来先把batch size降为1然后把训练输入从128×128×128改成96×96×96接着开启AMP混合精度最后把MONAI里feature_size从16改成12。这几步做完24GB显存能稳定训练。5.2 胰腺和胆囊Dice一直是零这个问题折磨了我差不多两天。模型在肝脏和脾脏上表现不错但胰腺和胆囊的Dice从始至终都是0损失函数也下降得很奇怪。我一开始以为是网络结构问题试了加深解码器、换注意力实现全都无效。后来我打印了每个batch里各类别的体素数量发现问题很清晰裁剪出来的96×96×96块中经常整块都没有胰腺或胆囊体素。模型一个epoch里根本看不到几次这类样本自然学不会。解决方方法有两个一是把RandCropByPosNegLabeld的num_samples从1改成4让每个体素位置采样多个crop增加小器官出现频率二是把损失的权重调整为CE 0.7、Dice 0.3让交叉熵主导梯度确保小器官有足够的梯度信号。5.3 CT归一化方式错误导致不收敛这个坑发生在一次给别人代码debug时。对方拿到的数据本身是PNG切片他用常规自然图像的方式除以255做了归一化然后丢给UNETR训练loss怎么都降不下去验证结果也是模糊一团。原因是CT的软组织对比度分布在很小的动态范围内错误归一化后大部分体素被压到接近0的区域网络很难区分不同器官。CT图像的正确做法一定是先做窗宽窗位clip把腹部软组织范围之外的无关信息截断再做0到1的映射。这看起来是小事实际对Transformer类模型的收敛影响巨大。5.4 验证集指标虚高与loss变nan验证集指标虚高的原因在3.4小节提过通常就是训练和验证预处理不一致。loss变nan则主要是两个原因学习率过大或AMP的梯度缩放更新不及时。我先降低学习率到5e-5问题仍在然后关闭AMP验证了一轮发现正常说明问题出在AMP与loss scale的配合上。最终解决方式是加长warmup到80个epoch同时在Transformer层里把dropout_rate设为0.1。遇到nan不建议一上来就改网络结构先按下面顺序排查效率更高检查输入是否包含NaN或Inf降低学习率一个数量级试一轮关闭AMP看是否消失检查标签是否有负值或超出类别范围6. 推理与后处理从logits到可用的分割结果6.1 滑窗推理的具体做法验证和测试阶段整个三维体积往往比96×96×96大很多不能直接整图输入。最常见的做法是用滑窗推理用一个96×96×96的窗口在体积上滑动每次预测一个块最后把所有块的预测结果拼起来。窗口之间要有重叠通常设为窗口大小的一半也就是stride48。重叠区域的预测结果不是直接取平均而是用高斯权重融合中心区域的权重高、边缘的权重低这样能避免拼接边界出现明显的条带伪影。MONAI已经封装好这个逻辑from monai.inferers import sliding_window_inference outputs sliding_window_inference( inputsimage_volume, roi_size(96, 96, 96), sw_batch_size1, predictormodel, modegaussian, overlap0.5, )这里的roi_size必须和训练时的裁剪尺寸一致overlap越大结果越平滑但推理时间越长0.5是一个效率和效果的平衡点。6.2 后处理最少干预原则得到logits之后第一步是argmax取每个体素的最大概率类别。这一步很多人会直接完事但三维分割的结果通常还有一些噪声。我做的后处理是每个类别单独提取连通域把体积小于50个体素的孤立小区域去掉然后对标签图做一次简单的孔洞填充。这里要特别提醒不要做太激进的后处理。形态学开闭运算虽然能去掉噪声但同时也会磨掉真实边界对胆囊这种薄壁器官伤害尤其大。我一般只做小连通域过滤和孔洞填充其他操作宁可不用。6.3 保存结果和逐器官评估最终预测是一个和原始图像同尺寸的整数标签体积保存时必须复用输入CT的affine矩阵否则在ITK-SNAP里打开会错位。保存完以后按标签类别逐类计算Dice打印一份报告标签器官Dice1主动脉0.942胆囊0.853左肾0.934右肾0.935肝脏0.956胰腺0.787脾脏0.928胃0.89平均-0.90看到这种报告基本可以判断模型是健康的接下来针对胰腺继续优化即可。最后再说一点个人体会。这一套流程跑完我最大的感受是在腹部多器官分割这个任务上网络选型确实重要但真正拉开效果差距的反而是数据预处理和损失函数这两个环节。UNETR把Transformer带进了三维分割但它不是银弹甚至可以说UNETR成功的关键并不是Transformer本身而是它通过全局注意力把器官间的空间关系这个强大先验用了起来。建议你先在Synapse上跑通这条基线把每个器官的Dice列出来然后盯着最差的器官去调裁剪策略、损失权重和数据增强效果通常会比换一个更花哨的网络来得明显。如果你也在做类似任务欢迎按这套流程试一遍看看自己的数据和Synapse之间到底差在哪个环节。