☰
3D医学图像分割数据管线详解:从NIfTI到PyTorch的预处理与踩坑指南
2026/10/1 13:58:50 网站建设 项目流程

简介:基于Pytorch的3D图像分割完整实践资源,以Luna16 CT肺结节数据为案例,覆盖UNet3d与VNet3d两种CNN结构,面向医学图像处理、深度学习方向的研究者与开发者,尤其适合希望从零搭建3D分割训练流程的入门及进阶用户。资源共92个文件,约61.66MB,主体为49个Python脚本,涵盖数据预处理、模型定义、训练验证、推断评估等模块;另有npy格式的预处理数组、nii格式的医学图像样本、png格式的训练曲线与预测可视化结果,以及若干工程配置文件,便于快速复现实验。目前已有543人学习使用。通过该资源可系统掌握CT结节数据重采样、mask与bbox标注生成、patch采样、多类别与单类别训练等关键步骤,并配有后处理、预测图像融合及评价指标计算脚本,代码思路与作者公开系列文章对应,适合边读边练,显著降低3D分割任务的入门门槛。

1. 基于 Pytorch 的 3D 图像分割任务,数据准备为什么是第一道坎

基于 PyTorch 做 3D 图像分割,很多人以为把 .nii 文件读进来、转成 numpy 数组就能直接训练了。真上手会发现,一个 CT 序列动不动就是 512×512×400 的体数据,加上标注文件里的 mask,直接塞进 U-Net,轻则 loss 不降,重则显存爆掉、mask 和原图错位。问题大多不出在网络上,而出在这条数据准备链路上:格式解析、重采样、归一化、滑窗切块、标签对齐,每一环都藏着能让模型静默翻车的细节。这篇文章就沿着这条链路讲,从文件格式到 PyTorch 的 Dataset 实现,把代码思路、参数选择和踩坑点一次说透。适合正在做医学影像分析、工业 CT 检测,或者其他体数据分割任务,并且不想在数据环节反复返工的人。

2. 从 NIfTI 到 PyTorch 张量:读入、重采样与归一化的完整管线

2.1 读入库选型:SimpleITK、NiBabel 还是 pydicom

3D 分割任务里最常见的数据格式是 NIfTI(.nii.gz)和 DICOM 序列。读入这一步不建议自己写文件解析,直接用成熟库。三个库的定位完全不同,选错后面处处别扭。

库适用场景优点短板
SimpleITKNIfTI / DICOM 序列 / 预处理完整保留 spacing、origin、direction,内置重采样和滤波接口偏底层,需要理解 ITK 的坐标系概念
NiBabel神经影像 NIfTI 为主轻量,转 numpy 方便,社区资源多对 DICOM 序列支持弱
pydicom单张 DICOM 或序列解析能精细访问每个 tag拼 3D 卷、重采样都要自己搭

我一般优先用 SimpleITK。原因是医学图像重采样、方向对齐这些操作它都是现成的,而 NiBabel 读出来只是纯 numpy 数组,spacing 和 direction 等信息要另外维护,容易丢。pydicom 通常只在需要读取 DICOM 自定义 tag 时才会用到。

import SimpleITK as sitk image = sitk.ReadImage("case_001_ct.nii.gz") print("size:", image.GetSize()) # (x, y, z) 顺序 print("spacing:", image.GetSpacing()) # 体素间距,单位 mm print("origin:", image.GetOrigin()) print("direction:", image.GetDirection())

这里有一个高频坑点:SimpleITK 的 GetSize 返回的是 (x, y, z),而 sitk.GetArrayFromImage 转换出的 numpy 数组维度顺序是 (z, y, x)。也就是说,arr.shape[0] 对应的是最慢的那一维,也就是切片方向。写代码时一旦搞混轴序,后面所有重采样、切块、可视化都会跟着错,表现为 mask 旋转了 90 度或者整体转置。读入后第一件事就是打印这些元数据,确认轴序和物理坐标范围。

2.2 重采样:统一体素间距,否则模型学到的是假形状

不同 CT 设备扫描参数差别很大,常见的有 512×512×400 配 0.5×0.5×1.0 mm 的 spacing,也有 256×256×200 配 1.0×1.0×2.5 mm 的。如果不做重采样,模型看到的是被拉伸或压扁的器官形态,同一个肝脏在不同病例里的尺寸和形状分布会被体素网格扭曲,分割精度直接受影响。

