- 人工智能
- 深度学习
- 计算机视觉
- 医疗健康
【免费下载链接】nnUNet
导读:本文基于 nnU-Net 开源仓库中 DKFZ 团队对 Toothfairy2 挑战赛(口腔 CBCT 全景牙/颌分割)的官方提交方案文档,完整还原其技术路线:以 ResEnc L 残差编码器 U-Net 为骨干,将 Patch 尺寸放大至 160×320×320、禁用左右镜像并延长训练至 1500 epoch;同时讲解针对 grand-challenge 平台"单例 10 分钟推理时限"所做的规划器(torch 重采样)、双模型集成与体积阈值后处理等工程优化。读者将掌握从 mha 数据集转换、plans 文件手工扩写、定制 Trainer 训练,到双模型集成推理与后处理的端到端复现能力,并理解每个环节背后的源码级原理。
一、方案概览:为 Toothfairy2 挑战赛定制的 nnU-Net
Toothfairy2 挑战赛要求算法在口腔 CBCT 影像上分割牙齿与颌骨结构。本仓库 Toothfairy2 提交文档(作者为 DKFZ 医学图像计算部门)记录的官方方案,其核心并非重新设计网络,而是对 nnU-Net v2 的标准流程做四处关键改造:
- 放大 Patch:将 ResEnc L 配置的 Patch 尺寸上采样到
160×320×320体素,让网络看到更大的解剖上下文; - 加深网络:架构比标准 ResEnc L 多一个 stage(7 个 stage,多一次池化与残差块);
- 调整训练策略:禁用左右镜像(left/right mirroring),训练轮数从默认 1000 提升到 1500;
- 面向推理时限优化:用 torch 实现的重采样替代默认重采样(更快但略欠精确),训练两个模型并以 cross-validation 集成策略合并,以满足 grand-challenge 单例 10 分钟的时间上限。
训练硬件为 2×A100 40GB 或单块 GH200 96GB,模型从零开始训练(不加载预训练权重)。
二、数据集转换:mha 转 NIfTI 与标签重映射
2.1 转换脚本与标签语义
官方在 Dataset119_ToothFairy2_All.py 中完成数据集转换。该脚本的作用正如文档所述:"只是将 mha 文件转换为 nifti(更小的文件体积)并移除未使用的标签 id"。
从源码看,脚本同时完成了三类工作:
- 格式转换:
image_to_nifi使用 SimpleITK 将*.mha图像直接写出为*.nii.gz; - 标签重映射:
label_mapping通过查表把原始标签 id 映射为连续 id,并压缩为uint8(np.zeros_like(label_np, dtype=np.uint8)),丢弃不存在的"NA 类"; - 生成元数据:重写
dataset.json(file_ending改为.nii.gz、更新labels字典),并额外生成 70:30 划分的splits_final.json(随机种子固定为 42)。
2.2 标签映射表(mapping_DS119)
转换脚本中mapping_DS119()定义了核心映射规则("移除所有 NA 类并让类 id 连续"):
def mapping_DS119() -> Dict[int, int]: """Remove all NA Classes and make Class IDs continuous""" mapping = {} mapping.update({i: i for i in range(1, 19)}) # [1-10]->[1-10] | [11-18]->[11-18] mapping.update({i: i - 2 for i in range(21, 29)}) # [21-28]->[19-26] mapping.update({i: i - 4 for i in range(31, 39)}) # [31-38]->[27-34] mapping.update({i: i - 6 for i in range(41, 49)}) # [41-48]->[35-42] return mapping即原始标签 1–18 保持不变,21–28 映射到 19–26,31–38 映射到 27–34,41–48 映射到 35–42,最终得到 42 个前景类别(标签 1–42)。脚本中还保留了mapping_DS120(仅保留牙与颌类)与mapping_DS121(仅保留牙类)两种变体(默认注释掉),说明同源数据可衍生出多个 nnU-Net 数据集。
转换入口(需按实际路径修改root):
root = "/media/l727r/data/Teeth_Data/ToothFairy2_Dataset" process_ds(root, "Dataset112_ToothFairy2", "Dataset119_ToothFairy2_All", mapping_DS119(), None)运行后将得到符合 nnU-Net v2 数据集规范的Dataset119_ToothFairy2_All(含imagesTr/、labelsTr/、dataset.json)。
三、实验规划与预处理:指纹提取、torch 重采样规划与 plans 文件扩写
3.1 提取数据集指纹
nnUNetv2_extract_fingerprint -d 119 -np 48该命令并行(-np 48)分析 Dataset 119 的体素间距、强度分布、形状统计等信息,生成数据指纹,是规划阶段的前提。
3.2 使用 torch 重采样规划器进行规划
nnUNetv2_plan_experiment -d 119 -pl nnUNetPlannerResEncL_torchres这里的核心是自定义规划器nnUNetPlannerResEncL_torchres,定义于 resample_with_torch.py。它继承自nnUNetPlannerResEncL(见 residual_encoder_unet_planners.py,目标显存约 24GB,UNet_reference_val_3d = 2100000000),并做了两点改动:
- 替换默认重采样方案:
determine_resampling()返回resample_torch_fornnunet(实现在 resample_torch.py,基于torch.nn.functional.interpolate),而非默认的基于 scipy 的方案——更快但精度略低; - 修改 plans 标识:
plans_name默认为nnUNetResEncUNetLPlans_torchres,generate_data_identifier()返回plans_identifier + '_' + configuration_name,保证不同 plans 文件产生的预处理数据互不混淆。
如文档所强调,由于挑战赛训练/测试图像本就统一为0.3×0.3×0.3 mm间距,理论上无需重采样,改用快速方案纯粹是"安全冗余";真正的动机是推理端的速度——grand-challenge 平台对每个病例限时 10 分钟。
3.3 编辑 plans 文件:手工追加超大 Patch 配置
规划完成后,需在生成的 plans 文件(nnUNetResEncUNetLPlans_torchres)中追加以下配置块。除修改 Patch 尺寸外,该配置还将网络加深了一个 stage(多一次池化 + 残差块),从而让网络更充分地利用更大输入:
"3d_fullres_torchres_ps160x320x320_bs2": { "inherits_from": "3d_fullres", "data_identifier": "nnUNetPlans_3d_fullres_torchres_ctnorm", "patch_size": [ 160, 320, 320 ], "normalization_schemes": [ "CTNormalization" ], "architecture": { "network_class_name": "dynamic_network_architectures.architectures.unet.ResidualEncoderUNet", "arch_kwargs": { "n_stages": 7, "features_per_stage": [ 32, 64, 128, 256, 320, 320, 320 ], "conv_op": "torch.nn.modules.conv.Conv3d", "kernel_sizes": [ [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3], [3, 3, 3] ], "strides": [ [1, 1, 1], [2, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2], [2, 2, 2], [1, 2, 2] ], "n_blocks_per_stage": [ 1, 3, 4, 6, 6, 6, 6 ], "n_conv_per_stage_decoder": [ 1, 1, 1, 1, 1, 1 ], "conv_bias": true, "norm_op": "torch.nn.modules.instancenorm.InstanceNorm3d", "norm_op_kwargs": { "eps": 1e-05, "affine": true }, "dropout_op": null, "dropout_op_kwargs": null, "nonlin": "torch.nn.LeakyReLU", "nonlin_kwargs": { "inplace": true } }, "_kw_requires_import": [ "conv_op", "norm_op", "dropout_op", "nonlin" ] } }关键点逐一解读:
| 字段 | 值 | 含义 |
|---|---|---|
inherits_from | 3d_fullres | 继承标准 3D 全分辨率配置的其余属性 |
patch_size | [160, 320, 320] | 相比 ResEnc L 默认(约 128×128×128 量级)大幅放大,利用各向同性 0.3mm 间距数据 |
n_stages | 7 | 比标准 ResEnc L 的 6 个 stage 多一级下采样 |
features_per_stage | (32,64,128,256,320,320,320) | 末三级特征数封顶在 320 |
strides末项 | [1, 2, 2] | 最后一个 stage 仅在 Y/Z 轴池化(第 0 轴不池化),适配 160 这一较小轴 |
normalization_schemes | ["CTNormalization"] | 与data_identifier中的ctnorm对应,CT 强度归一化 |
3.4 预处理
nnUNetv2_preprocess -d 119 -c 3d_fullres_torchres_ps160x320x320_bs2 -plans_name nnUNetResEncUNetLPlans_torchres -np 48注意-plans_name必须与规划阶段一致(nnUNetResEncUNetLPlans_torchres),预处理产物按nnUNetPlans_3d_fullres_torchres_ctnorm数据标识落盘,供训练时按配置名读取。
四、训练:双模型、1500 epoch 与仅 0/1 轴镜像
4.1 训练命令
官方在两个模型上训练全部训练病例(allfold):
nnUNetv2_train 119 3d_fullres_torchres_ps160x320x320_bs2 all -p nnUNetResEncUNetLPlans_torchres -tr nnUNetTrainer_onlyMirror01_1500ep nnUNet_results=${nnUNet_results}_2 nnUNetv2_train 119 3d_fullres_torchres_ps160x320x320_bs2 all -p nnUNetResEncUNetLPlans_torchres -tr nnUNetTrainer_onlyMirror01_1500ep两行命令的唯一差别在于第二行通过环境变量覆盖nnUNet_results(追加_2后缀),从而在不互相覆盖结果的前提下,把同一模型训练两遍——这是后续双模型集成的基础。模型均从零开始训练。
4.2 定制 Trainer 的源码原理
nnUNetTrainer_onlyMirror01_1500ep定义于 nnUNetTrainerNoMirroring.py,继承链为nnUNetTrainer_onlyMirror01_1500ep → nnUNetTrainer_onlyMirror01 → nnUNetTrainer:
class nnUNetTrainer_onlyMirror01(nnUNetTrainer): """ Only mirrors along spatial axes 0 and 1 for 3D and 0 for 2D """ def configure_rotation_dummyDA_mirroring_and_inital_patch_size(self): rotation_for_DA, do_dummy_2d_data_aug, initial_patch_size, mirror_axes = \ super().configure_rotation_dummyDA_mirroring_and_inital_patch_size() patch_size = self.configuration_manager.patch_size dim = len(patch_size) if dim == 2: mirror_axes = (0, ) else: mirror_axes = (0, 1) self.inference_allowed_mirroring_axes = mirror_axes return rotation_for_DA, do_dummy_2d_data_aug, initial_patch_size, mirror_axes class nnUNetTrainer_onlyMirror01_1500ep(nnUNetTrainer_onlyMirror01): def __init__(self, plans: dict, configuration: str, fold: int, dataset_json: dict, device: torch.device = torch.device('cuda')): super().__init__(plans, configuration, fold, dataset_json, device) self.num_epochs = 1500- 镜像策略:
mirror_axes = (0, 1)表示训练时只沿 3D 空间的前两个轴(对应 sagittal/coronal 方向)做镜像增强,禁用了左右(轴向,axis 2)镜像;同时inference_allowed_mirroring_axes也被设为(0, 1),保证推理时测试时增强(TTA)的镜像轴与训练一致。这一点对牙齿这类左右高度对称的结构很有意义——避免网络被镜像引入的"左右位置先验"干扰。 - 训练长度:
self.num_epochs = 1500覆盖默认的 1000 epoch。
同文件还提供了nnUNetTrainerNoMirroring(完全禁用镜像)与nnUNetTrainer_onlyMirror01_DA5、nnUNetTrainer_onlyMirror01_DASegOrd0等变体,可在其他项目按需复用。
4.3 数据增强的 CPU 瓶颈与nnUNet_n_proc_DA
文档特别提醒:建议提高数据增强进程数,否则容易遭遇 CPU 瓶颈:
export nnUNet_n_proc_DA=32(系统允许时可继续调高。)原因在于训练所用增强管线(旋转、缩放、弹性形变、噪声、低分辨率模拟、镜像等,实现在上述 Trainer 的get_training_transforms中)均需 CPU 逐 batch 执行,而 160×320×320 的超大 Patch 使单 batch 的增强计算量显著增大;若 DA 进程数不足,GPU 将被迫空转等待。
五、推理:双模型集成的限时优化
5.1 集成策略:把双模型伪装成交叉验证 fold
grand-challenge 平台对推理限时 10 分钟/病例,因此官方没有采用"跑两次推理再平均"的朴素集成,而是复用 nnU-Net 内置的交叉验证集成机制:
技术上,把两个模型的
fold_all目录复制到同一个训练输出目录下,分别重命名为fold_0与fold_1,从而启用 nnU-Net 的 cross-validation ensembling 策略(计算效率更高,满足平台时限)。
这样推理时两个模型共享同一份预处理后的滑窗高斯权重、同一批 patch 采样,逐个 fold 加载权重并把 softmax 概率累加,避免重复搬运大张量。
5.2 推理脚本整体流程
官方推理脚本为 inference_script_semseg_only_customInf2.py,其main流程如下:
if __name__ == '__main__': os.environ['nnUNet_compile'] = 'f' parser = argparse.ArgumentParser() parser.add_argument('-i', '--input_folder', type=Path, default="/input/images/cbct/") parser.add_argument('-o', '--output_folder', type=Path, default="/output/images/oral-pharyngeal-segmentation/") parser.add_argument('-sem_mod', '--semseg_trained_model', type=str, default="/opt/app/_trained_model/semseg_trained_model") parser.add_argument('--semseg_folds', type=str, nargs='+', default=[0, 1]) args = parser.parse_args() args.output_folder.mkdir(exist_ok=True, parents=True) semseg_folds = [i if i == 'all' else int(i) for i in args.semseg_folds] semseg_trained_model = args.semseg_trained_model rw = SimpleITKIO() input_files = list(args.input_folder.glob('*.nii.gz')) + list(args.input_folder.glob('*.mha')) for input_fname in input_files: output_fname = args.output_folder / input_fname.name # load test image im, prop = rw.read_images([input_fname]) with torch.no_grad(): semseg_pred = predict_semseg(im, prop, semseg_trained_model, semseg_folds) torch.cuda.empty_cache() gc.collect() # now postprocess semseg_pred = postprocess(semseg_pred, np.prod(prop['spacing']), True) semseg_pred = map_labels_to_toothfairy(semseg_pred) # now save rw.write_seg(semseg_pred, output_fname, prop)默认路径已按 grand-challenge 的容器规范设定:输入为/input/images/cbct/,输出为/output/images/oral-pharyngeal-segmentation/;脚本同时支持*.nii.gz与*.mha输入。单例流程为:读取 → 集成预测 → 体积后处理 → 标签反向映射 → 写盘。
5.3 CustomPredictor:面向显存与速度的定制
predict_semseg使用CustomPredictor(nnUNetPredictor)(源码同文件内定义),初始化参数为tile_step_size=0.5, use_mirroring=True, use_gaussian=True。相比基类,它重写了三个关键方法:
initialize_from_trained_model_folder:从model_training_output_dir读取dataset.json、plans.json与checkpoint_final.pth,用trainer_class.build_network_architecture(...)恢复网络(enable_deep_supervision=False),并支持nnUNet_compile环境变量触发torch.compile;predict_preprocessed_image:这是最核心的提速/省显存手段——- 每个 fold 的权重按需加载(
self.network.load_state_dict(torch.load(p, ...)['network_weights'])),而不是一次性加载全部参数,"每个参数集在测试集上只用一次,运行时间几乎不变,但显存占用大幅下降"; - 滑窗预测在
torch.autocast下以 float16 进行,高斯权重(compute_gaussian(tuple(patch_size), sigma_scale=1./8, value_scaling_factor=10),见 sliding_window_prediction.py)同样以 float16 计算,且数据与 logits 驻留 CPU、仅计算时搬运到cuda:0; - 每个 patch 的预测先做
pred /= (pred.max() / 100)归一化,再乘高斯权重累加,保证不同 fold 贡献可比较;
- 每个 fold 的权重按需加载(
convert_predicted_logits_to_segmentation_with_correct_shape:用configuration_manager.resampling_fn_probabilities把 logits 重采样回裁剪前原始间距(注意这里同样使用 torch 重采样),argmax取类别、回填到裁剪 bbox,最后按transpose_backward还原轴序。
单线程相关细节:预测与后处理中分别torch.set_num_threads(7),在受限容器环境中平衡 CPU 吞吐。
5.4 标签反向映射
由于训练标签是连续 id(1–42),提交结果必须映射回 ToothFairy2 官方标签体系。map_labels_to_toothfairy完成该逆映射(mapping数组长度为 43,仅对 19–42 重新映射):
def map_labels_to_toothfairy(predicted_seg: np.ndarray) -> np.ndarray: max_label = 42 mapping = np.arange(max_label + 1) remapping = {19: 21, 20: 22, 21: 23, 22: 24, 23: 25, 24: 26, 25: 27, 26: 28, 27: 31, 28: 32, 29: 33, 30: 34, 31: 35, 32: 36, 33: 37, 34: 38, 35: 41, 36: 42, 37: 43, 38: 44, 39: 45, 40: 46, 41: 47, 42: 48} for k, v in remapping.items(): mapping[k] = v return mapping[predicted_seg]这与数据集转换时的mapping_DS119恰好互逆,确保提交结果符合挑战赛的标签定义。
六、后处理:基于五折交叉验证的体积截断
文档明确后处理规则:若某个类别的预测连通体积小于对应 cutoff,则将其移除(替换为背景)。
6.1 Cutoff 的确定方法
- cutoff 值在 Toothfairy2 训练数据上通过五折交叉验证优化得到;
- 分别针对HD95与Dice两个指标优化出两组 cutoff;
- 每个类别的最终 cutoff 取两组中的较小值(更保守、宁可多删),因此对假阳性小碎块有较强的抑制作用。
6.2 推理脚本中的postprocess实现
官方把最终体积 cutoff 直接写死在推理脚本的postprocess函数中(按类 id 1–42 组织):
def postprocess(prediction_npy, vol_per_voxel, verbose: bool = False): cutoffs = {1: 0.0, 2: 78411.5, 3: 0.0, 4: 0.0, 5: 2800.0, 6: 1216.5, 7: 0.0, 8: 6222.0, 9: 1573.0, 10: 946.0, 11: 0.0, 12: 6783.5, 13: 9469.5, 14: 0.0, 15: 2260.0, 16: 3566.0, 17: 6321.0, 18: 4221.5, 19: 5829.0, 20: 0.0, 21: 0.0, 22: 468.0, 23: 1555.0, 24: 1291.5, 25: 2834.5, 26: 584.5, 27: 0.0, 28: 0.0, 29: 0.0, 30: 0.0, 31: 1935.5, 32: 0.0, 33: 0.0, 34: 6140.0, 35: 0.0, 36: 0.0, 37: 0.0, 38: 2710.0, 39: 0.0, 40: 0.0, 41: 0.0, 42: 970.0} vol_per_voxel_cutoffs = 0.3 * 0.3 * 0.3 for c in cutoffs.keys(): co = cutoffs[c] if co > 0: mask = prediction_npy == c pred_vol = np.sum(mask) * vol_per_voxel if 0 < pred_vol < (co * vol_per_voxel_cutoffs): prediction_npy[mask] = 0 if verbose: print( f'removed label {c} because predicted volume of {pred_vol} is less than the cutoff {co * vol_per_voxel_cutoffs}') return prediction_npy实现要点:
- 传入的
vol_per_voxel = np.prod(prop['spacing'])是该病例真实体素体积(mm³),而 cutoff 以"0.3mm 间距下的体素数"为单位,比较前乘以vol_per_voxel_cutoffs = 0.3×0.3×0.3换算为 mm³,从而对实际间距与标称 0.3mm 的微小偏差保持稳健; - cutoff 为 0 的类别(如 1、3、4、7、11 等)直接跳过,不做删除;
- 该后处理是类别级的(按
prediction_npy == c统计整个类别的总体积),而非连通域级,因此实现极快,几乎不占用推理时限预算。
七、端到端复现路线图与关键文件索引
把整条流水线汇总为可执行的步骤清单:
| 阶段 | 命令/动作 | 关键产物 |
|---|---|---|
| 1. 数据转换 | 运行 Dataset119_ToothFairy2_All.py | Dataset119_ToothFairy2_All(NIfTI + 连续标签) |
| 2. 指纹提取 | nnUNetv2_extract_fingerprint -d 119 -np 48 | dataset_fingerprint |
| 3. 规划 | nnUNetv2_plan_experiment -d 119 -pl nnUNetPlannerResEncL_torchres | nnUNetResEncUNetLPlans_torchres |
| 4. 编辑 plans | 追加3d_fullres_torchres_ps160x320x320_bs2配置(见 3.3) | 含超大 Patch 与 7-stage 架构 |
| 5. 预处理 | nnUNetv2_preprocess -d 119 -c 3d_fullres_torchres_ps160x320x320_bs2 -plans_name nnUNetResEncUNetLPlans_torchres -np 48 | 预处理后的训练数据 |
| 6. 训练 ×2 | 两条nnUNetv2_train ... -tr nnUNetTrainer_onlyMirror01_1500ep(第二条改nnUNet_results) | 两个fold_all模型 |
| 7. 集成准备 | 复制两个fold_all→fold_0/fold_1 | 单目录双 fold 模型 |
| 8. 推理 + 后处理 | 运行 inference_script_semseg_only_customInf2.py | 提交用分割结果 |
复现所需的源码锚点:
- 定制 Trainer:
nnUNetTrainer_onlyMirror01与nnUNetTrainer_onlyMirror01_1500ep定义于 nnUNetTrainerNoMirroring.py(镜像轴(0,1)、num_epochs=1500); - torch 重采样规划器:
nnUNetPlannerResEncL_torchres定义于 resample_with_torch.py,底层重采样实现见 resample_torch.py; - ResEnc L 家族规划器(显存目标与参考值)见 residual_encoder_unet_planners.py;
- 滑窗高斯权重
compute_gaussian(sigma_scale=1/8、value_scaling_factor=10)见 sliding_window_prediction.py; - 官方提交文档原文:Toothfairy2/readme.md。
八、结语:一套"标准方案 + 限时优化"的范式
Toothfairy2 的这套提交方案的价值不仅在于挑战赛成绩,更在于它示范了 nnU-Net v2 如何在"标准自动化流程"与"平台级工程约束"之间取得平衡:通过手工扩写 plans 文件把 ResEnc L 的 Patch 放大一个量级、增加一个 stage,通过定制 Trainer 微调镜像与训练长度,通过 torch 重采样与双模型集成把推理时间压进 10 分钟窗口,最后用五折交叉验证训出的体积 cutoff 清洗假阳性。对任何需要"在有限显存与推理时限内把 nnU-Net 推到极限"的任务,这套从 数据集转换脚本 到 推理脚本 的完整链路,都是一份可直接迁移参考的实战范本。
- 人工智能
- 深度学习
- 计算机视觉
- 医疗健康
【免费下载链接】nnUNet
相关推荐
OpenCore Legacy Patcher完整教程:4步让老Mac重获新生的终极指南
OpenCore Legacy Patcher完整教程:4步让老Mac重获新生的终极指南 还在为你的老Mac无法升级最新macOS而烦恼吗?看着2012年的Ma
人工智能深度学习计算机视觉医疗健康nnU-Net v2 入门指南:从安装、数据准备到训练与推理的完整工作流
nnU Net v2 入门指南:从安装、数据准备到训练与推理的完整工作流 本指南是面向首次使用 nnU Net v2 用户的完整入门路线图:它把官方文档中分散的
人工智能深度学习计算机视觉医疗健康eldarion-ajax实战:构建现代化CRUD应用的最佳实践
eldarion ajax实战:构建现代化CRUD应用的最佳实践 在当今Web开发领域,构建响应式、用户友好的CRUD(创建、读取、更新、删除)应用已成为标配需
人工智能深度学习计算机视觉医疗健康
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考