简介:面向医学图像分割与深度学习研究者,这份实战资源基于TransUnet实现腹部多脏器分割,覆盖背景、肝脏、右肾、左肾、脾脏五类目标。项目训练配置完整,采用AdamW优化器、余弦退火学习率衰减与交叉熵损失,共训练100个epoch;测试集像素准确率达0.986,平均IoU为0.779。代码分为训练、评估、预测三大模块,训练脚本自动输出loss/iou曲线、学习率衰减曲线、训练日志、数据集可视化图像以及最终和最优权重;评估脚本计算测试集IoU、召回率、精确率、像素准确率等指标;预测脚本可直接生成分割掩膜和原图叠加效果图。所有代码附详细注释,按README说明即可训练自定义数据,操作门槛较低。资源共1031个文件,以986张png数据图像、18个Python脚本与2个pth权重文件为主体,另有说明文档与日志文本,压缩包约200.83MB。目前已有536人学习下载,适合需要完整跑通多脏器分割流程或深入实践TransUnet的开发者参考。
1. 基于 TransUnet 的腹部多脏器分割:一套能跑通的数据到模型方案
拿到一批腹部 CT,要一次性把肝脏、脾脏、肾脏、胰腺、胆囊、十二指肠都勾出来,这是腹部多脏器分割最典型的诉求。基于 TransUnet 做这件事,意味着保留 U-Net 的解码器骨架,把编码器换成“CNN + Transformer”混合结构,让模型既能抓住器官边缘的局部纹理,又能感知整个腹腔的解剖布局。很多人第一次看到 TransUnet 总觉得它笨重、显存消耗大,但换来的收益是:在肝脏与邻近组织密度接近、胰腺形态千差万别的场景里,分割结果比纯 U-Net 稳定不少。这篇笔记就围绕“代码、数据集、训练结果”三件事展开:先讲网络结构和选型理由,再把原始 CT 处理成模型能吃的格式,给出可改着跑的训练骨架,最后把踩过的坑按现象列清楚。适合正在做医学图像分割、想复现 TransUnet 又不想被论文空转耗掉时间的同学。
2. TransUnet 结构拆解与选型理由:为什么腹部多脏器分割要用 Transformer 编码器
2.1 多脏器分割的难点:边界模糊、器官形态差异,纯 CNN 的瓶颈在哪
腹部 CT 里的器官分割,难点不在“看得见”,而在“分得清”。肝脏和周围的肌肉、肠道在 CT 值上经常非常接近,边界由血管、筋膜和脂肪的细微信号决定;胰腺形态随扫描体位和呼吸运动变化很大,胰腺头部又紧贴十二指肠,两者在灰度上几乎融为一体。这种任务如果只靠 U-Net 这类卷积网络,会面临一个现实约束:卷积核的感受野是有限的,下采样到 1/16 或 1/32 之后,高层特征里的每个点虽然能覆盖大范围,但位置信息已经被压缩得很模糊,模型容易把“和肝脏纹理很像的胃肠道”也判成肝脏,或者把胰腺头的一部分切给十二指肠。
多脏器分割还有一个被不少人忽略的点:器官之间的相对位置关系是极强的先验。肝脏基本在右上腹,脾在左侧膈下,左肾和右肾被脊柱隔开,胃和胰腺在中腹部。这些空间约束对卷积网络来说是“学得慢”的信息,因为每个卷积核都只盯着局部。即使 U-Net 通过跳连接恢复了分辨率,编码器最高层特征里也缺少一个明确的“全局坐标参考系”。换句话说,模型知道这里有纹理像肝脏,但不确定这里是不是肝脏应该出现的位置。
所以腹部多脏器分割选型时,需要一种能同时处理细节和全局定位的架构。纯 CNN 派的分割模型(U-Net、DeepLabV3+)在单器官分割上成绩很好,但多器官同时分割时,器官间重叠区域和相似纹理带来的混淆往往要靠大量后处理才能压下去。TransUnet 的思路是用 Transformer 的自注意力机制,在 CNN 特征图上建立一个跨整个图像范围的依赖关系,让每一个位置都能“看到”所有其他位置,从而把解剖相对位置编码进网络。
2.2 TransUnet 的编码器:ResNet 特征图与 ViT 序列建模怎么合并
TransUnet 的编码器分成两条路径。第一条是 CNN 路径:输入图像先经过一个 ResNet 骨干网络,比如 ResNet-50,前几层产生不同分辨率的特征图;第二条是 Transformer 路径:取自 ResNet 较后阶段的特征图(通常是输入图像的 1/16 分辨率),把它按空间位置展开成一组 token,送入 Transformer Encoder。
这里和原始 ViT 有一个关键差别:ViT 是把原始图像划分成 16x16 的 patch,输入就是像素级 token;TransUnet 则是在 ResNet 已经学会的语义特征图上做 token 化。这样做有几个好处:ResNet 提供了一定的平移等变性和局部纹理提取能力,Transformer 不需要从零学习边缘检测这类底层算子;同时特征图尺寸比原图小很多,token 数量更可控,训练时的显存占用和收敛难度都会降低。
一个常见的配置是:ResNet-50 的 stage3(也就是 layer3)输出 2048 通道的特征图,空间尺寸是 H/16 × W/16。对这个特征图做 flatten,得到序列长度为 H/16 × W/16、每 token 维度 2048。为了送入 Transformer,一般先用一个线性投影把维度压缩到一个较小的 embedding 维度,比如 768 或 1024,再叠加可学习的位置编码(position embedding)。Transformer Encoder 的层数通常在 6 到 12 之间,注意力头数 8 或 12。我一般用 12 层、embedding 维度 768、头数 12,显存占用中等,效果贴近论文报告。
需要特别说明的是,位置编码对医学图像分割影响比自然图像更大。自然图像的物体大致居中,位置编码只是辅助;腹部 CT 中器官绝对位置非常稳定,位置编码几乎变成了“解剖坐标”。如果你自己实现 TransUnet,位置编码的类型(二维还是二维)和是否保留原始坐标信息,直接关系到胰腺、十二指肠这类小器官的召回率。常见实现里采用 2D 位置编码,因为它保留了行和列的独立性,比纯 1D 位置编码更符合 CT 图像的空间结构。
2.3 解码器与分割输出:上采样路径和跳跃连接的恢复细节
Transformer Encoder 输出的序列会重新 reshape 回与输入特征图相同的高和宽,然后进入解码器。解码器沿用 U-Net 的经典设计:逐级上采样,并和 ResNet 早期阶段产生的不同分辨率特征图做跳跃连接。每一级上采样通常先做双线性插值或者转置卷积,把分辨率翻倍,然后与编码器的特征图沿着通道维度拼接,再做两层 3x3 卷积和 ReLU 激活。这种设计的价值在于:Transformer 部分提供了语义和全局信息,而 ResNet 浅层特征图保留了高分辨率的边缘细节,两者拼接后,解码器能同时利用“这是什么器官”和“边界在哪里”两种信息。
从通道数配置上看,Transformer 输出的 embedding 维度通常比 CNN 特征通道更大,拼接之前需要做对齐。常见做法是在 Transformer 输出后接一个 1x1 卷积,把通道数降到与当前解码器层匹配的通道数,再和对应的 CNN 特征图拼接。例如第一级上采样前,先把 Transformer 输出从 768 降到 512,再和 ResNet 的 layer2 特征图(512 通道)拼接,得到 1024 通道,之后用卷积压缩到 512,继续往上采样。这个过程很像是给 U-Net 的编码器加了一个“全局重标定”模块,实际上就是 TransUnet 名称里“Trans”的那部分。
解码器最后一级输出通道数等于分割类别数。对于腹部多脏器分割,常见设定是背景 + 8 个器官,也就是 9 类;如果在自己的数据集上只标注了肝脏、脾脏、肾脏、胰腺,那输出通道就是 5 类。这里有一个容易被忽略的坑:多类别分割的最后一层卷积只输出 logits,需要配合 softmax(而不是 sigmoid)来归一化,因为一个体素只能属于一个器官类别。有些做惯了二分类的同学把最后一层换成 sigmoid,训练时 loss 也能降,但推理结果经常出现一个体素同时属于好几个器官的假阳性。
Transformer Encoder 与 U-Net 解码器之间的信息流,是整个网络性能的关键。很多复现结果差,问题并不在 Transformer 层,而在解码器上采样时通道对齐和拼接方式错了。比如有些简化实现直接把高维 Transformer 特征和低维 CNN 特征强行相加而不是拼接,导致浅层细节被淹没。如果你自己搭模型,务必保留拼接和 1x1 对齐,不要贪省事。
3. 准备腹部多脏器分割数据集:从 nii.gz 到模型输入的规范化流程
3.1 数据来源与器官类别:Synapse 数据集的常见设定
腹部多脏器分割最常用的公开数据集是 Synapse 多器官数据集,来自 MICCAI 2015 的腹部多器官分割挑战赛。里面是腹部 CT 的 NIfTI 文件(nii.gz),每个病例包含原始 CT 体积和对应的分割 label 体积,label 里的每个整数代表一个器官。不同版本对器官的定义略有差异,常见的有 8 个目标器官:肝脏、脾脏、肾脏、胰腺、胃、胆囊、食管、十二指肠。有些版本还包含主动脉和下腔静脉,所以拿到数据后第一件事不是写代码,而是检查 label 值到器官名称的映射,确认哪些类别要保留、哪些要合并。
这一点特别重要,因为网上流传的预处理脚本常常写死“9 类”或“8 类”,而你的数据可能是 10 类或只有 6 类。如果直接套用别人的代码,模型最后一层输出通道数不匹配,训练会立刻报错。我通常会把 label 的取值分布先打印出来,统计每个整数的体素数,再决定哪些类别用于训练。对于体素数很少的类别,比如只有几百个体素的胆囊,如果直接参与训练,基本会被模型忽略;要么增加该类的采样权重,要么干脆去掉,不要硬凑器官数量。
3.2 预处理四步:窗宽裁剪、归一化、重采样和切片
腹部 CT 的原始值域是 CT 值(HU),范围可以到 -1000 到 +3000 以上,但软组织器官主要集中在 -125 到 +275 之间。直接把这个范围的数据喂给网络,模型会把大量精力花在区分空气、骨骼和软组织上,分割的目标器官反而被压缩到很窄的灰度区间。所以第一步是窗宽裁剪:只保留 [-125, 275] 这个窗口内信息,窗口外值统一截断到边界。这一步能让肝脏和胰腺的纹理对比度明显增强。
第二步是归一化。裁剪后的数据通常线性变换到 [0, 1] 区间,避免数值波动干扰梯度。归一化必须是在裁剪之后做,而不是直接对整张 CT 图做 min-max,否则个别高亮骨骼会把软组织灰度压到接近 0,损失对比度。这个顺序不要反过来,这是预处理里最容易出错的细节。
第三步是重采样。CT 数据来自不同设备,X、Y、Z 方向的像素间距(spacing)可能不同。如果直接训练,同一器官在不同样本里的尺度特征全乱了。常见做法是把所有体积重采样到固定的目标 spacing,比如 1.0 × 1.0 × 1.0 mm 或者 0.8 × 0.8 × 1.5 mm。重采样使用三线性插值对图像本身,但分割 label 必须用最近邻插值,防止插值产生新的整数标签。我见过有人统一用scipy.ndimage.zoom处理 image 和 label,结果 label 边缘出现 2.5、3.7 这些非法值,训练时类别数直接爆掉。
第四步是把 3D 体积沿轴向切成一叠 2D 切片。TransUnet 的常见使用方式是把 3D 体积当成多个 2D 切片独立处理,因为一次性输入整个 3D 体积的显存开销太大,Transformer 的 token 数也会爆炸。切片之后,每一张切片就是训练样本,shape 为 (H, W, 1),通道数通常是 1(灰度),也有用三通道复制来适配 ImageNet 预训练的 ResNet。如果你打算加载 ImageNet 预训练权重,必须把单通道灰度重复成三通道,并且注意归一化方式要与预训练一致。
以下是预处理的核心代码,我按实战中跑通过的方式写:
import nibabel as nib import numpy as np from scipy import ndimage def preprocess_volume(nii_img_path, nii_label_path, target_spacing=(1.0, 1.0, 1.0), window=(-125, 275)): # 加载原始 CT 和 label img = nib.load(nii_img_path).get_fdata().astype(np.float32) label = nib.load(nii_label_path).get_fdata().astype(np.int16) # 1. 窗宽裁剪,把 HU 值限制到 [-125, 275] img_clipped = np.clip(img, window[0], window[1]) # 2. 线性归一化到 [0, 1] img_norm = (img_clipped - window[0]) / (window[1] - window[0]) # 3. 读取当前 spacing,并按比例计算缩放因子 spacing = nib.load(nii_img_path).header.get_zooms()[:3] # (x, y, z) zoom_factor = [cur / target for cur, target in zip(spacing, target_spacing)] # 图像用三线性插值,label 用最近邻插值 img_resample = ndimage.zoom(img_norm, zoom_factor, order=3) label_resample = ndimage.zoom(label, zoom_factor, order=0) # 4. 沿轴向(depth)切出 2D 切片 slices = [] labels = [] for d in range(img_resample.shape[2]): slices.append(img_resample[:, :, d]) labels.append(label_resample[:, :, d]) return np.array(slices), np.array(labels)这段代码的关键参数:target_spacing决定模型看到的器官绝对尺寸。对腹部多脏器分割,1mm 的等向性 spacing 是最稳妥的选择,器官解剖结构完整,但数据量会变大;如果显存受限,Z 轴 spacing 放宽到 1.5mm 也能接受。window值直接影响器官对比度,我曾试过更窄的窗口比如 [-50, 200],胰腺和肝脏的对比度更高,但肠管内气体被完全压黑,边界反而容易断;[-125, 275] 是绝大多数医学分割竞赛采用的默认值,先不用改。ndimage.zoom的order=3是三次样条插值,会产生范围溢出,由于前面已经做了归一化,溢出不会太严重,但保险做法是在缩放后再 clip 回 [0, 1]。
3.3 目录结构与训练验证划分:代码直接可用的组织方式
预处理完成后,你需要把数据整理成固定目录结构,并划分训练集和验证集。目录结构要能直接支持 PyTorch 的Dataset类读取,我一般这样组织:
data/ images/ case1_slice_0.npy case1_slice_1.npy ... labels/ case1_slice_0.npy case1_slice_1.npy train_list.txt val_list.txtimages和labels下是两个一一对应的 npy 文件,每个文件是一张 2D 切片。train_list.txt和val_list.txt记录训练、验证用的文件名前缀。划分时一定要按“病例”划分,而不是按“切片”划分。也就是同一个 case 的全部切片要么都在训练集,要么都在验证集,不能交叉。否则同一病人相邻切片信息高度重复,验证集 Dice 虚高,换了新病人立刻大幅下降。
下面这段脚本完成按病例划分和文件列表生成:
import os import numpy as np from glob import glob img_dir = "data/images" label_dir = "data/labels" output_dir = "data" # 获取所有 case 前缀(假设文件名是 case1_slice_0.npy) all_files = glob(os.path.join(img_dir, "*.npy")) prefixes = sorted(set([os.path.basename(f).rsplit("_slice", 1)[0] for f in all_files])) # 按 8:2 划分病例,而不是切片 num_val = max(1, int(len(prefixes) * 0.2)) val_cases = set(prefixes[-num_val:]) # 固定取后 20%,保证可复现 train_cases = [c for c in prefixes if c not in val_cases] train_list = [] val_list = [] for p in train_cases: for f in glob(os.path.join(img_dir, p + "_slice_*.npy")): train_list.append(os.path.basename(f).replace(".npy", "")) for p in val_cases: for f in glob(os.path.join(img_dir, p + "_slice_*.npy")): val_list.append(os.path.basename(f).replace(".npy", "")) with open(os.path.join(output_dir, "train_list.txt"), "w") as fout: fout.write("\n".join(train_list)) with open(os.path.join(output_dir, "val_list.txt"), "w") as fout: fout.write("\n".join(val_list)) print("train slices:", len(train_list), "val slices:", len(val_list))这里有一个隐藏参数容易被忽略:rsplit("_slice", 1)。如果文件名里本身包含“slice”字样,切割位置可能会错。更稳妥的做法是文件名格式固定为case001_12.npy,然后用rsplit("_", 1)去掉最后一段数字。我实际用的时候会把病例编号单独存成前缀,避免文件名解析的玄学问题。划分比例方面,腹部多脏器公开数据集很小,总共也只有几十到一百多个病例,验证集比例 20% 是合理的;如果你的数据更少,可以用五折交叉验证,而不是硬切一个验证集。
训练时读取 npy 文件比每次从 nii.gz 现场处理快得多,建议提前把所有训练切片保存为 npy 或者内存映射。如果数据集太大,也可以只在__getitem__里读取对应 npy,避免一次性全loaded。接下来进入训练环节。
4. 训练 TransUnet 的 PyTorch 骨架:模型定义、损失函数与关键参数
4.1 最小可用模型定义:把论文结构落成可训练的代码
完整从零手写 TransUnet 的代码量很大,实际项目里我用的方式是:用 PyTorch 搭一个“结构正确”的轻量版,重点保证 ResNet 特征、Transformer token 化、解码器拼接三段逻辑清楚,再根据显存调整宽度和深度。下面的代码是模型核心骨架,去掉了 ResNet 内部重复层,只保留关键流程:
import torch import torch.nn as nn from torchvision import models class TransUNet(nn.Module): def __init__(self, n_classes=9, embed_dim=768, depth=12, heads=12): super().__init__() # 使用 ResNet-50 作为 CNN 编码器,取 layer1/layer2/layer3 作为跳连接特征 resnet = models.resnet50(weights=models.ResNet50_Weights.IMAGENET1K_V2) self.conv1 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu) self.maxpool = resnet.maxpool self.layer1 = resnet.layer1 # 256通道,stride 4 self.layer2 = resnet.layer2 # 512通道,stride 8 self.layer3 = resnet.layer3 # 1024通道,stride 16 self.layer4 = resnet.layer4 # 2048通道,stride 32 # 把 layer4 输出投影到 embed_dim,得到 Transformer 输入 self.proj = nn.Conv2d(2048, embed_dim, kernel_size=1) # 位置编码:序列长度等于特征图 H*W,这里设为动态 self.pos_embed = nn.Parameter(torch.zeros(1, (256 // 16) * (256 // 16), embed_dim)) # Transformer Encoder encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=heads, dim_feedforward=4 * embed_dim, activation="gelu", dropout=0.1, batch_first=True) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=depth) # 解码器:从 1/32 分辨率上采样到 1/16,再拼接 layer4 输出 self.up1 = nn.ConvTranspose2d(embed_dim, 1024, kernel_size=2, stride=2) # 2x上采样 self.conv_up1 = self._conv_block(1024 + 1024, 1024) # 拼接 layer4 # 再上采样到 1/8,拼接 layer3 self.up2 = nn.ConvTranspose2d(1024, 512, kernel_size=2, stride=2) self.conv_up2 = self._conv_block(512 + 1024, 512) # 继续上采样到 1/4,拼接 layer2 self.up3 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2) self.conv_up3 = self._conv_block(256 + 512, 256) # 上采样到 1/2,拼接 layer1 self.up4 = nn.ConvTranspose2d(256, 64, kernel_size=2, stride=2) self.conv_up4 = self._conv_block(64 + 256, 64) # 上采样到原分辨率,输出 n_classes 通道 self.up5 = nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2) self.seg_head = nn.Conv2d(32, n_classes, kernel_size=1) def _conv_block(self, in_c, out_c): return nn.Sequential( nn.Conv2d(in_c, out_c, 3, padding=1), nn.BatchNorm2d(out_c), nn.ReLU(inplace=True), nn.Conv2d(out_c, out_c, 3, padding=1), nn.BatchNorm2d(out_c), nn.ReLU(inplace=True)) def forward(self, x): # 输入为 B,3,H,W,要求 H、W 能被 32 整除 f1 = self.conv1(x) # 1/2 f1 = self.maxpool(f1) # 1/4 f2 = self.layer1(f1) # 1/4, 256 f3 = self.layer2(f2) # 1/8, 512 f4 = self.layer3(f3) # 1/16, 1024 f5 = self.layer4(f4) # 1/32, 2048 # Transformer 输入 B, C, H, W = f5.shape tokens = self.proj(f5).flatten(2).transpose(1, 2) # (B, H*W, embed_dim) tokens = tokens + self.pos_embed[:, :H * W, :] tokens = self.transformer(tokens) # 把 token 还原成特征图 feat = tokens.transpose(1, 2).reshape(B, -1, H, W) # (B, embed_dim, H, W) # 解码器逐级上采样 d1 = self.up1(feat) # (B,1024,2H,2W) d1 = self.conv_up1(torch.cat([d1, f4], dim=1)) # f4 是 1/16 d2 = self.up2(d1) # (B,512,4H,4W) d2 = self.conv_up2(torch.cat([d2, f3], dim=1)) d3 = self.up3(d2) # (B,256,8H,8W) d3 = self.conv_up3(torch.cat([d3, f2], dim=1)) d4 = self.up4(d3) # (B,64,16H,16W) d4 = self.conv_up4(torch.cat([d4, f1], dim=1)) # f1 是 1/4 d5 = self.up5(d4) # (B,32,32H,32W) logits = self.seg_head(d5) # (B,n_classes,32H,32W) return logits这段代码里有一个需要特别留意的参数:pos_embed的尺寸我写死了(256 // 16) * (256 // 16),对应输入图像 256x256、特征图 16x16。如果你输入尺寸不是 256x256,运行时会因为 token 数量不匹配直接报错。常见做法是把pos_embed定义成更大的尺寸,比如 512×16,然后在 forward 里用self.pos_embed[:, :H * W, :]裁剪。但裁剪频谱可能会损失位置精度。更稳妥的做法是在数据加载阶段统一 resize 到 256x256,这样位置编码可以固定,Transformer 也不需要处理变长序列。腹部 CT 原图尺寸通常在 512x512 左右,resize 到 256x256 会丢失一些细小边界,但换来的是显存大幅下降和稳定的训练行为。如果你对边缘质量要求高,可以改用 384x384,并同步把pos_embed扩到对应尺寸。
解码器部分,我在代码里用nn.ConvTranspose2d做上采样,也可以用nn.Upsample(scale_factor=2, mode="bilinear")配合卷积做。转置卷积有可学习参数,表达能力更强,但容易在高频区域产生棋盘格伪影;Upsample没有参数,训练更稳。实际跑腹部多脏器分割,我一般用双线性上采样,棋盘格伪影对分割边界的干扰值得警惕。
4.2 损失函数和评估指标:Dice Loss 与 Hausdorff 距离怎么配
多脏器分割里每个器官的体素数差异很大:肝脏可能占几万个像素,胆囊可能只有几千个。直接用交叉熵损失,小器官基本被大器官淹没,模型只需要把肝脏分割好,整体 loss 就很好看。所以训练时必须引入基于区域的损失函数,最常见的是 Dice Loss 的变体。多分类任务里,一个典型组合是:
loss = 0.5 * CrossEntropyLoss(logits, label) + 0.5 * DiceLoss(softmax(logits), label)交叉熵提供像素级梯度,Dice Loss 提供类别平衡的全局梯度。Dice Loss 在二分类里非常直接,多分类实现上需要把每个类别单独计算 Dice 再取平均。以下是一个可供参考的多分类 Dice Loss 实现:
import torch import torch.nn.functional as F def multiclass_dice_loss(logits, labels, eps=1e-6): # logits: (B, C, H, W) # labels: (B, H, W) 且值为 0..C-1 probs = F.softmax(logits, dim=1) # 转成概率 n_classes = probs.shape[1] dice_sum = 0.0 for c in range(1, n_classes): # 跳过背景 0,通常不计算 pred = probs[:, c] # (B, H, W) true = (labels == c).float() intersection = (pred * true).sum(dim=(1, 2)) union = pred.sum(dim=(1, 2)) + true.sum(dim=(1, 2)) dice = (2.0 * intersection + eps) / (union + eps) dice_sum += (1.0 - dice.mean()) # loss = 1 - dice return dice_sum / (n_classes - 1)这段代码对每个类别独立计算 Dice,然后对非背景类别取平均。eps是平滑项,防止某器官在验证集中完全没出现时除以零;但这治标不治本,真正原因是数据读取时漏了某几个类别。如果发现 loss 是 NaN,第一步检查 label 里是否有n_classes之外的值,而不是调大eps。另一个容易被忽略的细节:背景类别不参与 Dice 计算,不然本来占比就高的背景会让 loss 虚低,模型对小器官的惩罚被稀释。
评估指标上,论文里习惯报告两个值:Dice 系数(DC)和 95% Hausdorff 距离(HD95)。Dice 衡量体积重叠比例,而 HD95 衡量边界最大偏差的全 95 百分位数。多脏器分割中,胰腺的 Dice 可能不错但 HD95 很夸张,因为胰腺尾部细长,一个小的误判点会把最大边界距离拉得极大。HD95 可以用 SimpleITK 计算,HausdorffDistanceImageFilter得到的是最大距离,要计算 95 百分位需要保存距离图后手动排序。OpenCV 或者medpy库有现成函数,但要注意不同实现之间的距离单位(体素还是毫米)不一致,评估时务必统一。
4.3 训练循环与参数表:学习率、batch size、epoch 和数据增强的取舍
下面是训练的主循环骨架,包含关键的梯度累积和验证逻辑。2D 切片训练时,一般不用把整个体积灌进模型,逐切片喂入即可。
from torch.utils.data import Dataset, DataLoader import torch.optim as optim class SliceDataset(Dataset): def __init__(self, img_dir, label_dir, filelist_path, augment=False): with open(filelist_path, "r") as f: self.samples = [line.strip() for line in f if line.strip()] self.img_dir = img_dir self.label_dir = label_dir self.augment = augment def __len__(self): return len(self.samples) def __getitem__(self, idx): name = self.samples[idx] img = np.load(f"{self.img_dir}/{name}.npy") # 形状 H,W label = np.load(f"{self.label_dir}/{name}.npy") # 转成三通道,适配 ResNet 预训练 img = np.stack([img] * 3, axis=0).astype(np.float32) label = label.astype(np.long) # 可在此处做随机裁剪、翻转等增强 return torch.from_numpy(img), torch.from_numpy(label) def train_one_epoch(model, loader, optimizer, criterion, device, grad_accum_steps=2): model.train() running_loss = 0.0 optimizer.zero_grad() for step, (img, label) in enumerate(loader): img = img.to(device) label = label.to(device) logits = model(img) loss = criterion(logits, label) loss = loss / grad_accum_steps # 累积梯度 loss.backward() if (step + 1) % grad_accum_steps == 0: optimizer.step() optimizer.zero_grad() running_loss += loss.item() * grad_accum_steps return running_loss / len(loader)这段代码里的grad_accum_steps是显存不足时的后悔药。当 batch size 设为 4 也会爆显存时,可以把 batch size 降到 1,然后用累积步数模拟更大的 batch。注意梯度累积时 loss 要除以累积步数,否则实际学习率被放大,训练容易震荡。我还习惯在 optimizer.step() 之后做一次torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=12),这能防止 Transformer 在初期突然出现梯度爆炸。
实际操作中,我推荐的参数表如下:
| 参数 | 推荐值 | 说明 |
|---|---|---|
| 输入尺寸 | 256 × 256 | 平衡显存与细节,适配固定位置编码 |
| batch size | 8(2D切片) | 根据显存调整,配合 grad_accum_steps |
| epoch 数 | 100 | 带早停,验证指标连续 15 个 epoch 不升就停 |
| 优化器 | AdamW | weight_decay 用 1e-5,比 SGD 更稳 |
| 初始学习率 | 1e-4 | Transformer 部分偏大容易训崩 |
| LR 调度 | cosine decay | 从头训练用 warmup 10 个 epoch |
| 数据增强 | 随机翻转、随机缩放、灰度扰动 | 不要用强剪裁,防止破坏解剖结构 |
| 类别不平衡 | 小器官权重乘 2 | 或者在采样时对含小器官的切片加权 |
关于数据增强,多脏器分割和自然图像分割的场景不同。腹部 CT 的解剖方向相对固定,左右翻转是可以的,上下翻转绝对不要用——否则肝脏跑到右下腹,模型会被位置编码彻底搞乱。随机缩放也要保守,缩放系数控制在 0.9 到 1.1 之间,缩多了会改变器官相对位置关系。灰度扰动可以做,但要限制幅度,CT 窗宽已经把灰度归一化到 [0,1],扰动超过 ±0.1 就可能让组织纹理失真。
Transformer 部分的初始化非常关键。如果你用了 ImageNet 预训练的 ResNet,但 Transformer Encoder 是随机初始化,训练一开始模型会在 CNN 路径上表现良好,Transformer 却输出混乱的信息。常见做法是给 Transformer 部分一个更小的学习率,比如让 Transformer 的lr = 1e-5,而 CNN 部分保持1e-4。可以用 PyTorch 的param_group按模块名设置不同学习率。如果不区分,头几个 epoch 的 loss 会反复横跳,这是 Transformer 随机初始化权重和预训练 CNN 特征“抢方向”造成的,不是你的代码有 bug。
5. TransUnet 训练避坑:五个让模型翻车的典型场景与排查过程
5.1 显存不足:三维体数据直接灌进模型是第一个翻车点
现象:代码写完后,训练脚本一跑,很快爆出CUDA out of memory,显存直接打满,训练中断。
原因:最常见是有人试图把 CT 三维体积(1, 1, D, H, W)直接喂给 TransUnet。即便网络内部只处理 2D 特征图,PyTorch 的卷积和 Transformer 对 5D 输入会隐式展平,token 数量瞬间变成D×H×W,注意力矩阵大小随 token 数平方增长,显存必然爆炸。即便切成 3D patch,TransUnet 的原始设计也是 2D 网络,强行 3D 化会引入大量参数和训练难度。
解决:坚持 2D 切片输入。把每个体积沿 Z 轴切为几十到上百张 2D 切片,固定输入尺寸 256 × 256。如果显存仍然不够,先把 batch size 降到 2 或 1,用梯度累积弥补。另外检查 ResNet 部分是否真的加载了预训练权重,加载预训练权重虽然不省显存,但能显著减少训练轮数,间接降低显存压力。如果你用的 Transformer 层数较多,比如 12 层以上,可以把dim_feedforward从 4 倍 embedding 维度改成 2 倍,显存立刻降很多,精度损失有限。
5.2 Dice 高但边界破损:窗宽与分辨率匹配问题
现象:训练结束后,验证集平均 Dice 有 0.85,但把预测结果叠加到原始 CT 上看,器官边缘参差不齐,有的地方像被啃了一口,尤其是肝脏和脾脏的包膜边界。
原因:分割网络输出一个软概率图,边界处的像素本来就容易混淆;但如果边缘大面积破损,先检查预处理。CT 窗宽裁剪范围如果太窄或太宽,会直接把器官边界的信息压平。另一种原因是输入分辨率过低,256 × 256 在 512 × 512 的原始 CT 上等于每个像素代表约 2mm,而肝脏包膜厚度不到 1mm,边界自然只能用粗粒度逼近。
解决:调整预处理窗口为[-125, 275],这是多数通用模型的折中。如果特定器官边界仍差,可以在推理阶段增加多尺度测试:把输入切片缩放为 0.75、1.0、1.25 倍,分别推理后对概率图加权平均,再取 argmax。多尺度能明显改善边界连续性,代价是推理时间翻几倍。如果只是肝脏边界差,也可以单独对肝脏类别做个后处理,提取连通域后用形态学闭运算填补凹陷,但要控制结构元素大小,避免把相邻器官粘连。
5.3 小器官分割接近失效:样本不平衡带来的验证集假象
现象:训练日志里每个 epoch 的 loss 在下降,验证集平均 Dice 到了 0.80,但单独看胆囊和十二指肠的 Dice 只有 0.10 甚至 0。整体指标被肝脏、脾脏这些大器官拉得很高。
原因:小器官在数据集里占的体素数太少。一个 512×512×200 的 CT 中,肝脏可能占 300 万体素,胆囊只有 3 万体素,差了 100 倍。在用平均 Dice Loss 时,每个类别的 Dice 对 loss 的贡献是均等的,理论上没问题;但训练过程中梯度由每个像素的误差累加,大器官的误差数量占优,小器官的类别不出现在多数 batch 中。如果 DataLoader 随机采样切片,很多 2D 切片里根本没有胆囊,模型在这些切片上的 loss 完全由肝脏和背景决定,梯度方向会把模型推向忽略小器官。
解决:有两个实际办法。第一,在采样阶段做类别偏好采样,确保每个 batch 里包含至少一张含有小器官的切片。实现时给每个样本分配权重,凡是 label 中出现胆囊或十二指肠的切片,采样概率乘 2.5。第二,对损失函数做类别加权,小器官的 Dice 在 loss 计算中乘一个加权系数,比如胆囊加权 2 倍、胰腺 1.5 倍。我还会把每个 epoch 的验证结果按器官分别打印,而不是只打印平均 Dice,这样能实时观察到小器官是否在改善,而不是等到整个训练结束才发现胆囊没学出来。
5.4 推理结果层间闪烁:2D 切片模型的伪影来源
现象:模型训练和验证指标都正常,但拿来预测一个完整 3D 体积时,逐层观看轴向切片发现前后两层的分割结果很不稳定,同一器官在相邻层里一会儿多一会儿少,形成“拉链”状伪影。三维重建后表面充满凹凸不平。
原因:TransUnet 是 2D 分割模型,逐层独立预测时,层与层之间没有任何上下文约束。CT 扫描的 Z 方向采样间距通常比 X/Y 方向稀疏,跨层解剖结构变化更大,模型在每一层上只能猜测器官在这个截面上的位置,相邻两层的猜测如果置信度差不多,就可能出现边界抖动。
解决:推荐两种方式。第一是做重叠切片推理:沿 Z 轴以步长 2 或 3 滑动,同一位置被多个邻近预测覆盖,最后对重叠区域取平均概率。这会消除大部分层间抖动,代价是推理时间线性增加。第二是对预测结果做三维中值滤波,以 Z 方向一个 3×3×3 的窗口对概率体素做平滑,再用平滑后的概率图取 argmax。这个方法不增加计算量,只会稍微模糊边界,但能大幅提升三维重建的平滑度。我习惯两者都用:重叠推理只用于测试阶段,训练时不用;中值滤波在最终提交前做一次,效果肉眼可见。
5.5 训练结果与论文差距过大:预处理和评估口径不一致
现象:按公开网络结构、公开数据集训练完,自己跑出来的平均 Dice 和论文报告的相差 0.05 以上。反复调参也追不上,怀疑代码有 bug。
原因:论文里的评估指标往往有隐藏前提。常见口径差异包括:只评估有标注的体素范围还是全图;是否去掉肝脏等大器官再计算平均;验证集是每个病例固定切片还是全部切片;是否排除了与训练集重合的病例。另一个可能是预处理差异:论文可能使用了特定方向的插值、固定裁剪区域(比如只保留腹部中央区域)、或者对每一例做了直方图匹配。这些细节论文里只会写一句“we preprocess all images to 256×256”,实际执行差异很大。
解决:遇到指标差距,先把预处理对齐到和论文一致。具体做法是检查公开的预处理脚本,对比窗宽范围、输入尺寸、是否对 mask 也做了同样的缩放、验证集划分随机种子。如果这些都一致,再检查模型加载预训练权重的方式。很多人加载 ResNet 预训练权重时只加载了 CNN 部分,却忽略了位置嵌入和 Transformer Encoder 的随机初始化,这会让模型需要多训练 50 个 epoch 才能达到论文水平。我通常会在训练完成后,用相同 test set 用不同随机种子跑三次,取平均作为最终结果,避免一次训练的运气成分被误判为模型差距。
6. 从训练结果到落地验证:可视化重叠图、按器官 Dice 与一组后处理习惯
6.1 推理脚本与可视化重叠图
训练结束后,第一个想看的不是指标数字,而是预测结果和真实标注叠加在原图上是什么样。下面这段代码读取一张验证切片,输出彩色分割覆盖图:
import torch import numpy as np import matplotlib.pyplot as plt model.eval() with torch.no_grad(): img_np = np.load("val_slice.npy") # H,W img_tensor = torch.from_numpy(np.stack([img_np] * 3, axis=0)).unsqueeze(0).to(device) logits = model(img_tensor) # 1, C, H, W pred = torch.argmax(logits, dim=1).squeeze(0).cpu().numpy() # H, W # 创建彩色标签图 color_map = { 0: [0, 0, 0], 1: [255, 0, 0], # 肝脏 2: [0, 255, 0], # 脾 3: [0, 0, 255], # 肾 4: [255, 255, 0], # 胰腺 5: [255, 0, 255], # 胆囊 6: [0, 255, 255], # 胃 } rgb = np.zeros((*pred.shape, 3), dtype=np.uint8) for idx, color in color_map.items(): rgb[pred == idx] = color img_gray = (img_np * 255).astype(np.uint8) img_gray = np.stack([img_gray] * 3, axis=2) alpha = 0.5 overlay = (alpha * rgb + (1 - alpha) * img_gray).astype(np.uint8) plt.imsave("overlay.png", overlay)这段代码里alpha控制原图和分割图的透明度,0.5 能同时看清器官边界和周围解剖。如果你要验证层间连续性,把同一个病例的所有切片推理结果按 Z 顺序堆叠,保存为三维 npy 后用Slicer或者ITK-SNAP查看三维重建,比一张张看图更直接。推理时记得把模型切换到eval()并关闭梯度,否则每个 batch 都会累积计算图,显存和速度都会不可控。
6.2 按器官评估与后处理习惯
验证阶段不能只看平均 Dice,要单独输出每个器官的指标。腹部多脏器数据集里,器官类别数量有限,用表格记录最直观。我一般会生成下面这样的表:
| 器官 | Dice | HD95/mm |
|---|---|---|
| 肝脏 | 0.961 | 4.3 |
| 脾脏 | 0.947 | 2.8 |
| 左肾 | 0.923 | 5.1 |
| 右肾 | 0.915 | 4.9 |
| 胰腺 | 0.861 | 8.7 |
| 胆囊 | 0.742 | 12.4 |
| 胃 | 0.893 | 6.2 |
| 十二指肠 | 0.812 | 11.0 |
这张表能直接告诉你模型在哪些器官上还有余量。如果胆囊或十二指肠 Dice 明显低于其他器官,优先怀疑样本不平衡而不是模型结构。后处理方面,我复现 TransUnet 时最终固定用三个动作:第一步是去掉小于500体素的孤立连通域,这个值按 CT 层厚调整,层厚 1mm 时 500 体素约为小指头大小,不会伤到真实小器官;第二步是对每个器官类别单独做一次条件膨胀,只在原 label 概率高于某个阈值(比如 0.1)的邻域内扩展,避免把邻近器官粘在一起;第三步是保存预测概率图而不是硬标签,因为后续要用重叠推理精细分割边界时,概率图能提供更多信息。
最后说一个让我印象深刻的教训。我第一次用 TransUnet 跑腹部多脏器分割,从头到尾盯着平均 Dice 从 0.70 涨到 0.85,以为模型没问题,直到三维可视化时才发现十二指肠区域全是噪声,原来我用的一个公开数据版本里十二指肠标签本身就很小,而且切片方向重采样后 label 几乎被最近邻插值破坏。后来我把预处理脚本里所有对 label 的插值都改成order=0,并单独过滤掉没有目标器官的切片,结果各项指标才稳定上涨。医学图像分割没有玄学,绝大多数翻车都出在数据读取和预处理这一环,模型架构反而是最不容易出问题的地方。希望这些做法能帮你把 TransUnet 真正跑起来,看到一份可靠的分割结果。
本文还有配套的精品资源,点击获取