重采样的核心思路是:保持图像覆盖的物理区域不变,改变体素网格的密度。目标 spacing 的选择取决于任务器官的大小和网络下采样的深度。肝脏、肺这类大器官常用 1.0×1.0×1.0 mm 或 1.5×1.5×1.5 mm,细小结构比如血管、神经则需要 0.5 mm 级别的各向同性分辨率,但数据量和显存开销会成倍增长。

def resample_to_spacing(itk_image, new_spacing=(1.0, 1.0, 1.0), is_label=False): original_spacing = itk_image.GetSpacing() original_size = itk_image.GetSize() new_size = [ int(round(size * spacing / target)) for size, spacing, target in zip(original_size, original_spacing, new_spacing) ] resampler = sitk.ResampleImageFilter() resampler.SetOutputSpacing(new_spacing) resampler.SetSize(new_size) resampler.SetOutputOrigin(itk_image.GetOrigin()) resampler.SetOutputDirection(itk_image.GetDirection()) if is_label: resampler.SetInterpolator(sitk.sitkNearestNeighbor) else: resampler.SetInterpolator(sitk.sitkLinear) resampler.SetDefaultPixelValue(0) return resampler.Execute(itk_image)

这段代码里两个关键参数:interpolator 和 defaultPixelValue。图像用线性插值,标注永远用最近邻插值,这是不能商量的;对 label 做线性插值会出现 0.3、0.7 这种非整数标签,交叉熵损失直接报错。defaultPixelValue 是重采样后落在原物理区域之外的体素的值,图像填 0 不影响归一化,label 填 0 表示背景,前提是你的标注里背景确实编码为 0。

new_size 计算用了 round,四舍五入后可能与原始物理区域差一个体素。如果任务对坐标对齐要求极严,更稳的做法是保留原图的 origin 和 direction,仅重采样网格,这也是上面代码的思路。

2.3 强度归一化:CT 用窗宽窗位,MRI 用 z-score

归一化直接影响网络训练的稳定性。CT 图像的体素值是亨氏单位(HU),有明确的物理含义,空气约 -1000,水约 0,骨骼可达 +1000 以上。如果直接拿原始 HU 值训练,数据范围跨度太大,模型很难收敛;但也不能简单做全局 z-score,因为 CT 扫描时背景空气占了大半个体积,全局均值和方差会被空气拉偏。

常见做法是先做窗宽窗位截断,再做 z-score。截断范围按目标组织选:肝脏、肾脏等软组织用 [-200, 200],肺实质建议 [-1200, 600],骨骼相关任务用更宽的 [-500, 1500]。截断后体素值基本集中在目标组织范围内,再算均值和方差就合理得多。

import numpy as np def ct_normalize(volume, clip_min=-200, clip_max=200): volume = np.clip(volume, clip_min, clip_max) mean = volume.mean() std = volume.std() return (volume - mean) / (std + 1e-8)

MRI 没有统一的物理量纲,同一序列不同病例的强度分布差异很大。对 MRI 我一般不做固定窗口,而是按百分位截断,比如 0.5% 到 99.5% 分位,再做 z-score。要注意的是 preprocessing 里用到的 mean 和 std 必须来自训练集并保存下来,推理阶段沿用同一组统计量,否则数据集间的数据分布会对不齐,验证指标的可靠性会打折扣。

2.4 把预处理串成一条管线:从 nii.gz 到可训练的 numpy

实际项目中每个病例要做的处理不止重采样和归一化,还可能包括裁剪背景、去除极端体素、根据 body mask 统计归一化参数等。我习惯写一个load_and_preprocess函数把读入、重采样、归一化串在一起,并且把计算好的统计量落到磁盘,方便训练时快速复用。

def load_and_preprocess(nii_path, target_spacing=(1.0, 1.0, 1.0)): image = sitk.ReadImage(nii_path) image = resample_to_spacing(image, target_spacing, is_label=False) volume = sitk.GetArrayFromImage(image).astype(np.float32) volume = ct_normalize(volume) return volume

