3D图像分割Pytorch工程实战:从Luna16数据预处理到VNet/UNet训练
发布时间:2026/10/1 3:10:27来源:尧图网络
简介一套基于 PyTorch 的三维图像分割完整方案以肺部 CT 结节公开数据集 Luna16 为例结合 UNet3d 与 VNet3d 两种卷积神经网络结构系统梳理了从原始数据重采样、掩膜生成、感兴趣区域裁剪、数据增强到训练、验证、测试、评估、可视化与结果后处理的代码思路适合医学影像与三维深度学习方向的中高级学习者参考。资源包共 92 个文件其中 49 个 Python 源码是核心涵盖预处理、数据集加载、模型定义、损失函数、推理脚本和可视化工具16 个 pyc 为编译缓存7 个 npy 存储预处理后的数组数据5 张 png 展示训练损失与验证指标变化曲线另有 xml、gz、nii 等配置与医学图像文件整体大小约 61.66MB。目前已有 543 人学习下载。压缩包目录按数据预处理、训练主程序、模型库、后处理与预测结果等模块化组织并提供多个版本的训练入口和 bbox 标签转换脚本便于逐模块复现实验也可作为扩展新分割任务的工程模板。1. 一份把 3D 图像分割数据准备写到极致的 Pytorch 工程如果你和我一样拿到 Luna16 的第一反应是直接丢进训练脚本那你大概率会在第一个 epoch 就翻车——CT 体素间距不一致、病灶边缘错位、显存瞬间爆掉。这份基于 Pytorch 的 3D 图像分割工程包是我见过少数把「数据准备 → 训练 → 验证 → 推理 → 评估 → 后处理」整条链路讲完整的开源代码案例数据就是 CT 结节的 Luna16模型是 CNN 结构的 VNet3d 和 UNet3d。它的价值不在于模型多先进而在于把 3D 分割最折磨人的数据预处理环节抽成了三个 step 脚本每一步都有对应输出适合正在做医学影像分割、想快速跑通一个 3D baseline 的从业者和研究生。这篇笔记我会按实际拆包顺序把每个脚本的用法、参数和踩过的坑讲清楚。2. 把 CT 原始数据做成训练集三个 step 脚本与 patch 裁剪策略2.1 为什么 3D 分割的数据准备比模型更费心思Luna16 的原始数据是 .mhd/.raw 格式每个病例的体素间距不一样有的层厚 1mm有的是 2.5mm。如果直接拿原始分辨率去训练同一个结节在不同病例里的大小和形状在体素空间里完全对不上模型学到的是「间距伪影」而不是「结节特征」。这就是 3D 分割和 2D 分类最大的差别2D 分类可以靠归一化把尺寸抹掉3D 分割的坐标、间距、方向必须显式处理。工程包里的 preProcess 目录下有三个脚本作者把数据准备拆成了三步这个设计我很认同——每一步生成中间产物出问题能定位到具体环节。脚本作用主要输出step1.generate_resample_image_and_mask.py把原始 CT 和 mask 重采样到统一间距重采样后的 image 和 maskstep2.mask2bbox_centerCoor.py从 mask 中提取每个结节的 bbox 和中心坐标bbox 坐标文件step3.generateNewBboxLabel_save_csv.py生成带外扩边界的结节标签并保存 csv新的 bbox 标签 csv2.2 step1 重采样CT 图像和 mask 必须用不同插值方式第一个脚本做的是重采样。常见做法是用 SimpleITK 读取原始 .mhd拿到 spacing 后按目标间距重新采样。这里有一个关键区别图像用线性插值mask 必须用最近邻插值。import SimpleITK as sitk def resample_image(itk_image, new_spacing[1.0, 1.0, 1.0], is_labelFalse): original_spacing itk_image.GetSpacing() original_size itk_image.GetSize() new_size [ int(round(orig_sz * orig_spa / new_spa)) for orig_sz, orig_spa, new_spa in zip(original_size, original_spacing, new_spacing) ] resampler sitk.ResampleImageFilter() resampler.SetOutputSpacing(new_spacing) resampler.SetSize(new_size) resampler.SetOutputDirection(itk_image.GetDirection()) resampler.SetOutputOrigin(itk_image.GetOrigin()) if is_label: resampler.SetInterpolator(sitk.sitkNearestNeighbor) else: resampler.SetInterpolator(sitk.sitkLinear) return resampler.Execute(itk_image)这段代码的核心是new_size的计算——原始尺寸乘以原始间距再除以目标间距。参数is_label控制插值方式图像用sitkLinear保留灰度连续性mask 用sitkNearestNeighbor防止标签灰阶出现 0.5 这种混叠值。我见过太多人在这一步偷懒图像和 mask 都用线性插值结果 loss 里出现一堆非 0 非 1 的预测目标训练直接不收敛。2.3 step2 与 step3从 mask 到 bbox再到带外扩的标签重采样完成后第二步是从 mask 里提取每个结节的 bbox 和中心坐标。Luna16 的 mask 是稀疏的一个 CT 序列 200 多层结节可能只占 5 层。用连通域分析把每个结节的包围盒提出来这一步相当于给后续的 patch 裁剪提供锚点。第三步的关键设计是外扩。原作者把 bbox 外扩了固定体素数比如前后左右各扩 5 个体素然后生成新的标签 mask 存成 csv。这个外扩的意义在于结节和周围组织是有粘连的只裁剪结节本体让模型学不到边界上下文外扩后模型能同时看到结节和邻近组织分割边界更干净。2.4 patch 裁剪为什么必须围绕中心点切而不是滑窗切全图数据集相关脚本里作者放了 datasets.py、datasets_patch.py、datasets_v2.py、datasets_v3.py 好几个版本核心逻辑都在datasets_patch.py里。它做的事情是读 csv 里的 bbox 和中心坐标以中心点为中心裁一个固定大小的 patch比如 96×96×96同时把对应的 mask patch 一起切出来。class DatasetPatch(Dataset): def __init__(self, csv_path, image_dir, mask_dir, patch_size96, is_trainTrue): self.annotations pd.read_csv(csv_path) self.image_dir image_dir self.mask_dir mask_dir self.patch_size patch_size self.is_train is_train def __getitem__(self, idx): row self.annotations.iloc[idx] center_z, center_y, center_x row[center_z], row[center_y], row[center_x] half self.patch_size // 2 z_start max(0, center_z - half) z_end min(max_z, center_z half) image_patch load_patch(self.image_dir, row[case_id], z_start, z_end) mask_patch load_patch(self.mask_dir, row[case_id], z_start, z_end) return torch.from_numpy(image_patch).float(), torch.from_numpy(mask_patch).float()这里要注意两个参数patch_size决定显存占用和感受野的平衡96 在 16GB 显存下 batch_size 也就只能给到 2 左右center_*是 step2 算出来的中心坐标裁剪时只用它定位不额外做对齐。为什么选中心点裁剪而不是全图滑窗因为 Luna16 里一个结节约占整个 CT 体素的千分之一都不到全图训练正负样本严重失衡模型会倾向于把所有 voxel 预测为背景。围绕中心点裁 patch每个样本都保证包含一个完整的目标学习效率高很多。3. 训练主脚本怎么读VNet3d 与 UNet3d 的选型逻辑和 loss 设计3.1 模型文件里的四个版本分别解决什么问题工程包的 models 目录下有 vnet3d.py、unet3d.py、unet3d_bn.py、unet3d_bn_activate.py 四个文件。VNet3d 是 2016 年 Milletari 提出的体素分割网络核心是残差连接和编码-解码结构每个 stage 内部用残差块堆叠UNet3d 是把经典 U-Net 的 2D 卷积换成 3D 卷积后得到的 baseline。两个模型在同一个工程里并列存在说明作者在拿它们做对照实验。实际使用时我一般这样选如果没有特殊需求先用 UNet3d_bn 起步因为它加了 BatchNorm训练稳定性好显存占用比 VNet 小一截如果 DICE 卡在 0.7 上不去换 VNet3d 试试残差连接对梯度的传递更友好在结节这类小目标上往往能多涨 2~3 个点。unet3d_bn_activate.py里的 activate 是作者对激活函数位置做的调整哪个版本在你的数据上更好只能跑实验验证这部分说实话有点玄学。3.2 train_main 系列脚本为什么有 2 到 6 五个版本train_main_2.py到train_main_6.py加上train_main_multCls.py作者把训练入口拆成多个副本每个脚本对应一次实验配置变体。这种做法在医学影像圈很常见——每次改一个关键变量就复制一个脚本方便回溯对比。你不需要全部看懂只要抓config.py里的配置项就行。# config.py 中的关键参数按常见工程配置整理 learning_rate 1e-4 batch_size 2 patch_size [96, 96, 96] num_epochs 300 loss_name dice_bce # 可选: dice, bce, dice_bce save_model_dir ./models_savedlearning_rate初始给 1e-4这是 3D 医疗分割的常用起点batch_size 2不是因为不想调大而是 96 的 patch 在 3D 卷积下显存开销约等于 2D 图像的几十倍强行调大直接 OOM。loss_name参数对应loss.py里的不同损失实现建议用dice_bce组合。logs目录下的lr.png、train_loss.png、valid_dice.png就是训练曲线的可视化输出每跑完一个 epoch 自动保存。我拿到一份别人的训练代码第一件事就是看valid_dice.png的曲线如果验证集 Dice 在震荡而不是上升多半是 loss 组合或者学习率调度出了问题省得自己再跑一遍才发现。3.3 loss.py 里的组合逻辑单用 BCE 会怎样loss.py 里作者写了多种损失核心的组合是 DiceLoss BCE。为什么不能单用 BCE因为结节在 patch 里的占比可能只有 2%~3%BCE 对每个 voxel 是平等对待的背景 voxel 的梯度会淹没前景。DiceLoss 按区域重叠计算损失天然对类别不平衡不敏感。但单用 DiceLoss 也有问题——在训练初期预测全为背景时梯度不稳定。组合 loss 的常见做法是def dice_bce_loss(pred, target): bce F.binary_cross_entropy_with_logits(pred, target) dice dice_loss(pred, target) # 返回 1 - Dice return bce dice参数上注意pred是未经过 sigmoid 的 logitsBCE 用with_logits版本Dice 计算时再套 sigmoid避免数值溢出。如果发现训练集 Dice 高但验证集 Dice 低可以先检查 loss 是不是只用了bce——那种配置下模型很容易把背景学得很好但结节的边界完全模糊。3.4 训练脚本里另一个值得抄的点global_annos.py 和 global_.py这两个文件看着像命名不规范的工具脚本打开看其实是把标注信息做全局管理的模块。global_annos.py负责把 csv 标注加载成全局字典global_.py里放的是跨脚本共享的常量。它的设计思路比具体代码更值得借鉴3D 分割的数据量大每个 epoch 都重新读 csv 不现实一次性加载到内存字典里训练过程中用索引访问能省掉大量 IO 时间。如果你的数据集标注文件很大可以直接复用这个模式。4. 从推理到评估再到可视化把四类脚本串成一条流水线4.1 inference_main.py 与 predict_noTargetCropPatch_4.py 的分工训练完成后拿到权重接下来是推理。inference_main.py是主推理脚本它读取测试 CT用滑动窗口生成 patch逐个输入网络得到概率图最后把概率图拼回原始尺寸。predict_noTargetCropPatch_4.py是另一种推理变体从文件名可以看出它不依赖预先的 bbox 标注适合在测试集没有 mask 的情况下做预测。这里有个推理策略的细节3D 推理时窗口的 stride 一般设置为 patch 大小的一半重叠区域取平均概率这样能避免拼接边界处的预测断层。# 推理主循环逻辑按常见工程写法整理 for z in range(0, depth - patch_d 1, stride_d): for y in range(0, height - patch_h 1, stride_h): for x in range(0, width - patch_w 1, stride_w): patch volume[z:zpatch_d, y:ypatch_h, x:xpatch_w] with torch.no_grad(): pred model(patch.unsqueeze(0).cuda()) prob_map[z:zpatch_d, y:ypatch_h, x:xpatch_w] pred.squeeze().cpu().numpy()stride设为 patch 一半的目的前面说了重叠区域取平均能抑制边缘伪影。patch输入前记得unsqueeze(0)补 batch 维输出后squeeze()去掉。这段代码里最容易错的是坐标遍历范围——最后一个窗口如果超出边界需要截断而不是报错跳出。4.2 evaluate_pred_mask.py用哪些指标评判一次分割效果推理产物是每个病例的预测 mask存成 NIfTI 或 NRRD 格式后evaluate_pred_mask.py负责和标准 mask 对比。医学分割任务最常用的指标是 Dice 系数但只报 Dice 容易被边界噪声蒙骗所以作者还加了其他评估维度。查看脚本里的实现通常会计算指标计算方式说明Dice2×A∩BVOE1 - Dice体积重叠误差敏感性TP / (TPFN)漏检率低更关键特异性TN / (TNFP)假阳性控制能力作者在工程里单独放了一个「评价方法.py」且标注是原作者给的说明这份工程对评估环节是认真的。评估 mask 前记得确认预测和标准 mask 的 spacing 是否一致——很多翻车现场都是这里没对齐导致 Dice 奇低。4.3 showRes_fromNrrd.py 和 crop_merge 脚本结果可视化与后处理showRes_fromNrrd.py负责把 NRRD 格式的预测结果叠加到 CT 原图上进行可视化便于肉眼确认分割边界是否贴合。crop_merge_fromNII_one.py和crop_merge_fromNrrd_one.py是配套的后处理脚本——推理阶段是按 patch 预测的所以要按 patch 的索引关系重新合并回整图这两个脚本就是干这个的。后处理还有一个几乎必做的环节用连通域分析把预测 mask 中面积小于设定阈值的孤立区域删除。结节是实体目标不太可能出现孤立小斑点这些小区域基本都是假阳性。工具函数可以在 tools.py 里找一般会用 scipy.ndimage 的 label 函数实现。5. 避坑与排查跑这套工程最容易翻车的五个地方5.1 维度顺序错乱导致结节位置偏移半个肺现象训练正常推理也能出 mask但把预测结果叠加到 NIfTI 原图上结节位置整体偏移了几十个 voxel像是被平移过。原因NIfTI 数据的内部存放顺序是 z, y, x 或者说 D, H, W而 Pytorch 3D 输入要求 B, C, D, H, W两个顺序在代码里没有严格统一。别的工具预处理时用了 x, y, z 顺序训练时又按 z, y, x 读入维度就错位了。解决在 path.py 和最外层脚本里强制约定顺序我一般会在读取 NIfTI 后立刻打印 shape 和 spacing确认每一维的含义再往下走。这个检查花不了两分钟能省下好几个小时的排查时间。5.2 mask 用线性插值重采样标签出现 0.5 的混叠值现象loss 前期下降正常但验证集 Dice 始终在 0.1 以下徘徊。原因重采样脚本对 mask 也用了 sitkLinear导致掩码边缘出现 0.2、0.5 这类非整数灰度值模型在学一个「模糊边界」。解决回到 step1确认 is_labelTrue 时强制用 sitkNearestNeighbor并检查重采样后的 mask 唯一值是否只有 0 和 1。这条是最隐蔽的因为重采样结果用肉眼看不出问题必须靠代码检查。5.3 裁剪 patch 越界导致训练中断现象训练到一半抛 IndexError报错在数据加载部分。原因结节的中心坐标靠近图像边缘裁剪 patch 时下界为负或上界超出图像尺寸。解决在裁剪逻辑里加边界截断常见做法是z_start max(0, center - half)并记得同时调整 mask 的裁剪起点保证 image 和 mask 偏移一致。更好的做法是提前计算所有 patch 的索引范围对越界的病例做 padding 而不是丢弃因为边缘区域正是肿瘤常出现的位置靠近肺壁。5.4 GPU 显存溢出batch_size 调到 1 还是崩现象CUDA out of memory把 batch_size 降到 1 后仍然报错。原因显存瓶颈不一定在 batch 维度patch_size128 的 3D 输入在 UNet3d 深层特征图上的显存开销是 2D 的几十倍。解决优先调小 patch_size 而不是 batch_size比如从 96 降到 64观察 Dice 是否有明显变化另一个办法是启用梯度累积每 4 个 step 更新一次参数等效于 batch_size 保持不变但显存压力减半。如果还 OOM检查 num_workers 是不是调得太高DataLoader 的 worker 进程也会吃内存。5.5 csv 路径在跨平台迁移后全部失效现象在 Windows 上调试好的代码放到 Linux 服务器上训练直接报文件找不到。原因csv 里存的路径是绝对路径且用反斜杠分隔Windows 上能跑Linux 上不认识。解决改用一个统一的路径配置模块比如作者工程里的 path.py把所有数据路径在入口处重新拼接一次不要直接读 csv 里的路径字段。更稳的方案是 csv 里只存相对路径和 case_id加载时用 os.path.join 拼绝对路径。6. 把这套流程搬到你自己的 CT 数据改 config 和 path 的五个关键点不少人下载这份工程后第一反应是拿自己的 CT 数据直接跑然后卡在第一步。我拆完这套代码后总结出一个可复用的迁移清单按顺序检查这五个位置基本能一次跑通。第一重采样目标间距不必照抄 1.0mm要根据你的数据实际分布定。Luna16 是肺部 CT结节在 1mm 间距下足够清晰如果你的数据是肝脏肿瘤或脑部病灶建议先统计所有病例的 spacing 范围取中位数或众数作为目标间距不要无脑统一到 1mm。第二patch_size 要参考你的目标尺寸——结节直径普遍在 5~30mm96 的 patch 足够覆盖如果你的目标是整个器官或大病灶patch 至少要 128 起步否则模型永远看不到完整边界。第三数据增强要在 datasets_patch.py 的__getitem__里加主流的 3D 增强是随机翻转和随机旋转 90 度弹性形变在 CT 分割里慎用容易把解剖结构拉变形。第四类别数变了记得检查 loss 和模型输出维度。工程里有 train_main_multCls.py说明原作者也做了多分类扩展如果你要分三类模型输出通道改成 3loss 用 CrossEntropy 或对应多分类 Dice不能再套 sigmoid。第五评估指标要根据任务补充结节分割看 Dice 就行但器官分割建议加 HD95 豪斯多夫距离这指标对边界质量更敏感。工程里的online.py是我建议你最后读的脚本——它把整条推理链路封装成了在线预测流程输入一个 CT 序列输出分割结果适合部署和批量测试。从那以后我拿到任何一份新的 3D 分割数据都会强制先跑一遍 step1 到 step3 的预处理脚本看中间产物没问题再碰训练代码这个习惯帮我避掉了一大半的坑希望帮到你。本文还有配套的精品资源点击获取
网站建设高端定制企业官网