简介:这份资源围绕3D U-Net在三维医学图像分割中的应用展开,面向医学影像处理方向的研究者、学生与算法从业者,帮助其理解并复现基于CT、MRI等体数据的器官与病灶分割流程。压缩包共17个文件,约9KB,以4个Python脚本为核心,涵盖模型定义、训练入口与nii、yaml等工具模块,另有9个xml配置文件及iml、gitignore、txt、md等辅助文件,可用于环境配置与项目说明。内容涉及三维卷积编解码结构、Dice或Jaccard损失、数据增强与后处理等关键环节,并配有README与依赖清单,便于快速搭建训练与验证环境。目前已有1149人学习下载,适合希望从二维分割过渡到三维体数据实践、需要可运行代码骨架与排错参考的读者。
1. 3D U-Net 医学图像分割源码包:从体积数据到体素级掩膜
拿到一个 CT 或 MRI 序列,你面对的不是一张图,而是一摞切片堆成的三维体积。二维分割模型逐片推理再堆叠,层间连续性全靠运气,Z 轴上的解剖结构经常被切得七零八落。3D U-Net 就是冲着这个问题来的——它把卷积、池化、上采样全部搬到三维空间,让网络在体积上直接学习空间上下文。这次拆的3DUNET-simply_3dunet分割_3DU-Net_3dUnet_recordydn_医学图像分割源码包,核心就是一份能跑通三维医学图像分割的训练框架,包含train.py、model.py、utils工具集和requirements.txt依赖清单。它适合已经具备 PyTorch 基础、手头有 NIfTI 格式标注数据、想快速验证 3D U-Net 在自己数据集上表现的从业者。下面按「结构怎么搭 → 数据怎么喂 → 训练怎么调 → 坑怎么避」的顺序拆开讲。
2. 3D U-Net 结构拆解:编码器、解码器与跳跃连接的三维实现
2.1 为什么必须是三维卷积而不是二维堆叠
二维 U-Net 在医学图像分割里统治了很多年,但它的根本假设是「切片之间独立」。实际 CT 数据层厚 1mm 到 5mm 不等,病灶在相邻切片上的形态变化是连续的,二维模型学不到这种连续性。3D U-Net 用Conv3d替代Conv2d,卷积核在 D×H×W 三个方向上同时滑动,感受野天然覆盖层间信息。
代价也很直接:参数量和显存占用大约按核尺寸的立方增长。一个kernel_size=3的 3D 卷积,参数量是同等通道数 2D 卷积的 3 倍左右。所以 3D U-Net 的通道基数通常比 2D 版本小,常见做法是首层 16 或 32 通道起步,而不是 64。
源码包里model.py是结构定义的核心文件。我一般会先确认它用的是标准 3D U-Net 还是带残差连接的变体。标准结构长这样:
import torch import torch.nn as nn class ConvBlock3D(nn.Module): """3D U-Net 的基础卷积单元:两次 3x3x3 卷积 + BN + ReLU""" def __init__(self, in_ch, out_ch): super().__init__() self.block = nn.Sequential( nn.Conv3d(in_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm3d(out_ch), nn.ReLU(inplace=True), nn.Conv3d(out_ch, out_ch, kernel_size=3, padding=1), nn.BatchNorm3d(out_ch), nn.ReLU(inplace=True), ) def forward(self, x): return self.block(x)这段定义了两个连续的三维卷积,每个卷积后面跟 BatchNorm 和 ReLU。padding=1保证卷积后空间尺寸不变,这样跳跃连接时编码器和解码器的特征图尺寸才能对齐。BatchNorm3d在三维数据上做归一化,对医学图像这种强度分布差异大的输入尤其重要。
2.2 编码器下采样与解码器上采样的对称设计
编码器负责逐层提取语义特征,每经过一个卷积块就用MaxPool3d(2)把空间尺寸减半、通道数翻倍。解码器反过来,用ConvTranspose3d或插值上采样恢复空间分辨率,再把编码器对应层的特征图拼接过来。
class Down3D(nn.Module): """编码器下采样:最大池化 + 卷积块""" def __init__(self, in_ch, out_ch): super().__init__() self.pool = nn.MaxPool3d(2) self.conv = ConvBlock3D(in_ch, out_ch) def forward(self, x): return self.conv(self.pool(x)) class Up3D(nn.Module): """解码器上采样:转置卷积 + 跳跃连接拼接 + 卷积块""" def __init__(self, in_ch, out_ch): super().__init__() self.up = nn.ConvTranspose3d(in_ch, out_ch, kernel_size=2, stride=2) self.conv = ConvBlock3D(out_ch * 2, out_ch) # 拼接后通道翻倍 def forward(self, x, skip): x = self.up(x) x = torch.cat([x, skip], dim=1) # 沿通道维拼接 return self.conv(x)Up3D里的torch.cat是 U-Net 的灵魂操作。编码器在浅层保留了大量空间细节,解码器在深层拥有强语义信息,跳跃连接把两者拼在一起,让网络同时具备「看得准」和「看得懂」的能力。dim=1是通道维度,拼接后通道数翻倍,所以后面的ConvBlock3D输入通道要写成out_ch * 2。
注意:如果输入体积的某个维度不是 16 的倍数,经过 4 次下采样后尺寸可能变成奇数,上采样时
ConvTranspose3d的输出和跳跃连接的特征图尺寸会对不上。常见做法是在数据预处理阶段把体积裁剪或填充到 16 的倍数。
2.3 输出层与损失函数的选择逻辑
输出层通常是一个Conv3d把通道数降到类别数,二分类就是 1,多分类就是 N。激活函数二分类用 Sigmoid,多分类用 Softmax。
损失函数方面,医学图像分割最头疼的是类别极度不平衡——病灶体素可能只占整个体积的百分之几。纯交叉熵在这种情况下会被背景体素主导,模型倾向于全预测为背景。Dice Loss 直接优化预测掩膜和真实掩膜的重叠度,对不平衡数据更鲁棒。实践中常见做法是Dice Loss + CrossEntropy Loss加权组合,比如0.5 * Dice + 0.5 * CE。
class DiceLoss(nn.Module): """Dice 损失:直接优化分割重叠度,缓解类别不平衡""" def __init__(self, smooth=1e-6): super().__init__() self.smooth = smooth def forward(self, pred, target): pred = torch.sigmoid(pred) # 二分类输出转概率 pred = pred.view(-1) target = target.view(-1) intersection = (pred * target).sum() dice = (2. * intersection + self.smooth) / (pred.sum() + target.sum() + self.smooth) return 1 - dicesmooth是防止分母为零的平滑项,view(-1)把三维输出展平成一维向量再算全局 Dice。这个实现是「整批算一个 Dice」,也有按样本分别算再平均的写法,区别在于对难样本的权重不同。
3. 数据管线搭建:NIfTI 读取、归一化与三维 Patch 切分
3.1 NIfTI 格式读取与强度归一化
医学图像最常见的格式是 NIfTI(.nii或.nii.gz),源码包的utils/nii_utils.py就是干这个的。读取用nibabel,核心就一行nib.load(path).get_fdata()。
原始 CT 的 HU 值范围从 -1000 到 3000 不等,直接喂给网络会导致梯度爆炸或收敛极慢。CT 数据常见做法是先做窗宽窗位裁剪,比如腹部 CT 把 HU 限制在 [-100, 200],然后归一化到 [0, 1]。MRI 没有标准 HU 值,通常按均值和标准差做 Z-Score 归一化。
import nibabel as nib import numpy as np def load_nifti(path): """读取 NIfTI 文件并返回 numpy 数组""" img = nib.load(path) data = img.get_fdata().astype(np.float32) return data def normalize_ct(volume, hu_min=-100, hu_max=200): """CT 体积窗宽窗位裁剪 + 归一化到 [0,1]""" volume = np.clip(volume, hu_min, hu_max) volume = (volume - hu_min) / (hu_max - hu_min) return volume def normalize_mri(volume): """MRI 体积 Z-Score 归一化""" mean = volume.mean() std = volume.std() if std < 1e-8: return volume - mean return (volume - mean) / stdnp.clip把超出窗宽窗位的值截断,避免极端值影响归一化。MRI 的 Z-Score 里加了std < 1e-8的保护,防止全黑切片导致除零。
3.2 三维 Patch 切分策略与正负样本平衡
完整 CT 体积动辄 512×512×300,直接塞进网络显存扛不住。标准做法是切 Patch。训练时随机采样固定大小的三维块,比如 128×128×128 或 64×64×64,推理时用滑窗加权重叠拼接。
切 Patch 有个关键问题:如果纯随机采样,大部分 Patch 可能全是背景,正样本比例极低。常见做法是「前景优先采样」——先定位所有包含病灶的体素坐标,以这些坐标为中心采样一部分 Patch,再随机采样一部分背景 Patch,比例控制在 1:1 到 1:3 之间。
def extract_patches(volume, label, patch_size=(64, 64, 64), pos_ratio=0.5, num_patches=100): """前景优先的三维 Patch 采样""" patches, labels = [], [] fg_coords = np.argwhere(label > 0) # 前景体素坐标 num_pos = int(num_patches * pos_ratio) num_neg = num_patches - num_pos for _ in range(num_pos): if len(fg_coords) == 0: break center = fg_coords[np.random.randint(len(fg_coords))] patch, lbl = crop_at_center(volume, label, center, patch_size) patches.append(patch) labels.append(lbl) for _ in range(num_neg): center = [np.random.randint(0, s) for s in volume.shape] patch, lbl = crop_at_center(volume, label, center, patch_size) patches.append(patch) labels.append(lbl) return np.stack(patches), np.stack(labels)pos_ratio=0.5表示一半 Patch 以病灶为中心采样。crop_at_center是自定义裁剪函数,需要处理边界情况——当中心点靠近体积边缘时,Patch 会超出范围,常见做法是镜像填充或直接跳过。
3.3 DataLoader 与数据增强的工程实现
PyTorch 的Dataset和DataLoader负责把上面的采样逻辑串起来。数据增强在三维场景下比二维更需要注意——旋转、缩放、弹性形变都要在三个方向上同步操作,否则会破坏解剖结构的连续性。
from torch.utils.data import Dataset, DataLoader import torch class MedicalVolumeDataset(Dataset): def __init__(self, volume, label, patch_size=(64, 64, 64), num_patches=200): self.patches, self.labels = extract_patches( volume, label, patch_size=patch_size, num_patches=num_patches ) def __len__(self): return len(self.patches) def __getitem__(self, idx): x = torch.from_numpy(self.patches[idx]).unsqueeze(0).float() # 加通道维 y = torch.from_numpy(self.labels[idx]).unsqueeze(0).float() return x, y # 使用示例 dataset = MedicalVolumeDataset(volume, label, num_patches=200) loader = DataLoader(dataset, batch_size=2, shuffle=True, num_workers=4)unsqueeze(0)在通道维插入一个维度,因为Conv3d期望输入形状是(N, C, D, H, W)。batch_size=2是 3D 分割的常见起点,显存够可以往上加。num_workers=4加速数据加载,但 Windows 下有时会有多进程问题,设成 0 可以排查。
4. 训练脚本配置与调参:从 train.py 到收敛判据
4.1 train.py 核心流程与超参数设置
源码包的train.py是训练入口。典型流程是:加载数据 → 实例化模型 → 定义损失和优化器 → 循环训练 → 验证 → 保存最优模型。
import torch import torch.optim as optim from model import UNet3D from utils.nii_utils import load_nifti, normalize_ct # 超参数 LR = 1e-4 EPOCHS = 200 BATCH_SIZE = 2 PATCH_SIZE = (64, 64, 64) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = UNet3D(in_ch=1, out_ch=1).to(device) optimizer = optim.Adam(model.parameters(), lr=LR, weight_decay=1e-5) scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, patience=10, factor=0.5) criterion = DiceLoss() for epoch in range(EPOCHS): model.train() epoch_loss = 0 for x, y in loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() pred = model(x) loss = criterion(pred, y) loss.backward() optimizer.step() epoch_loss += loss.item() scheduler.step(epoch_loss) print(f"Epoch {epoch+1}, Loss: {epoch_loss:.4f}")lr=1e-4是 3D U-Net 的稳妥起点,太大容易震荡,太小收敛慢。weight_decay=1e-5是轻量 L2 正则,防止过拟合。ReduceLROnPlateau在损失不再下降时把学习率减半,patience=10表示连续 10 个 epoch 没改善才触发。
4.2 学习率调度与早停策略
固定学习率在 3D 分割里几乎不够用。前期需要较大学习率快速下降,后期需要小学习率精细调整。除了ReduceLROnPlateau,CosineAnnealingLR也是常见选择,它按余弦曲线平滑衰减,不需要手动设 patience。
早停策略是另一道保险。验证集 Dice 连续 N 个 epoch 不提升就停止训练,保存验证集上最优的模型权重。这个逻辑在train.py里通常用一个best_dice变量跟踪。
best_dice = 0.0 patience_counter = 0 EARLY_STOP_PATIENCE = 30 for epoch in range(EPOCHS): # ... 训练代码 ... val_dice = evaluate(model, val_loader, device) if val_dice > best_dice: best_dice = val_dice torch.save(model.state_dict(), "best_model.pth") patience_counter = 0 else: patience_counter += 1 if patience_counter >= EARLY_STOP_PATIENCE: print(f"Early stop at epoch {epoch+1}") breakEARLY_STOP_PATIENCE=30给模型足够的探索空间,太小容易在损失还在波动时误停。
4.3 显存不够时的降级方案
3D U-Net 最常翻车的地方就是显存。CUDA out of memory一出来,训练直接中断。按优先级排列的降级方案:
| 方案 | 操作 | 影响 |
|---|---|---|
| 减小 Patch | 64³ → 32³ | 空间上下文减少,小病灶可能漏检 |
| 减小 Batch | 2 → 1 | 梯度噪声增大,收敛可能变慢 |
| 减少通道 | 首层 32 → 16 | 模型容量下降,欠拟合风险 |
| 混合精度 | torch.cuda.amp | 几乎无损,显存省 30%-40% |
| 梯度累积 | 累积 4 步等效 batch=4 | 训练变慢,但等效 batch 增大 |
混合精度是性价比最高的方案,改动量小,效果立竿见影:
scaler = torch.cuda.amp.GradScaler() for x, y in loader: x, y = x.to(device), y.to(device) optimizer.zero_grad() with torch.cuda.amp.autocast(): pred = model(x) loss = criterion(pred, y) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()autocast自动把部分运算降到 float16,GradScaler防止梯度下溢。这套组合在 3D 分割里基本是标配。
5. 避坑与排查:3D U-Net 训练中最容易翻车的五个地方
5.1 损失不下降,Dice 一直在 0.1 以下
现象:训练几十个 epoch,loss 几乎不动,验证集 Dice 极低。
原因:最常见的是归一化没做对。CT 的 HU 值没裁剪直接送进网络,或者 MRI 用了 CT 的归一化方式。另一个可能是学习率太大,梯度直接炸了。
解决:先打印输入数据的均值和方差,确认在合理范围。CT 应该在 [0,1] 之间,MRI 应该在均值 0 附近。然后检查学习率,从 1e-5 开始试,确认 loss 有下降趋势再往上调。
5.2 验证集 Dice 很高但推理结果全是背景
现象:训练时验证 Dice 能到 0.8,但拿新数据推理,输出全是 0。
原因:数据泄露。训练集和验证集来自同一个病人的相邻切片,模型记住了病人特征而不是病灶特征。换一个病人的数据就失效。
解决:按病人划分训练集和验证集,而不是按切片随机划分。源码包里如果有split_by_patient之类的函数,确认它被正确调用。
5.3 显存溢出但 batch_size 已经降到 1
现象:batch_size=1仍然 OOM,但 GPU 显存看起来够。
原因:PyTorch 的缓存分配器会保留已释放的显存,nvidia-smi显示的占用不等于实际可用。另外,验证阶段的torch.no_grad()如果忘了加,验证也会建计算图。
解决:验证循环包在with torch.no_grad():里。训练前调torch.cuda.empty_cache()。如果还不行,用混合精度或减小 Patch 尺寸。
5.4 上采样后尺寸对不上,报错 size mismatch
现象:torch.cat时报错,编码器特征图和上采样后的尺寸不一致。
原因:输入体积的某个维度不是 16 的倍数,经过 4 次下采样后变成奇数,上采样回来差 1 个像素。
解决:在 Dataset 里把体积裁剪或填充到 16 的倍数。常见做法是np.pad到最近的 16 倍数,推理后再裁回来。
5.5 训练 loss 震荡剧烈,Dice 忽高忽低
现象:loss 曲线像心电图,Dice 在 0.3 到 0.7 之间反复横跳。
原因:学习率太大,或者 batch_size 太小导致梯度噪声大。另外,Dice Loss 本身在预测和真实掩膜完全无重叠时梯度不稳定。
解决:降低学习率到 1e-5,或者用Dice + CE组合损失,CE 提供稳定的梯度信号。增大 batch_size 或使用梯度累积也能平滑梯度。
6. 推理与后处理:滑窗拼接、连通域过滤与 Dice 验证
训练完模型只是第一步,推理阶段同样有讲究。完整体积推理不能直接整块送进网络,要用滑窗加权重叠拼接。窗口大小和训练时的 Patch 一致,步长通常设为窗口的 1/2 或 1/4,重叠区域取平均或高斯加权。
def sliding_window_inference(model, volume, patch_size=(64,64,64), stride=32): """滑窗推理:重叠区域取平均""" model.eval() D, H, W = volume.shape output = np.zeros_like(volume, dtype=np.float32) count = np.zeros_like(volume, dtype=np.float32) with torch.no_grad(): for d in range(0, D - patch_size[0] + 1, stride): for h in range(0, H - patch_size[1] + 1, stride): for w in range(0, W - patch_size[2] + 1, stride): patch = volume[d:d+patch_size[0], h:h+patch_size[1], w:w+patch_size[2]] inp = torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).float().cuda() pred = torch.sigmoid(model(inp)).squeeze().cpu().numpy() output[d:d+patch_size[0], h:h+patch_size[1], w:w+patch_size[2]] += pred count[d:d+patch_size[0], h:h+patch_size[1], w:w+patch_size[2]] += 1 output = output / np.maximum(count, 1) # 避免除零 return outputstride=32是 Patch 尺寸的一半,保证重叠区域足够平滑。np.maximum(count, 1)防止边缘区域除零。
推理完得到的是概率图,需要二值化。阈值通常取 0.5,但医学图像里常见做法是扫一遍 0.3 到 0.7 的阈值,看哪个在验证集上 Dice 最高。二值化之后用连通域分析去掉小面积噪声——scipy.ndimage.label标记连通区域,把体素数小于某个阈值的区域置零。
from scipy import ndimage def postprocess(pred_mask, min_size=100): """连通域过滤:去掉小于 min_size 的孤立区域""" labeled, num = ndimage.label(pred_mask) for i in range(1, num + 1): if (labeled == i).sum() < min_size: pred_mask[labeled == i] = 0 return pred_maskmin_size=100是体素数阈值,具体值取决于病灶的实际大小。太小去不掉噪声,太大会把真实小病灶也删掉。我一般会先统计验证集上所有连通域的体素数分布,取第 5 百分位数作为参考。
验证 Dice 的时候有个细节容易忽略:Dice 要在原始分辨率上算,而不是在 Patch 上算。Patch 级别的 Dice 会被采样比例影响,不能反映真实性能。另外,如果有多类,每一类分别算 Dice 再平均,不要混在一起算。
从那以后我每次跑完推理都会强制走一遍「滑窗拼接 → 阈值扫描 → 连通域过滤 → 原始分辨率 Dice」这个流程,少一步都可能被假象骗过去。希望帮到你。
本文还有配套的精品资源,点击获取