简介:面向有一定PyTorch基础、希望上手3D医学图像分割的开发者,这套以Luna16 CT肺结节数据为案例的实战代码包,完整覆盖基于UNet3d与VNet3d两种模型的训练、验证、测试、评估、可视化与后处理全过程,重点讲解数据准备与代码思路,适合算法工程师和研究生作为入门参考。压缩包共92个文件、61.66MB,其中49个Python脚本按预处理、数据集、模型、推理、后处理等模块组织,覆盖原始数据重采样、结节标注生成、patch采样、模型搭建与推理评估等关键环节;另有16个pyc编译文件、7个npy数组、5张png训练曲线图、3个nii医学影像、xml工程配置及csv标注信息等辅助资料,帮助理解各文件的用途。已有543人学习,借助该案例可系统掌握3D分割任务从数据准备到后处理的完整流程,包括UNet3d/VNet3d网络设计、损失函数与Dice评估、可视化调参与结果后处理等具体技巧,有效降低复现门槛。
1. 3D图像分割为什么先卡在数据上
跑过 2D 分割的人第一次接触基于 Pytorch 的 3D 图像分割任务,十有八九不是先调网络,而是先断在数据上。一张自然图像是 3 通道的 H×W,一块 CT 或 MRI 体数据是单通道的 D×H×W,单个样本从几 MB 到上百 MB,显存根本撑不起整图直接进网络。于是「怎么切 patch、怎么读文件、怎么对齐标签、增强写在哪个环节」全成了绕不开的事,这些统称数据准备,恰恰是决定 3D 分割项目能不能跑通的第一步。这篇文章写给两类人:刚把 Pytorch 环境搭好、准备第一个 3D 分割实战的初学者,以及被 NIfTI 读取、随机裁剪、多进程加载反复搞崩的老手。我按自己做医学影像分割的习惯,把从原始文件到 (1, D, H, W) 张量的整个过程拆开讲,重点落在代码思路和参数选择上。
2. 数据准备第一步:把NIfTI变成可训练张量的三件事
3D 分割的数据准备不是简单调一个torchvision.transforms就完事了。医学影像常见的 NIfTI、NRRD、MHD 格式自带空间信息,图像数组的排列顺序和放射学坐标不一定一致,间距也千差万别。我一般会把「格式解析、重采样、归一化」当成三个独立步骤来做,每一步单独验证,而不是写一个大而全的预处理脚本一路跑到底。
2.1 认识NIfTI文件:affine矩阵与方向对齐的坑
NIfTI 文件(.nii / .nii.gz)的核心不是像素数组,而是那个 4×4 的 affine 矩阵。它负责把体素坐标映射到解剖坐标,也决定了你在numpy里看到的数组第 0 维到底是左右方向还是前后方向。用nibabel加载时,img.shape返回的是 (D, H, W),其中 D 是层数,H 和 W 是平面内尺寸,这和 2D 图像 (H, W) 的直觉完全不同,也是很多人第一次看到数据「躺倒」的原因。
我踩过的坑是混合使用 SimpleITK 和 nibabel。同一批数据,SimpleITK 的GetArrayFromImage读出来是 (D, H, W),nibabel 的get_fdata读出来顺序一样,但方向矩阵的处理逻辑不同,一旦混用,图像和标签就会出现镜像或旋转错位。常见做法是建立项目后统一走 nibabel,只在重采样这种 SimpleITK 更顺手的环节用它,读进来后立刻把数组np.ascontiguousarray转一次,避免后续切片时出现奇怪的 strides。方向对齐这个问题没有玄学,就是「一个项目只用一套读取约定」,并且在进入 Dataset 前打印一次 affine 和 shape 留档。
2.2 体素间距重采样:为什么0.5mm和1.0mm不能混在一个batch
医院的 CT 扫描间距不统一,同一个数据集里可能有 0.5mm、0.75mm、1.0mm 三种体素间距。如果不重采样,同一个网络在物理世界里看到的感受野就不一样:间距 0.5mm 的数据里,一个 7×7×7 的卷积核覆盖 3.5mm 的组织;间距 1.0mm 的数据里,同样大小的卷积核覆盖 7mm 组织。混在一个 batch 里,BatchNorm 的统计量会被拉偏,最后模型在验证集上忽高忽低。这就是 3D 分割数据准备里最容易被忽略的「数据口径不一致」。
我一般把目标间距设在 1.0mm 或 1.5mm,具体看 GPU 内存和任务需求。重采样用 SimpleITK 的ResampleImageFilter,图像和标签用同一套目标间距,但插值方式不同:图像用线性插值,标签用最近邻插值,否则器官边缘会插出不属于任何类别的中间值。关键参数如下:
| 参数 | 建议值 | 说明 |
|---|---|---|
target_spacing | (1.0, 1.0, 1.0) | 兼顾细节与显存;病灶小就设 0.5 |
interpolator | 图像 Linear / 标签 NearestNeighbor | 标签插值必须最近邻 |
output_size | 按原尺寸和新间距反算 | 不要让 SimpleITK 自动给尺寸 |
default_pixel_value | 0 | 标签背景为 0,图像越界填 0 |
重采样代码并不复杂,核心是算好output_size:new_size = ceil(original_size * original_spacing / target_spacing)。这里有个值得注意的边界:如果数据是各向异性严重的核磁,比如层厚 3mm、平面内 0.5mm,硬压到 1mm 会丢失层间细节。我通常按层厚的 3 倍作为容忍上限,超过就降级到 1.5mm 而不是 1.0mm。
2.3 归一化与窗位:CT和MRI的数据口径不一样
归一化是 3D 分割最不缺争议的环节,场景不同,方法完全不同。CT 的 CT 值(Hounsfield Unit)有物理含义,空气是 -1000,水是 0,软组织在 -100 到 100 之间,所以业界通行做法是加窗:设定窗位(window center)和窗宽(window width),把关心范围内的灰度映射到 [0,1]。比如肺结节任务常设窗位 -600、窗宽 1500,肝部任务常设窗位 40、窗宽 250。超出窗范围的值被截断,这比全局 min-max 归一化稳定得多,因为 CT 里的高密度金属伪影会把全局最大值拉到几千,直接压缩软组织对比度。
MRI 没有绝对密度单位,同一个序列不同扫描仪之间的数值范围差异很大。我一般做百分位归一化:先算 2% 和 98% 分位数,把x = (x - p2) / (p98 - p2)再裁剪到 [0,1]。少用 z-score,因为 3D 数据往往只有稀疏前景,全局均值和标准差会被大面积背景带偏。这套逻辑必须写进预处理脚本,而不是在 Dataset 里临时算,否则每个 epoch 都重算一遍分位数,纯属浪费算力。把归一化参数写成字典,随数据一起缓存,后面做推理时用同一组参数还原,这也是可复现性的基本要求。
3. 写Dataset类之前先想清楚:懒加载、缓存与返回结构
很多人的 3D 分割项目死不是死在网络结构上,而是死在 Dataset 写得像一锅粥。Pytorch 的 Dataset 不只是一个「返回样本的类」,它是数据管线的核心调度单元。我写 Dataset 之前习惯先回答四个问题:什么时候真正读磁盘、一次往内存放多少、返回什么结构、增强在哪里做。想清楚了再写代码,基本一次跑通;不想清楚,后面调参时会被各种奇怪问题拖一个星期。
3.1 懒加载与元数据缓存:为什么__init__里不要调get_fdata
最快搞崩 3D 分割训练的做法,就是在__init__里遍历所有 NIfTI 文件并调用get_fdata()把体数据全部读进内存。一个 512×512×300 的 uint16 数据大约 150MB,解压成 float32 再乘个 3 倍,100 个训练样本就是 45GB,笔记本直接不死也卡。但完全不读又会面临训练时第一轮很慢的问题。常见做法是「懒加载 + 元数据预缓存」两个层级分开:__init__只扫描文件路径,读取每个文件的shape、affine和像素间距这种轻量头信息,缓存成一个 Python list;真正读取体数据的操作推迟到__getitem__,并用一个 LRU 缓存把最近用过的样本留在内存里,命中时跳过磁盘 I/O。
我一般这样设计:nib.load(path)返回的是一个懒加载对象,并不会立刻读盘;get_fdata()才真正解压。所以__init__里可以放心nib.load拿 head 信息,体数据留给get_fdata。如果数据集不大,干脆手动做一个缓存字典,key 是文件路径,value 是已经转成 float32 的数组。这样内存开销可控,训练脚本也不会因为频繁读盘而每 epoch 慢三五倍。
3.2 返回dict还是tuple:多模态、标签对齐与内存拷贝
单模态分割任务,__getitem__返回(volume, label)的 tuple 完全没有问题。但只要涉及多模态输入,比如 PET-CT 联合分割,我强烈建议返回 dict,键设计成{"image": ..., "label": ..., "path": ..., "index": ...}。理由很简单:tuple 模式下模态一多,位置顺序稍微写错一次,模型就把 CT 当 PET 读了,而且很难查。dict 虽然每次多一点开销,但在代码可读性和排错效率上非常值。
还有一个总被忽略的点:nibabel 通过dataobj拿到的底层数据是内存映射,只读且非连续。如果直接对这个数组做torch.from_numpy,轻则张量不可写,重则 Pytorch 的某些算子因为非连续内存直接崩。所以从get_fdata()拿到数据后,我一般紧接着做np.ascontiguousarray(arr),再考虑转 tensor。这一行代码几乎不值钱,但能省掉后面一连串「为什么 loss 是 nan」「为什么 DataLoader 报错」的排查时间。对了,标签读进来之后立刻astype(np.int64)或者astype(np.uint8),别在 loss 函数里截断。
3.3 显存口径:先算patch尺寸再写训练循环
3D 分割和 2D 最大的不同是「一个 patch 就是几个 MB 到几十个 MB」。训练一个 128³ 的 patch,输入张量本身就占 8MB,中间的卷积特征图按通道数再膨胀几倍到几十倍,所以显存口径必须在写训练循环之前算清楚。经验公式是:单卡可用显存除以 2,这个结果大概是能同时装下的输入张量大小(forward/backward 中间量按同等体积算)。比如 24GB 显存,留出 2GB 给框架和优化器,剩下 22GB,除以 2 后约 11GB 可用给输入,那么 128³ float32 的 patch 是 8MB,一个 batch 设为 8 就在合理范围;想用 256³ 的大 patch,一次就只能放 1 个,梯度累积到 8 步等效 batch。这不是精确模型,但能帮你避免「代码写完了,一启动就 OOM」的尴尬。
实际项目里我更倾向先定 patch 再定 batch:把 patch 尺寸定在「能覆盖目标器官且可以被 8 整除」的前提下,再根据显存回推 batch。常见组合我列在下面:
| 显存 | Patch 尺寸 | Batch 大小 | 备注 |
|---|---|---|---|
| 12GB | 96³ | 4 | 适合小器官 |
| 24GB | 128³ | 8 | 通用肺/肝/脑 |
| 40GB+ | 160³ | 6-8 | 需要配合梯度累积 |
| 多卡 | 128³ | 8×N | 用DistributedDataParallel |
3.4 数据增强的归属:放进Dataset还是交给外部回调
3D 分割的增强必须写在 Dataset 内部,直接在__getitem__返回前执行,而不是像 2D 任务那样放在训练循环外部。原因在于 3D 图像和标签必须共享同一套几何变换参数:如果先翻转图像、再单独翻转标签,随机数没对上,标签就错位了。放在 Dataset 里可以让「图像-标签」组合在一个函数内完成变换,天然保证一致性。
有人会把增强挪到 DataLoader 外面,用transform回调钩子实现,但这会让数据管线分叉,调试时很难判断某个样本到底有没有做过翻转。我一般把增强拆成三个纯函数:random_flip_3d(volume, label)、random_rotate_90(volume, label)、random_shift(volume, label),它们都接收一堆 numpy 数组,返回变换后的数组对。__getitem__里按一定概率依次调用。这样增强代码可单测,也方便以后替换成 MONAI 的增强套件。不过要提醒一句,3D 的旋转和 2D 不一样,任意角度的旋转会引入插值,对标签的最近邻插值处理稍有不慎边界就出现伪标签,所以我的默认策略是「翻转 + 90°旋转 + 随机裁剪 + 轻微噪声」,先不碰大角度旋转,等基线跑通了再增量加。
4. 基于Pytorch的Dataset实现:从文件路径到(1, D, H, W)张量的完整代码
这一章直接把上一章的决策落成代码。我不会用 torchio 这类封装库替你把活干完,而是用 nibabel 加 numpy 手写一个最小可跑的 3D Dataset,这样你才能理解每一步在干什么,换到自己的数据格式时也知道改哪里。
4.1 文件扫描与元数据缓存:先有清单再写迭代器
第一步是构建文件清单。假设目录结构是这样的:images/放 NIfTI 图像,masks/放同名标签文件。我习惯写一个独立的scan_cases函数,返回一个 list,每个元素是(image_path, label_path)的元组,同时断言两个文件的形状信息一致。注意,断言 shape 一致这一步要放在扫描阶段,不是训练阶段。
import pathlib import nibabel as nib def scan_cases(image_dir, label_dir, pattern="*.nii.gz"): image_dir = pathlib.Path(image_dir) label_dir = pathlib.Path(label_dir) image_paths = sorted(image_dir.glob(pattern)) label_paths = [] for img_path in image_paths: # 假设标签文件名与图像文件名相同,仅目录不同 label_path = label_dir / img_path.name if not label_path.exists(): raise FileNotFoundError(f"missing label: {label_path}") label_paths.append(label_path) # 顺手核对每个样本的 shape,避免训练到一半才炸 meta = [] for img_path, label_path in zip(image_paths, label_paths): img = nib.load(str(img_path)) lab = nib.load(str(label_path)) if img.shape != lab.shape: raise ValueError(f"shape mismatch: {img_path} {img.shape} vs {lab.shape}") meta.append({ "image_path": str(img_path), "label_path": str(label_path), "shape": img.shape[:3], "spacing": img.header.get_zooms()[:3], }) return meta这段代码的逻辑说明:nib.load只读取文件头和元数据,不加载体素数据,所以 100 个样本的扫描过程非常快。get_zooms()返回每个维度的体素间距,放进 meta 是为了后面判断是否需要重采样。这里返回的 meta 列表可以直接用作 Dataset 初始化的输入参数,而不是让 Dataset 自己再去解析目录。这样数据来源变更时,只需要替换scan_cases一个函数。
4.2 Dataset核心实现:随机裁剪、统一增强与边界断言
下面是 Dataset 的核心。三个要点:__init__只存 meta,不读体数据;_load_volume负责真正读盘并做类型转换;__getitem__先随机裁剪再增强,最后转成张量。所有随机操作都用同一个 numpy RandomState 实例,保证图像和标签的随机行为绑定在一起。
import numpy as np import torch from torch.utils.data import Dataset import nibabel as nib class VolumeSegDataset(Dataset): def __init__(self, meta, patch_size=(96, 96, 96), normalize="percentile", seed=0): self.meta = meta self.patch_size = patch_size self.normalize = normalize self.rng = np.random.RandomState(seed) self._cache = {} def __len__(self): return len(self.meta) def _load_volume(self, idx): # 简单的样本缓存:同一个样本只在第一次真正读盘 if idx in self._cache: return self._cache[idx] image_path = self.meta[idx]["image_path"] label_path = self.meta[idx]["label_path"] image = nib.load(image_path).get_fdata().astype(np.float32) label = nib.load(label_path).get_fdata().astype(np.int64) # 标签去重,防止遇到 0/255 这种输入 label = (label > 0).astype(np.int64) self._cache[idx] = (image, label) return image, label def _random_crop(self, image, label): d, h, w = image.shape pd, ph, pw = self.patch_size if d < pd or h < ph or w < pw: raise ValueError( f"patch {self.patch_size} larger than volume {image.shape}" ) # 随机起点,图像和标签共用同一组偏移 start_d = self.rng.randint(0, d - pd + 1) start_h = self.rng.randint(0, h - ph + 1) start_w = self.rng.randint(0, w - pw + 1) image_crop = image[start_d:start_d + pd, start_h:start_h + ph, start_w:start_w + pw] label_crop = label[start_d:start_d + pd, start_h:start_h + ph, start_w:start_w + pw] return image_crop, label_crop def _augment(self, image, label): # 翻转:图像和标签同时翻,方向共用 if self.rng.rand() < 0.5: image = np.flip(image, axis=0).copy() label = np.flip(label, axis=0).copy() if self.rng.rand() < 0.5: image = np.flip(image, axis=1).copy() label = np.flip(label, axis=1).copy() return image, label def _normalize(self, image): if self.normalize == "percentile": lo, hi = np.percentile(image, [2, 98]) else: lo, hi = image.min(), image.max() image = (image - lo) / (hi - lo + 1e-8) return np.clip(image, 0.0, 1.0) def __getitem__(self, idx): image, label = self._load_volume(idx) image, label = self._random_crop(image, label) image, label = self._augment(image, label) image = self._normalize(image) # unsqueeze(0) 得到 (1, D, H, W) volume_tensor = torch.from_numpy(image).unsqueeze(0).float() label_tensor = torch.from_numpy(label).long() return { "volume": volume_tensor, "label": label_tensor, "path": self.meta[idx]["image_path"], }代码逻辑说明:_load_volume里的idx in self._cache实现了一个简单的样本缓存。如果训练集只有几十个样本,缓存能省掉反复解压的开销;如果数据集大,把self._cache替换成functools.lru_cache并限定 maxsize 即可。_random_crop的关键是三个随机起点只生成一次,图像和标签切片共用,这是标签对齐的根基。_augment里翻转操作要加.copy(),因为np.flip返回的是视图,后续torch.from_numpy可能会遇到非连续内存的警告。
参数说明:patch_size我默认 96³,这是 24GB 显存卡的稳妥选择。normalize支持 percentile 和 minmax 两种,CT 场景建议改成窗宽窗位,直接在_normalize里加一个 branch 即可。seed参数听起来没用,但其实决定了你每次实验的随机裁剪分布,建议固定下来,方便复现实验。label = (label > 0).astype(np.int64)这行最容易被嫌弃,但它是防御多类别标签被误存成 0/255 的有效手段;如果任务是多类别分割,把它改成np.clip(label, 0, n_classes - 1)。
4.3 collate_fn与DataLoader:多进程参数怎么配才不白开
Dataset 写好之后,DataLoader 的配置同样有门道。3D 数据的单样本体积大,num_workers开少了加载跟不上 GPU,开多了又会内存暴涨。我的默认配置是num_workers=4配合persistent_workers=True和prefetch_factor=4。persistent_workers=True让 worker 进程在多个 epoch 之间存活,不重复 fork;prefetch_factor=4意味着每个 worker 预取 4 个 batch,够 GPU 忙一阵子。内存不够时把prefetch_factor降到 2,优先保护住缓存。
from torch.utils.data import DataLoader def collate_3d(batch): volumes = torch.stack([item["volume"] for item in batch]) labels = torch.stack([item["label"] for item in batch]) paths = [item["path"] for item in batch] return volumes, labels, paths loader = DataLoader( dataset, batch_size=4, shuffle=True, num_workers=4, pin_memory=True, persistent_workers=True, prefetch_factor=4, collate_fn=collate_3d, ) # 训练循环里直接解包 for volumes, labels, paths in loader: volumes = volumes.cuda(non_blocking=True) labels = labels.cuda(non_blocking=True)collate_3d的作用是把一个 batch 里的字典整理成三个独立元素。这里必须用torch.stack,因为每个样本已经带了unsqueeze(0)的通道维,stack 后会得到 (B, 1, D, H, W) 的张量。pin_memory=True配合non_blocking=True能减少 CPU 到 GPU 的搬运时间,在 3D 任务上收益比 2D 更明显,因为单个张量大。注意:如果某个 batch 里样本尺寸不齐,stack 会报错,所以在 Dataset 的_random_crop里保证裁剪后尺寸恒定是硬要求。
5. 数据准备避坑指南:5个最常见的翻车现场
这一章写的是我真实踩过的坑,按「现象 → 原因 → 解决」的流水来。前面章节的代码思路是理想路径,这里补上现实世界的补丁。
5.1 图像和标签方向不一致:NIfTI方向读法不统一的代价
现象:训练 loss 能正常下降,但验证集 Dice 始终在 0.1 左右徘徊,可视化 label 叠在 image 上,发现标签的边缘永远偏了几个像素或直接镜像。 原因:数据预处理阶段有人用 SimpleITK 读了一批数据,有人用 nibabel 读了另一批,两批数据的体素顺序在方向矩阵上不统一。NIfTI 文件本身带 qform/sform,两个库对它的解释侧重点不同,直接按数组索引对齐,就会得到空间错位的 pair。 解决:全项目统一用 nibabel,读入后打印一次img.affine和img.shape留档。如果发现数据集本身存在方向不一致,用nibabel.probtensor或nibabel.affines把标签的 affine 与图像对齐,核心代码是label = nib.Nifti1Image(label_data, img.affine, img.header)。最彻底的方案是在预处理阶段重采样成固定方向后再存盘,让进入网络的每一份数据方向一致。
5.2 标签是0和255不是0和1:loss爆炸的常见前奏
现象:训练第一个 batch 时 loss 是几十甚至上百,随后出现 nan,或者 DICE Loss 报除零错误。 原因:标注工具导出的标签往往不是二值掩码,而是 0/255 的 uint8 图像,有些甚至存成了 0/1/2/255 的稀疏格式。直接用CrossEntropyLoss处理,255 会被当成一个合法类别,模型为了拟合它学出一堆无效特征。 解决:在_load_volume里加一道标签清洗逻辑。我的习惯是先打印np.unique(label),确认类别数;二值任务就label = (label > 0).astype(np.int64),多类别任务就做label = np.clip(label, 0, num_classes - 1)。这个步骤必须放在缓存之前,避免把脏标签缓存起来。
5.3 随机裁剪抽不到前景:用质心引导采样自救
现象:训练了 20 个 epoch,模型输出全是背景,Dice 为 0,验证集上预测结果只有黑块。 原因:3D 体数据中目标器官可能只占整体体积的 2% 以下,纯随机裁剪 96³ patch,大概率裁到全是背景的区域,模型根本没有见过前景样本,梯度被背景主导。 解决:先统计 foreground ratio,如果小于 5%,改用质心引导采样。做法是找到标签中前景体素的质心,在该质心附近随机偏移一定范围作为裁剪起点,保证 patch 里总有前景。下面的代码片段可以直接替换_random_crop:
def _centroid_crop(self, image, label): d, h, w = image.shape pd, ph, pw = self.patch_size coords = np.argwhere(label > 0) if len(coords) == 0: return self._random_crop(image, label) centroid = coords.mean(axis=0).astype(int) # 在质心周围 32 个体素范围内随机偏移 offset = self.rng.randint(-32, 33, size=3) start_d = centroid[0] + offset[0] - pd // 2 start_h = centroid[1] + offset[1] - ph // 2 start_w = centroid[2] + offset[2] - pw // 2 start_d = min(max(start_d, 0), d - pd) start_h = min(max(start_h, 0), h - ph) start_w = min(max(start_w, 0), w - pw) image_crop = image[start_d:start_d + pd, start_h:start_h + ph, start_w:start_w + pw] label_crop = label[start_d:start_d + pd, start_h:start_h + ph, start_w:start_w + pw] return image_crop, label_crop逻辑说明:质心偏移量 32 是一个宽泛的经验值,patch 是 96³ 时,偏移 32 既能让前景出现在 patch 内,又保留一定随机性。如果病灶特别小且分散,偏移量可以缩小到 16,甚至直接用质心硬剪裁。注意min(max(start_d, 0), d - pd)这行边界处理必须写,否则质心靠近边缘时会越界。
5.4 worker越多越快吗:num_workers与内存飙升的真实关系
现象:num_workers从 4 调到 16,显存没涨,CPU 内存却从 20GB 一路涨到 90GB,训练速度反而更慢了。 原因:每个 DataLoader worker 都会把__getitem__里读到的数据保存在自己的进程内存里,3D 单样本几十 MB,16 个 worker 加上prefetch_factor的放大效应,内存翻倍是正常的。而且如果self._cache存了大量解压后的体数据,每个 worker 复制一份,内存直接爆。 解决:限制num_workers=4,prefetch_factor=4改成 2。如果训练集几十个样本且全部能放进单进程内存,可以把self._cache做成全量缓存后,直接把num_workers=0,省掉进程拷贝,速度反而稳定。我的血泪经验是 3D 分割里 worker 的边际收益衰减很快,IO 瓶颈通常不在进程数上,而在磁盘随机读和 NIfTI 解压上,与其加 worker,不如把预处理结果存成npy或zarr格式,省掉每次都解压.nii.gz的 CPU 开销。
5.5 同一次旋转没作用到标签:随机种子不一致导致错位
现象:训练过程看起来正常,但验证集可视化发现标签偶尔旋转了 90°,边界错乱,Dice 波动剧烈。 原因:增强代码里分别对图像和标签调用np.random,两个调用的随机数状态不在同一个流上;或者用了torch的随机 API 和numpy的随机 API 混用,种子状态没有同步。更隐蔽的情况是翻转操作没有.copy(),导致标签的 stride 异常,后续切片时索引错位。 解决:上文 Dataset 里的做法是核心规范——所有随机数都从self.rng这个唯一的RandomState获取,图像和标签共用同一组决策。增强函数一律写成(image, label) -> (image, label)的纯函数,不要在同一个函数里用np.random.xxx全局 API。我再在__getitem__末尾加一道断言:assert image.shape == label.shape,出错时第一时间报错,而不是训练几百步后模型不收敛。
6. 让数据准备和推理闭环:patch联动与可视化验证
数据准备做得对不对,光看 train loss 不足以说明问题。我每次跑新数据集,都先单独执行一遍__getitem__逻辑,把裁剪后的 image 和 label 叠加保存成 PNG,再开始训练。
6.1 滑窗推理的patch尺寸要和训练对齐
推理阶段的 patch 通常不等于体积整幅,常见做法是滑窗(sliding window)切 patch 预测,再把结果拼回完整体积。这里最容易犯的错是推理 patch 尺寸和训练不一致。训练用 96³,推理换 128³,网络的感受野和 padding 行为都可能变化,导致边缘出现接缝伪影。我的做法是训练定 patch 时,就把推理 patch 定成相同的值,至少保持一个维度上的连续性。滑窗时overlap设为 0.25 或 0.5,overlap 区域用线性权重叠加,或者用汉宁窗让靠近 patch 中心的体素贡献更高,这样接缝处不会出现明显的分块痕迹。一个值得记住的细节:如果最后要导出 onnx 做部署,推理输入尺寸最好就是训练 patch 尺寸,否则它不支持动态多维 shape 时还得在外面包一层 pad。
6.2 可视化与分布统计:数据准备做完先看这两样
最后分享我现在的习惯:新数据集到手,先写一个几行的检查脚本,统计每个样本的 shape、spacing、label 百分比,按增序排一遍。shape 差异过大的样本要么重采样没做,要么方向有问题;label 百分比为 0 的样本直接查,多半是标签读取路径错了。然后把第一个 patch 的中层三个切面叠加保存成图,盯着看 10 秒:图像和标签边缘是否贴合,有没有旋转错位。这两步做完,再动网络结构。这套流程看起来慢,实际比训练 50 个 epoch 后发现数据有问题再返工快得多。希望帮到你。
本文还有配套的精品资源,点击获取