IEEE TMI IF=9.8 | UNet-DeformSA与TransDeformer:基于注意力机制实现腰椎脊柱MR图像无伪影几何重建

引言

腰椎间盘退变是导致腰背痛的关键因素,而磁共振成像(MRI)是评估其形态学变化的主要手段。然而,现有基于图像分割的自动化方法常产生带有伪影的分割结果或无结构的点云,难以转换为可用于精准医学参数测量的结构化几何模型,这阻碍了临床定量分析的效率与一致性。

发表在《IEEE Transactions on Medical Imaging》上的一项研究提出了突破性解决方案。由迈阿密大学、耶鲁大学等机构组成的研究团队,开发了两种名为UNet-DeformSA和TransDeformer的新型注意力神经网络。这些网络能够从腰椎MRI中直接、无伪影地重建出具有患者间网格对应关系的高精度几何模型,并首次实现了对重建误差的估计,为腰椎疾病的快速、标准化定量评估铺平了道路。

基本信息

文章标题:Attention-Based Shape-Deformation Networks for Artifact-Free Geometry Reconstruction of Lumbar Spine From MR Images
期刊:IEEE TRANSACTIONS ON MEDICAL IMAGING
影响因子:9.8
发表时间:2025年7月15日
研究单位:迈阿密大学计算机科学系;耶鲁大学医学院急诊医学系;迈阿密大学米勒医学院神经外科系;迈阿密大学机械与航空航天工程系
Github地址:https://github.com/linchenq/TransDeformer-Mesh
论文地址:https://ieeexplore.ieee.org/document/3588831
算力描述:训练每个模型在单张配备48GB VRAM的NVIDIA RTX A6000 GPU上耗时24小时;推理时使用单张NVIDIA RTX A6000 GPU。

研究内容与方法

1. 数据集构建

  • SSMSpine数据集
    • 样本组成:包含7000个训练样本、250个验证样本、2500个测试样本,均来自100名患有腰椎退行性疾病的患者
    • 标注规范:由3名专家对矢状位MR图像进行标注,生成覆盖6个椎体(L1-L5、S1)与5个椎间盘的11个腰椎组件的分割掩码,并构建具有点对对应关系的网格(固定1082个节点、11个元素)
  • MRSpineSeg数据集
    • 样本组成:包含150个训练样本、22个验证样本、43个测试样本,补充了S1椎体的标注
    • 标注规范:将原始分割掩码转换为统一拓扑结构的网格,确保样本间的点对对应关系
  • 模板生成:对每个数据集,取训练集所有网格的平均值作为变形模板
    【腰椎脊柱网格表示,包含11个组件、1082个节点,展示椎体与椎间盘的拓扑连接】
    在这里插入图片描述

2. 带相对位置嵌入的注意力机制(核心基础组件)

  • 设计目标:捕捉形状特征与图像特征之间的长-range依赖关系,提升模板变形的鲁棒性与准确性

  • 对应代码片段(来自github/attention.py):

    def forward(self, x, y, pos_x, pos_y):
        q = self.w_q(x)  # [N, E],形状特征作为查询
        k = self.w_k(y)  # [M, E],图像特征作为键
        v = self.w_v(y)  # [M, E],图像特征作为值
    
        # 相对位置编码生成
        pos_encoding_x = self.pos_enc(pos_x)
        pos_encoding_y = self.pos_enc(pos_y)
    
        # 融合位置编码的注意力分数计算
        q = q.unsqueeze(1) * pos_encoding_x[:, :self.embed_dim].unsqueeze(1)
        k = k.unsqueeze(0) * pos_encoding_y[:, :self.embed_dim].unsqueeze(0)
        attn_score = torch.matmul(q, k.transpose(-1, -2)) / math.sqrt(self.embed_dim)
        attn_score += torch.matmul(q, self.r2.unsqueeze(0).transpose(-1, -2)) / math.sqrt(self.embed_dim)
        attn = F.softmax(attn_score, dim=-1)
    
        # 注意力值输出
        out = torch.matmul(attn, v)
        out = self.out_proj(out.squeeze(1))
    
        # 添加相对位置差的投影项
        pos_diff = pos_x.unsqueeze(1) - pos_y.unsqueeze(0)
        pos_out = torch.matmul(attn, pos_diff)
        pos_out = self.pos_proj(pos_out.squeeze(1))
        out += pos_out
        return out
    

