ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

脑部血管分割实战:3D U-Net与预处理全流程解析

脑部血管分割实战:3D U-Net与预处理全流程解析 简介面向医学图像处理与深度学习初学者这份资源围绕脑部血管分割任务完整演示了从数据预处理到U-Net模型训练与评估的全流程。压缩包共196个文件以gif、png、tif格式的血管影像、掩膜标注与分割过程可视化为主另有8个Python脚本及xml配置文件对应图像增强、归一化、裁剪、数据扩增等预处理程序以及U-Net模型搭建、损失函数与优化器配置、评估指标计算等核心代码整体大小约42.69MB。已有272人学习。通过对照脚本与可视化结果读者可深入理解医学图像分割流程掌握跳跃连接、转置卷积、Dice损失等关键知识点并可修改脚本迁移至其他分割任务是医疗AI入门与实战的实用参考资料。1. 脑部血管分割最值得复用的部分其实是预处理脑部血管分割是从 TOF-MRA 或 CTA 体数据中把脑血管标记为前景的语义分割任务目标结构非常细长末梢分支常常只有 12 个像素或体素宽。整个任务里最麻烦的不是网络选型而是血管体素占比极低、个体差异大、伪影多。不少以“脑部血管分割”命名的工程包结构基本都是预处理脚本加 U-Net 模型定义再加训练与推理脚本真正决定 Dice 高低的往往是预处理那一半。本文按这个顺序讲先把体数据预处理和 patch 采样做扎实再看 3D U-Net 怎么设计然后说损失函数与训练配参最后落到后处理和血管连通性验证。2. 脑部血管影像预处理体素统计、ROI 裁剪与 patch 采样2.1 TOF-MRA 影像的体素分布特点脑部血管分割最常用的输入是 TOF-MRA时间飞跃法 MR 血管成像它对流动血液敏感静态脑组织信号被饱和压制因此血管与背景有天然对比。但实际数据并不干净颅底有高信号伪影头皮脂肪信号强扫描协议不同还会导致体素值整体偏移。查看一组数据的体素直方图会发现典型长尾分布绝大多数体素集中在低值区间少量伪影体素拖出很长的右尾。如果直接做 min-max 归一化伪影会被放大血管的动态范围反而被压缩。这里需要的是分位数截断而不是固定窗宽因为不同医院扫描参数差异很大。2.2 分位数截断、包围盒裁剪与 Z-score 归一化我通常按固定顺序做三步预处理对非零体素计算 2% 和 98% 分位数做截断统计非零区域的包围盒裁剪掉纯背景只用裁剪后非零体素计算均值和标准差做 Z-score顺序不能颠倒。如果先做 Z-score 再做截断极端值会拉高方差标准化后的数据仍不稳定如果先统计全图均值背景占比过高会把血管信号压得很低。import SimpleITK as sitk import numpy as np def preprocess_volume(path): image sitk.ReadImage(path) arr sitk.GetArrayFromImage(image).astype(np.float32) # 步骤1非零体素分位数截断抑制长尾伪影 non_zero arr[arr 0] lower, upper np.percentile(non_zero, [2, 98]) arr np.clip(arr, lower, upper) # 步骤2包围盒裁剪去除纯背景区域 mask arr 0 coords np.argwhere(mask) z0, z1 coords[:, 0].min(), coords[:, 0].max() 1 y0, y1 coords[:, 1].min(), coords[:, 1].max() 1 x0, x1 coords[:, 2].min(), coords[:, 2].max() 1 arr arr[z0:z1, y0:y1, x0:x1] # 步骤3仅用前景体素统计Z-score fg arr[arr 0] mean, std fg.mean(), fg.std() arr (arr - mean) / (std 1e-8) arr[arr -3] -3 return arr这段逻辑的重点在于SimpleITK 返回的数组顺序是 (z, y, x)包围盒裁剪后后续所有处理都基于新的坐标。第 2% 和 98% 分位数的选择来自经验保留绝大部分真实信号同时切掉图像边缘的异常高亮。Z-score 的下界压到 -3 是为了避免纯背景区域在卷积时产生过大的负激活。如果血管边缘特别弱分位数可以放宽到 1% 和 99%但注意不要低于 1%否则高亮伪影会重新占据动态范围。报告里如果发现背景 Dice 虚高而血管 Dice 偏低多半是归一化时背景体素参与了统计回读一下这段代码的 fg 取值就能定位问题。2.3 血管中心采样不随机均匀采 patch 的原因预处理完的体数据无法直接整块输入 3D U-Net显存放不下所以训练时要采 patch。常规做法是随机均匀采样但血管分割里这会带来一个很直接的问题血管体素占比可能低于 5%均匀采样得到的 patch 绝大多数是纯背景每个 epoch 有效样本太少。我采用血管中心采样以标注 mask 的血管体素集合作为候选点随机选一个作为 patch 中心再叠加一个随机的偏移量。这样每个 patch 至少覆盖一段血管同时保留部分背景上下文。import torch import numpy as np from torch.utils.data import Dataset class VesselPatchDataset(Dataset): def __init__(self, volume, label, patch_size64): self.volume volume self.label label self.ps patch_size self.vessel_pos np.argwhere(label 0) self.z_dim, self.y_dim, self.x_dim volume.shape def sample_center(self): idx np.random.randint(len(self.vessel_pos)) z, y, x self.vessel_pos[idx] # 随机偏移到血管周围保留背景上下文 offset int(self.ps * 0.4) z np.random.randint(-offset, offset 1) y np.random.randint(-offset, offset 1) x np.random.randint(-offset, offset 1) return z, y, x def __getitem__(self, idx): z, y, x self.sample_center() half self.ps // 2 z0 np.clip(z - half, 0, self.z_dim - self.ps) y0 np.clip(y - half, 0, self.y_dim - self.ps) x0 np.clip(x - half, 0, self.x_dim - self.ps) patch self.volume[z0:z0self.ps, y0:y0self.ps, x0:x0self.ps] lbl self.label[z0:z0self.ps, y0:y0self.ps, x0:x0self.ps] # 随机左右翻转 if np.random.rand() 0.5: patch patch[:, :, ::-1] lbl lbl[:, :, ::-1] # 灰度扰动乘性噪声模拟设备差异 scale 1.0 np.random.uniform(-0.1, 0.1) shift np.random.uniform(-0.1, 0.1) patch patch * scale shift return (torch.from_numpy(patch).unsqueeze(0).float(), torch.from_numpy(lbl).unsqueeze(0).float())这段代码的采样核心是sample_center每个 patch 中心来自血管体素偏移量为 patch 边长的 40%让 patch 不会完全压在血管主干上保证网络能同时看到血管与周围组织。灰度扰动里乘性噪声模拟不同扫描协议的信号幅度差异加性偏移模拟偏置场。前面已经做过全局 Z-scorepatch 内部就不再重复归一化否则会破坏血管与背景的相对对比。增强策略还有一个细节不要用大角度旋转。血管是细长结构旋转超过 10 到 15 度后体素网格上的离散化会切断血管。脑部近似左右对称水平翻转足够安全。弹性形变如果要用sigma 控制在 24 个体素内并且只对小部分样本生效否则主干会被拉变形。2.4 数据预处理与采样的参数汇总操作推荐参数说明分位数截断2% / 98%伪影多时可放宽到 1% / 99%包围盒非零体素注意保持 z/y/x 轴顺序Z-score前景体素统计背景不参与计算patch 大小64³显存够可提到 96³中心采样偏移patch 边长 40%过大则丢失血管上下文旋转角度±10°防止细长结构断裂灰度扰动±10% 乘性模拟设备差异这套流程做完数据才真正适合 U-Net 训练。换个网络结构效果可能变化不大但预处理顺序错了后面的训练基本都是在补充预处理欠下的债。3. U-Net 网络选型与血管分割的 3D 实现3.1 3D U-Net 还是 2D U-Net血管在三维空间里是连续管状结构单张轴向切片只能看到血管的截面或一段投影。2D U-Net 的预测在单张切片上可能很完整但重建回三维后血管经常在切片之间断裂或明显变窄这是因为网络没有见过相邻切片的上下文。3D U-Net 直接处理体数据感受野覆盖 z/y/x 三个方向血管连通性的学习就有了基础。代价是显存开销大。以 64³ patch、32 基础通道、batch_size2 为例显存占用约 68 GB如果显存小于 6 GB退而求其次用 2D U-Net 时可以在轴向、冠状位和矢状位三个平面分别训练再融合预测能部分缓解切片断裂问题。3.2 深度、通道基数和残差连接的选择血管分割不需要特别深的 U-Net。下采样 4 次足够第 5 次下采样后特征图降到 4×4×4空间信息几乎消失对细血管没有任何帮助。通道基数我一般选 32显存充裕可以提到 48但没必要继续往上加血管分割只有前景和背景两类不需要大量语义类别。上采样方式推荐用双线性或三线性插值加卷积而不是 ConvTranspose。转置卷积在细长结构上容易产生棋盘伪影表现为血管边缘出现周期性的亮度条纹。残差连接值得加在编码器每层输出前做一次恒等映射对收敛速度和稳定性都有帮助尤其是在血管体素很稀疏、多数卷积激活接近零的情况下。3.3 血管分割 3D U-Net 的 PyTorch 实现这里给一个可直接改用的实现输入单通道体数据输出 sigmoid 概率图。import torch import torch.nn as nn class ConvBlock(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv3d(in_ch, out_ch, 3, padding1), nn.BatchNorm3d(out_ch), nn.ReLU(inplaceTrue), nn.Conv3d(out_ch, out_ch, 3, padding1), nn.BatchNorm3d(out_ch), nn.ReLU(inplaceTrue) ) self.shortcut ( nn.Conv3d(in_ch, out_ch, 1) if in_ch ! out_ch else nn.Identity() ) def forward(self, x): return self.conv(x) self.shortcut(x) class UNet3D(nn.Module): def __init__(self, in_ch1, out_ch1, base32): super().__init__() self.enc1 ConvBlock(in_ch, base) self.enc2 ConvBlock(base, base * 2) self.enc3 ConvBlock(base * 2, base * 4) self.enc4 ConvBlock(base * 4, base * 8) self.pool nn.MaxPool3d(2) self.bottleneck ConvBlock(base * 8, base * 16) self.up4 nn.Upsample(scale_factor2, modetrilinear, align_cornersFalse) self.dec4 ConvBlock(base * 16 base * 8, base * 8) self.up3 nn.Upsample(scale_factor2, modetrilinear, align_cornersFalse) self.dec3 ConvBlock(base * 8 base * 4, base * 4) self.up2 nn.Upsample(scale_factor2, modetrilinear, align_cornersFalse) self.dec2 ConvBlock(base * 4 base * 2, base * 2) self.up1 nn.Upsample(scale_factor2, modetrilinear, align_cornersFalse) self.dec1 ConvBlock(base * 2 base, base) self.out nn.Conv3d(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 torch.sigmoid(self.out(d1))这个实现的核心设计包括三处。残差连接让梯度在深层网络中能跳过中间卷积对体素稀疏的血管数据更友好。上采样统一用Upsample加卷积避免 ConvTranspose 在细血管边缘产生棋盘伪影。输出端使用 sigmoid可以直接对接 BCE 或 Dice 损失。显存不足时优先砍 batch_size其次砍 base 通道数不建议把 patch 缩到 48³ 以下。patch 太小会让网络看不到血管主干的全貌末梢分支的上下文信息也不够。3.4 训练前用前向反向验证模型结构跑全量训练之前先验证模型定义没有维度错误model UNet3D(in_ch1, out_ch1, base16) x torch.randn(1, 1, 64, 64, 64) y model(x) loss y.sum() loss.backward() print(y.shape) # torch.Size([1, 1, 64, 64, 64])输出形状与输入一致反向传播能通说明模型基础没问题。base 设为 16 只是为了快速验证正式训练再改回 32。这一步能避免在训练跑了几小时后才发现解码器拼接维度错误。4. 训练配置与损失函数血管分割的调参要点4.1 BCE 与 Dice 的组合损失血管分割的正负样本极不平衡血管体素通常只占 1% 到 5%。只用 BCE 训练网络会倾向于把所有体素预测为背景因为这样损失已经很低。Dice 损失直接把前景重合度作为优化目标对类不平衡有天然鲁棒性但它对边缘细节不敏感容易产生过度平滑的分割边界。常见做法是将两者加权组合L 0.3 * BCE 0.7 * DiceBCE 提供像素级密集梯度帮助早期收敛Dice 主导后期细化和不平衡处理。权重可以按验证集表现调整血管过细时提高 Dice 比重背景噪声多时提高 BCE 比重。计算训练损失时预测概率值直接参与 Dice 计算不要先做阈值化。阈值化会截断梯度导致网络无法学习。阈值只在评估和后处理阶段使用。4.2 优化器、学习率与推荐超参参数推荐值说明优化器AdamWweight_decay1e-5防止过拟合初始学习率5e-43D U-Net 常用 3e-4 到 1e-3学习率调度CosineAnnealing配合 510 个 epoch 的 warmupbatch_size24显存不足优先减这个patch_size64³视血管尺度调整训练轮数200300血管分割不需要太长训练损失权重BCE 0.3Dice 0.7血管更细可调到 0.2/0.84.3 训练循环与模型保存训练循环里常见的坑有三个验证集没有做与训练集相同的归一化保存模型时只存了model.state_dict()而没保存预处理参数续训时没有恢复优化器和调度器状态。import torch import torch.nn.functional as F from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR def dice_soft(pred, target, smooth1.0): pred pred.contiguous().view(pred.size(0), -1) target target.contiguous().view(target.size(0), -1) intersection (pred * target).sum(dim1) return ((2.0 * intersection smooth) / (pred.sum(dim1) target.sum(dim1) smooth)) def loss_fn(pred, target): bce F.binary_cross_entropy(pred, target) dice 1 - dice_soft(pred, target) weights 0.3 * bce 0.7 * dice return weights, bce, dice model UNet3D(in_ch1, out_ch1, base32).cuda() optimizer AdamW(model.parameters(), lr5e-4, weight_decay1e-5) scheduler CosineAnnealingLR(optimizer, T_max200) checkpoint { model: model.state_dict(), optimizer: optimizer.state_dict(), scheduler: scheduler.state_dict(), epoch: epoch } torch.save(checkpoint, funet3d_epoch{epoch}.pth)这里dice_soft直接对 sigmoid 输出做计算没有阈值化梯度可以正常回流。保存 checkpoint 时把优化器和调度器状态一起存进去续训时才能保持学习率曲线一致。如果只做推理只用model字段加载即可。损失中 BCE 与 Dice 的权重并不需要频繁调整先用 0.3/0.7 跑 20 个 epoch 观察训练集 Dice。如果 Dice 在后期震荡明显尝试把 BCE 权重提高到 0.4。如果血管末梢大量丢失把 Dice 权重提到 0.8。5. 后处理与连通性验证不只盯着 Dice 看5.1 滑窗推理与概率图合并推理时整图输入显存放不下需要滑动窗口采样。窗口重叠部分取平均值能有效减少拼接边缘的预测突变。重叠率一般用 1/2推理速度敏感时可以降到 1/4。def sliding_window_infer(model, volume, patch_size64, overlap0.5): model.eval() stride int(patch_size * (1 - overlap)) z_dim, y_dim, x_dim volume.shape output np.zeros_like(volume, dtypenp.float32) count np.zeros_like(volume, dtypenp.float32) with torch.no_grad(): for z in range(0, z_dim - patch_size 1, stride): for y in range(0, y_dim - patch_size 1, stride): for x in range(0, x_dim - patch_size 1, stride): patch volume[z:zpatch_size, y:ypatch_size, x:xpatch_size] inp torch.from_numpy(patch).unsqueeze(0).unsqueeze(0).float().cuda() pred model(inp).squeeze().cpu().numpy() output[z:zpatch_size, y:ypatch_size, x:xpatch_size] pred count[z:zpatch_size, y:ypatch_size, x:xpatch_size] 1 output / np.maximum(count, 1) return output滑窗输出的是概率图不是二值 mask。保存概率图后再做阈值处理方便后续对不同阈值做评估和调优。5.2 连通域筛选与形态学修正概率图转二值 mask 后先做一次连通域分析去掉体积过小的组件。这些孤立小块通常是伪影或噪声而非真实血管分支。体积阈值可设为 50 个体素数据分辨率不同需按实际调整。from scipy import ndimage import numpy as np def remove_small_components(binary_mask, min_volume50): labeled, num_features ndimage.label(binary_mask) sizes ndimage.sum(binary_mask, labeled, range(1, num_features 1)) remove [i 1 for i, s in enumerate(sizes) if s min_volume] for label_id in remove: binary_mask[labeled label_id] 0 return binary_mask形态学闭运算能填补血管断面间的细小间隙但内核不能太大3³ 足够。内核太大会把相邻但不相连的血管错误地连接起来。这一步做完再算 Dice 通常会比直接用 0.5 阈值的结果高 1 到 3 个点更重要的是血管连通性会明显改善。5.3 用 clDice 评估血管连通性Dice 只能反映体素重叠程度两个预测可能有相同的 Dice但一个血管完整连通另一个断成碎片。clDice 是专门评估管状结构连通性的指标核心思想是预测结果的骨架是否被真实标注覆盖以及真实标注的骨架是否被预测结果覆盖。对预测 mask 和真实 mask 分别做三维骨架化然后计算两个骨架之间的 Dice。骨架提取可以使用 scikit-image 提供的skeletonize_3d之后用与 Dice 相同的方式计算 clDicefrom skimage.morphology import skeletonize_3d pred_skel skeletonize_3d(pred_mask) tru_skel skeletonize_3d(true_mask) cldice (2 * (pred_skel * tru_mask).sum()) / (pred_skel.sum() tru_skel.sum())如果 Dice 很高但 clDice 明显偏低说明预测血管虽然和真实血管有大量体素重叠但结构不连续可能是 2D 网络或未加后处理的典型表现。血管分割项目建议同时报告 Dice 和 clDice前者反映体素精度后者反映拓扑完整性。预处理时保留的包围盒信息也需要在这一步保留便于把评价结果映射回原始影像坐标但纯分割评估在裁剪坐标系下计算即可。本文还有配套的精品资源点击获取
返回列表