☰
基于UNet的CBCT牙齿分割实战:从数据预处理到3D后处理全流程
2026/10/11 6:44:58 网站建设 项目流程

简介:这份资源面向医学图像处理方向的深度学习学习者与牙科影像研究者,提供使用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.2~0.4 mm,而层内像素间距可能只有 0.15~0.3 mm,不同设备出来的体数据尺寸差异很大。更麻烦的是 CBCT 没有 CT 那样的 CT 值标定,灰度是相对值,同一个病人在不同机器上拍出来的灰度分布能差出一大截。所以拿到数据第一步不是直接归一化,而是先看直方图,确认骨组织和软组织的灰度峰在哪里。常见做法是用百分位裁剪,把 1% 和 99% 分位之外的灰度截掉,再线性映射到 0~1,这样能压掉金属伪影带来的极端亮斑。金属伪影在 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, axis=0): """ 将 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_ok=True) 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, f"img_{i:04d}.nii.gz")) nib.save(nib.Nifti1Image(m, np.eye(4)), os.path.join(out_dir, f"mask_{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 通常能涨 1~2 个点,但参数量和显存也会涨。如果显存紧张,可以只在最上面两层跳跃连接加,下面两层保持直接拼接。另一个思路是把普通卷积换成残差块,缓解深层梯度消失,这个改动对训练稳定性帮助明显,尤其是 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, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, x): return self.conv(x) class UNet(nn.Module): def __init__(self, in_ch=1, out_ch=2, base=64): 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, stride=2) self.dec4 = DoubleConv(base*16, base*8) self.up3 = nn.ConvTranspose2d(base*8, base*4, 2, stride=2) self.dec3 = DoubleConv(base*8, base*4) self.up2 = nn.ConvTranspose2d(base*4, base*2, 2, stride=2) self.dec2 = DoubleConv(base*4, base*2) self.up1 = nn.ConvTranspose2d(base*2, base, 2, stride=2) 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], dim=1)) d3 = self.dec3(torch.cat([self.up3(d4), e3], dim=1)) d2 = self.dec2(torch.cat([self.up2(d3), e2], dim=1)) d1 = self.dec1(torch.cat([self.up1(d2), e1], dim=1)) return self.out(d1)

这个实现里base=64控制第一层通道数,显存不够就降到 32。DoubleConv里加了 BatchNorm,CBCT 数据 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.9~1.1,模拟不同设备的分辨率差异。弹性形变对牙齿这种硬组织要慎用,形变太强会把牙根弯成不合理的形状。灰度增强用随机 Gamma 校正,Gamma 范围 0.8~1.2,模拟不同设备的灰度差异。验证集不做增强,但要做和训练集一致的归一化。

import torch from torch.utils.data import Dataset, DataLoader import numpy as np class CBCTSliceDataset(Dataset): def __init__(self, img_dir, mask_dir, augment=True): 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并设order=1做双线性插值,掩码用order=0保持标签整数。训练时num_workers设 4 到 8,CBCT 切片读取是 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 或 4,BatchNorm 的 running mean 和 variance 估计不准。解决办法是把 BatchNorm 换成 GroupNorm 或 InstanceNorm,GroupNorm 的组数设 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, sigma=1.0): """ 三个方向的概率图做高斯平滑后平均 prob_*: 3D 数组,形状一致,值域 0-1 """ p_a = gaussian_filter(prob_axial, sigma=sigma) p_s = gaussian_filter(prob_sagittal, sigma=sigma) p_c = gaussian_filter(prob_coronal, sigma=sigma) 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] = 0

sigma控制平滑强度,设 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 数据上就能到临床可用的水平,剩下的提升靠后处理和针对性的难例挖掘。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询