3. UNet-DeformSA模型架构

  • 整体架构:由预训练UNet骨干网络 + 3个串联的几何变形模块组成
  • 组件细节
    • UNet骨干网络:预训练用于腰椎分割,冻结后提取4种空间分辨率(512×512、128×128、32×32、16×16)的特征图
    • 几何变形模块:每个模块包含图像采样层与形状自注意力(SSA)层
      • 图像采样层:在模板节点位置双线性采样UNet特征图的局部特征
      • 形状自注意力层:基于上述带相对位置嵌入的注意力机制,捕捉模板节点间的依赖关系,预测节点位移
  • 对应代码片段(来自github/unet_deformsa.py):
    class GeometryDeformationModule(nn.Module):
        def __init__(self, in_channels, out_channels):
            super().__init__()
            self.image_sampler = BilinearSampler()
            self.ssa = ShapeSelfAttention(in_channels, out_channels)
            self.disp_mlp = nn.Sequential(
                nn.Linear(out_channels, out_channels),
                nn.ReLU(),
                nn.Linear(out_channels, 2)  # 预测2D位移
            )
    
        def forward(self, unet_feat, template_points):
            # 采样模板节点对应的图像特征
            sampled_feat = self.image_sampler(unet_feat, template_points)
            # 形状自注意力融合节点特征
            ssa_feat = self.ssa(sampled_feat, template_points)
            # 预测节点位移并更新模板
            disp = self.disp_mlp(ssa_feat)
            new_points = template_points + disp
            return new_points, ssa_feat
    

【UNet-DeformSA整体架构图,展示UNet骨干与3个几何变形模块的串联关系】在这里插入图片描述

4. TransDeformer模型架构

  • 整体架构:由图像特征提取器 + 形状变形块 + 形状精修块组成
  • 各组件细节:
    • 图像特征提取器
      • CNN层:生成2种空间分辨率(512×512、128×128)的初始特征图
      • 图像自注意力(ISA)模块:将特征图划分为非重叠patch,通过自注意力捕捉图像长-range依赖
        对应代码片段(来自github/transdeformer.py):
        class ImageSelfAttention(nn.Module):
            def __init__(self, patch_size=4, embed_dim=256):
                super().__init__()
                self.patch_embed = PatchEmbed(patch_size=patch_size, in_chans=64, embed_dim=embed_dim)
                self.attn_blocks = nn.ModuleList([SelfAttentionBlock(embed_dim) for _ in range(2)])
        
            def forward(self, x):
                # 将图像特征图转换为patch序列
                x = self.patch_embed(x)  # [B, L, E], L=H*W/P^2
                # 位置嵌入
                pos_embed = self.pos_embed(torch.arange(x.shape[1], device=x.device))
                # 多轮自注意力融合
                for attn in self.attn_blocks:
                    x = attn(x, pos_embed)
                return x
        

【图像自注意力(ISA)模块架构图,展示patch划分与自注意力层的串联】在这里插入图片描述
图像采样模块:在模板节点位置采样CNN特征图,结合节点坐标生成初始形状特征
形状变形块:由形状自注意力(SSA)模块与形状-图像注意力(S2IA)模块交替串联
形状自注意力(SSA)模块:捕捉模板节点间的依赖关系,基于未变形模板的位置嵌入计算注意力
对应代码片段(来自github/transdeformer.py):

     class ShapeSelfAttention(nn.Module):
         def __init__(self, embed_dim=256):
             super().__init__()
             self.attn = RelativeAttention(embed_dim)
             self.norm = nn.LayerNorm(embed_dim)
     
         def forward(self, shape_feat, template_points):
             # 归一化模板节点位置作为位置嵌入
             pos_embed = self.normalize_points(template_points)
             # 自注意力融合形状特征
             attn_feat = self.attn(shape_feat, shape_feat, pos_embed, pos_embed)
             return self.norm(attn_feat + shape_feat)  # 残差连接

【形状自注意力(SSA)模块架构图,展示节点特征的自注意力融合】在这里插入图片描述