这条管线的顺序是有讲究的:先重采样再归一化,因为重采样会改变体素数量,但每个体素的物理值不变;先归一化再切 patch 也没有问题,但如果你用的是 per-sample 的归一化,切 patch 后再归一化会让每个 patch 的均值和方差不一致,训练时的数据分布被人为改动了。我习惯在整个 volume 上完成归一化再切 patch,保证每个 patch 共享同一分布。

到这里,原始图像已经变成了规范的 float32 数组。下一步要处理的是标注文件,以及把整块 volume 切成网络能吃的 patch。

3. 标注处理与滑窗切块:从 mask 到监督信号,从整图到 patch

3.1 mask 读取、类别合并与重映射

标注文件的读取同样用 SimpleITK,但要单独处理。多类别分割任务里,标注有两种常见形态:一种是单个文件里 label 值分别为 1、2、3,另一种是每个类别一个独立的二值 mask 文件。第二种形态需要先合并,合并时注意不能简单相加,否则重叠区域会变成 2 甚至 3,正确做法是按优先级覆盖或取最大值。

def load_label(label_path, target_classes=(1, 2, 3)): label_itk = sitk.ReadImage(label_path) label_np = sitk.GetArrayFromImage(label_itk) if target_classes: remapped = np.zeros_like(label_np, dtype=np.uint8) for new_id, old_id in enumerate(target_classes, start=1): remapped[label_np == old_id] = new_id return remapped return label_np.astype(np.uint8)

这里做了一次隐式的类别重映射:把原始 label 值映射到连续的 1、2、3,背景保持 0。这一映射非常有用,尤其是标注文件里类别编号有跳号的情况,比如只有 1 和 3 却没有 2,交叉熵损失会认为 2 也是目标类,导致莫名其妙的误分割。映射后网络要预测的类别就是 {0,1,2,...,K},简单干净。

另一个容易忽略的点是 mask 和 image 是否在同一坐标系下。标注文件通常与对应的 CT 图像对齐,但不同来源的数据集之间偶尔会有 origin 或 direction 不一致的情况。处理方式很简单:读取 label 后检查 GetSpacing、GetOrigin、GetDirection 是否与 image 一致,不一致就对 label 做重采样,参照 image 的空间参数,插值器用最近邻。

3.2 滑窗切块:patch size、stride(overlap)与边界处理

整张 512×512×400 的卷直接送进 3D 网络,显存根本扛不住,所以需要滑窗切块。patch size 的选取首先受制于网络的下采样次数:假设 U-Net 做了 4 次下采样,patch 每个维度就必须能被 16 整除,经典选择是 128×128×64 对应 stride 64、64、32,保证 overlap 的同时不浪费感受野。

def extract_patches(volume, label, patch_size=(128, 128, 64), stride=(64, 64, 32)): D, H, W = volume.shape pd, ph, pw = patch_size sd, sh, sw = stride patches, labels = [], [] for z in range(0, D - pd + 1, sd): for y in range(0, H - ph + 1, sh): for x in range(0, W - pw + 1, sw): patches.append(volume[z:z + pd, y:y + ph, x:x + pw]) labels.append(label[z:z + pd, y:y + ph, x:x + pw]) return np.stack(patches)[:, None], np.stack(labels)[:, None]

这段代码返回的 shape 是 (N, 1, D, H, W),第一维是 patch 数量,第二维是通道。训练阶段我通常不设 overlap 或只设很小的 overlap,相当于用 stride 等于 patch size,这样数据增广相当于给了不同位置的采样;推理阶段则相反,stride 要设为 patch size 的一半甚至更小,确保器官边界处的体素至少落在多个 patch 里,拼回整图时不容易出现漏检。

边界部分如果 patch 滑出图像范围,做法是丢弃还是补边需要看任务。如果器官经常贴到图像边缘,我推荐镜像填充而不是补零——零填充会在边界制造一个不存在的强边缘,网络学到的边界特征在推理时会对不上。镜像填充的代码是np.pad(volume, pad_width, mode="reflect"),开销小且效果好。

3.3 前景引导采样:解决背景占比过高的问题

3D 分割里背景体素占比通常超过 95%,肝脏可能只占体数据的 5%,肿瘤只有 0.5%。纯随机采 patch 的结果是绝大多数 patch 里连一个目标体素都没有,模型训练时看到的全是背景,分割精度自然上不去。解决办法是前景引导采样:以标注中的目标体素为锚点生成 patch 中心,同时混入一部分纯背景 patch 保持模型的判别力。

