
简介这份资源围绕3D U-Net在三维医学图像分割中的应用展开面向具备一定深度学习基础、希望将U-Net从二维扩展到三维的医学影像研究者与工程实践者可用于CT、MRI等体积数据的分割实验与代码复现。压缩包共17个文件约9KB以4个Python脚本为核心涵盖模型定义、训练流程与nii、yaml等工具模块另含9个xml配置、1个iml工程文件及README、requirements等说明文档便于快速配置环境并跑通训练验证。资源重点呈现三维卷积、池化与上采样构成的对称编解码结构以及数据预处理、Dice损失训练、超参调优与连通组件后处理等关键环节读者可据此理解体积分割的完整链路并在此基础上开展肿瘤检测、脑结构分割等方向的实验与改进。目前已有1149人学习下载适合作为三维医学图像分割的入门与参考项目。1. 3D U-Net 医学图像分割从体素堆里把病灶抠出来的最小闭环手里有一批 CT 或 MRI 体数据标注是三维的想直接跑通一个 3D U-Net 做医学图像分割却卡在「数据怎么喂进去、显存怎么不炸、输出怎么还原回原始空间」这三件事上——这是绝大多数人第一次碰 3DUNET 的真实处境。2D 分割那套ImageFolder DataLoader的经验搬到体数据上基本失效因为切片之间不是独立样本通道维和深度维会打架标签还常常是稀疏的。3D U-Net 的价值就在于它把下采样、上采样、跳跃连接全部搬到三维空间里做让网络能同时利用层内纹理和层间连续性对肝脏、肾脏、脑肿瘤、血管这类有体积结构的器官尤其管用。这篇不铺原理教科书按「数据准备 → 模型搭起来 → 训练跑通 → 推理还原 → 踩坑排查」的顺序把一条能复现的最小闭环讲清楚新手能照着敲熟手能对着参数和边界挑刺。2. 3D U-Net 到底在算什么体素级分类和 2D 版本的本质差别2.1 从 2D U-Net 到 3D多出来的那一维带来了什么2D U-Net 把每张切片当独立图像卷积核是3x3感受野只在平面内扩张。3D U-Net 把卷积核换成3x3x3特征图是D×H×W的体块一次卷积同时看相邻三层。这个改动看着小影响却很大层间连续性被显式建模器官边界在 Z 方向上的断裂会被平滑掉对层厚较薄1~3mm的 CT 收益明显。代价是参数量和显存按立方增长一个3x3x3卷积核的权重是3x3的 3 倍特征图显存也是 3 倍量级。所以 3D U-Net 从来不是「把 2D 代码维度改一改」就能跑patch 大小、batch size、混合精度这三件事必须一起调。常见做法是训练时不做整卷输入而是从体数据里随机裁patch比如128×128×128或96×96×96推理时再用滑窗sliding window拼回整卷。这样显存可控还能靠随机裁剪做数据增强。我一般会把 patch 的 Z 深度设成 16 的倍数因为 3D U-Net 通常下采样 4 次2^416深度不能被 16 整除时跳跃连接的两端尺寸会对不上直接报维度错误。2.2 各向异性体素3D U-Net 最容易被忽略的输入问题医学体数据几乎都是各向异性的——层内分辨率0.7mm层厚可能5mm。直接送进网络等于告诉模型「Z 方向和 XY 方向一样重要」这是错的。两种处理方式一是重采样到各向同性如全部1mm代价是插值引入模糊、数据量膨胀二是保留原始 spacing但在 patch 采样时对 Z 方向单独控制。我倾向第二种因为重采样对小结节的形状破坏不可逆。具体做法是在数据加载时记录每例的spacing推理还原时按原 spacing 把预测结果映射回去保证输出和原始影像对齐。提示如果数据集里 spacing 差异很大有的 0.6mm 有的 3mm务必在训练前统计一遍必要时按 spacing 分层采样否则模型会偏向层厚大的样本。3. 数据准备把 NIfTI 体数据变成能喂进 3D U-Net 的 patch3.1 用 nnU-Net 风格的数据指纹先摸清家底动手写模型前先花十分钟把数据统计清楚这一步能省掉后面一半的玄学问题。核心是每例的 shape、spacing、标签类别分布、强度范围。下面这段脚本读 NIfTI 并打印指纹依赖nibabel和numpy。import nibabel as nib import numpy as np from pathlib import Path def dataset_fingerprint(img_dir, lbl_dir): rows [] for img_path in sorted(Path(img_dir).glob(*.nii.gz)): case img_path.name.replace(.nii.gz, ) img nib.load(str(img_path)) lbl_path Path(lbl_dir) / f{case}.nii.gz lbl nib.load(str(lbl_path)) img_data img.get_fdata() lbl_data lbl.get_fdata() # spacing 来自 affine 的对角线单位 mm spacing nib.affines.voxel_sizes(img.affine) # 标签类别去掉背景 0 classes np.unique(lbl_data) classes classes[classes 0] rows.append({ case: case, shape: img_data.shape, spacing: tuple(round(s, 2) for s in spacing), intensity: (round(float(img_data.min()), 1), round(float(img_data.max()), 1)), classes: classes.tolist(), fg_ratio: round(float((lbl_data 0).mean()), 4), }) return rows for r in dataset_fingerprint(./imagesTr, ./labelsTr): print(r)逻辑说明nib.affines.voxel_sizes从 affine 矩阵直接算体素物理尺寸比手动读 header 稳。fg_ratio是前景体素占比这个值低于 1% 时普通 Dice loss 会被背景淹没需要换损失或加采样权重。参数上img.get_fdata()返回 float64大体积数据建议改get_fdata(dtypenp.float32)省一半内存。3.2 patch 采样与前景过采样别让模型只学背景体数据里病灶往往只占几个百分点随机裁 patch 大概率裁到纯背景。常见做法是前景过采样以一定概率强制 patch 中心落在前景体素上。下面是一个可复用的 patch 采样器。import numpy as np def random_patch(img, lbl, patch_size(96, 96, 96), fg_prob0.7): d, h, w img.shape pd, ph, pw patch_size # 保证 patch 不越界 assert d pd and h ph and w pw, patch 比体积还大 if np.random.rand() fg_prob: # 前景过采样随机挑一个前景体素作为中心 fg_idx np.argwhere(lbl 0) cz, cy, cx fg_idx[np.random.randint(len(fg_idx))] else: cz, cy, cx d // 2, h // 2, w // 2 # 以中心反推起点并夹到合法范围 z0 int(np.clip(cz - pd // 2, 0, d - pd)) y0 int(np.clip(cy - ph // 2, 0, h - ph)) x0 int(np.clip(cx - pw // 2, 0, w - pw)) img_p img[z0:z0pd, y0:y0ph, x0:x0pw] lbl_p lbl[z0:z0pd, y0:y0ph, x0:x0pw] return img_p, lbl_p逻辑说明fg_prob0.7表示七成 patch 以病灶为中心三成随机兼顾正负样本。np.clip处理边界避免中心靠近边缘时越界。参数上patch_size的每个维度都应是 16 的倍数对应 4 次下采样如果显存吃紧优先砍 Z 深度而不是 XY因为 XY 承载的纹理信息更多。前景体素索引np.argwhere对超大体积会慢可以预先算好每例的前景坐标缓存起来。3.3 强度归一化CT 和 MRI 要分开处理CT 的 HU 值有物理意义通常做clip(-1000, 1000)再线性映射到[0,1]MRI 没有绝对单位一般按前景体素的均值和标准差做 z-score。混用会翻车把 MRI 当 CT 做 HU 窗模型学到的全是噪声。归一化参数要按数据集统计不要拍脑袋写mean0.5。4. 搭一个能跑的 3D U-Net通道数、下采样和跳跃连接怎么定4.1 网络结构的关键参数3D U-Net 的结构本身不复杂难的是通道数和深度的取舍。下面给一份我常用的配置输入单通道输出类别数可调。层级操作输入通道输出通道特征图尺寸以 96³ 输入为例enc12×Conv3d(3³)BNReLU13296³enc2MaxPool3d(2) 2×Conv3d326448³enc3MaxPool3d(2) 2×Conv3d6412824³enc4MaxPool3d(2) 2×Conv3d12825612³bottleneckMaxPool3d(2) 2×Conv3d2565126³dec4ConvTranspose3d(2) concat 2×Conv3d51225625612³dec3ConvTranspose3d(2) concat 2×Conv3d25612812824³dec2ConvTranspose3d(2) concat 2×Conv3d128646448³dec1ConvTranspose3d(2) concat 2×Conv3d64323296³outConv3d(1³)32num_classes96³首层 32 通道是显存和表达力的平衡点。通道翻倍到 64 起步96³patch 在 16G 显存上基本跑不动 batch1。下采样 4 次是常规选择再深一层到3³的特征图空间信息损失太大对小病灶不友好。4.2 用 PyTorch 把结构写出来import torch import torch.nn as nn class DoubleConv3d(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.block nn.Sequential( nn.Conv3d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm3d(out_ch), nn.ReLU(inplaceTrue), nn.Conv3d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm3d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.block(x) class UNet3D(nn.Module): def __init__(self, in_ch1, num_classes2, base32): super().__init__() chs [base, base*2, base*4, base*8, base*16] self.enc1 DoubleConv3d(in_ch, chs[0]) self.enc2 DoubleConv3d(chs[0], chs[1]) self.enc3 DoubleConv3d(chs[1], chs[2]) self.enc4 DoubleConv3d(chs[2], chs[3]) self.pool nn.MaxPool3d(2) self.bottleneck DoubleConv3d(chs[3], chs[4]) self.up4 nn.ConvTranspose3d(chs[4], chs[3], 2, stride2) self.dec4 DoubleConv3d(chs[3]*2, chs[3]) self.up3 nn.ConvTranspose3d(chs[3], chs[2], 2, stride2) self.dec3 DoubleConv3d(chs[2]*2, chs[2]) self.up2 nn.ConvTranspose3d(chs[2], chs[1], 2, stride2) self.dec2 DoubleConv3d(chs[1]*2, chs[1]) self.up1 nn.ConvTranspose3d(chs[1], chs[0], 2, stride2) self.dec1 DoubleConv3d(chs[0]*2, chs[0]) self.out nn.Conv3d(chs[0], num_classes, 1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) e4 self.enc4(self.pool(e3)) b self.bottleneck(self.pool(e4)) d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out(d1)逻辑说明DoubleConv3d是 3D U-Net 的基本砖块两次3³卷积加 BN 和 ReLU。biasFalse是因为后面跟 BN偏置冗余。跳跃连接用torch.cat在通道维拼接这是 U-Net 的核心让解码器能拿到编码器的高分辨率细节。参数上base32控制整体宽度显存不够就降到 16但低于 16 表达力会明显下降。num_classes包含背景二分类任务设 2。4.3 损失函数Dice 和 CE 的组合怎么配医学分割里前景占比低纯交叉熵会被背景主导。常见做法是Dice Loss CrossEntropy加权求和权重0.5:0.5起步。Dice 对类别不平衡鲁棒但梯度在训练初期不稳定CE 提供稳定的逐体素梯度。两者互补。class DiceCELoss(nn.Module): def __init__(self, dice_weight0.5): super().__init__() self.dice_weight dice_weight self.ce nn.CrossEntropyLoss() def forward(self, logits, target): # logits: [B, C, D, H, W], target: [B, D, H, W] long ce_loss self.ce(logits, target) probs torch.softmax(logits, dim1) # one-hot 化 target target_oh torch.zeros_like(probs).scatter_(1, target.unsqueeze(1), 1) dims (0, 2, 3, 4) inter (probs * target_oh).sum(dims) union probs.sum(dims) target_oh.sum(dims) dice (2 * inter 1e-5) / (union 1e-5) # 跳过背景类只对前景算 dice dice_loss 1 - dice[1:].mean() return self.dice_weight * dice_loss (1 - self.dice_weight) * ce_loss逻辑说明scatter_把整型标签转 one-hot和 softmax 概率逐体素相乘。dice[1:]跳过背景类因为背景 Dice 通常接近 1算进去会稀释前景信号。1e-5是平滑项防止空 patch 除零。参数上dice_weight在 0.3~0.7 之间调前景极稀疏时调高到 0.7。5. 训练与推理显存、混合精度和滑窗拼接的实操5.1 训练循环与混合精度3D 网络显存吃紧混合精度AMP几乎是必开项能省 30%~40% 显存速度也快。下面是一个最小训练循环。from torch.cuda.amp import autocast, GradScaler def train_one_epoch(model, loader, optimizer, criterion, scaler, device): model.train() total_loss 0.0 for img, lbl in loader: img img.float().to(device) # [B, 1, D, H, W] lbl lbl.long().to(device) # [B, D, H, W] optimizer.zero_grad() with autocast(): logits model(img) loss criterion(logits, lbl) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss loss.item() return total_loss / len(loader)逻辑说明autocast让卷积在 fp16 下算BN 和 loss 仍在 fp32兼顾速度和数值稳定。GradScaler处理 fp16 梯度下溢。参数上优化器用AdamW学习率1e-3起步配合CosineAnnealingLR衰减到1e-5。batch size 在96³patch 下通常只能设 1~2这时用梯度累积模拟大 batch每 4 步optimizer.step()一次。5.2 滑窗推理与结果还原整卷推理不能直接送网络得用滑窗。窗口大小和训练 patch 一致步长取窗口的 1/2 做重叠重叠区取平均降低拼接缝。import numpy as np import torch torch.no_grad() def sliding_window_inference(model, volume, patch_size(96,96,96), overlap0.5): model.eval() d, h, w volume.shape pd, ph, pw patch_size sd, sh, sw int(pd*(1-overlap)), int(ph*(1-overlap)), int(pw*(1-overlap)) prob np.zeros((2, d, h, w), dtypenp.float32) count np.zeros((d, h, w), dtypenp.float32) 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): patch volume[z:zpd, y:yph, x:xpw] t torch.from_numpy(patch[None, None]).float() logits model(t) p torch.softmax(logits, dim1)[0].numpy() prob[:, z:zpd, y:yph, x:xpw] p count[z:zpd, y:yph, x:xpw] 1 prob / np.maximum(count[None], 1e-5) return prob.argmax(0)逻辑说明overlap0.5表示相邻窗口重叠一半重叠区概率累加后取平均能有效消除边界伪影。count记录每个体素被覆盖次数避免边缘除零。参数上步长越小越平滑但越慢0.5重叠是速度和质量的平衡点。推理完的预测是 patch 空间还要按原始 affine 写回 NIfTI保证和输入影像对齐。6. 避坑与排查3D U-Net 训练里最常见的五个翻车现场6.1 现象loss 一直不降Dice 卡在 0.1 以下原因通常是标签和图像没对齐或者标签值不是从 0 开始的连续整数。有些数据集标签用 1 表示背景、2 表示前景直接送进CrossEntropyLoss会越界或错位。解决训练前打印np.unique(lbl)确认类别是[0, 1, 2, ...]不是就做一次映射。6.2 现象训练正常推理结果全是背景多半是滑窗推理时归一化参数和训练不一致。训练用了 z-score推理忘了做输入分布偏移模型输出全塌到背景类。解决把归一化封装成一个函数训练和推理共用同一份代码别两处各写一遍。6.3 现象显存溢出报CUDA out of memory除了降 patch 和 batch还有一个隐蔽原因验证阶段没关梯度。验证时忘了torch.no_grad()中间激活全存着显存直接翻倍。解决验证和推理一律套torch.no_grad()并用model.eval()切 BN 到推理模式。6.4 现象Dice 在训练集很高验证集很低典型过拟合但 3D 数据量小的时候更隐蔽。除了加数据增强随机翻转、旋转、弹性形变还要检查验证集和训练集是不是同一批病人。同一病人的不同扫描如果分到两边等于变相泄漏。解决按病人 ID 划分数据集不是按切片或按文件。6.5 现象预测边界锯齿严重小病灶整块丢失原因可能是下采样太深小目标在6³的特征图上只剩一两个体素解码器恢复不回来。解决减少一次下采样改成 3 次或者引入深监督在中间层也加辅助 loss。另一个办法是把 patch 的 Z 深度调大让网络看到更多上下文。7. 把 Dice 从 0.7 推到 0.85后处理和深监督的两个具体技巧训练跑通只是及格线真正拉开差距的是后处理和结构微调。先说后处理3D U-Net 的输出常有孤立的小连通域这些多半是假阳性。用scipy.ndimage.label做连通域分析去掉体素数小于阈值的块Dice 通常能涨 1~3 个点。阈值按最小病灶体积定比如50体素。另一个是保留最大连通域适合单器官分割任务。from scipy import ndimage import numpy as np def remove_small_objects(mask, min_size50): labeled, num ndimage.label(mask) if num 0: return mask sizes ndimage.sum(mask, labeled, range(1, num 1)) keep np.zeros_like(mask, dtypebool) for i, s in enumerate(sizes, start1): if s min_size: keep | (labeled i) return keep.astype(mask.dtype)逻辑说明ndimage.label给每个连通域编号ndimage.sum统计每个域的体素数小于阈值的丢弃。参数min_size要按数据集的最小真实病灶来定设太大反而会误删真阳性。再说深监督。在解码器的dec3和dec2输出上各接一个1x1x1卷积上采样到原尺寸算辅助 loss权重设 0.3 和 0.2。这样梯度能更直接地传到浅层缓解深层网络梯度消失对小目标尤其有效。代价是显存和计算量增加约 15%显存紧的话只加一层深监督。我自己的习惯是每次改结构或损失只动一个变量跑完对比验证集 Dice 再决定留不留。3D 训练一轮动辄几小时同时改三处出了问题根本不知道是哪处的锅。这套最小闭环我在几批腹部 CT 上跑过从数据指纹到滑窗推理按上面的参数走单卡 16G 能稳定复现。希望帮到你。本文还有配套的精品资源点击获取