ARTICLE DETAIL

资讯详情

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

基于TransUnet的腹部多脏器分割实战:从CT预处理到模型训练

基于TransUnet的腹部多脏器分割实战:从CT预处理到模型训练 简介面向医学图像分割与深度学习研究者这份实战资源基于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输入就是像素级 tokenTransUnet 则是在 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, order3) label_resample ndimage.zoom(label, zoom_factor, order0) # 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的order3是三次样条插值会产生范围溢出由于前面已经做了归一化溢出不会太严重但保险做法是在缩放后再 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_classes9, embed_dim768, depth12, heads12): super().__init__() # 使用 ResNet-50 作为 CNN 编码器取 layer1/layer2/layer3 作为跳连接特征 resnet models.resnet50(weightsmodels.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_size1) # 位置编码序列长度等于特征图 H*W这里设为动态 self.pos_embed nn.Parameter(torch.zeros(1, (256 // 16) * (256 // 16), embed_dim)) # Transformer Encoder encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadheads, dim_feedforward4 * embed_dim, activationgelu, dropout0.1, batch_firstTrue) self.transformer nn.TransformerEncoder(encoder_layer, num_layersdepth) # 解码器从 1/32 分辨率上采样到 1/16再拼接 layer4 输出 self.up1 nn.ConvTranspose2d(embed_dim, 1024, kernel_size2, stride2) # 2x上采样 self.conv_up1 self._conv_block(1024 1024, 1024) # 拼接 layer4 # 再上采样到 1/8拼接 layer3 self.up2 nn.ConvTranspose2d(1024, 512, kernel_size2, stride2) self.conv_up2 self._conv_block(512 1024, 512) # 继续上采样到 1/4拼接 layer2 self.up3 nn.ConvTranspose2d(512, 256, kernel_size2, stride2) self.conv_up3 self._conv_block(256 512, 256) # 上采样到 1/2拼接 layer1 self.up4 nn.ConvTranspose2d(256, 64, kernel_size2, stride2) self.conv_up4 self._conv_block(64 256, 64) # 上采样到原分辨率输出 n_classes 通道 self.up5 nn.ConvTranspose2d(64, 32, kernel_size2, stride2) self.seg_head nn.Conv2d(32, n_classes, kernel_size1) def _conv_block(self, in_c, out_c): return nn.Sequential( nn.Conv2d(in_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), nn.Conv2d(out_c, out_c, 3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue)) 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], dim1)) # f4 是 1/16 d2 self.up2(d1) # (B,512,4H,4W) d2 self.conv_up2(torch.cat([d2, f3], dim1)) d3 self.up3(d2) # (B,256,8H,8W) d3 self.conv_up3(torch.cat([d3, f2], dim1)) d4 self.up4(d3) # (B,64,16H,16W) d4 self.conv_up4(torch.cat([d4, f1], dim1)) # 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_factor2, modebilinear)配合卷积做。转置卷积有可学习参数表达能力更强但容易在高频区域产生棋盘格伪影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, eps1e-6): # logits: (B, C, H, W) # labels: (B, H, W) 且值为 0..C-1 probs F.softmax(logits, dim1) # 转成概率 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, augmentFalse): 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, axis0).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_steps2): 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_norm12)这能防止 Transformer 在初期突然出现梯度爆炸。实际操作中我推荐的参数表如下参数推荐值说明输入尺寸256 × 256平衡显存与细节适配固定位置编码batch size82D切片根据显存调整配合 grad_accum_stepsepoch 数100带早停验证指标连续 15 个 epoch 不升就停优化器AdamWweight_decay 用 1e-5比 SGD 更稳初始学习率1e-4Transformer 部分偏大容易训崩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 patchTransUnet 的原始设计也是 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, axis0)).unsqueeze(0).to(device) logits model(img_tensor) # 1, C, H, W pred torch.argmax(logits, dim1).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), dtypenp.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, axis2) 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要单独输出每个器官的指标。腹部多脏器数据集里器官类别数量有限用表格记录最直观。我一般会生成下面这样的表器官DiceHD95/mm肝脏0.9614.3脾脏0.9472.8左肾0.9235.1右肾0.9154.9胰腺0.8618.7胆囊0.74212.4胃0.8936.2十二指肠0.81211.0这张表能直接告诉你模型在哪些器官上还有余量。如果胆囊或十二指肠 Dice 明显低于其他器官优先怀疑样本不平衡而不是模型结构。后处理方面我复现 TransUnet 时最终固定用三个动作第一步是去掉小于500体素的孤立连通域这个值按 CT 层厚调整层厚 1mm 时 500 体素约为小指头大小不会伤到真实小器官第二步是对每个器官类别单独做一次条件膨胀只在原 label 概率高于某个阈值比如 0.1的邻域内扩展避免把邻近器官粘在一起第三步是保存预测概率图而不是硬标签因为后续要用重叠推理精细分割边界时概率图能提供更多信息。最后说一个让我印象深刻的教训。我第一次用 TransUnet 跑腹部多脏器分割从头到尾盯着平均 Dice 从 0.70 涨到 0.85以为模型没问题直到三维可视化时才发现十二指肠区域全是噪声原来我用的一个公开数据版本里十二指肠标签本身就很小而且切片方向重采样后 label 几乎被最近邻插值破坏。后来我把预处理脚本里所有对 label 的插值都改成order0并单独过滤掉没有目标器官的切片结果各项指标才稳定上涨。医学图像分割没有玄学绝大多数翻车都出在数据读取和预处理这一环模型架构反而是最不容易出问题的地方。希望这些做法能帮你把 TransUnet 真正跑起来看到一份可靠的分割结果。本文还有配套的精品资源点击获取
返回列表