形状-图像注意力(S2IA)模块:建立形状特征与图像特征的跨模态注意力,解决模板初始化偏差问题
对应代码片段(来自github/transdeformer.py):

     class ShapeToImageAttention(nn.Module):
         def __init__(self, embed_dim=256):
             super().__init__()
             self.attn = RelativeAttention(embed_dim)
             self.norm = nn.LayerNorm(embed_dim)
     
         def forward(self, shape_feat, image_feat, shape_points, image_patch_pos):
             # 分别归一化形状节点与图像patch的位置
             pos_shape = self.normalize_points(shape_points)
             pos_image = self.normalize_patches(image_patch_pos)
             # 跨模态注意力融合图像特征到形状特征
             attn_feat = self.attn(shape_feat, image_feat, pos_shape, pos_image)
             return self.norm(attn_feat + shape_feat)

【形状-图像注意力(S2IA)模块架构图,展示形状特征与图像特征的跨模态融合】在这里插入图片描述

形状精修块:对形状变形块输出的中间形状进行精修,预测最终节点位移得到完整腰椎网格
【TransDeformer整体架构图,展示图像特征提取、形状变形、形状精修的流程】在这里插入图片描述

5. 损失函数与训练/推理策略

  • 损失函数设计
    • UNet-DeformSA损失:融合分割损失与几何损失
      L = L s e g + L g e o m \mathcal{L} = \mathcal{L}_{seg} + \mathcal{L}_{geom} L=Lseg+Lgeom
      分割损失为Dice损失与面积加权交叉熵损失的组合:
      L s e g = − ∑ i = 0 11 2 × ∑ ( M ^ i ⊙ M i ) ∑ M ^ i + ∑ M i − ∑ i = 0 11 ∑ ( w i M i ⊙ log ⁡ ( M ^ i ) ) \mathcal{L}_{seg} = -\sum_{i=0}^{11} \frac{2 \times \sum (\hat{M}_i \odot M_i)}{\sum \hat{M}_i + \sum M_i} - \sum_{i=0}^{11} \sum (w_i M_i \odot \log(\hat{M}_i)) Lseg=i=011M^i+Mi2×(M^iMi)i=011(wiMilog(M^i))
      几何损失:融合所有中间变形输出与最终输出的MSE损失
    • TransDeformer损失:仅使用几何损失,同样融合中间输出与最终输出
      L g e o m = ∑ m L g e o m ( t ) ( S ^ ( m ) , S ) \mathcal{L}_{geom} = \sum_m \mathcal{L}_{geom}^{(t)}(\hat{S}^{(m)}, S) Lgeom=mLgeom(t)(S^(m),S)
  • 三模式训练策略
    • 模式1:仅预测腰椎质心,损失为质心坐标的MSE
      L g e o m ( 1 ) ( S ^ , S ) = ∥ C ^ − C ∥ 2 2 \mathcal{L}^{(1)}_{geom}(\hat{S}, S) = \|\hat{C} - C\|_2^2 Lgeom(1)(S^,S)=C^C22
      训练时模板随机放置在图像中
    • 模式2:预测完整腰椎形状,损失为所有节点坐标的MSE
      L g e o m ( 2 ) ( S ^ , S ) = 1 N p o i n t s ∥ S ^ − S ∥ 2 2 \mathcal{L}^{(2)}_{geom}(\hat{S}, S) = \frac{1}{N_{points}} \|\hat{S} - S\|_2^2 Lgeom(2)(S^,S)=Npoints1S^S22
      训练时模板初始化为接近真实腰椎形状
    • 模式3:同模式2的损失计算,训练时模板先经过非线性变换再初始化
  • 两阶段推理策略
    • 阶段1:预测模板质心位移,将模板移动到腰椎质心位置
    • 阶段2:预测每个模板节点的位移,得到最终腰椎网格
  • 对应代码片段(来自github/loss.py):
    def geometric_loss(pred_points, gt_points, mode='full'):
        if mode == 'centroid':
            # 质心预测损失
            pred_centroid = torch.mean(pred_points, dim=1)
            gt_centroid = torch.mean(gt_points, dim=1)
            return F.mse_loss(pred_centroid, gt_centroid)
        elif mode == 'full':
            # 完整形状预测损失
            return F.mse_loss(pred_points, gt_points, reduction='mean') / pred_points.shape[1]
    

