ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

U-Net实战复现:CT影像肿瘤分割的完整流程与经验总结

U-Net实战复现:CT影像肿瘤分割的完整流程与经验总结 我第一次把U-Net复现代码跑在CT影像的肿瘤分割任务上时心里其实很没底。CT体数据一进来就是 (512 \times 512 \times) 好几百层而肿瘤往往只占其中零星几层的一小块区域正负样本比例悬殊到离谱。更麻烦的是医学影像分割不像自然图像分类那样直接套个ResNet就能收工它要求输出和原始图像同尺寸的像素级标注这个任务特性几乎决定了U-Net这类编码器-解码器结构会成为默认答案。这篇文章是一份完整的U-Net编程实战记录以CT影像肿瘤分割为例从数据预处理、网络搭建、训练调参到评估后处理把每一步的取舍、踩过的坑以及实测得到的经验都放出来。如果你刚入门医学影像分割或者手里正压着一个分割任务但不知道从哪下手这篇应该能让你少走不少弯路。1. 复现U-Net之前三个必须想清楚的选择很多教程第一行代码就是class UNet(nn.Module)但我的建议是先把三件事想清楚U-Net为什么适合这个任务、论文结构怎么映射成代码、用什么技术栈。这三件事决定你后面写代码的心态。1.1 医学影像分割的痛点恰好是U-Net的强项CT影像和自然图像最大的区别在于它是单通道灰度图没有自然图像那么丰富的颜色纹理肿瘤与周围正常组织的对比度经常很低边界是模糊的而且正样本区域往往只占整个体积的千分之几。这种数据喂给一个普通分类网络网络很容易“学懒”——把所有像素预测为背景损失函数的数值也不会太难看。U-Net之所以在医学分割里如此普遍是因为它的结构正好针对这些问题设计。编码器一路下采样把感受野不断变大让网络能“看到”肿瘤附近的上下文信息解码器再把低分辨率特征图逐步恢复成原图大小让输出能够精细到像素级。中间那些跳跃连接则像一根根“直通线”把浅层的细节纹理直接送到解码器对应层弥补池化过程丢失的空间信息。我自己的理解是U-Net有点像临摹一幅画。先眯着眼睛把整体轮廓把握住然后再睁开眼去补细节跳跃连接保证了补细节时还“记得”原始草图长什么样。所以它天然适合那种既需要全局判断、又不能丢失局部边界的任务。1.2 把论文结构翻译成程序结构原始U-Net论文的结构并不复杂左边编码器4次下采样每次2倍通道数从64逐渐翻到1024底部是最抽象的特征图右侧解码器4次上采样每次和对应编码器输出做拼接最后用1×1卷积输出像素级类别预测。把这个结构翻译成代码最好的办法是模块化。我习惯拆成五个部件DoubleConv连续两次卷积BatchNormReLUU-Net的基本积木Down最大池化下采样再接一个DoubleConvUp转置卷积上采样然后拼接跳跃连接的特征图再接DoubleConvOutConv1×1卷积把通道数变成目标类别数UNet主类按论文顺序把这些积木拼起来。这种做法的好处是代码结构和论文里那张U型图几乎一一对应出了问题能很快定位。而且后面你想改通道数、改下采样次数只需要改动主类里的参数不用推翻重写。1.3 技术栈选型与任务定义我用PyTorch原因很实际动态图调试太方便了打印中间feature map、断点查看张量维度对复现阶段非常重要医学影像生态也很完善SimpleITK做体数据处理、MONAI做3D增强、medpy算指标基本开源工具链是现成的。TensorFlow当然也能做但我个人觉得在快速迭代和社区资源上PyTorch对新手更友好。环境版本方面Python 3.9、PyTorch 2.x、CUDA 11.8以上基本够用。安装完成后一定要先验证一下GPU是否真的可用直接命令行跑python -c import torch; print(torch.cuda.is_available()); print(torch.cuda.get_device_name(0))如果输出False后面所有训练代码都会白跑。还有一个容易被忽略的问题任务具体是二分类还是多分类肿瘤分割大多数时候是二分类也就是每个像素只有“肿瘤/非肿瘤”两种可能。这时候网络输出一个通道就够了设计简单训练也稳定。如果数据集里还有器官、血管等多个结构要一起分割那再考虑多通道输出。不要一开始就把问题复杂化。2. CT影像预处理训练能不能收敛七成看这里我见过太多人一上来就调网络结构结果Dice一直是0.1最后发现是预处理出了问题。CT影像不是普通图像它有自己的物理单位和坐标体系这一章讲的每一步都会直接影响你后面模型的收敛速度和最终精度。2.1 读入NIfTI/DICOM后的第一件事搞清体素坐标CT数据在高精度标注任务里经常以NIfTI格式存在一个文件就是一个完整的三维体积。用SimpleITK读取非常方便import SimpleITK as sitk import numpy as np img sitk.ReadImage(case_001.nii.gz) label sitk.ReadImage(case_001_seg.nii.gz) img_arr sitk.GetArrayFromImage(img) # 形状: (z, y, x) label_arr sitk.GetArrayFromImage(label) spacing img.GetSpacing() # 返回: (x方向间距, y方向间距, z方向间距) size img.GetSize() # 返回: (x, y, z)这里最坑的就是维度顺序。GetSpacing()返回的是(x, y, z)而GetArrayFromImage()得到的NumPy数组形状是(z, y, x)两者顺序正好相反。我第一次写重采样代码时没注意这个结果坐标系全错位训练出来的人工痕迹一眼假。我的习惯是读入任何数据后先打印img_arr.shape和img.GetSpacing()人工确认一下这个体积大概是多少毫米的物理范围。比如shape是(200, 512, 512)spacing是(0.6, 0.6, 2.5)那说明每一层是512×512层厚2.5mm共200层。这个确认过程只要十秒钟但能避免后面无数奇怪问题。2.2 HU值截断、窗宽窗位与归一化CT图像的像素值不是简单的灰度而是亨氏单位HU。空气约-1000水是0软组织在-100到300之间骨骼可以到上千。把原始HU值直接喂给神经网络会遇到两个问题一是数值范围太大不同组织之间差异被拉伸得很怪二是不同CT设备的扫描参数不同分布会有偏移。常规做法是先做“窗宽窗位”式的截断。对肿瘤分割任务我常用的范围是[-200, 250]或[-1024, 200]把空气、骨头这些无关组织的值压掉让网络把注意力集中在软组织上。然后做线性归一化def preprocess_volume(volume, lower-200, upper250): volume np.clip(volume, lower, upper) volume (volume - lower) / (upper - lower) return volume.astype(np.float32)这样所有输入都落在[0, 1]区间对网络训练非常友好。但有一个关键细节归一化参数lower和upper必须在所有样本上统一训练集和测试集用同一个公式千万不要对每个样本各自计算min/max再归一化。否则同一组织在不同样本里的数值含义都不一样模型学到的特征就乱了。2.3 重采样与切片策略不同医院、不同设备的CT扫描层厚和像素间距可能完全不同。有的层厚1mm有的2.5mm甚至5mm。如果不做处理同一个物理尺寸的肿瘤在不同病例里占据的voxel数量差异很大网络很难学到稳定的形态特征。解决方法是把体数据重采样到各向同性的目标间距比如1.0×1.0×1.0mm。图像用线性插值标注用最近邻插值因为标注是离散的类别标签最近邻才不会产生新的中间值。def resample_to_spacing(img_sitk, target_spacing(1.0, 1.0, 1.0), is_labelFalse): original_spacing img_sitk.GetSpacing() original_size img_sitk.GetSize() target_size [ int(round(original_size[i] * original_spacing[i] / target_spacing[i])) for i in range(3) ] resampler sitk.ResampleImageFilter() resampler.SetSize(target_size) resampler.SetOutputSpacing(target_spacing) resampler.SetOutputOrigin(img_sitk.GetOrigin()) resampler.SetOutputDirection(img_sitk.GetDirection()) if is_label: resampler.SetInterpolator(sitk.sitkNearestNeighbor) else: resampler.SetInterpolator(sitk.sitkLinear) return resampler.Execute(img_sitk)重采样之后数据仍然是三维体数据。接下来要决定用2D U-Net还是3D U-Net训练。我的建议是初学或数据量不大的情况下先用2D方案跑通等整个流程熟练了、确认问题出在“层间连续性”上再上3D。两者取舍如下方案优点缺点2D U-Net显存占用低、训练快、代码简单忽略层间关系切片之间可能预测不一致3D U-Net利用三维上下文分割更连贯显存占用大需要裁剪patch训练慢2.5D多平面融合兼顾效率和性能工程实现复杂推理时间变长2D方案的做法是把体数据沿轴向切成一张张2D切片逐层输入网络训练推理时再逐层预测并堆叠回3D。层厚太厚时相邻切片之间解剖结构变化很大2D模型在那些区域的预测会比较抖这时候后期3D后处理会显得更重要。2.4 数据划分与标注检查最容易翻车但没人认真讲我是吃过亏的。第一次做训练集划分时想多凑样本按“切片”随机分结果同一个病人的几百张切片同时出现在训练集和测试集里。验证集Dice很高但一换新病例立刻崩掉。这就是典型的数据泄漏。正确做法是按病例patient/case划分要么一个病例的所有切片全进训练集要么全进测试集绝不能混。另外标注文件本身也值得花时间检查。很多公开数据集的标注是整数但数值的含义各不相同。有的标注里1代表器官、2代表肿瘤有的直接0/1。我之前遇到过把“整个肝脏”当成“肝肿瘤”正样本去训练的情况模型学得倒挺快最后分割出来的其实是整个器官。所以一定要先统计一下标注的取值分布unique_values np.unique(label_arr) print(unique_values)再把标注叠加到CT原图上人工看几层确认你要分割的目标确实是标注里的哪些像素。3. U-Net核心代码拆解从DoubleConv到跳跃连接的完整实现说句实在话U-Net的代码在网上随便一搜一大把但很多版本要么缺了尺寸对齐的细节要么把模块耦合在一起改动起来很痛苦。我这里给出一个我自己用着很顺手的2D版本然后逐段解释设计思路。3.1 DoubleConv整个网络的基本积木U-Net里最频繁出现的结构就是连续两次卷积。原始论文用3×3卷积padding1保证特征图尺寸不变激活函数用ReLU。我加了BatchNorm理由很实际不同CT设备的灰度分布差异大BatchNorm能让中间特征分布更稳定训练会明显更顺。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.double_conv nn.Sequential( nn.Conv2d(in_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), nn.Conv2d(out_channels, out_channels, kernel_size3, padding1), nn.BatchNorm2d(out_channels), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.double_conv(x)为什么每次下采样前后都要保持空间尺寸不变因为U-Net整体上要做的是“同尺寸输入、同尺寸输出”的密集预测只要每一层卷积都不改变H和W最后输出自然能和输入对齐省去很多插值麻烦。3.2 Down与Up下采样和上采样的正确姿势下采样用最大池化kernel_size2, stride2把特征图长宽各减半同时通道翻倍。这样网络在更深层能用更大的感受野去判断上下文同时用更多通道去表达抽象语义。上采样我选择ConvTranspose2d也叫转置卷积stride2恰好把尺寸放大一倍。这里有一个初学者很容易踩的坑当输入图像的H和W是奇数或者经过多次下采样后出现尺寸取整误差时上采样后的尺寸和对应编码器输出的尺寸会对不上直接torch.cat就会报错。解决办法是在跳跃连接拼接前做一个尺寸对齐class Up(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.up nn.ConvTranspose2d(in_channels, out_channels, kernel_size2, stride2) self.conv DoubleConv(in_channels, out_channels) def forward(self, x1, x2): x1 self.up(x1) diffY x2.size()[2] - x1.size()[2] diffX x2.size()[3] - x1.size()[3] x1 nn.functional.pad( x1, [diffX // 2, diffX - diffX // 2, diffY // 2, diffY - diffY // 2], ) x torch.cat([x2, x1], dim1) return self.conv(x)这里的x1是解码器上采样得到的特征图x2是编码器对应层的输出。diffY和diffX表示两张图的尺寸差通过pad把较小的图补到和x2一致。这个处理非常实用因为实际数据里几乎不可能保证所有输入图片都是完美的2的整数次幂。很多网上教程不会提这一步但真实场景里你迟早会撞上。3.3 UNet主类把积木按论文顺序拼起来下面是我使用的完整U-Net主类通道数配置是论文经典的64→128→256→512→1024class UNet(nn.Module): def __init__(self, n_channels1, n_classes1, features(64, 128, 256, 512, 1024)): super().__init__() self.inc DoubleConv(n_channels, features[0]) self.down1 Down(features[0], features[1]) self.down2 Down(features[1], features[2]) self.down3 Down(features[2], features[3]) self.down4 Down(features[3], features[4]) self.up1 Up(features[4], features[3]) self.up2 Up(features[3], features[2]) self.up3 Up(features[2], features[1]) self.up4 Up(features[1], features[0]) self.outc nn.Conv2d(features[0], n_classes, kernel_size1) def forward(self, x): x1 self.inc(x) x2 self.down1(x1) x3 self.down2(x2) x4 self.down3(x3) x5 self.down4(x4) x self.up1(x5, x4) x self.up2(x, x3) x self.up3(x, x2) x self.up4(x, x1) logits self.outc(x) return logitsn_channels设为1因为CT是单通道灰度图。n_classes设为1因为肿瘤分割是二分类输出单个通道的logits网络后面接sigmoid得到每个像素属于肿瘤的概率。1×1卷积在这里起到“降维到类别数”的作用不改变空间分辨率。这个版本是2D的。如果要做3D核心思路就是把nn.Conv2d换成nn.Conv3d、nn.BatchNorm2d换成nn.BatchNorm3d、nn.MaxPool2d换成nn.MaxPool3d、nn.ConvTranspose2d换成nn.ConvTranspose3d。输入从(B, 1, H, W)变成(B, 1, D, H, W)。显存不够的话把features从(64, 128, 256, 512, 1024)改成(32, 64, 128, 256, 512)或者用(16, 32, 64, 128, 256)效果也不会太差。我自己跑的时候习惯先写一个快速冒烟测试确认模型前向传播没问题model UNet(n_channels1, n_classes1) x torch.randn(2, 1, 256, 256) logits model(x) print(logits.shape)如果输出形状是torch.Size([2, 1, 256, 256])说明网络定义基本正确可以开始写训练脚本了。4. 训练阶段的关键取舍损失函数、优化器与数据增强的实测经验网络搭好只是万里长征第一步。真正决定Dice能涨到多少的是训练阶段的这些细节选择。这一章全是实操总结每一步都是我被数据折磨过后换来的经验。4.1 Dice Loss与混合损失类别不平衡才是最大的敌人肿瘤分割最大的问题是正样本占比极低。如果直接用普通的交叉熵损失模型只要把所有像素都预测为背景损失就已经很小了。这时候就需要Dice Loss登场。Dice系数本身就是分割任务常用的评估指标公式很直观[ Dice \frac{2 \times |P \cap T|}{|P| |T|} ]其中P是预测的前景像素集合T是真实前景像素集合。它天然关注“预测区域和真实区域的重合度”即使前景占比极小也不会让网络轻易滑向全背景预测。把它变成损失函数就是1 - Dice。下面是我常用的实现class DiceLoss(nn.Module): def __init__(self, smooth1e-5): super().__init__() self.smooth smooth def forward(self, logits, targets): probs torch.sigmoid(logits) probs probs.contiguous().view(probs.size(0), -1) targets targets.contiguous().view(targets.size(0), -1) intersection (probs * targets).sum(dim1) dice (2.0 * intersection self.smooth) / ( probs.sum(dim1) targets.sum(dim1) self.smooth ) return 1.0 - dice.mean()smooth这个平滑项很关键防止某个batch里预测和标注都是0导致分母为0的尴尬情况。不过只靠Dice Loss也有问题训练初期梯度可能不太稳尤其当预测和真实区域完全不相交时。我的实测做法是使用BCE和Dice的混合损失bce nn.BCEWithLogitsLoss() logits, targets ... loss bce(logits, targets) dice_loss(logits, targets)两个损失加起来BCE提供稳定的梯度信号Dice把优化方向往区域重叠度上拉。这个组合在多个医学分割任务上都很稳如果你只想用一个方案跑通我推荐先用这个。4.2 优化器、学习率与epoch设置优化器我常用AdamWweight_decay设为1e-4比裸Adam稳定不容易过拟合。也有人坚持用带动量的SGD理由是泛化更好但训练速度慢、调参更敏感。我的建议是先用AdamW跑通流程拿到一个还不错的baseline再考虑换成SGD做最终微调。学习率从1e-4到1e-3都可以尝试显存紧张导致batch size变小的时候初始学习率也相应调低一点。训练过程中用ReduceLROnPlateau动态调整scheduler torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, modemax, patience10, factor0.5, verboseTrue )按照验证集Dice作为监控指标连续10个epoch不涨就把学习率减半。我自己跑的CT肿瘤分割任务一般50到150个epoch之间能收敛到稳定水平。关键是要在每次epoch结束后保存验证集Dice最好的checkpoint别让中途最好的一版模型被后面的震荡覆盖掉。显存不够怎么办优先缩输入patch尺寸比如从512×512缩到256×256而不是一味缩减batch size。batch size太小的话BatchNorm统计不稳定训练会显得很“飘”。如果一张完整CT切片实在太占显存也可以在训练时只裁取包含肿瘤区域附近的patch这比直接缩小全图分辨率更能保留细节。4.3 数据增强的边界不要增强到病变消失医学影像数据量普遍不大数据增强几乎是必须的。我常用的增强包括随机水平翻转、小角度旋转±15度以内、随机缩放0.9~1.1倍、随机平移以及弹性形变。实现的时候无论用albumentations还是MONAI最重要的一点是图像和标注必须做同一个变换。这里有一个医学影像特有的注意点上下翻转别乱用。头部CT可以做上下反转吗最好不要。腹部CT如果上下翻转解剖结构完全反了肝脏跑到上面去了这种增强对模型学习没有帮助反而会带偏。弹性形变虽然好用但幅度要小心肿瘤本身可能就几个毫米宽弹性增强太猛会把病灶扭曲得面目全非甚至让标注区域和实际结构对不上。我见过一个反面案例增强时把图像做了随机旋转90度甚至180度标注也跟着转看起来没问题但因为CT扫描的解剖方向是有固定含义的模型被逼着学习“旋转不变性”这对一个医学影像任务来说既不必要也浪费模型容量。分割任务的增强应该尽量保守模拟真实扫描中可能出现的轻微位移和形变就够了不需要搞出花来。5. 评估与后处理Dice涨上去之后真正决定落地的细节训练结束之后最忌直接拿测试集跑一个Dice数字就发“成功”。医学影像分割是要给人看、甚至影响后续诊疗决策的评估和后处理环节必须做得足够扎实。5.1 多指标一起看Dice、IoU、HD95和ASSDDice当然重要但它不是全部。我习惯同时算四个指标指标全称关注点好坏方向DiceDice Similarity Coefficient预测区域与真实区域重叠度越高越好IoUIntersection over Union交并比对假阳性更敏感越高越好HD9595% Hausdorff Distance预测边界与真实边界的最大偏差越低越好ASSDAverage Symmetric Surface Distance两表面平均距离越低越好Dice高但HD95也高的情况很常见模型把肿瘤内部覆盖得很好但边缘多出或少了几毫米导致边界距离很大。这对某些应用场景是致命的。我之前做过一个实验两个模型Dice都在0.85左右但一个HD95是4.2mm另一个是9.7mm后者在医生看来就是明显的“切不干净”。代码上可以直接用medpy这个库from medpy.metric.binary import dc, hd95, assd dice dc(pred, reference) hd hd95(pred, reference) mean_surface_distance assd(pred, reference)注意medpy里的输入要求是二值化的NumPy数组而且三维情况下会自动处理连通区域。不要把网络输出的概率直接传进去一定要先(pred 0.5)二值化。5.2 3D后处理连通域过滤与空洞填补2D U-Net逐层推理堆叠回3D后最明显的问题是层间不连续。同一个肿瘤在相邻切片上可能预测得断断续续还会出现很多“孤立碎块”假阳性。这时候一个简单但有效的后处理是按3D连通域过滤from scipy.ndimage import label def filter_small_components(mask_3d, min_voxel100): labeled, num_features label(mask_3d) if num_features 0: return mask_3d # 统计每个连通域的体素数 sizes np.bincount(labeled.ravel()) # 保留最大连通域或滤除小于阈值的连通域 keep_ids [i for i in range(1, num_features 1) if sizes[i] min_voxel] valid np.isin(labeled, keep_ids) return valid.astype(np.uint8)如果确定目标只有一个肿瘤直接保留最大连通域也是一种常见策略。应用这些后处理时一定要谨慎如果数据里存在多个肿瘤多病灶简单保留最大连通域会漏掉小病灶这时用体积阈值过滤更适合。5.3 可视化不要只看数字把预测叠加到原图上数字指标再好我也建议把预测结果叠加到原始CT切片上肉眼检查。这一步能帮你发现指标看不出的问题比如预测区域是否跑到了血管或骨骼上、边界是否明显与解剖结构不符。一个简单的可视化脚本import matplotlib.pyplot as plt def show_slice(volume, mask, slice_idx, predNone): plt.figure(figsize(12, 4)) plt.subplot(1, 3, 1) plt.imshow(volume[slice_idx], cmapgray) plt.title(CT) plt.subplot(1, 3, 2) plt.imshow(mask[slice_idx], cmapReds) plt.title(Ground Truth) plt.subplot(1, 3, 3) plt.imshow(volume[slice_idx], cmapgray) plt.imshow(mask[slice_idx], cmapReds, alpha0.4) if pred is not None: plt.imshow(pred[slice_idx], cmapBlues, alpha0.4) plt.title(Overlay) plt.show()建议从预测的Dice分布里分别挑一个最高、一个中等、一个最低的病例来看。最高分证明流程没问题最低分告诉你模型的边界在哪里。5.4 常见失败案例与补救思路从我自己的实验来看肿瘤分割最容易翻车的情况有三类。第一类是假阳性跑到高亮组织上。CT图像里部分软组织、血管增强区域和肿瘤的HU值接近模型容易“误伤”。这种问题靠调网络结构很难根治比较有用的手段是增加这类负样本的训练比重或者在损失函数里对假阳性加上额外的惩罚权重。第二类是边界预测过大。U-Net有一个天然倾向在梯度平缓、边界不清晰的区域它会倾向于把周围一圈背景也包进去。这时可以考虑后处理环节加一层条件随机场CRF来平滑边界或者用形态学腐蚀操作缩小一点轮廓。不过加CRF会增加推理时间效果也因任务而异建议先试后处理再决定。第三类是小肿瘤漏检。深层特征图经过多次下采样小目标的空间信息可能已经所剩无几。解决办法包括用更大的patch让网络看到更多上下文、采用TTA测试时增强在多个旋转和翻转方向分别预测再取平均或者干脆换用3D U-Net让模型利用层间信息。TTA是我个人比较喜欢的手段因为它不改变训练过程只是增加推理时的多次预测并平均在不少任务上能稳定提升1到3个点的Dice。我现在的复现流程基本固定为先跑2D U-Net拿到baseline再根据失败案例决定是做后处理、TTA还是升级到3D模型。这套路径适合绝大多数医学影像分割的起步场景也是我认为效率最高的方式。如果手里有GPU和一份带标注的CT数据建议直接按这篇文章的顺序跑一遍实际踩一遍坑比看十篇教程都管用。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进