def sample_patch_around_foreground(volume, label, patch_size=(128, 128, 64), fg_ratio=0.7): pd, ph, pw = patch_size D, H, W = volume.shape fg_voxels = np.argwhere(label > 0) if fg_ratio > 0 and len(fg_voxels) > 0 and np.random.rand() < fg_ratio: z, y, x = fg_voxels[np.random.randint(len(fg_voxels))] z = min(max(z - pd // 2, 0), D - pd) y = min(max(y - ph // 2, 0), H - ph) x = min(max(x - pw // 2, 0), W - pw) else: z = np.random.randint(0, max(D - pd, 1)) y = np.random.randint(0, max(H - ph, 1)) x = np.random.randint(0, max(W - pw, 1)) return volume[z:z + pd, y:y + ph, x:x + pw], label[z:z + pd, y:y + ph, x:x + pw]

fg_ratio 设置在 0.6 到 0.8 之间比较稳。比例过高会让模型对背景区域的判断变弱,推理时容易把背景误判成目标;比例过低又采不到足够的前景。实际工程里更细致的做法是先计算前景连通域,以每个连通域的质心或随机内部点为锚点,再按连通域体积加权采样,避免只盯着最大的那一个器官。

这里我再说一句:训练时的 patch 采样建议在 Dataset 的__getitem__里在线做,而不是像上一节那样离线把整张卷的所有 patch 全部切好存盘。离线切会放大几十倍的存储,而且无法在线做数据增强。在线采样每次随机取一个位置,相当于天然带来了无限的位置增广。

4. 基于 PyTorch 的 Dataset 与 DataLoader:数据准备代码的核心思路

4.1 自定义 Dataset 的三个必须方法

PyTorch 的数据准备最终落到自定义 Dataset 类上。核心思路只有三个方法:__init__只存文件列表和配置参数,__len__返回样本数,__getitem__在每次被调用时加载一个样本并返回张量。最容易犯的错误是在__init__里把全量数据读进内存,一次 50 个病例还能扛,几百个病例就直接内存爆炸。

import torch from torch.utils.data import Dataset class Seg3DDataset(Dataset): def __init__(self, file_list, patch_size=(128, 128, 64), fg_ratio=0.7): self.file_list = file_list self.patch_size = patch_size self.fg_ratio = fg_ratio def __len__(self): return len(self.file_list) def __getitem__(self, idx): volume = load_and_preprocess(self.file_list[idx]) label = load_label(self.file_list[idx].replace("_ct.nii.gz", "_label.nii.gz")) volume_patch, label_patch = sample_patch_around_foreground( volume, label, self.patch_size, self.fg_ratio ) image_tensor = torch.from_numpy(volume_patch).float().unsqueeze(0) label_tensor = torch.from_numpy(label_patch).long() return image_tensor, label_tensor

这段代码里有个细节:unsqueeze(0)加的维度是通道维,3D 分割网络的输入约定是 (batch, channel, depth, height, width),单样本返回 (1, D, H, W) 让 DataLoader 自动堆成 (N, 1, D, H, W)。label 用.long()是因为交叉熵损失要求 target 是整型张量,dtype 必须是 int64。

如果任务需要多通道输入,比如同时输入 CT 和 MRI,或者 CT 加上一个先验概率图,就在load_and_preprocess里把多个模态沿通道维拼接,__getitem__里unsqueeze(0)改成一个循环拼接。多模态数据的对齐是另一大门类,核心原则仍然是在同一坐标网格下重采样。

4.2 3D 数据增强:随机翻转、旋转和弹性形变的实现顺序

3D 数据增强和 2D 最大的区别在于:几何变换必须对 image 和 label 用同一套随机参数,稍有不对称,标注就错位了。增强的顺序也有讲究,先做几何变换,再做强度变换,最后转 tensor。下面的代码实现了两个最常用的几何增强。

import random import numpy as np class RandomFlip3D: def __init__(self, axis=0, prob=0.5): self.axis = axis self.prob = prob def __call__(self, volume, label): if random.random() < self.prob: volume = np.flip(volume, axis=self.axis) label = np.flip(label, axis=self.axis) return np.ascontiguousarray(volume), np.ascontiguousarray(label) class RandomRotate3D: def __call__(self, volume, label): k = random.randint(0, 3) if k: volume = np.rot90(volume, k, axes=(1, 2)) label = np.rot90(label, k, axes=(1, 2)) return np.ascontiguousarray(volume), np.ascontiguousarray(label)

旋转这里只用了 90 度的整数倍,原因是np.rot90不会产生任何插值误差,label 的类别值完全不会被破坏。如果要做小角度旋转比如 ±10 度,必须用scipy.ndimage.rotate,对 image 用 order=1 的线性插值,对 label 必须 order=0 最近邻。我在任务里一般只用翻转和 90 度旋转,医学图像里器官的朝向有解剖学意义,过度旋转反而会干扰模型。

弹性形变是医学分割的标志性增强,但 3D 实现开销不小。常见做法是用 scipy 的map_coordinates配合高斯滤波生成平滑变形场,alpha 控制形变幅度,sigma 控制平滑程度。alpha=3、sigma=0.5 是轻柔的形变,适合肿瘤这类结构;大器官可以适当加大 alpha,但要防止解剖结构扭曲到不可识别。注意实现的时候 label 必须用 order=0,然后 round 回整数再转回 uint8。

4.3 DataLoader 参数配置:num_workers、pin_memory 与 prefetch_factor

Dataset 写完后,DataLoader 的参数直接决定训练时数据供给是否跟得上 GPU。3D 数据一个 nii.gz 解压后有 100 MB 到几百 MB,IO 和解析的开销远大于 2D,worker 太少会让 GPU 饿着等数据,worker 太多内存会先爆。

from torch.utils.data import DataLoader train_loader = DataLoader( train_dataset, batch_size=8, shuffle=True, num_workers=4, pin_memory=True, prefetch_factor=2, persistent_workers=True, )

我一般从 num_workers=4 起步,然后观察训练时的 CPU 占用和内存涨幅。每个 worker 会预取 prefetch_factor 个 batch,3D patch 本身不大,但预取的 nii 文件经过 preprocess 后是一整个 volume,内存放大效应明显。如果内存占用超过物理内存的 60%,先把 prefetch_factor 调成 1,再考虑降 num_workers。pin_memory=True 在 GPU 训练时几乎总是值得开的,它把数据拷到页锁定内存,减少 GPU 拷贝时间。

persistent_workers=True让 worker 在每一轮 epoch 结束后不销毁,省掉反复启动的开销,前提是 Dataset 里的状态不依赖 epoch,比如我们的随机采样状态不保存在 Dataset 里,就可以放心开。collate_fn 这里不需要自定义,因为所有 patch 的 shape 都一致,DataLoader 默认的堆叠行为就能正常工作。

5. 避坑汇总:3D 图像分割数据准备的六个典型翻车现场

5.1 翻车一:mask 和 image 错位,模型学了鬼影

现象:训练时把 image 和 label 叠加可视化,发现标注轮廓与组织边缘对不上,整体平移或者旋转了一个角度。原因:image 和 label 文件来自不同处理流程,origin、spacing、direction 三者至少有一个不一致,或者读取后用错了轴序。解决:先打印两者的元数据对比,代码里加一个断言,shape 和 spacing 不等直接抛异常。重采样 label 时用 image 作为 reference image,保证两者落在同一个物理网格上。

5.2 翻车二:重采样后标签变成非整数,损失函数当场报错

现象:交叉熵损失报错,提示 target 的 dtype 不是 long,或者 Dice loss 出现 NaN。原因:对 label 用了线性插值,0 和 1 之间插出了 0.6;偶尔是增强代码里把 label 转成 float 后忘了 round。解决:label 的所有几何变换一律用最近邻插值,重采样用 sitkNearestNeighbor,旋转用 order=0。转 numpy 后可以加一句label = np.round(label).astype(np.uint8)兜底。

5.3 翻车三:num_workers 开太高,训练直接内存爆炸

现象:num_workers 设成 12,训练刚跑一个 step,内存占用一路冲高,机器卡死甚至被系统杀掉。原因:每个 worker 都加载并缓存了一整批 nii 文件的解析结果,12 个 worker 同时预取,内存成倍放大。解决:先把 num_workers 降到 4,prefetch_factor 降到 1,观察内存稳定后再逐步加。如果 Dataset 内部有缓存 dict,一定要控制缓存上限,否则多个 worker 之间还会互相叠加,这也是 3D 分割比 2D 更容易吃内存的原因。

5.4 翻车四:全局归一化把软组织对比度压没了

现象:CT 归一化后的切片整体灰蒙蒙的,肝脏和周围组织的边界看不清,训练时 loss 降得很慢。原因:整幅 CT 超过一半体素是背景空气,这些 -1024 的 HU 值把全局标准差拉得很大,软组织的小差异被稀释了。解决:先做窗宽窗位截断,再算统计量。截断到 [-200, 200] 后背景空气变成 -200 的常数,均值和方差才能真正反映目标组织的分布。MRI 同理,先按百分位截断再 z-score。

5.5 翻车五:增强引入的错误标签语义

现象:加了旋转增强后训练集 loss 正常,但验证集效果忽高忽低,特别不稳定。原因:小角度旋转产生的黑色填充区域被当成背景类别,但填充区域在语义上既不是真正的空域,旋转后的器官边界与黑色填充之间也没有明确界限,模型学到了错误的边界特征。解决:旋转时填充模式用 nearest 而不是 constant,或者限制旋转角度在 ±10 度以内,并在增强后对 label 做一次腐蚀膨胀,把边缘噪声清理干净。弹性形变同样要限制幅度,sigma 太小时形变场不光滑,器官结构会被拉成一个怪异的形状。

5.6 翻车六:滑窗推理拼回整图时边界出现块状伪影

现象:推理阶段把 patch 预测结果拼回整幅 3D 卷,器官边界处出现一道道明显的接缝,尤其肿瘤这类小目标,patch 交界处的分割结果断断续续。原因:推理时 stride 太小导致同一个体素被多个 patch 预测,而代码里直接取了某一个 patch 的输出,没有做重叠区域的融合。解决:拼接时对重叠区域取平均值,更讲究的做法是加权融合——patch 中心的预测概率置信度高,边缘低,用一个与 patch 同尺寸的高斯权重做加权平均。

def stitch_with_overlap(pred_patches, positions, volume_shape, patch_size): pd, ph, pw = patch_size prob_map = np.zeros((num_classes, *volume_shape)) weight_map = np.zeros(volume_shape) for pred, (z, y, x) in zip(pred_patches, positions): prob_map[:, z:z + pd, y:y + ph, x:x + pw] += pred weight_map[z:z + pd, y:y + ph, x:x + pw] += 1 return prob_map / np.maximum(weight_map, 1)

6. 数据管线写完后怎么验证:三个让代码思路落地的手段

数据准备写完了,不要急着全量训练。先用三个手段验证管线,确认数据本身没有静默错误,再让模型进场。第一个手段是可视化叠加检查。把 image 和 label 沿轴向每隔几十个切片输出一张叠加图,人眼扫一遍就能发现错位、轴序颠倒、标注空洞这类大问题。

import matplotlib.pyplot as plt for z in range(0, label.shape[0], 20): fig, axes = plt.subplots(1, 2, figsize=(10, 5)) axes[0].imshow(volume[z], cmap="gray") axes[1].imshow(label[z]) plt.savefig(f"check_layer_{z:04d}.png", dpi=150) plt.close()

第二个手段是检查一个 batch 的结构。跑next(iter(train_loader)),打印张量 shape、dtype、min、max,以及 label 里的类别。image 应该落在 0 附近,label 应该只在 {0,1,2,...,K} 内取值。这一步能抓出数据类型错误、归一化失效、类别映射错误。

第三个手段是单样本过拟合测试。只取几个 patch 放到一个小网络里训练 50 个 step,如果 loss 能一路降到接近 0,说明梯度和数据链路都是通的;如果 loss 纹丝不动,问题大概率出在数据端,比如 label 全是背景、patch 里根本没有前景、或者归一化把输入变成了常量。我自己每拿到一个新数据集,都会先用这三个手段验证一遍,尤其单样本过拟合,一次能省下好几天排查时间。数据准备做到这个程度,后面模型训练和调参才不会被看不见的数据问题反复打断。希望这套思路和踩坑记录能帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询