6. 形状误差估计模型

  • 模型架构:基于TransDeformer修改而来
    • 输入:MR图像 + 任意输入形状
    • 修改点:移除中间形状预测层,添加误差预测头,输出每个节点的非负标量(代表该节点与真实形状的距离估计)
  • 训练方式:输入形状添加随机噪声,以真实误差(节点与真实形状的距离)为监督,使用MSE损失训练
  • 对应代码片段(来自github/error_estimation.py):
    class ShapeErrorEstimator(nn.Module):
        def __init__(self, transdeformer_backbone):
            super().__init__()
            self.image_encoder = transdeformer_backbone.image_encoder
            self.shape_encoder = transdeformer_backbone.shape_encoder
            self.error_head = nn.Sequential(
                nn.Linear(256, 128),
                nn.ReLU(),
                nn.Linear(128, 1),
                nn.Sigmoid()  # 输出0-1的误差比例,再缩放为实际距离范围
            )
    
        def forward(self, mr_image, input_shape):
            # 提取图像特征
            image_feat = self.image_encoder(mr_image)
            # 采样输入形状对应的图像特征
            sampled_feat = self.image_sampler(image_feat, input_shape)
            # 编码形状特征
            shape_feat = self.shape_encoder(sampled_feat, input_shape)
            # 预测每个节点的误差
            error = self.error_head(shape_feat) * 10  # 缩放至0-10mm的误差范围
            return error
    

【形状误差估计模型架构图,展示输入、特征编码与误差预测的流程】在这里插入图片描述

实验结果分析

腰椎几何重建模型的性能评估

通过平均点对点距离(APPD)Dice相似系数DSC)和豪斯多夫距离HD)等指标,比较了所提模型与其他基线模型在两个数据集上的整体性能。

  • 几何重建精度:在SSMSpine和MRSpineSeg测试集上,我们的模型在APPD主要指标)上均取得了最低的平均值,表明其重建的几何形状与真实标注最为接近。例如,在SSMSpine数据集上,TransDeformer的APPD为0.59 mm,显著优于UNet-GCN(0.77 mm)和UNet-Disp(0.81 mm)。
  • 分割性能对比:尽管分割模型(如DGMSNet)在DSC指标上表现良好,但其输出是分割掩码,无法直接获得具有点对应关系的结构化网格,因此不适用于基于模板的医学参数测量。我们的模型在保证网格对应关系的同时,其DSC和HD指标也极具竞争力。
  • 最坏情况分析:在最坏情况指标(如Q95,即95%分位数)下,我们的模型同样表现稳健。例如,在SSMSpine数据集上,TransDeformer的APPD Q95值为1.31 mm,优于其他所有基于模板变形的模型,表明其在具有挑战性的病例中也能产生可靠输出

模型对模板初始化的鲁棒性

本部分评估了模型对模板初始位置扰动的鲁棒性。实验通过在不同半径的圆内随机平移初始模板中心,并比较开启/关闭推理第一阶段(中心预测)时的性能变化来完成。

  • 鲁棒性优势:当模板初始化存在较大偏差时(例如半径40像素),我们的模型性能下降幅度远小于对比模型。在SSMSpine数据集上,即使关闭第一阶段推理,TransDeformer在40像素扰动下的APPD仅从0.59 mm增至0.86 mm,而UNet-Disp则从0.81 mm恶化至1.74 mm。
  • 中心预测的贡献:启用推理第一阶段(即先预测腰椎中心)能显著提升所有模型对初始化扰动的鲁棒性。我们的模型在该阶段也表现出更准确的中心预测能力,这为后续的精细变形奠定了良好基础。
  • 位置编码的影响:实验还比较了不同位置编码方法。我们提出的结合正弦函数的相对位置编码方法,在TransDeformer中取得了最佳性能,尤其是在模板初始化存在扰动时,其鲁棒性优于绝对位置编码和经典相对位置编码方法。

医学参数测量与误差估计的应用

本部分展示了所提模型在最终应用目标——医学参数测量上的准确性,以及所构建的形状误差估计模型在质量控制中的潜力。

在这里插入图片描述

图注: 展示了基于网格模板定义的15个与椎间盘退变相关的几何参数。

  • 参数测量精度:使用我们模型重建的几何形状进行医学参数测量时,其相对误差最低。例如,在测量椎间盘前高(ADH)和中高(MDH)等关键参数时,TransDeformer的相对误差中位数分别为3.21%和3.58%,优于其他对比模型。
  • 误差估计与质量控制:基于TransDeformer修改的形状误差估计模型,能够有效预测重建几何的误差。如表该模型预测的误差与真实误差之间存在高度相关性(Pearson相关系数最高达0.85)。这使其能够对重建结果进行可靠排序,将估计误差较大的案例优先提交给人工审核,从而在实际应用中辅助质量控制流程

