ARTICLE DETAIL

资讯详情

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

基于PyTorch的3D医学图像分割:从NIfTI数据预处理到DataLoader全流程

基于PyTorch的3D医学图像分割:从NIfTI数据预处理到DataLoader全流程 简介基于Pytorch的3D图像分割任务配套资源面向医学影像分析与深度学习入门者以Luna16 CT结节数据为案例完整覆盖UNet3d与VNet3d两种CNN结构从数据准备到后处理的全流程。资源共92个文件以49个Python脚本为核心配合16个编译缓存pyc、7个npy预处理数据、5张训练曲线图及5个xml配置另有少量nii/gz原始数据与csv标注文件压缩包总计61.66MB。按数据预处理、训练、推理、后处理等模块组织便于对照学习。代码思路参考作者系列文章详细讲解了重采样、掩码生成、bbox坐标提取、patch采样、损失函数、训练验证、模型评估与预测结果裁剪合并等环节可视化脚本输出loss、dice曲线及分割示例。目前已有543人学习下载适合想要系统掌握3D医学图像分割工程实现的读者。1. 3D 图像分割的第一步不是网络是数据这篇笔记对标的是什么很多做 2D 分割跑得很顺的人第一次转 3D 分割任务时会懵掉——不是模型不会写而是数据根本喂不进去。2D 的 jpg/png 拿到就能读3D 的医学影像或者工业 CT 往往是一个.nii文件、一组几百层的切片还带着 spacing、orientation、窗宽窗位这些额外信息。你如果直接把 2D 的 ResNet 套上去大概率会在第一次torch.Tensor转换时翻车。这篇笔记就是围绕基于 Pytorch 的 3D 图像分割任务把数据准备过程和代码思路从零拆开从原始体数据怎么读、怎么预处理到怎么切 patch、怎么写 DataLoader再到那些不翻一次车根本记不住的边界条件。内容面向两种人一种是想复现 3D U-Net / V-Net 但卡在数据上的人另一种是已经跑通 2D 分割、想迁移到 3D 但不确定数据流怎么设计的工程师。我不会讲太深的数学只讲能落地、能跑通的那条路。2. 先解决“数据长什么样”3D 医学影像的格式、坐标系与重采样2.1 NiFTI 不是一张图读取 .nii 之前先弄懂三个字段3D 分割数据集最常见的是 NiFTI 格式.nii或.nii.gz由医学影像社区主导。和普通图片不同一个.nii文件除了体素数组还带着affine、spacing、orientation三个核心元数据。有人觉得“反正都是数组”直接np.array(img)拿过来用结果训练出来的模型换个数据集就崩——这大概率是没对齐坐标系和物理间距。读取.nii文件常见做法是用SimpleITK或nibabel。我一般用SimpleITK因为它处理 spacing 和重采样更顺手而且 API 更接近医学影像工程师的习惯import SimpleITK as sitk import numpy as np def load_nii(file_path): # 读取整个体数据返回图像对象 img sitk.ReadImage(file_path) # 取出像素数组shape 是 (depth, height, width) arr sitk.GetArrayFromImage(img) # 获取体素间距单位通常是毫米 spacing img.GetSpacing() # 获取仿射矩阵4x4把体素坐标映射到物理坐标 affine img.GetDirection() origin img.GetOrigin() return arr, spacing, affine, origin # 使用示例 arr, spacing, affine, origin load_nii(case_001.nii.gz) print(f体数据 shape: {arr.shape}) print(f体素间距 spacing: {spacing})逻辑说明GetArrayFromImage返回的 shape 是(z, y, x)也就是先深度后行列这和nibabel的get_fdata()默认返回顺序(x, y, z)不一样。新手最容易在这一步踩坑——自己写代码时读出来(512, 512, 300)觉得没问题但后面做 patch 切分的时候shape顺序错了整个训练就全乱了。我习惯在数据加载的第一行就固定统一为(depth, height, width)后面所有函数都按这个约定写。参数说明spacing是一个三维向量表示每个体素在三个方向上的物理间距。大部分公开数据集是各向异性数据比如spacing(0.7, 0.7, 3.0)意味着层间间距远大于平面内间距。这个值不能丢后面做重采样和归一化都要用它。2.2 重采样到各向同性WHY 和 HOW大多数 3D 分割模型比如 3D U-Net假设体素是各向同性的。如果你的数据spacing(0.7, 0.7, 3.0)直接送进网络模型会认为第三维和第二维的空间尺度一样结果训练时损失权重被深层切片噪声干扰分割边界在层间方向糊成一片。处理方向有两个一是重采样到各向同性比如都变成 1.0mm 或 1.5mm二是保留各向异性但网络结构里用非对称卷积。对绝大多数团队重采样到各向同性是更稳的路因为你换网络结构意味着重新设计模型。def resample_to_iso(img, target_spacing(1.0, 1.0, 1.0), is_labelFalse): 将体数据重采样到目标 spacing is_labelTrue 时使用最近邻插值避免标签被平滑 original_spacing img.GetSpacing() original_size img.GetSize() # 计算目标尺寸原始尺寸 * 原始间距 / 目标间距四舍五入取整 target_size [ int(round(orig_sz * orig_sp / target_sp)) for orig_sz, orig_sp, target_sp in zip(original_size, original_spacing, target_spacing) ] resampler sitk.ResampleImageFilter() resampler.SetSize(target_size) resampler.SetOutputSpacing(target_spacing) resampler.SetOutputDirection(img.GetDirection()) resampler.SetOutputOrigin(img.GetOrigin()) if is_label: resampler.SetInterpolator(sitk.sitkNearestNeighbor) else: resampler.SetInterpolator(sitk.sitkLinear) return resampler.Execute(img)逻辑说明这个函数里最关键的是target_size的计算公式——目标尺寸不是随便定的必须由原始 size、原始 spacing、目标 spacing 三者推导否则几何信息会错位。is_label参数控制插值方式这是无数人踩过的坑标签图如果用了线性插值原本是 1 的体素可能在边界处变成 0.6损失计算直接爆炸。标签图一律用最近邻插值这个规则永远不要破。参数说明target_spacing设为(1.0, 1.0, 1.0)是把所有数据变成各向同性但注意重采样后体积可能会变大比如层间距 3mm 变成 1mm深度方向尺寸膨胀 3 倍。显存不够的机器建议先设(1.5, 1.5, 1.5)试跑。2.3 HU 值截断与归一化窗宽窗位是分割任务的隐藏参数CT 影像的原始值叫 HUHounsfield Unit范围从 -1024 到 3071 甚至更高。不同部位的组织 HU 范围差异极大空气约 -1000脂肪约 -120水 0软组织 4080骨骼超过 400。如果你直接 Min-Max 归一化全图背景和低密度组织会占据大部分数值区间目标器官反而被压缩。正确做法是先做 HU 截断把无关范围砍掉再归一化。def ct_preprocess(arr, lower-200, upper300): CT 数据预处理先截断 HU 值范围再线性归一化到 [0,1] 默认窗宽窗位适合腹部软组织肝/脾/肾类任务 # 截断到 [lower, upper] 区间 arr_clipped np.clip(arr, lower, upper) # 线性映射到 [0, 1] arr_normalized (arr_clipped - lower) / (upper - lower) return arr_normalized.astype(np.float32) # 标签不需要截断但需要保证为整数类型 def preprocess_label(arr): return arr.astype(np.int64)参数说明lower-200, upper300是腹部器官分割的常见窗宽窗位。如果你做骨分割窗口要放到lower200, upper2000做肺结节窗口可能是lower-1350, upper150。这里没有绝对标准我一般直接看数据的直方图取目标组织所在的峰段。别在窗口上花太多时间玄学调参先跑一版试试分割效果差再回头调窗口效率最高。3. 3D 数据怎么进模型patch 切分、滑动窗口与前景采样3.1 为什么 3D 分割几乎必须用 patch-based 方法3D 体数据通常很大。一个典型的腹部 CT 重采样后是(300, 256, 256)甚至更大直接整体送进 3D U-Net 会显存溢出。即便是高端显卡batch size 为 1 也未必能塞下一个完整体积。业界标准做法是 patch-based 训练从整个体数据里切出固定大小的小块比如(96, 96, 96)或(128, 128, 64)用这些 patch 来训练。patch 大小的选择有讲究。切得太小感受野不足分割目标内部容易出现空洞切得太大显存装不下不说batch size 被迫变小训练稳定性差。我常见的策略是先看目标器官在数据集里的体积分布取能包裹 90% 目标的最小 patch 尺寸。另外要保证 patch 里包含边界背景否则模型会学不到“目标外就是背景”这个基础分类信号。3.2 随机采样 vs 滑动窗口采样训练和推断要分开设计训练阶段用随机采样推断阶段用滑动窗口拼接这两者不能混用。随机采样的思路是在每个训练 epoch 里随机从体数据中抽 patch。如果目标器官体积小纯随机采样会让大量 patch 落在背景区域模型训练半天学不到东西。此时要加一个“前景采样比例”——比如 50% 的 patch 强制落在标注区域附近。def sample_patch_foreground(img_arr, label_arr, patch_size(96, 96, 96), foreground_ratio0.5, rngNone): 训练阶段随机采样 patch按比例混合前景采样和全图随机采样 img_arr: (D, H, W) 的预处理后图像 label_arr: (D, H, W) 的标签0 为背景非 0 为目标 patch_size: 采样块大小 foreground_ratio: 前景采样比例0~1 之间 D, H, W label_arr.shape pD, pH, pW patch_size if rng is None: rng np.random.default_rng(42) # 找到标签里所有前景体素的位置 foreground_indices np.argwhere(label_arr 0) for _ in range(5): # 最多尝试 5 次防止越界 if len(foreground_indices) 0 and rng.random() foreground_ratio: # 从前景体素里随机取一个点作为 patch 中心 center foreground_indices[rng.integers(0, len(foreground_indices))] else: # 全图随机取中心点 center np.array([rng.integers(0, D), rng.integers(0, H), rng.integers(0, W)]) # 根据 patch 尺寸计算起始坐标并限制在图内 start_d max(0, min(center[0] - pD // 2, D - pD)) start_h max(0, min(center[1] - pH // 2, H - pH)) start_w max(0, min(center[2] - pW // 2, W - pW)) img_patch img_arr[start_d:start_d pD, start_h:start_h pH, start_w:start_w pW] label_patch label_arr[start_d:start_d pD, start_h:start_h pH, start_w:start_w pW] return img_patch, label_patch # 极端情况兜底全图随机再试一次 start_d rng.integers(0, D - pD 1) start_h rng.integers(0, H - pH 1) start_w rng.integers(0, W - pW 1) return img_arr[start_d:start_d pD, start_h:start_h pH, start_w:start_w pW], \ label_arr[start_d:start_d pD, start_h:start_h pH, start_w:start_w pW]逻辑说明这段代码的核心逻辑是“先取中心点再算起始坐标最后越界裁剪”。越界裁剪不是直接clip到边界就完事而是要保证切割窗口始终在体数据内部——所以起始坐标的计算同时受控于center - patch_size // 2和D - pD这两个边界条件。rng.integers用的是numpy.random.Generator比老式np.random.randint更推荐种子可控且线程安全好一些。参数说明foreground_ratio一般取 0.5意思是训练过程中半数 patch 围绕目标区域。如果目标器官极小小于整个体数据的 1%可以抬到 0.8 甚至 0.9但不要直接设 1.0否则模型会把所有背景区域都判断成目标附近泛化能力会下降。3.3 滑动窗口推断重叠区怎么合并train 阶段你随机采样没问题但 test 阶段必须覆盖完整体积。常见做法是滑动窗口 重叠区加权平均。窗口滑过整个体数据每个位置切一块 patch 送进模型得到概率图再把重叠部分的概率加权平均。def sliding_window_infer(model, img_arr, patch_size(96, 96, 96), stride_ratio0.5): 推断阶段滑动窗口拼接概率图 stride_ratio0.5 表示步长为 patch 尺寸的一半重叠率 50% import torch import torch.nn.functional as F D, H, W img_arr.shape pD, pH, pW patch_size stride_d int(pD * stride_ratio) stride_h int(pH * stride_ratio) stride_w int(pW * stride_ratio) # 输出概率图num_classes 从模型输出推断 model.eval() with torch.no_grad(): # 先粗略估计 class 数用一次临时 forward dummy_patch torch.from_numpy(img_arr[:pD, :pH, :pW]).float().unsqueeze(0).unsqueeze(0) dummy_out model(dummy_patch) num_classes dummy_out.shape[1] # 概率累加器和计数累加器 prob_acc np.zeros((num_classes, D, H, W), dtypenp.float32) count_acc np.zeros((D, H, W), dtypenp.float32) # 滑动窗口遍历 for d in range(0, D - pD 1, stride_d): for h in range(0, H - pH 1, stride_h): for w in range(0, W - pW 1, stride_w): patch img_arr[d:dpD, h:hpH, w:wpW] patch_tensor torch.from_numpy(patch).float().unsqueeze(0).unsqueeze(0) logits model(patch_tensor) # (1, C, pD, pH, pW) probs F.softmax(logits, dim1).squeeze(0).cpu().numpy() prob_acc[:, d:dpD, h:hpH, w:wpW] probs count_acc[d:dpD, h:hpH, w:wpW] 1.0 # 重叠区取平均 count_acc[count_acc 0] 1.0 # 防止除零 prob_mean prob_acc / count_acc[np.newaxis, ...] return prob_mean逻辑说明代码里用count_acc做计数累加是有讲究的——重叠区域内每个体素被覆盖的次数可能不一样边界区域覆盖次数少、中心区域覆盖次数多直接用概率累加不除次数边界处亮度会明显暗一截argmax 后容易出现锯齿状伪影。这个滑动窗口方案是标准做法重叠率越高结果越平滑但推理时间成正比增长。stride_ratio0.5是常见折中方案重叠率 50% 能让边界区域至少被两个 patch 覆盖足够平滑。参数说明如果你的显存允许更大的 patchstride_ratio可以降到 0.25 获取更平滑的边界反之显存紧张只能用小 patch 时重叠率务必保持 50% 以上否则拼接缝隙会很明显。3.4 显存不够的备选思路分块级联与伪 3D如果 patch 切到(64, 64, 64)依然 OOM有两个备选方案。第一个是分块级联先用低分辨率跑整个体积得到粗糙分割图再把粗糙分割图里的目标区域放大到原始分辨率精修。这个思路在很多医学影像竞赛里拿过名次代价是代码复杂度翻倍。第二个是伪 3D把 3D 卷积拆成三路 2D 卷积分别处理轴向、冠状位、矢状位三个视角最后融合。这个方案显存消耗只有真 3D 的三分之一左右但模型要自己写没有现成的预训练权重可以抄。上策还是先从数据下手——检查自己 CT 数据的 z 轴层厚是不是太大了层厚 5mm 的数据重采样到 1mm 是没有意义的信息量根本不够纯属浪费显存。4. 3D 数据增强与类别不均衡不能直接抄 2D 的增强库4.1 哪几种几何增强在 3D 上能用且值得用2D 分割常用的翻转、旋转、缩放在 3D 里大部分保留但有细节差别。翻转方面轴向翻转永远不要开——医学图像尤其是 CT天然具有“头在上脚在下”的空间一致性翻转轴向等于把解剖结构上下颠倒会让模型学到错误的空间先验。我最常用的是绕 z 轴旋转轴向旋转角度取 90°/180°/270°因为腹部 CT 的冠状位和矢状位方向没有必须保持的方向性旋转不影响诊断价值。弹性形变在 3D 上用得少因为三维弹性形变的控制点数量巨大计算代价高而且形变场如果处理不当标签和图像的对齐会崩。我倾向于少用或不用。def augment_3d(img_patch, label_patch, flip_prob0.3, rotate_prob0.3): 3D 数据增强轴向翻转绕 z 轴 绕 z 轴旋转 90/180/270 img_patch: (D, H, W) 图像 label_patch: (D, H, W) 标签类别为整数 import random # 翻转只翻 H 和 W 维度不翻 D 维度 if random.random() flip_prob: img_patch np.flip(img_patch, axis1) # 翻转 H label_patch np.flip(label_patch, axis1) if random.random() flip_prob: img_patch np.flip(img_patch, axis2) # 翻转 W label_patch np.flip(label_patch, axis2) # 旋转绕 z 轴旋转 k * 90 度 k random.choice([0, 1, 2, 3]) if k ! 0: img_patch np.rot90(img_patch, kk, axes(1, 2)) label_patch np.rot90(label_patch, kk, axes(1, 2)) return img_patch.copy(), label_patch.copy()注意这里返回值用了.copy()原因是np.flip和np.rot90返回的是原数组的视图不是新数组。如果直接返回后续对 patch 的原地修改会连带影响原数据DataLoader 多进程时甚至会引发数据竞争。.copy()写在这里是血的教训。4.2 类别不均衡3D 分割的 Dice Loss 和采样策略怎么配合3D 分割的类别不均衡比 2D 更严重。一个肝分割数据里背景体素可能占 99%前景只占 1%。用纯 Cross Entropy Loss 会得到“全预测背景”的模型指标上还特别好看——Dice 直接变 0。常见做法是联合Dice Loss Cross EntropyDice Loss 天然处理不均衡因为它直接优化重叠区域比例CE 则提供梯度稳定性。import torch import torch.nn.functional as F class DiceCE3DLoss(torch.nn.Module): def __init__(self, num_classes, ce_weight0.4, smooth1.0): super().__init__() self.num_classes num_classes self.ce_weight ce_weight self.smooth smooth def forward(self, logits, targets): # logits: (B, C, D, H, W) # targets: (B, D, H, W), 值为类别索引 probs F.softmax(logits, dim1) # 转为概率 # 将 targets 转为 one-hot targets_onehot F.one_hot(targets, num_classesself.num_classes).permute(0, 4, 1, 2, 3).float() # targets_onehot: (B, C, D, H, W) # 计算每个类别的 Dice dice_loss 0.0 for c in range(self.num_classes): pred_c probs[:, c] target_c targets_onehot[:, c] intersection (pred_c * target_c).sum(dim(1, 2, 3)) denominator pred_c.sum(dim(1, 2, 3)) target_c.sum(dim(1, 2, 3)) dice (2.0 * intersection self.smooth) / (denominator self.smooth) dice_loss (1.0 - dice).mean() dice_loss dice_loss / self.num_classes ce_loss F.cross_entropy(logits, targets) return ce_loss * self.ce_weight dice_loss * (1 - self.ce_weight)逻辑说明Dice Loss 的一个微妙之处是它对小目标很苛刻——如果一张大 patch 里只有一小块目标预测稍微偏移几个体素Dice 会从 0.9 崩到 0.3梯度信号很大。反过来如果目标占了 patch 的一半Dice 就非常钝感梯度近乎为零。所以代码里保留了一部分 CE Loss 来提供稳定的梯度。ce_weight0.4是经验值如果训练初期 Dice 完全不下降把 CE 权重调到 0.5 以上如果模型出现过拟合降到 0.2 左右。4.3 预处理顺序不要乱归一化和增强的先后关系我见过有人先做归一化再做裁剪也见过先裁剪再归一化两者结果差异不大但有一个原则窗口截断和黄窗操作必须在最前面归一化其次增强最后。原因是弹性形变和旋转这类几何增强会对灰度值做插值插值后的数值范围可能会超出 [0, 1]如果归一化在增强之后截断和缩放的效果会被插值破坏数据分布会漂移。反过来归一化在前增强只改变空间位置不改变数值分布训练更稳定。代码结构上把增强放在 Dataset 的__getitem__里预处理放在数据加载阶段统一跑好存成 npy能省很多时间。5. 避坑3D 分割数据准备阶段最常见的 5 类翻车现场5.1 坑一imgaug/torchvision的增强库直接套 3Dshape 对不上现象调用torchvision.transforms.RandomRotation处理 3D 数据时报错或者直接运行通过但输出形状不对。原因这类增强库是给 2D 图像设计的内部假设输入是(C, H, W)遇到(D, H, W)或(C, D, H, W)会误把D当作C或者当成 batch 维处理。解决3D 增强一律自己写。上面 4.1 节里的augment_3d已经覆盖了翻转和旋转这两个核心需求如果要做弹性形变用scipy.ndimage.map_coordinates配合随机形变场生成不要依赖 2D 库。5.2 坑二label 里类别编号不连续one-hot 之后总维数爆炸现象数据注释时标签是[0, 1, 4]0 背景、1 器官 A、4 器官 BF.one_hot默认会生成 5 个通道训练时损失计算多出两个空类。原因标注人员习惯用原始编号不关心类别之间的空洞。模型输出通道数只能由最大编号 1 决定空类别让训练计算量白白变多还可能导致模型在这两个空类上产生预测噪声。解决在预处理阶段重新映射标签。写一个函数做类别压缩def remap_labels(label_arr, original_ids, target_idsNone): 把原始标签编号映射为连续编号 original_ids: [0, 1, 4] - target_ids: [0, 1, 2] if target_ids is None: target_ids list(range(len(original_ids))) mapping {orig: target for orig, target in zip(original_ids, target_ids)} remapped np.zeros_like(label_arr) for orig, target in mapping.items(): remapped[label_arr orig] target return remapped.astype(np.int64)逻辑说明映射时必须遍历original_ids里的每个类别把原值替换成连续编号。注意np.zeros_like初始化是为了覆盖未标注区域原数组里等于 0 或未出现的编号保证这些位置在映射后仍然是背景。original_ids建议直接从数据集的np.unique(label_arr)获取不要写死。5.3 坑三spacing 没有统一就训练结果模型在不同设备采集的数据上漂移现象训练集是 A 医院 1.0mm 层厚验证集是 B 医院的 3.0mm 层厚。训练时 Dice 达到 0.85验证集直接掉到 0.3。原因模型学到的是“每个体素的语义”没有学到“物理空间距离”。不同厚度的数据进同一网络感受野覆盖的真实物理范围不同模型学到一半的尺度信息就乱了。解决在数据准备阶段强制统一 spacing。所有训练和验证数据resample_to_iso到同一目标 spacing这步不能跳过。如果硬件显存限制不能到 1.0mm 各向同性可以统一到 1.5mm 或 2.0mm但必须全数据集一致。没有统一 spacing 之前不要碰模型训练。5.4 坑四滑动窗口推断时 padding 方式不对导致边界出现黑带现象推理结果里整个体积的边缘出现明显的低概率区域分割掩膜在边界处收缩。原因滑动窗口遍历时窗口超出体数据的边界代码里用了零填充。零填充的 patch 里有大量 0 值模型认为这些“背景”可信度很高输出概率偏向背景导致边缘目标预测被压低。解决不要用零填充用边缘反射填充或者干脆限制窗口不要越过边界。限界方式见 3.3 节的实现起始坐标range(0, D - pD 1, stride_d)保证了窗口不超出边界不需要任何填充。如果非要用填充用np.pad(img_arr, pad_width, modereflect)——反射填充出的内容在语义上和邻近组织更接近模型不会把它们误判成空气。5.5 坑五验证时用随机采样的 patch 评估指标虚高又翻车现象验证集上用随机采样的 patch 跑 Dice每次运行结果都不一样两次之间的波动超过 5 个点。原因随机采样 patch 带来的评估方差。今天采到 500 个富目标 patchDice 高明天采到 500 个背景 patchDice 低。验证集评估必须用完整体积的滑动窗口推理不能图省事。解决验证阶段调用sliding_window_infer做全量推理然后基于原始体素空间计算 Dice 和 IoU。如果验证集太大导致推理缓慢可以只评估固定 seed 下抽样的 35 个完整体积保证每次评估的数据一致。验证阶段不要使用任何随机性。6. 把数据流串成 Pytorch Dataset/DataLoader一个可直接改的完整骨架6.1 Dataset 的__getitem__要返回什么3D 分割的 Dataset 和 2D 的差别在于__getitem__返回的是五维张量(C, D, H, W)而且必须在返回前把增强和采样逻辑封装好。这里我给一个完整的骨架内部把采样、增强、tensor 转换都串起来from torch.utils.data import Dataset, DataLoader import torch class Seg3DDataset(Dataset): def __init__(self, img_paths, label_paths, patch_size(96, 96, 96), foreground_ratio0.5, augmentTrue, target_spacing(1.0, 1.0, 1.0)): self.img_paths img_paths self.label_paths label_paths self.patch_size patch_size self.foreground_ratio foreground_ratio self.augment augment self.target_spacing target_spacing # 读取并预处理所有数据缓存到内存 self.images [] self.labels [] for img_path, label_path in zip(img_paths, label_paths): img sitk.ReadImage(img_path) lb sitk.ReadImage(label_path) # 统一 spacing img resample_to_iso(img, self.target_spacing, is_labelFalse) lb resample_to_iso(lb, self.target_spacing, is_labelTrue) img_arr sitk.GetArrayFromImage(img) lb_arr sitk.GetArrayFromImage(lb) # 类别重映射假设类别 ID 列表从数据里获取 original_ids np.unique(lb_arr) if len(original_ids) 1: lb_arr remap_labels(lb_arr, original_ids) # CT 预处理如果是 CT 数据 img_arr ct_preprocess(img_arr) self.images.append(img_arr.astype(np.float32)) self.labels.append(lb_arr.astype(np.int64)) self.num_classes len(np.unique(np.concatenate([l.ravel() for l in self.labels]))) def __len__(self): return len(self.images) * 20 # 每个样本每个 epoch 抽样 20 个 patch def __getitem__(self, idx): # 根据 patch 计数决定用哪个原始体积 vol_idx idx // 20 img_arr self.images[vol_idx] label_arr self.labels[vol_idx] # 随机采样 patch img_patch, label_patch sample_patch_foreground( img_arr, label_arr, self.patch_size, self.foreground_ratio ) # 数据增强 if self.augment: img_patch, label_patch augment_3d(img_patch, label_patch) # 转换为 tensor 并加 channel 维 img_tensor torch.from_numpy(img_patch).float().unsqueeze(0) # (1, D, H, W) label_tensor torch.from_numpy(label_patch).long() # (D, H, W) return img_tensor, label_tensor参数说明__len__返回len(self.images) * 20是我常用的做法——每个 epoch 里每份体数据抽样 20 个 patch。这个 20 不是固定标准样本少的时候可以调到 50样本多的时候 5 就够了。要注意foreground_ratio参数被传给了sample_patch_foreground确保每个 patch 都能包含到目标区域。num_classes是在初始化时从所有标签里统计算出来的后面损失函数要用它做 one-hot。6.2 DataLoader 的坑num_workers和内存占用3D 数据比 2D 大得多DataLoader 的num_workers设置不当会导致内存暴涨或启动失败。train_loader DataLoader( train_dataset, batch_size2, shuffleTrue, num_workers4, # 可根据 CPU 核数和内存调整 pin_memoryTrue if torch.cuda.is_available() else False, drop_lastTrue, # 3D 分割 batch 通常不会整除丢弃最后不完整的 batch persistent_workersTrue, # 减少多 epoch 重复 spawn worker 的开销 )参数说明batch_size2对 3D 分割是保守选择(2, 1, 96, 96, 96)的输入体量在 12GB 显存上跑 3D U-Net 差不多是极限。你如果显存紧张先跑batch_size1确认 OOM 界限再逐步往上加。num_workers4不是越高越好每个 worker 都会把一份 patch 数据复制到内存worker 太多内存直接吃满。Linux 系统下persistent_workersTrue能省掉每个 epoch 重建 worker 的开销Windows 下可能不稳定建议设False。6.3 验证整个数据管线用一次前向传播检查五件事新写好数据管线后不要急着训练。先在 CPU 上跑一次完整的train_loader迭代检查五个关键点# 快速验证脚本跑通数据管线 for i, (img_tensor, label_tensor) in enumerate(train_loader): print(f图像 tensor shape: {img_tensor.shape}) # 期望 (B, 1, D, H, W) print(f标签 tensor shape: {label_tensor.shape}) # 期望 (B, D, H, W) print(f图像 range: [{img_tensor.min().item():.3f}, {img_tensor.max().item():.3f}]) print(f标签唯一值: {torch.unique(label_tensor).tolist()}) assert img_tensor.shape[2:] label_tensor.shape[1:], 图像和标签空间尺寸不一致 assert img_tensor.shape[0] label_tensor.shape[0], batch 维度不一致 break # 只跑一个 batch 做验证 # 用一个简单 3D 卷积验证 forward 能通 import torch.nn as nn dummy_model nn.Sequential( nn.Conv3d(1, 8, kernel_size3, padding1), nn.BatchNorm3d(8), nn.ReLU(inplaceTrue), ).eval() with torch.no_grad(): out dummy_model(img_tensor) print(f3D 卷积输出 shape: {out.shape}) # 期望 (B, 8, D, H, W)注意脚本里的两个断言图像和标签的空间尺寸必须一致batch 维度必须一致。这两个检查过不了后面训练每一步都是错的。我第一次做 3D 分割时就是漏了空间尺寸检查图像是 96³、标签是 94³跑了一个 epoch 才发现 Dice 始终不降最后定位到是预处理里裁剪边界差了 2 个像素。这种问题训练时几乎不可能看出来数据管线验证阶段必须拦下来。检查完后看一眼标签类别数是否小于模型输出通道数如果模型通道数大于实际类别数最后一层输出里会有空通道损失计算里它们是常量分母会让反向传播出现无效梯度——我通常直接让模型out_channels len(torch.unique(label_tensor))这样最省事。最后提一个习惯3D 分割项目的数据准备阶段永远比模型结构调参花更多时间。不少团队在模型上折腾了两周最后发现问题出在 spacing 没对齐、标签重采样插值方式错了、验证集 patch 采样有随机性。我自己的流程是数据准备占六成时间模型结构占两成训练调参占两成——数据稳了模型换什么结构差别都不大。希望帮到你。本文还有配套的精品资源点击获取
返回列表