简介这份资源围绕3D U-Net在三维医学图像分割中的实现展开面向具备一定深度学习基础、希望将U-Net从二维扩展到三维的医学影像研究者与工程实践者可用于CT、MRI等体积数据的器官与病灶分割任务。压缩包共17个文件约9KB以4个Python脚本为核心涵盖模型定义、训练流程与nii、yaml等工具模块另含9个xml配置、1个iml工程文件及README、requirements等说明文档便于快速配置环境并复现训练。资源已积累1152人学习下载具备一定参考热度。内容侧重三维卷积、池化与上采样构成的编解码结构读者可据此理解体积数据的预处理、Dice损失训练、超参调整与后处理思路并借助工程目录组织方式搭建自己的分割实验流程适合作为三维医学图像分割的入门与对照实践材料。1. 从一张 512×512×128 的 CT 说起3DUNET 到底在解决什么手里有一批腹部增强 CT层厚 1 mm单例数据 512×512×128 体素。要做肝脏和肿瘤的逐体素分割。如果按 2D 思路把每一层切出来单独送进 U-Net会立刻撞上两个问题层与层之间的解剖连续性被切断肿瘤上下缘的判定全靠单层纹理边界抖动明显同时 128 张图各自推理显存和耗时都不划算。3DUNET 就是冲着这个场景来的——把 U-Net 的编码器-解码器结构整体搬到三维卷积上让卷积核在深度方向也滑动直接吃 (D, H, W) 的体数据输出同样尺寸的概率体。它适合有体数据、有逐体素标注、且对空间连续性敏感的任务肝脏、胰腺、脑肿瘤、血管、肺结节。代价也直白显存吃紧、标注昂贵、训练慢。这篇笔记按「结构怎么立住 → 数据怎么喂 → 模型怎么写 → 坑在哪 → 怎么验证」的顺序把 3DUNET 从标题落到能跑通的代码。2. 3DUNET 的结构选型为什么是三维卷积而不是逐层 2D2.1 三维卷积核到底多算了什么2D 卷积核是 k×k3D 是 k×k×k。一个 3×3×3 的核在单次前向里覆盖 27 个体素而 3×3 只覆盖 9 个。这意味着同样的层数3DUNET 的有效感受野在深度方向天然延展肿瘤的上下缘、血管的走行方向都能被编码进去。代价是参数量和计算量按 k 的倍数增长3×3×3 相对 3×3参数量约为 1.5 倍27/9 再考虑通道但实际显存占用往往涨 3 到 5 倍因为中间特征图也是三维的。我一般会先算一笔账输入 patch 取 128×128×128第一层 32 通道单精度浮点仅这一层激活就约 128³×32×4 Byte ≈ 268 MB。所以 3DUNET 落地第一件事不是写网络是决定 patch 大小和 batch size 的取舍。2.2 编码器-解码器在三维下的对应关系标准 3DUNET 沿用 U-Net 的对称结构编码器每级两次 3×3×3 卷积加 ReLU接 2×2×2 最大池化降采样解码器每级先 2×2×2 转置卷积上采样与编码器同级特征在通道维拼接再两次 3×3×3 卷积。最后一层 1×1×1 卷积把通道压到类别数。关键差异在跳跃连接2D 里拼的是同层特征图3D 里拼的是同尺寸的体块。如果输入深度不是 16 的整数倍降采样到某一级会出现奇数尺寸拼接时对不齐这是 3DUNET 最常见的翻车点之一后面避坑章会细说。2.3 选型对照3DUNET、2.5D、逐层 2D 怎么选方案输入深度信息显存标注成本适用逐层 2D U-Net单层无低低层间独立、大层厚2.5D多通道堆叠相邻 3~5 层弱中低层厚较大、算力有限3DUNET体块强高高各向同性、边界敏感判断标准很简单层厚小于 2 mm、且标注是逐体素的优先 3DUNET层厚 5 mm 以上、标注只在少数层2.5D 性价比更高。别为了用 3D 而用 3D。3. 数据准备把 NIfTI 变成能喂进 3DUNET 的 patch3.1 读取、重采样与强度归一化医学图像分割的输入通常是 NIfTI.nii/.nii.gz。第一步不是切 patch是统一体素间距。不同设备层厚不同直接切 patch 会让网络学到「层厚」这个伪特征。import nibabel as nib import numpy as np from scipy.ndimage import zoom def load_and_resample(img_path, lbl_path, target_spacing(1.0, 1.0, 1.0)): img nib.load(img_path) lbl nib.load(lbl_path) # 原始体素间距顺序对应 (x, y, z) src_spacing img.header.get_zooms()[:3] # 计算每个轴的缩放因子 factors [s / t for s, t in zip(src_spacing, target_spacing)] img_data img.get_fdata().astype(np.float32) lbl_data lbl.get_fdata().astype(np.uint8) # 图像用三线性标签必须用最近邻否则会出现不存在的类别 img_data zoom(img_data, factors, order1) lbl_data zoom(lbl_data, factors, order0) return img_data, lbl_data逻辑说明get_zooms()返回体素物理尺寸zoom的order1是线性插值适合灰度图标签是整数类别必须order0最近邻否则插值会造出 0.5 这种非法标签。参数target_spacing我一般设 (1.0, 1.0, 1.0)如果显存紧张可以放宽到 (1.5, 1.5, 1.5)。归一化用 z-score按前景非零区域统计均值和标准差别用全图否则大量空气体素会把分布拉偏。def normalize_foreground(img_data): mask img_data 0 mean img_data[mask].mean() std img_data[mask].std() 1e-8 img_data (img_data - mean) / std img_data[~mask] 0 return img_data3.2 滑窗切 patch 与前景采样整例数据太大必须切 patch。两种策略随机采样和滑窗。训练用随机采样保证每个 batch 里前景占比不低于阈值推理用滑窗加重叠最后按权重融合。def random_patch(img, lbl, patch_size(128, 128, 128), pos_ratio0.7): d, h, w img.shape pd, ph, pw patch_size if np.random.rand() pos_ratio: # 前景采样从非零标签里随机挑一个中心点 fg np.argwhere(lbl 0) if len(fg) 0: return random_patch(img, lbl, patch_size, pos_ratio0.0) cz, cy, cx fg[np.random.randint(len(fg))] else: cz, cy, cx np.random.randint(d), np.random.randint(h), np.random.randint(w) # 以中心点反推 patch 起点并做边界裁剪 z0 np.clip(cz - pd // 2, 0, d - pd) y0 np.clip(cy - ph // 2, 0, h - ph) x0 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参数说明pos_ratio0.7表示 70% 的 patch 中心落在前景这是应对类别极不平衡的关键。如果肿瘤只占体数据的 1%纯随机采样会让网络几乎只看到背景Dice 直接躺平。patch_size要能被 16 整除对应 4 次降采样128 是常用值显存不够就降到 96 或 64。4. 用 PyTorch 写一个能跑通的 3DUNET4.1 基础卷积块与下采样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, kernel_size3, padding1, biasFalse), nn.BatchNorm3d(out_ch), nn.ReLU(inplaceTrue), nn.Conv3d(out_ch, out_ch, kernel_size3, padding1, biasFalse), nn.BatchNorm3d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.block(x)padding1保证 3×3×3 卷积不改变空间尺寸这样跳跃连接才能对齐。biasFalse是因为后面接了 BatchNorm偏置会被抵消省参数。BatchNorm3d 在小 batch 下统计不稳如果 batch size 只能开到 1 或 2换成 InstanceNorm3d 更稳这是血泪经验。4.2 编码器、解码器与跳跃连接class UNet3D(nn.Module): def __init__(self, in_ch1, num_classes2, base16): 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.head 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.head(d1)逻辑说明torch.cat沿通道维拼接所以解码器卷积的输入通道是chs[i]*2。base16是显存和表达力的折中显存够可以上 32。输入尺寸必须能被 16 整除因为经过 4 次 pool 后要能整除回来否则cat时尺寸对不上报错信息通常是 size mismatch很难一眼看出是输入深度的问题。4.3 损失函数与优化器配置医学分割类别极不平衡纯交叉熵会被背景主导。常用组合是 Dice CE。class DiceCELoss(nn.Module): def __init__(self, ce_weight0.5): super().__init__() self.ce nn.CrossEntropyLoss() self.ce_weight ce_weight def forward(self, logits, target): ce_loss self.ce(logits, target) # softmax 后取前景通道计算 Dice probs torch.softmax(logits, dim1)[:, 1] target_fg (target 1).float() inter (probs * target_fg).sum() dice_loss 1 - (2 * inter 1e-5) / (probs.sum() target_fg.sum() 1e-5) return self.ce_weight * ce_loss (1 - self.ce_weight) * dice_loss参数说明ce_weight0.5是起点前景占比极低时可以降到 0.3让 Dice 主导。优化器用 Adam初始学习率 1e-4配合CosineAnnealingLR从 1e-4 退到 1e-6。batch size 在 128³ patch 下通常只能开到 2这时把 BatchNorm 换 InstanceNorm学习率再降一半。5. 3DUNET 训练与推理的避坑清单5.1 显存爆炸现象、原因、解决现象训练第一个 batch 就 OOM或跑到第 3 个 epoch 突然爆。原因3D 特征图随层数立方增长且 PyTorch 默认保留中间激活用于反向。解决开启混合精度torch.cuda.amppatch 从 128 降到 96base 通道从 32 降到 16必要时用梯度累积模拟大 batch。我一般先跑一个 batch 的 forward 看显存峰值再决定 patch 大小别凭感觉设。5.2 输入尺寸不整除拼接报错的排查现象cat时 size mismatch或上采样后尺寸比编码器特征大 1。原因输入深度不是 16 的整数倍某次 pool 后出现奇数。解决切 patch 时强制尺寸为 16 的倍数或在解码器上采样后做一次裁剪对齐。最省事的做法是在forward里对up的输出做F.interpolate到编码器特征尺寸但会引入插值误差不如从数据源头保证。5.3 前景太少的空 patchDice 不升反降现象训练几个 epoch 后 Dice 卡在 0.1 不动或验证集全预测背景。原因随机采样导致大量 patch 无前景网络学会「全预测背景」这个局部最优。解决前景采样比例提到 0.7 以上损失里 Dice 权重加大并在验证时按前景体素统计指标别用全图准确率自欺欺人。5.4 归一化用错统计范围验证集分布漂移现象训练集指标很好换一台设备的验证集直接崩。原因归一化用了全图均值和标准差而不同设备空气体素比例不同导致前景分布被拉偏。解决统一按前景非零 mask统计并把均值和标准差固定为训练集的值推理时直接套用别在验证集上重新统计。5.5 滑窗推理重叠不够边界出现拼接缝现象推理结果在 patch 边界出现明显台阶或断裂。原因滑窗步长等于 patch 尺寸没有重叠边界体素只被一个 patch 覆盖上下文不足。解决步长设为 patch 的 1/2 或 1/4重叠区域用高斯权重融合边界处权重低、中心权重高拼接缝基本消失。6. 验证与进阶把 Dice 从 0.7 推到 0.85 的几个具体动作训练跑通只是起点。验证阶段我习惯先做一次「过拟合测试」拿 2 例数据关掉所有增强让网络死记硬背如果 Dice 上不到 0.95说明网络或损失有问题别急着上全量数据。这个动作能省下大量调参时间。指标上Dice 和 HD95 一起看。Dice 高但 HD95 大说明内部填得满、边界毛糙这时加边界损失或做后处理。下面是一个按前景体素统计 Dice 的验证片段def compute_dice(pred, target, num_classes2): dice_list [] for c in range(1, num_classes): # 跳过背景 p (pred c) t (target c) inter (p t).sum() dice (2 * inter 1e-5) / (p.sum() t.sum() 1e-5) dice_list.append(dice) return sum(dice_list) / len(dice_list)参数说明1e-5防止空前景时除零跳过背景是因为背景占比过高会虚高指标。进阶方向有三个。一是深监督在解码器每级加辅助输出缓解梯度消失对小目标尤其有效。二是注意力门控在跳跃连接处加注意力模块抑制无关背景。三是各向异性卷积如果数据层厚明显大于层内间距把深度方向的核设为 1 或 3层内保持 3能省算力又贴合物理特性。后处理别忽视连通域分析去掉小于 50 体素的孤立预测再做一次形态学闭运算填内部空洞Dice 通常能涨 1 到 2 个点。这些动作不玄学都是可复现的。我自己的习惯是每改一个变量只跑一次对照记录 patch 大小、base 通道、损失权重、学习率四个值别一次改一堆否则出了问题根本不知道是谁的锅。3DUNET 训练一次动辄几小时后悔药很贵。希望帮到你。本文还有配套的精品资源点击获取