优势与局限

优势

高精度与无伪影重建:提出的UNet-DeformSA和TransDeformer模型通过新型注意力模块,实现了从腰椎MR图像中高空间精度的几何重建,输出无伪影的网格,优于现有分割和模板变形方法。
跨患者网格对应性强:基于模板变形的重建方法确保了所有患者输出网格具有相同的拓扑结构,建立了点对点对应关系,支持在统一参考解剖结构上一致地定义和测量医学参数。
支持误差估计与质量控制:基于TransDeformer改进的误差估计网络能够预测重建几何的误差,为临床应用的输出质量评估和优先级排序提供了有效工具,有助于质量控制。

局限

模板适用性受限:当前模板主要针对腰椎退行性病变相关参数测量设计,对于其他病理形态(如椎体内部骨折模式)的表征能力有限,可能不适用于所有临床场景。
对中矢状面图像的依赖:模型训练和验证主要基于中矢状面MR图像,对于存在严重脊柱侧弯或定位不佳的病例,模型性能可能下降,泛化能力有待在更广泛图像平面上验证。
数据驱动方法的固有不确定性:模型本质是数据驱动的,其输出精度依赖于训练数据分布,无法提供理论保证,仍需人工核查以确保临床使用的可靠性。

参考文献

  1. Evaluation of deep neural network models for instance segmentation of lumbar spine MRI Chen et al., 2024:该论文系统评估了包括UNet++、Attention U-Net、Swin-Unet等在内的15种深度学习模型在腰椎MRI实例分割任务上的表现,并构建了SSMSpine数据集。本研究以此为基线,指出了基于分割的方法在生成结构化网格和医学参数测量方面的固有缺陷,从而引出了基于模板变形的几何重建新范式
  2. Voxel2Mesh: 3D mesh model generation from volumetric data Wickramasinghe et al., 2020:本文提出了Voxel2Mesh模型,利用图卷积网络(GCN)从体数据中重建器官网格。本研究在相关工作部分将其作为基于模板变形的深度学习方法进行对比,并指出其生成的网格顶点数不确定,缺乏患者间网格对应关系,而本研究方法则确保了这种对应性。
  3. A deep-learning approach for direct whole-heart mesh reconstruction Kong et al., 2021:该论文提出了一个融合UNet与GCN的复合框架,用于心脏几何的直接网格重建。本研究将其方法适配到腰椎应用(即UNet-GCN模型)作为对比基准之一,实验表明其仍会产生几何伪影,从而凸显了本研究引入注意力机制以捕获长程依赖关系的优势。
  4. CorticalFlow: A diffeomorphic mesh transformer network for cortical surface reconstruction Lebrat et al., 2021:本文提出了CorticalFlow模型,利用微分同胚流变形球形模板来重建大脑皮层表面。本研究将其方法进行适配并作为对比模型之一,实验发现其对模板初始化非常敏感,而本研究提出的三阶段训练策略显著提升了模型对此的鲁棒性。
  5. SpineParseNet: Spine parsing for volumetric MR image by a two-stage segmentation framework with semantic image representation Pang et al., 2021:该论文提出了专门用于腰椎MR图像分割的先进模型SpineParseNet。本研究将其与另一个先进分割模型DGMSNet一同作为基于分割方法的代表进行性能对比,结果证实即使是最优的分割模型,其输出的掩码仍存在伪影且难以转换为适用于参数测量的结构化网格,从而反衬了本研究直接进行几何重建方法的必要性与优越性。
评论
添加红包

请填写红包祝福语或标题

红包个数最小为10个

红包金额最低5元

当前余额3.43前往充值 >
需支付:10.00
成就一亿技术人!
领取后你会自动成为博主和红包主的粉丝 规则
hope_wisdom
发出的红包
实付
使用余额支付
点击重新获取
扫码支付
钱包余额 0

抵扣说明:

1.余额是钱包充值的虚拟货币,按照1:1的比例进行支付金额的抵扣。
2.余额无法直接购买下载,可以购买VIP、付费专栏及课程。

余额充值