简介:基于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_label=False): """ 将体数据重采样到目标 spacing is_label=True 时使用最近邻插值,避免标签被平滑 """ 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 影像的原始值叫 HU(Hounsfield Unit),范围从 -1024 到 3071 甚至更高。不同部位的组织 HU 范围差异极大:空气约 -1000,脂肪约 -120,水 0,软组织 40~80,骨骼超过 400。如果你直接 Min-Max 归一化全图,背景和低密度组织会占据大部分数值区间,目标器官反而被压缩。正确做法是先做 HU 截断,把无关范围砍掉再归一化。
def ct_preprocess(arr, lower=-200, upper=300): """ 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, upper=300是腹部器官分割的常见窗宽窗位。如果你做骨分割,窗口要放到lower=200, upper=2000;做肺结节,窗口可能是lower=-1350, upper=150。这里没有绝对标准,我一般直接看数据的直方图,取目标组织所在的峰段。别在窗口上花太多时间玄学调参,先跑一版试试,分割效果差再回头调窗口,效率最高。
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_ratio=0.5, rng=None): """ 训练阶段随机采样 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_ratio=0.5): """ 推断阶段滑动窗口拼接概率图 stride_ratio=0.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), dtype=np.float32) count_acc = np.zeros((D, H, W), dtype=np.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:d+pD, h:h+pH, w:w+pW] patch_tensor = torch.from_numpy(patch).float().unsqueeze(0).unsqueeze(0) logits = model(patch_tensor) # (1, C, pD, pH, pW) probs = F.softmax(logits, dim=1).squeeze(0).cpu().numpy() prob_acc[:, d:d+pD, h:h+pH, w:w+pW] += probs count_acc[d:d+pD, h:h+pH, w:w+pW] += 1.0 # 重叠区取平均 count_acc[count_acc == 0] = 1.0 # 防止除零 prob_mean = prob_acc / count_acc[np.newaxis, ...] return prob_mean逻辑说明:代码里用count_acc做计数累加是有讲究的——重叠区域内每个体素被覆盖的次数可能不一样,边界区域覆盖次数少、中心区域覆盖次数多,直接用概率累加不除次数,边界处亮度会明显暗一截,argmax 后容易出现锯齿状伪影。这个滑动窗口方案是标准做法,重叠率越高结果越平滑,但推理时间成正比增长。stride_ratio=0.5是常见折中方案,重叠率 50% 能让边界区域至少被两个 patch 覆盖,足够平滑。
参数说明:如果你的显存允许更大的 patch,stride_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_prob=0.3, rotate_prob=0.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, axis=1) # 翻转 H label_patch = np.flip(label_patch, axis=1) if random.random() < flip_prob: img_patch = np.flip(img_patch, axis=2) # 翻转 W label_patch = np.flip(label_patch, axis=2) # 旋转:绕 z 轴旋转 k * 90 度 k = random.choice([0, 1, 2, 3]) if k != 0: img_patch = np.rot90(img_patch, k=k, axes=(1, 2)) label_patch = np.rot90(label_patch, k=k, 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 Entropy,Dice Loss 天然处理不均衡,因为它直接优化重叠区域比例;CE 则提供梯度稳定性。
import torch import torch.nn.functional as F class DiceCE3DLoss(torch.nn.Module): def __init__(self, num_classes, ce_weight=0.4, smooth=1.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, dim=1) # 转为概率 # 将 targets 转为 one-hot targets_onehot = F.one_hot(targets, num_classes=self.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_weight=0.4是经验值,如果训练初期 Dice 完全不下降,把 CE 权重调到 0.5 以上;如果模型出现过拟合,降到 0.2 左右。
4.3 预处理顺序不要乱:归一化和增强的先后关系
我见过有人先做归一化再做裁剪,也见过先裁剪再归一化,两者结果差异不大,但有一个原则:窗口截断和黄窗操作必须在最前面,归一化其次,增强最后。原因是弹性形变和旋转这类几何增强会对灰度值做插值,插值后的数值范围可能会超出 [0, 1],如果归一化在增强之后,截断和缩放的效果会被插值破坏,数据分布会漂移。反过来归一化在前,增强只改变空间位置不改变数值分布,训练更稳定。代码结构上,把增强放在 Dataset 的__getitem__里,预处理放在数据加载阶段统一跑好存成 npy,能省很多时间。
5. 避坑:3D 分割数据准备阶段最常见的 5 类翻车现场
5.1 坑一:imgaug/torchvision的增强库直接套 3D,shape 对不上
现象:调用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 器官 B),F.one_hot默认会生成 5 个通道,训练时损失计算多出两个空类。
原因:标注人员习惯用原始编号,不关心类别之间的空洞。模型输出通道数只能由最大编号 +1 决定,空类别让训练计算量白白变多,还可能导致模型在这两个空类上产生预测噪声。
解决:在预处理阶段重新映射标签。写一个函数做类别压缩:
def remap_labels(label_arr, original_ids, target_ids=None): """ 把原始标签编号映射为连续编号 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, mode='reflect')——反射填充出的内容在语义上和邻近组织更接近,模型不会把它们误判成空气。
5.5 坑五:验证时用随机采样的 patch 评估,指标虚高又翻车
现象:验证集上用随机采样的 patch 跑 Dice,每次运行结果都不一样,两次之间的波动超过 5 个点。
原因:随机采样 patch 带来的评估方差。今天采到 500 个富目标 patch,Dice 高;明天采到 500 个背景 patch,Dice 低。验证集评估必须用完整体积的滑动窗口推理,不能图省事。
解决:验证阶段调用sliding_window_infer做全量推理,然后基于原始体素空间计算 Dice 和 IoU。如果验证集太大导致推理缓慢,可以只评估固定 seed 下抽样的 3~5 个完整体积,保证每次评估的数据一致。验证阶段不要使用任何随机性。
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_ratio=0.5, augment=True, 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_label=False) lb = resample_to_iso(lb, self.target_spacing, is_label=True) 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_size=2, shuffle=True, num_workers=4, # 可根据 CPU 核数和内存调整 pin_memory=True if torch.cuda.is_available() else False, drop_last=True, # 3D 分割 batch 通常不会整除,丢弃最后不完整的 batch persistent_workers=True, # 减少多 epoch 重复 spawn worker 的开销 )参数说明:batch_size=2对 3D 分割是保守选择,(2, 1, 96, 96, 96)的输入体量在 12GB 显存上跑 3D U-Net 差不多是极限。你如果显存紧张,先跑batch_size=1确认 OOM 界限,再逐步往上加。num_workers=4不是越高越好,每个 worker 都会把一份 patch 数据复制到内存,worker 太多内存直接吃满。Linux 系统下persistent_workers=True能省掉每个 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_size=3, padding=1), nn.BatchNorm3d(8), nn.ReLU(inplace=True), ).eval() with torch.no_grad(): out = dummy_model(img_tensor) print(f"3D 卷积输出 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 采样有随机性。我自己的流程是数据准备占六成时间,模型结构占两成,训练调参占两成——数据稳了,模型换什么结构差别都不大。希望帮到你。
本文还有配套的精品资源,点击获取