基于UNet的CBCT牙齿分割实战:从数据预处理到3D后处理全流程

发布时间:2026/10/11 6:45:02
基于UNet的CBCT牙齿分割实战:从数据预处理到3D后处理全流程
简介这份资源面向医学图像处理方向的深度学习学习者与牙科影像研究者提供使用UNet对CBCT牙齿数据进行图像分割的完整项目源码帮助解决高噪声、低对比度CBCT影像中牙齿自动分割这一挑战性任务。压缩包共17个文件以16个Python脚本和1份README说明文档为主整体约32KB涵盖DICOM与nrrd转PNG、数据预处理、训练验证测试集划分、网络与数据加载模块以及train.py训练入口结构清晰便于按流程阅读。项目完整呈现从灰度归一化、去噪、增强对比度等预处理到损失函数与优化器选择、模型训练、验证调参及测试集泛化评估的全链路并可能附带预测结果可视化工具方便直观对比分割效果。目前已有997人学习下载适合希望掌握UNet原理与医学图像分割实战经验的初学者及进阶开发者参考。1. 牙齿分割从 CBCT 到 UNet一条能跑通的实战路径拿到一份 CBCT 数据想把上下颌牙齿逐颗分离出来这件事在口腔正畸、种植规划、颌面外科里几乎是绕不开的前置步骤。手工勾画一套全口牙大约要花掉一个熟练技师两三个小时而且不同人勾出来的边界差异肉眼可见。牙齿分割这个任务本质上是把 CBCT 体数据里每一颗牙的体素归到它自己的类别上难点在于牙根之间骨小梁密集、牙釉质和骨皮质灰度接近、相邻牙在咬合面处几乎贴在一起。UNet 之所以在这个场景里被反复提起是因为它的编码器-解码器加跳跃连接结构能在小样本医学数据上同时抓住全局位置和局部边界对牙齿这种“形状固定但个体差异大”的目标特别合适。这篇笔记面向已经会写 PyTorch、手头有 CBCT 数据、想用 UNet 把牙齿分割跑起来的从业者从数据准备一路讲到训练参数和踩坑源码结构也会按可复现的方式拆开讲。2. 数据准备CBCT 体数据怎么变成 UNet 能吃的切片2.1 CBCT 的物理特性决定了预处理不能照搬 CT 套路CBCT 和常规螺旋 CT 最大的区别在于它的体素是各向异性的层厚通常在 0.20.4 mm而层内像素间距可能只有 0.150.3 mm不同设备出来的体数据尺寸差异很大。更麻烦的是 CBCT 没有 CT 那样的 CT 值标定灰度是相对值同一个病人在不同机器上拍出来的灰度分布能差出一大截。所以拿到数据第一步不是直接归一化而是先看直方图确认骨组织和软组织的灰度峰在哪里。常见做法是用百分位裁剪把 1% 和 99% 分位之外的灰度截掉再线性映射到 01这样能压掉金属伪影带来的极端亮斑。金属伪影在 CBCT 里非常常见种植体、烤瓷冠周围会出现放射状亮暗条纹如果直接送进网络模型会把这些条纹当成牙齿边界血泪经验是预处理阶段就要用简单的阈值加形态学把明显伪影区域标记出来训练时给这些区域降权。2.2 从体数据到 2D 切片的三种切法UNet 原生是 2D 分割网络处理 3D 体数据有两条路一是把体数据按轴位、矢状位、冠状位三个方向切成 2D 切片分别训练推理时再融合二是把 UNet 的卷积换成 3D 卷积直接吃体数据块。前者显存友好、数据量翻三倍、预训练权重好找后者能利用层间连续性但显存吃紧。我一般会先走 2D 多方向切片这条路因为 CBCT 的轴位切片上牙齿排列最清晰矢状位能看到牙根的弯曲走向冠状位对判断牙根和上颌窦的关系有帮助。三个方向各训一个模型推理时把三个方向的概率图在 3D 空间里平均边界会比单方向稳不少。切片的步长建议设为层厚的 1 倍不要跳层否则牙根尖这种细小结构容易漏掉。2.3 标注格式转换与数据集划分CBCT 的标注常见有两种一种是每颗牙一个 label 的多类标注另一种是牙齿和背景的二分类标注。如果做全口逐颗分割建议先用二分类把牙齿整体分出来再用实例分割或分水岭做后处理拆颗这样训练难度低很多。标注文件如果是 NIfTI 格式用 nibabel 读进来是 3D 数组需要按切片方向转成 2D 图像和掩码。数据集划分要按病人划分不能按切片随机划分否则同一个病人的相邻切片会同时出现在训练集和验证集里验证指标虚高这是新手最容易翻车的地方。import nibabel as nib import numpy as np import os def volume_to_slices(vol_path, mask_path, out_dir, axis0): 将 3D CBCT 体数据和标注转成 2D 切片 axis: 0-轴位 1-矢状位 2-冠状位 vol nib.load(vol_path).get_fdata() mask nib.load(mask_path).get_fdata() # 百分位裁剪压掉金属伪影极端值 p1, p99 np.percentile(vol, (1, 99)) vol np.clip(vol, p1, p99) vol (vol - vol.min()) / (vol.max() - vol.min() 1e-8) # 按指定方向取切片 vol np.moveaxis(vol, axis, 0) mask np.moveaxis(mask, axis, 0) os.makedirs(out_dir, exist_okTrue) for i in range(vol.shape[0]): img (vol[i] * 255).astype(np.uint8) m (mask[i] 0).astype(np.uint8) * 255 # 跳过全黑切片节省训练时间 if img.max() 10: continue nib.save(nib.Nifti1Image(img, np.eye(4)), os.path.join(out_dir, fimg_{i:04d}.nii.gz)) nib.save(nib.Nifti1Image(m, np.eye(4)), os.path.join(out_dir, fmask_{i:04d}.nii.gz))这段代码做了三件事读取体数据和标注、按百分位裁剪并归一化、按指定方向切片保存。axis参数控制切片方向轴位切片适合看牙冠排列矢状位适合看牙根走向。img.max() 10这个判断用来跳过纯背景切片CBCT 边缘经常有大量全黑层不跳过会浪费大量训练时间。保存成 NIfTI 而不是 PNG 是为了保留空间信息后续做 3D 融合时不用再对齐。3. UNet 模型搭建编码器深度和跳跃连接怎么定3.1 经典 UNet 结构在牙齿分割上的适配经典 UNet 是 4 层下采样加 4 层上采样每层两个 3x3 卷积加 ReLU下采样用最大池化上采样用转置卷积跳跃连接把编码器同层特征拼到解码器。牙齿分割里这个结构基本够用但有两个地方要改一是输入通道CBCT 切片是单通道灰度图第一层卷积输入通道改成 1二是输出通道二分类牙齿分割输出 2 通道多类逐颗分割输出类别数加背景。编码器深度不建议加到 5 层以上CBCT 切片分辨率通常在 512x512再往下采两次牙根尖这种几个像素宽的结构就没了。我一般保持 4 层第一层 64 通道每下采样一次通道翻倍最深层 512 通道。3.2 跳跃连接上加注意力门控的取舍原始 UNet 的跳跃连接是直接拼接编码器浅层特征里包含大量背景噪声直接拼到解码器会让边界变糊。牙齿和牙槽骨交界处灰度差异小这个问题尤其明显。常见改进是在跳跃连接上加注意力门控让解码器根据当前语义特征去筛选编码器特征。加了注意力门控后边界 Dice 通常能涨 12 个点但参数量和显存也会涨。如果显存紧张可以只在最上面两层跳跃连接加下面两层保持直接拼接。另一个思路是把普通卷积换成残差块缓解深层梯度消失这个改动对训练稳定性帮助明显尤其是 batch size 只能开到 4 或 8 的时候。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_ch1, out_ch2, base64): super().__init__() # 编码器4 层下采样 self.enc1 DoubleConv(in_ch, base) self.enc2 DoubleConv(base, base*2) self.enc3 DoubleConv(base*2, base*4) self.enc4 DoubleConv(base*4, base*8) self.pool nn.MaxPool2d(2) # 瓶颈层 self.bottleneck DoubleConv(base*8, base*16) # 解码器转置卷积上采样 跳跃拼接 self.up4 nn.ConvTranspose2d(base*16, base*8, 2, stride2) self.dec4 DoubleConv(base*16, base*8) self.up3 nn.ConvTranspose2d(base*8, base*4, 2, stride2) self.dec3 DoubleConv(base*8, base*4) self.up2 nn.ConvTranspose2d(base*4, base*2, 2, stride2) self.dec2 DoubleConv(base*4, base*2) self.up1 nn.ConvTranspose2d(base*2, base, 2, stride2) self.dec1 DoubleConv(base*2, base) self.out nn.Conv2d(base, out_ch, 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)这个实现里base64控制第一层通道数显存不够就降到 32。DoubleConv里加了 BatchNormCBCT 数据 batch 小的时候 BN 的 running mean 会抖如果 batch size 小于 4建议换成 GroupNorm把nn.BatchNorm2d换成nn.GroupNorm(8, out_ch)。跳跃拼接用torch.cat沿通道维拼拼接前要确保上采样后的尺寸和编码器特征一致输入尺寸不是 16 的整数倍时会在拼接处报尺寸不匹配所以预处理时把切片 resize 到 512x512 或 256x256 这种 2 的幂次尺寸最省事。3.3 损失函数选型Dice 和 BCE 怎么配牙齿分割里前景像素占比通常不到 10%纯交叉熵会让模型倾向于全预测背景。常见做法是 Dice Loss 加 BCE 按权重相加Dice 管区域重叠BCE 管像素级分类。权重我一般设 Dice 0.7、BCE 0.3如果边界一直糊把 Dice 权重提到 0.8。牙根尖这种细长结构用 Dice 容易梯度不稳可以再加一个 Tversky Loss调 alpha 和 beta 让召回率优先因为漏掉牙根比多分一点骨组织后果严重。学习率用 1e-4 配 Adam训练 100 个 epoch 左右前 50 个 epoch 用余弦退火把学习率降到 1e-6。验证指标看 Dice 和 IoU但别只看这两个牙根尖的 Hausdorff 距离才是真正反映临床可用性的指标。4. 训练与推理从单切片到 3D 体数据的完整链路4.1 训练循环里必须加的几件事CBCT 数据量通常不大一个医院能拿到的标注病例可能就几十例切完片也就几千张。这种量级下数据增强是必须的但医学图像的增强不能照搬自然图像那套。随机旋转角度控制在 ±15 度以内因为牙齿排列有固定解剖方向转太多会破坏先验。随机缩放 0.91.1模拟不同设备的分辨率差异。弹性形变对牙齿这种硬组织要慎用形变太强会把牙根弯成不合理的形状。灰度增强用随机 Gamma 校正Gamma 范围 0.81.2模拟不同设备的灰度差异。验证集不做增强但要做和训练集一致的归一化。import torch from torch.utils.data import Dataset, DataLoader import numpy as np class CBCTSliceDataset(Dataset): def __init__(self, img_dir, mask_dir, augmentTrue): self.img_dir img_dir self.mask_dir mask_dir self.augment augment self.files sorted(os.listdir(img_dir)) def __getitem__(self, idx): img nib.load(os.path.join(self.img_dir, self.files[idx])).get_fdata() mask nib.load(os.path.join(self.mask_dir, self.files[idx])).get_fdata() img img.astype(np.float32) mask (mask 0).astype(np.float32) if self.augment: # 随机 Gamma 校正 gamma np.random.uniform(0.8, 1.2) img np.power(img, gamma) # 随机旋转 ±15 度 if np.random.rand() 0.5: k np.random.randint(1, 4) img np.rot90(img, k) mask np.rot90(mask, k) img torch.from_numpy(img).unsqueeze(0) mask torch.from_numpy(mask).unsqueeze(0).long() return img, mask def __len__(self): return len(self.files)这个 Dataset 里 Gamma 校正和 90 度旋转是最安全的增强90 度旋转不会引入插值误差对牙齿这种有方向性的目标来说旋转后牙冠朝向变了但形状没变模型能学到旋转不变性。如果要加小角度旋转用scipy.ndimage.rotate并设order1做双线性插值掩码用order0保持标签整数。训练时num_workers设 4 到 8CBCT 切片读取是 IO 瓶颈worker 少了 GPU 会等数据。4.2 推理阶段的多方向融合三个方向各训一个模型后推理时把每个方向的 2D 概率图按原方向叠回 3D 体数据然后在体素级别取平均。融合前要确保三个方向的体数据已经对齐到同一个空间用 nibabel 的 affine 矩阵做重采样。融合后做一次 3D 连通域分析去掉小于 100 体素的孤立区域这些通常是伪影或噪声。如果做逐颗分割在二分类结果上用分水岭算法以牙冠中心为种子点牙根处用距离变换找分界线。分水岭容易过分割可以在距离变换前做一次高斯平滑sigma 设 1.5 左右。4.3 评估指标怎么算才不骗自己Dice 和 IoU 是体素级指标对边界不敏感。牙齿分割真正要看的指标有三个一是牙根尖的 Hausdorff 距离反映最坏情况下的边界偏差二是每颗牙的 Dice不是整体 Dice因为整体 Dice 会被大牙冠主导小牙根的分割质量被掩盖三是牙根和下颌神经管、上颌窦的距离误差这个直接关系到手术规划安全。计算每颗牙 Dice 时要用连通域给预测和标注分别编号再按重叠面积匹配匹配不上的算漏检。验证集上如果整体 Dice 0.92 但某颗磨牙的 Dice 只有 0.7说明模型对多根牙的分支结构学得不好需要针对这类样本做重采样或加权重。5. 避坑与排查牙齿分割训练里最常见的五个翻车点5.1 验证集 Dice 很高但推理结果全是背景现象是训练日志里验证 Dice 从第 10 个 epoch 开始稳定在 0.9 以上但拿模型去推理新数据输出几乎全黑。原因通常是验证集和训练集来自同一个病人的相邻切片模型记住了这个病人的灰度分布和牙齿形状换个人就失效。解决方法是按病人划分数据集验证集病人和训练集病人完全不重叠如果病例数太少至少保证验证集病人不在训练集里出现。另一个可能是归一化方式不一致训练时用了百分位裁剪推理时忘了做灰度分布对不上。5.2 牙根尖分割断裂成几段现象是牙冠部分分割完整但牙根尖处预测结果断成几截连通域分析后牙根被拆成多个小区域。原因是牙根尖在 CBCT 里只有几个体素宽下采样 4 次后特征图上的响应已经非常弱解码器上采样时无法恢复。解决办法有两个一是把输入切片分辨率从 512 提到 768 或 1024让牙根尖占更多像素二是在损失函数里对牙根尖区域加权用距离变换生成权重图离牙根尖越近权重越高。如果显存不够提分辨率可以在解码器最后加一层额外的上采样把输出恢复到输入尺寸的两倍再做插值下采样。5.3 金属伪影导致种植体周围过分割现象是有种植体或烤瓷冠的病例种植体周围出现一圈被预测成牙齿的区域。原因是金属伪影在 CBCT 里表现为放射状亮条纹灰度值和牙釉质接近模型分不清。解决办法是在预处理阶段做金属伪影检测用阈值加形态学找出高密度区域膨胀后生成伪影掩码训练时把伪影区域的损失权重降到 0.1。推理时对伪影区域做后处理用周围正常区域的灰度分布做插值填充再送进模型。如果伪影太严重直接把这部分数据剔除不要硬训。5.4 Batch size 太小导致 BN 统计量失准现象是训练 loss 震荡剧烈验证指标忽高忽低同一份数据两次推理结果差异明显。原因是 CBCT 切片分辨率高显存只能开 batch size 2 或 4BatchNorm 的 running mean 和 variance 估计不准。解决办法是把 BatchNorm 换成 GroupNorm 或 InstanceNormGroupNorm 的组数设 8 或 16对 batch size 不敏感。如果坚持用 BN可以开梯度累积累积 4 个 batch 再更新一次等效 batch size 到 16但 BN 的统计量还是按实际 batch 算效果有限。另一个办法是冻结 BN 的 running 统计量用预训练模型的统计量但 CBCT 和自然图像分布差太远这个办法不推荐。5.5 多类逐颗分割时类别不平衡现象是逐颗分割时磨牙这种大牙的 Dice 很高但切牙和尖牙的 Dice 很低因为切牙体积小在损失函数里贡献的梯度少。解决办法是用类别加权的 Dice Loss每类的权重和它的体积成反比切牙权重设成磨牙的 3 到 5 倍。另一个办法是分阶段训练先训二分类把牙齿整体分出来再在牙齿区域内做逐颗分类这样小牙的梯度不会被背景淹没。如果某些牙位样本特别少比如智齿可以用数据增强做针对性过采样把含智齿的切片复制几份再训。6. 进阶技巧用 3D 一致性后处理把 Dice 再提两个点2D 切片训练出来的模型推理时逐切片预测层与层之间没有约束容易出现相邻切片预测结果跳变的情况。一个成本很低但效果明显的后处理是 3D 一致性滤波对每个体素看它在三个方向上的邻域预测概率如果某个体素在轴位切片上被预测成牙齿但在矢状位和冠状位上都是背景那它大概率是噪声把它的概率拉低。具体做法是把三个方向的概率图做高斯平滑sigma 设 1 左右然后取平均再阈值化。这个操作不需要重新训练推理后处理加几行代码就能做。from scipy.ndimage import gaussian_filter def fuse_3d_probability(prob_axial, prob_sagittal, prob_coronal, sigma1.0): 三个方向的概率图做高斯平滑后平均 prob_*: 3D 数组形状一致值域 0-1 p_a gaussian_filter(prob_axial, sigmasigma) p_s gaussian_filter(prob_sagittal, sigmasigma) p_c gaussian_filter(prob_coronal, sigmasigma) fused (p_a p_s p_c) / 3.0 return fused # 阈值化后做连通域分析去掉小于 100 体素的孤立区域 from scipy.ndimage import label binary (fused 0.5).astype(np.uint8) labeled, num label(binary) for i in range(1, num 1): if (labeled i).sum() 100: binary[labeled i] 0sigma控制平滑强度设 1.0 对 0.3 mm 层厚的 CBCT 大约对应 0.3 mm 的空间平滑不会把牙根尖抹掉。如果层厚更厚sigma 可以设到 1.5。连通域阈值 100 体素是按 0.3 mm 体素算的大约对应 2.7 立方毫米比牙根尖小不会误删真实结构。融合后 Dice 通常能比单方向提升 1.5 到 2.5 个点Hausdorff 距离下降更明显因为跳变被平滑掉了。还有一个技巧是测试时增强推理时对输入切片做水平翻转、小角度旋转每个变换各预测一次把概率图变换回原空间后平均。这个操作能让 Dice 再涨 0.5 到 1 个点代价是推理时间翻几倍。如果做临床规划推理时间不敏感值得加。如果做实时导航就只保留 3D 一致性滤波。我自己做 CBCT 牙齿分割这几年最大的教训是别一上来就堆模型复杂度。先把数据预处理和标注质量抓到位把按病人划分数据集这件事做对比换什么注意力机制都管用。很多次验证指标上不去回头查都是某个病人的标注把牙槽骨标成了牙齿或者归一化时百分位裁剪的参数写错了。模型结构用经典 UNet 加 GroupNorm 加 Dice 加权损失在几百例 CBCT 数据上就能到临床可用的水平剩下的提升靠后处理和针对性的难例挖掘。希望帮到你。本文还有配套的精品资源点击获取