U-Net实战:基于PyTorch的CT肿瘤分割完整流程与避坑指南
提起U-Net搞医学影像的几乎没有不知道的。CT影像里的肿瘤分割不管是肝脏、胰腺、肾脏还是肺结节很多开源方案和论文的baseline都离不开这个网络结构。我第一次完整复现U-Net是在一个肝脏肿瘤数据集上当时只是为了验证模型能不能跑通没想到数据预处理、调参、后处理加起来折腾了整整两周。这篇就把整个过程踩过的坑都翻出来按从数据准备、模型搭建、训练调参到推理后处理的顺序写一遍。如果你正打算做U-Net编程实战或者准备用CT影像做肿瘤分割可以参考这套完整流程要是已经跑过类似项目可以直接跳到后面的问题排查部分也许能帮你省几天时间。1. U-Net为什么是CT肿瘤分割的首选网络1.1 CT影像分割到底难在哪很多人第一次看到医学影像分割任务时会下意识觉得“这不就是语义分割吗”但上手之后才会发现完全不是一回事。CT影像和自然图像有本质区别它本质是单通道的灰度图反映的是人体组织对X射线的衰减系数像素值不是“亮度”而是有物理意义的CT值单位是亨斯菲尔德单位HU。空气约-1000HU水是0HU骨骼可以到1000HU以上肿瘤和正常组织的差值往往只有几十HU。这个特性带来两个直接问题。第一CT图像的颜色信息几乎没有只有灰度分布模型没法靠纹理颜色绕开几何边界必须理解组织本身的形态。第二肿瘤在整张图里占比通常非常小肝脏肿瘤可能只占整幅切片的1%到3%背景像素远远多于前景像素模型稍微偷懒一点输出全背景损失函数也不会太难堪。再加上肿瘤边界模糊、形态不规则、不同病例之间差异很大这些因素叠加在一起如果网络设计或者损失函数选得不好训练出来的模型很容易陷入局部最优。1.2 U-Net结构如何对症下药U-Net的核心设计非常直接左边是编码器通过卷积和池化不断缩小特征图尺寸提取越来越抽象的语义信息右边是解码器通过上采样逐步恢复分辨率最后输出和输入相同大小的分割概率图。关键动作是中间那条“U”形路径上的跳跃连接把编码器每一层的细节特征直接拼到解码器对应层上。这个设计刚好命中CT肿瘤分割的痛点。肿瘤边界依赖浅层特征里的边缘、纹理信息而判断“这个区域是不是肿瘤”又依赖深层语义特征。如果只用FCN那种先下采样再直接上采样的方式细节早就丢了如果完全不降分辨率计算量又扛不住。U-Net通过跳跃连接在保留细节的同时又拿到了高层语义等于把“看得清”和“认得准”合到了一起。再加上医学图像结构相对固定目标尺度不像自然图像那么悬殊U-Net这种编码器-解码器结构在这个场景下特别稳。1.3 为什么先自己复现U-Net而不是直接套现成框架现在nnU-Net、MONAI这些工具已经很成熟不少人会问直接调库不就行了吗为什么还要自己写一遍U-Net我的观点是如果是做课题、发论文、快速出结果确实可以直接用现成框架但如果你是刚接触医学影像分割想搞懂原理或者需要针对自己的数据做深度调优从头实现一遍的价值非常大。U-Net代码量不算大但真正动手写才能理解通道数变化、跳跃连接维度匹配、损失函数里sigmoid的坑后面遇到问题才具备排查能力。另外不建议一上来就直接上3D U-Net。3D模型对显存要求高数据预处理也更复杂很多公开数据集本身就是稀疏标注的三维volume处理不好很容易把训练时间拉长几倍。用2D的U-Net先在切片上跑通整个流程再平滑迁移到3D是性价比最高的路径。我身边不少做医学影像的同事第一版baseline都是2D U-Net起步效果稳定调试方便迭代速度快。2. 数据准备与CT预处理这一半的功夫全在这2.1 从DICOM到NIfTI先搞定影像格式再谈模型做CT影像的肿瘤分割数据格式绕不开DICOM和NIfTI。医院原始数据大多是DICOM一个病例就是一个文件夹里面是几十到几百张不同层面的二维图附带大量扫描参数。公开发布的数据集基本都转成了NIfTI格式也就是一个.nii.gz文件包含整个三维volume处理起来方便很多。读取NIfTI我习惯用SimpleITK或者nibabel。读取后先别急着训练一定要打印几个关键信息数据维度、体素间距spacing、方向、数值范围。CT的DICOM像素值通常需要乘以Rescale Slope再加Rescale Intercept才是标准HU值好在大多数公开数据集已经完成了这步转换。spacing这个参数非常关键不同病例扫描层厚可能不一样如果不做重采样直接按切片喂给模型同一个肿瘤在不同病例里的形态会有系统性偏差。最简单做法是统一重采样到各向同性比如1mm×1mm×1mm但这一步会影响显存2D切片分割时对平面内分辨率更敏感轴向可以放宽。2.2 窗宽窗位与归一化的正确姿势CT值的绝对数值对我们人类看片很重要但对神经网络来说直接输入原始值反而不利。不同扫描设备、不同参数下CT值范围波动很大需要先做裁剪再归一化到0到1。医学上调整CT显示效果依赖“窗宽窗位”这个概念。比如腹部CT常用窗宽350到400HU、窗位40到50HU意思是把窗位附近的CT值拉伸到全灰度范围超出范围的直接截断。代码实现很简单import numpy as np def ct_window(volume, window_width400, window_level40): lower window_level - window_width / 2 upper window_level window_width / 2 volume np.clip(volume, lower, upper) volume (volume - lower) / (upper - lower) return volume.astype(np.float32)看起来简单但选不对窗宽窗位会直接影响结果。我做肝脏肿瘤时先对比了几种方案发现直接裁剪到[-200, 250]再归一化比照搬腹部窗效果更稳定肺结节分割则适合用肺窗参数。这里没有统一答案建议在正式训练前抽几个病例用不同的窗宽窗位把图像显示出来看一遍重点观察肿瘤和周围组织的对比度。如果你标注的肿瘤区域在数据里本来就特别不清晰模型学到的特征也会很有限。2.3 切片、划分训练集和避免数据泄漏CT数据是一个三维volume2D U-Net要按轴向切成一堆二维切片。切割时别把三维信息全丢掉要记住这条切片原来属于哪个病例、位于volume的哪个位置推理时还需要拼回去。切片前我一般会根据标注的实际情况做筛选肿瘤分割任务里大量切片只有背景没有肿瘤如果全部进训练集模型会严重偏向预测背景。可以按比例保留一部分纯背景切片再和含肿瘤的切片混合避免正负样本失衡太夸张。训练集和验证集的划分一定要按“病例”维度切不能按“切片”随机划分。同一个病例的相邻切片高度相似如果验证集和训练集里混着同一个病例的不同切片验证集的Dice分数会虚高上线后实际效果一塌糊涂。按病例划分后我习惯再按80比20左右的比例拆训练集和验证集如果病例数量太少就做五折交叉验证。2.4 数据增强哪些该用哪些要小心数据增强在医学影像分割里非常重要但也要带着脑子用。我常用的增强包括水平翻转、垂直翻转、90度旋转、小角度旋转、缩放、弹性形变以及轻微的对比度扰动。这些操作对CT影像的语义没有破坏性能明显提高模型的鲁棒性。需要注意的一点是输入图像和标注mask必须用同一套变换。用albumentations库最省心它有专门的Compose机制可以保证image和mask用完全相同的参数变换。另一个容易翻车的点是随机亮度对比度增强CT值是有物理含义的大幅度改变亮度对比度会让模型学到错误规律我用的时候会把幅度控制得很小甚至干脆不用。裁剪也很重要真实腹部CT周围有大量黑色背景直接resize整张图会让目标占比进一步变小可以先裁剪到目标区域附近再resize到256×256或512×512。3. PyTorch手搭U-Net网络结构与关键代码3.1 网络整体结构与通道设计标准U-Net的输入是一张单通道CT切片输出是一张和输入分辨率相同的概率图。编码器部分由4个下采样阶段组成每个阶段先做两次3×3卷积加ReLU再做一次2×2最大池化。输出通道数从64开始逐层翻倍64、128、256、512到达最底部bottleneck时是1024。解码器部分则反着来先做一次反卷积上采样再把编码器同层的特征拼过来然后两次3×3卷积通道数逐层减半最后用1×1卷积输出单通道。这个配置在显存充裕的GPU上是比较标准的但如果你只有8G显存第一次跑实验就可以把起始通道改成32也就是32、64、128、256、512效果下降不明显显存压力小很多。输入尺寸我一般选256×256或者512×512。256×256对大部分CT切片分割任务足够512×512对小肿瘤细节更友好但训练时间会成倍增加。建议先在小尺寸上把流程跑通确认数据管线没有bug再放大输入和通道数。3.2 编码器、解码器与跳跃连接代码实现PyTorch实现U-Net非常直观。先写一个基础的双卷积模块再写编码器、解码器最后用跳跃连接拼接。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_channels, out_channels): super().__init__() self.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.conv(x) class UNet(nn.Module): def __init__(self, in_channels1, out_channels1, features(64, 128, 256, 512)): super().__init__() self.encoders nn.ModuleList() self.pool nn.MaxPool2d(kernel_size2, stride2) self.bottleneck DoubleConv(features[-1], features[-1] * 2) self.decoders nn.ModuleList() self.upconvs nn.ModuleList() in_ch in_channels for f in features: self.encoders.append(DoubleConv(in_ch, f)) in_ch f for i in range(len(features) - 1, 0, -1): self.upconvs.append(nn.ConvTranspose2d(features[i] * 2, features[i - 1], kernel_size2, stride2)) self.decoders.append(DoubleConv(features[i], features[i - 1])) self.out_conv nn.Conv2d(features[0], out_channels, kernel_size1) def forward(self, x): enc_outputs [] for encoder in self.encoders: x encoder(x) enc_outputs.append(x) x self.pool(x) x self.bottleneck(x) for i, (upconv, decoder) in enumerate(zip(self.upconvs, self.decoders)): x upconv(x) skip enc_outputs[-(i 1)] if x.shape ! skip.shape: x nn.functional.interpolate(x, sizeskip.shape[2:], modebilinear, align_cornersFalse) x torch.cat([x, skip], dim1) x decoder(x) return self.out_conv(x)这里有个细节容易踩坑解码器拼接前要确认上采样后的尺寸和跳跃连接的尺寸完全一致。如果输入尺寸是奇数池化会造成一个像素的偏差直接torch.cat会报错所以我在拼接前加了一步interpolate兜底。实际训练时大部分输入都是256×256这种偶数尺寸不会触发这个问题但写代码时保留这个判断会省事很多。3.3 损失函数BCE和Dice Loss的配合肿瘤分割中肿瘤区域占整张图的比例非常小单独用交叉熵损失会让模型更倾向于把一切预测成背景。常用做法是把BCE Loss和Dice Loss结合起来用两个损失共同优化模型。Dice Loss直接衡量预测掩膜和真实掩膜的重叠程度天然对类别不平衡更鲁棒。def dice_loss(pred, target, smooth1e-6): pred torch.sigmoid(pred) intersection (pred * target).sum() return 1 - (2.0 * intersection smooth) / (pred.sum() target.sum() smooth)组合时建议使用BCEWithLogitsLoss而不是普通BCELoss因为BCEWithLogitsLoss内部会做sigmoid数值上更稳定。如果在网络输出层提前加了sigmoid又把输出扔给BCEWithLogitsLoss相当于做了两次sigmoid损失会非常奇怪。bce_fn nn.BCEWithLogitsLoss() dice_weight 0.5 loss (1 - dice_weight) * bce_fn(logits, target) dice_weight * dice_loss(logits, target)dice_weight可以从0.5开始调如果发现模型预测结果偏保守、边界太细可以适当提高Dice的权重如果训练震荡明显就降低Dice权重。实际项目中我很少加Focal LossBCE加Dice的组合已经能解决大部分问题Focal Loss的alpha和gamma又多两个超参收益不一定成正比。3.4 评估指标Dice、IoU与边界距离训练过程中不仅看loss还要看分割质量指标。最常用的是Dice系数和IoU。Dice是两倍交集除以两个集合元素数之和直觉上约等于重叠比例的两倍IoU是交集除以并集。医学影像论文里还经常出现HD95即95%豪斯多夫距离衡量预测表面和真实表面的最大偏差对边界误差更敏感。做CT肿瘤分割时Dice和IoU一般足够如果任务特别关注边界质量再考虑HD95。建议在验证阶段用sigmoid后再计算指标不能直接在logits上算。还有个容易忽视的点Dice是逐病例计算后取平均还是把所有切片的像素累加在一起算全局Dice结果会有差异。我通常先按病例汇总再对病例取平均这样能避免大肿瘤病例主导整体指标让每个病例的权重更均衡。4. 训练实操超参、显存与曲线分析4.1 一个标准的训练循环长什么样训练循环本身不复杂难的是稳定复现。我会在训练脚本里固定随机种子保证每次实验结果可比较。训练阶段开model.train()推理阶段开model.eval()并包在torch.no_grad()里。每轮epoch结束后跑一次验证集保存验证集Dice最高的模型权重而不是保存最后一轮的权重。optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) scheduler torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, modemax, factor0.5, patience10) best_dice 0.0 for epoch in range(num_epochs): model.train() epoch_loss 0.0 for images, masks in train_loader: images images.cuda() masks masks.cuda() logits model(images) loss criterion(logits, masks) optimizer.zero_grad() loss.backward() optimizer.step() epoch_loss loss.item() val_dice evaluate(model, val_loader) scheduler.step(val_dice) if val_dice best_dice: best_dice val_dice torch.save(model.state_dict(), best_unet.pth)这段代码跑起来一般没问题但有几个细节要注意负样本比例过高时每个epoch的loss下降不明显要看Dice曲线而不是只看lossoptimizer.zero_grad()放在loss.backward()之前别漏AMP混合精度训练时要单独维护GradScaler。验证集的评估函数里记得把模型输出:sigmoid后再算Dice不然logits的范围会让指标失真。4.2 超参调优记录我试过的几组参数超参组合我推荐从一组比较稳的基线开始而不是一上来就用极端参数。以256×256输入、2D U-Net为例我用的基线是AdamW优化器初始学习率1e-4batch size 16训练100到150个epoch学习率在Dice指标停滞时减半。这组参数在多数医学分割任务上都能收敛到不错的结果。试过的其他组合里SGD加上momentum0.99在某些任务上能拿到更高精度但收敛很慢需要配合更长的训练周期而学习率调到3e-4以上显存不变但训练容易震荡尤其是在数据量不大的时候。batch size方面如果显存只允许batch size 4或8建议直接改用GroupNorm替代BatchNorm不然小batch下BatchNorm的统计量不稳定验证集效果会飘。分组归一化实现很简单把nn.BatchNorm2d(out_channels)替换成nn.GroupNorm(num_groups8, num_channelsout_channels)即可。4.3 显存不够的三种解决办法显存溢出是U-Net编程实战里最常见的报错尤其输入尺寸到512×512、通道数拉到64起步时8G显存基本扛不住。第一招是降低输入分辨率从512降到384或者256显存占用减少非常明显第二招是减少通道数起始通道从64改成32第三招是开混合精度训练。PyTorch的torch.cuda.amp使用起来成本很低实际训练中能节省30%到40%显存而且对精度影响极小。from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for images, masks in train_loader: images, masks images.cuda(), masks.cuda() optimizer.zero_grad() with autocast(): logits model(images) loss criterion(logits, masks) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()如果显存还是不够再考虑梯度累积也就是每隔若干个小batch更新一次梯度。这种做法不会减少显存占用但能把有效batch size顶上去。需要注意的是梯度累积后BN统计量还是按小batch计算的这也是我觉得不如直接换GroupNorm省心的原因。5. 推理和后处理从模型输出到干净的分割掩膜5.1 推理流程与图像恢复训练完的模型在推理阶段要注意输入数据的预处理必须和训练时完全一致。这里“一致”指的是窗宽窗位参数、归一化范围、resize尺寸、中心裁剪坐标都要一模一样。我见过很多人训练时做了裁剪推理时直接喂全图结果模型输出一团糟就是因为输入的图像分布变了。推理流程一般是加载最佳模型权重切换到model.eval()用torch.no_grad()包住前向过程得到logits后用sigmoid转换为概率图再按阈值二值化。阈值的默认值是0.5但实际调优时可以在验证集上扫一遍0.3到0.7选择一个Dice或IoU最高的阈值。对整个volume逐切片推理后如果之前做过中心裁剪和resize要用原始坐标把每个切片的预测结果填回全图尺寸的空白矩阵里这样才能和原始mask对齐。5.2 后处理连通域过滤、形态学操作与CRF二值化后的mask通常有很多散落的小噪声点直接交给下游任务会很头疼。我一般用scipy.ndimage做连通域分析把面积过小的区域删掉。肿瘤如果预期是单发病灶可以直接保留最大连通域如果是多灶性病变就按体积阈值过滤。形态学闭运算填补内部小空洞也很常用。from scipy import ndimage labeled, num ndimage.label(mask) sizes ndimage.sum(mask, labeled, range(1, num 1)) keep_labels [i 1 for i, size in enumerate(sizes) if size 50] mask_clean np.isin(labeled, keep_labels).astype(np.uint8) mask_clean ndimage.binary_closing(mask_clean, structurenp.ones((3, 3))).astype(np.uint8)这里要注意体积阈值别设太大。肝脏肿瘤小病灶可能就几十个像素阈值设高了会把真实的小目标全删掉。阈值50是个相对保守的起点具体数值要根据自己数据集的像素分辨率调整。CRF条件随机场是另一种可选后处理能让分割边界更平滑但参数多、速度慢在模型本身输出已经比较稳定的前提下提升有限我一般最后才考虑。5.3 结果评估与可视化不只看一个数字评估整个volume的分割效果时不要只盯着一个全局Dice。我更推荐按病例统计Dice、IoU和体积误差再看分布。有的模型平均Dice不低但某些小肿瘤病例完全漏检只报平均分会掩盖这个问题。可视化的价值也很大把CT灰度图、真实mask、预测mask叠加在同一张图上红绿对比一目了然。医学影像里常把预测轮廓画在原图上检查边界和真实肿瘤是否对得上。import matplotlib.pyplot as plt def visualize_case(image_slice, gt_mask, pred_mask, save_path): plt.figure(figsize(15, 5)) plt.subplot(1, 3, 1) plt.imshow(image_slice, cmapgray) plt.title(CT) plt.subplot(1, 3, 2) plt.imshow(gt_mask, cmapReds, alpha0.7) plt.title(GT) plt.subplot(1, 3, 3) plt.imshow(pred_mask, cmapGreens, alpha0.7) plt.title(Pred) plt.savefig(save_path, dpi150, bbox_inchestight)如果发现某些切片预测明显抖动比如上一层有肿瘤、下一层消失了说明逐切片预测缺乏轴向一致性。这时可以考虑对相邻切片的概率图做平均或者用滑动窗口重叠预测能减轻一部分z轴不连续问题。6. 踩坑记录与排查思路6.1 显存溢出与训练崩溃速查训练过程中可能遇到各种各样的问题这里整理一份速查表都是我用2D U-Net做CT肿瘤分割时实际碰到过的。现象可能原因解决思路CUDA out of memory输入尺寸大、batch大、通道数多降分辨率、减通道数、开AMP、加梯度累积Loss为NaN学习率过大、输入或标签含NaN/inf检查数据范围降低学习率用detect_anomaly定位验证Dice几乎为0正负样本严重不均、阈值设错、模型没收敛查看loss曲线确认输出是否sigmoid调阈值预测mask全背景数据预处理不一致、标签范围不是0/1检查测试集窗口参数打印标签unique训练收敛慢学习率太低、BN batch太小增大学习率或换GroupNorm遇到CUDA out of memory时最先看报错发生的位置。如果是在forward阶段优先减小输入尺寸如果是在backward阶段梯度爆炸也可能大量占用显存需要用torch.utils.clip_grad_norm_做梯度裁剪。排查NaN时我会在训练循环里打印每个batch输入的min、max以及标签的unique值数据里混入NaN的概率比想象中高。6.2 模型预测全黑或全白的排查思路预测结果全黑的常见原因有三个。第一个是测试数据和训练数据预处理不一致比如训练时用了窗宽窗位裁剪推理时忘了对CT值做同样的操作第二个是标签范围不是0和1有些公开数据集的标注背景是0、肿瘤是2直接计算二值交叉熵会导致损失异常第三个是阈值设得不对肿瘤区域的概率输出本身比较低0.5可能过滤掉了大部分有效区域。排查时先在验证集上统计模型输出的概率分布看肿瘤区域的概率峰值到底在哪个范围而不是急着调网络结构。全白的情况相对少见一般出现在损失函数写错、模型输出和标签对不齐的时候。如果你发现损失持续下降但预测mask是整个区域全1先把标签可视化出来确认读取的mask和CT图像是否对应同一层尤其注意数据增强阶段image和mask有没有同步。6.3 验证集指标高但真实场景效果差的原因一种典型情况是划分数据时按切片随机分同一个病例的相邻切片一部分进了训练集一部分进了验证集。模型相当于提前见过这个病人的“大致样子”验证集Dice虚高。解决办法就是前面强调过的按病例划分并且要把这个原则贯彻到交叉验证里。另一种常见情况是训练和推理的图像分布不一致。比如训练时统一重采样到了1mm间距推理时跳过重采样步骤直接用原始间距或者训练时裁剪到固定尺寸推理时忘记恢复。这类问题不会报错模型也能输出结果但指标和视觉效果都会明显变差。养成一个习惯把所有预处理逻辑封装成函数训练、验证、推理都调用同一个函数不要复制粘贴代码。6.4 小肿瘤漏检与边界过细的优化方向小肿瘤漏检通常是因为下采样次数太多或者输入分辨率太低。U-Net默认下采样4次256×256的输入在bottleneck层变成16×16一个直径只有10像素的小肿瘤在低分辨率下可能就剩一两个像素了。这时可以考虑减少下采样次数比如改成3次或者在解码器阶段融合更浅层的特征。如果目标是多发小病灶后处理的面积阈值也不能设太大否则等于人为清掉小目标。边界过细、预测结果比真实肿瘤小一圈大概率是损失函数里Dice权重偏高模型为了追求重叠更保守。可以适当降低Dice权重或者把阈值从0.5下调到0.4。实际项目里我会在验证集上多做几组阈值扫描观察不同阈值对Dice、IoU的影响选择一个综合指标最好的组合而不是死守0.5。我在做这个项目时最深的体会是U-Net代码本身并不复杂真正的坑全在数据和预处理上。很多人喜欢花大量时间调网络结构却忽略了窗宽窗位、病例划分、增强同步这些细节。等你把这部分做扎实即使模型只是最普通的2D U-Net结果也往往不会差。