跳转至

Unified Multi-plane Autoregressive Diffusion for 3D Multi-Contrast MRI Synthesis

会议: ECCV 2026
论文: ECCV 2026
领域: 医学图像
关键词: 多对比度MRI合成、隐式扩散模型、自回归生成、多平面先验、3D解剖一致性

一句话总结

提出统一多平面自回归扩散框架(MPAD),通过各向同性 3D 隐空间与 2D 掩码切片扩散训练,结合跨正交平面的自回归生成与先验传递机制,以 2D 扩散的高效率合成了具有严格 3D 解剖一致性的多对比度脑部 MRI 图像。

研究背景与动机

磁共振成像(MRI)在临床神经放射学中依赖采集多种对比度序列(如 T1 加权、T2 加权、质子密度加权 PD),不同对比度能够显影不同软组织的物理特性与病理病灶。然而,在临床实际采集完整序列组合面临极大的时间成本与患者不适,长时间扫描极易引入周期性运动伪影,降低成像质量,这在需要全空间覆盖的 3D 高分辨率扫描中尤为突出。因此,从已采集的单一源对比度序列精确合成缺失的目标对比度序列,成为加速临床扫描流程与标准化影像数据的重要研究途径。

早期多对比度 MRI 合成主要基于 2D GAN 或 2D 扩散模型进行逐切片跨模态图像转换(如 T1 到 T2)。尽管 2D 模型具备训练轻量、高分辨率生成稳定的优势,但因割裂了层间连续性,沿切片法向量重建时往往出现明显的阶梯状断层和解剖几何畸变。为了弥补这一缺陷,近期工作转向全 3D 生成网络(如 3D GAN 与 3D 隐式扩散模型 LDM-3D)。然而,3D 卷积与三维注意力的计算量和显存开销随体素体积呈三次幂立方级增长(\(\mathcal{O}(N^3)\)),不仅推理速度慢、训练峰值显存极高,且极难在一个统一模型内自适应扩展至多目标对比度合成(一对多转换),限制了其临床实用性。

解决该矛盾的核心挑战在于:如何在完全规避 3D 扩散骨干网络三次幂计算复杂度的同时,为切片生成引入强有力的全局 3D 空间结构约束与解剖连续性。核心 idea:将 3D 图像合成重构为 2D 隐式切片的掩码预测任务,通过各向同性 3D 隐空间消除平面几何偏差,并在推理时采用正交多平面自回归生成与切片先验传递(轴位、矢状位、冠状位循环约束),用纯 2D 扩散操作达成全局 3D 解剖结构的一致性与高保真合成。

方法详解

整体框架

MPAD 包含两阶段训练流程与正交多平面自回归推理流程。在第一阶段,采用 3D KL 正则化自编码器将高维 3D MR 图像体压缩为各向同性(isotropic)3D 隐表征空间,支持沿任意解剖平面无畸变切片。在第二阶段,利用 3D 多模态条件编码器(MCE)提取融合源对比度体素与部分掩码目标对比度的全局条件特征,并通过 SPADE 空间自适应归一化注入 2D 扩散模型,训练 2D 去噪网络预测目标平面的被掩码隐式切片。

在推理阶段,MPAD 执行多平面自回归合成(Multi-Plane Autoregressive Diffusion)。模型依次遍历冠状位、矢状位和轴位三个正交平面,首先在第一平面通过面内自回归生成(IAS)建立初始 3D 体素结构,随后在后续正交平面通过面间先验推理(IPI)以中间扩散时间步 \(\tau\) 快速去噪,最后执行第一平面的全局细化,四份候选体素通过体素级平均聚合,送入 3D 解码器重建完整 3D 目标对比度体数据。

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["输入: 源对比度 MRI 体数据 + 目标模态文本提示"] --> B["阶段 1: 各向同性 3D 隐表征学习<br/>3D KL-Autoencoder 压缩为 32×32×32 隐空间"]
    B --> C["阶段 2: 3D 多模态条件编码与 2D 切片掩码训练<br/>3D MCE 提取上下文 + SPADE 调制 2D UNet 扩散"]
    C --> D["面内自回归切片生成<br/>随机序列逐步去噪,维持面内解剖连续性"]
    D --> E["跨正交平面先验传递与细化<br/>正交平面共享体素先验 + 初始平面回环细化"]
    E --> F["多视角体素聚合与 3D 解码<br/>四路候选均值融合 + 3D Decoder 输出高保真体数据"]

关键设计

1. 各向同性 3D 隐表征学习:消除正交切片的几何尺度偏差 传统的逐切片 2D 扩散如果直接在图像空间进行多平面重投影,常因层厚与切片内分辨率不一致导致非正方形体素或严重拉伸。MPAD 首先训练一个由 3D 编码器 \(\mathcal{E}\) 和 3D 解码器 \(\mathcal{D}\) 组成的 KL 正则化自编码器,将原始体积 \(X \in \mathbb{R}^{H \times W \times D \times C}\) 压缩到各向同性隐空间 \(z \in \mathbb{R}^{h \times w \times d \times c}\)(论文中下采样至 \(32 \times 32 \times 32 \times 3\))。自编码器结合重建损失、KL 散度惩罚和 3D 对抗损失联合优化: $\(\mathcal{L}_{\text{ae}} = \mathcal{L}_{\text{recon}}(X, \hat{X}) + \lambda_{\text{kl}}\mathcal{L}_{\text{kl}}(z) + \lambda_{\text{adv}}\mathcal{L}_{\text{adv}}(X, \hat{X})\)$ 各向同性的隐空间保证了沿冠状位(coronal)、矢状位(sagittal)和轴位(axial)任意方向切片时,潜变量具有完全一致的空间维度与几何统计分布,为后续统一 2D 扩散网络切片处理奠定了无偏的几何基础。

2. 3D 多模态条件编码与 SPADE 调制:保留完整全局空间上下文 单纯将 3D 体切片输入 2D 扩散网络会导致层间断裂。为了使 2D 去噪网络能够感知全局 3D 解剖背景,MPAD 引入 3D 多模态条件编码器(MCE)。输入包含源对比度隐变量 \(z_{\text{src}}\)、被掩码的目标对比度隐变量 \(\tilde{z}_{\text{tar}}\) 以及通过预训练医学多模态模型 BioMedCLIP 提取的目标对比度文本嵌入 \(e_{\text{tar}}\)。MCE 采用 3D 卷积架构提取富含全局层间空间上下文的特征体 \(c_{\text{tar}}\): $\(c_{\text{tar}} = \text{MCE}(z_{\text{src}} \oplus \tilde{z}_{\text{tar}}, e_{\text{tar}})\)$ 在训练切片去噪时,根据当前抽取的平面方向 \(\pi \in \{\text{axial}, \text{coronal}, \text{sagittal}\}\) 和切片索引 \(n\),从 \(c_{\text{tar}}\) 中切出对应的 2D 条件切片 \(c_{\text{tar}}^{\pi, n}\),通过空间自适应归一化(SPADE)调制 2D UNet 的每一层特征: $\(\text{SPADE}(z_{\text{tar}}^{\pi, n}, c_{\text{tar}}^{\pi, n}) = \gamma(c_{\text{tar}}^{\pi, n}) \odot \text{Norm}(z_{\text{tar}}^{\pi, n}) + \beta(c_{\text{tar}}^{\pi, n})\)$ 此外,论文在训练时采用高斯噪声替代可学习掩码 token(掩码比例在 0.7 到 1.0 之间均匀采样),使未掩码区域和加噪区域保持在自然激活分布内,强迫 MCE 充分挖掘源模态全局解剖结构与目标模态上下文。

3. 面内自回归与跨正交平面先验推理:以 2D 计算量达成 3D 闭环约束 在推理合成阶段,MPAD 设计了两大协同机制:面内自回归生成(IAS)与面间先验推理(IPI)。 - 面内自回归合成(IAS):在单一平面方向 \(\pi\) 内,切片生成顺序采用随机排列,目标切片生成概率被因式分解为已生成切片的自回归条件分布: $\(p(\hat{z}_{\text{tar}}^\pi \mid c_{\text{tar}}) = \prod_{n=1}^{N_s} p(\hat{z}_{\text{tar}}^{\pi, n} \mid \hat{z}_{\text{tar}}^{\pi, <n}, c_{\text{tar}})\)$ 通过组生成机制(实验中将 \(g=4\) 个切片并行成组),每一步将已生成切片更新入 MCE 的目标输入端,维持单平面内的连续过渡。 - 面间先验推理(IPI):为了消除单一平面切片伪影,MPAD 采用四轮正交传播方案(\(\pi_1 \to \pi_2 \to \pi_3 \to \pi_1'\))。首个平面 \(\pi_1\)(如冠状位)从纯高斯噪声出发,在最大时间步 \(T\) 下以 10 步 DDIM 生成完整体数据;随后进入第二正交平面 \(\pi_2\)(矢状位)和第三平面 \(\pi_3\)(轴位)时,直接利用已完成的 3D 体素切片,通过加噪至中间时间步 \(\tau < T\) 作为结构先验,仅需 2 步 DDIM 去噪即可快速补全。最后,对初始平面 \(\pi_1\) 执行一次综合了三个视角的细化生成 \(\pi_1'\)。最终输出通过体素级平均聚合: $\(\hat{z}_{\text{tar}} = \frac{1}{4} \left( \hat{z}_{\text{tar}}^{\pi_1} + \hat{z}_{\text{tar}}^{\pi_2} + \hat{z}_{\text{tar}}^{\pi_3} + \hat{z}_{\text{tar}}^{\pi_1'} \right)\)$ 最后经 3D 解码器 \(\hat{X}_{\text{tar}} = \mathcal{D}(\hat{z}_{\text{tar}})\) 映射回体素空间,彻底消除方向性断层伪影。

损失函数 / 训练策略

2D 扩散网络 \(\epsilon_\theta\) 在隐空间中最小化条件切片去噪均方误差: $\(\mathcal{L}_{\text{diff}} = \mathbb{E}_{\epsilon, t, \pi, n} \left[ \| \epsilon - \epsilon_\theta(z_{\text{tar}, t}^{\pi, n}, t, c_{\text{tar}}^{\pi, n}, \pi) \|_2^2 \right]\)$ 其中训练时各个正交平面的切片混合成 batch 进行联合优化,使去噪网络对解剖切片方向具有完全的朝向不变性。训练使用 AdamW 优化器,学习率为 \(4 \times 10^{-6}\),训练 1000 个 epoch;扩散步数 \(T=1000\),噪声调度采用线性规划(\(\beta_1 = 0.0015\) 至 \(\beta_T = 0.0195\))。

实验关键数据

主实验

实验在两大公开脑部 MRI 数据集 ADNI(阿尔茨海默病患者数据,737 例 1.5T 扫描)和 IXI(健康人群多中心数据,577 例 1.5T/3T 扫描)上进行评估,涵盖 T1、T2、PD 之间的全部 6 种双向转换任务。对比基线包括 3D GAN 模型(CycleGAN-3D、EaGAN)与 3D 扩散模型(LDM-3D、cWDM、ALDM)。

表 1: ADNI 数据集 3D 体积合成定量对比

方法 T1 → T2 (PSNR / SSIM / NMSE) T1 → PD (PSNR / SSIM / NMSE) T2 → T1 (PSNR / SSIM / NMSE) T2 → PD (PSNR / SSIM / NMSE) PD → T1 (PSNR / SSIM / NMSE) PD → T2 (PSNR / SSIM / NMSE)
CycleGAN-3D 24.88 / 0.810 / 0.130 20.09 / 0.774 / 0.179 21.49 / 0.807 / 0.197 19.63 / 0.761 / 0.204 19.08 / 0.751 / 0.248 19.23 / 0.751 / 0.450
EaGAN 21.72 / 0.826 / 0.273 19.80 / 0.815 / 0.203 21.14 / 0.829 / 0.199 20.32 / 0.834 / 0.184 19.54 / 0.779 / 0.269 20.42 / 0.784 / 0.362
LDM-3D 22.89 / 0.803 / 0.214 23.59 / 0.816 / 0.085 21.63 / 0.817 / 0.161 25.29 / 0.843 / 0.068 19.82 / 0.772 / 0.261 22.07 / 0.785 / 0.257
cWDM 22.79 / 0.814 / 0.221 23.58 / 0.763 / 0.072 21.26 / 0.825 / 0.158 23.44 / 0.779 / 0.081 18.03 / 0.747 / 0.399 14.69 / 0.651 / 1.337
ALDM 20.77 / 0.770 / 0.352 22.96 / 0.797 / 0.099 20.15 / 0.789 / 0.236 24.00 / 0.816 / 0.085 19.15 / 0.754 / 0.307 20.15 / 0.751 / 0.412
MPAD (本文) 23.10 / 0.830 / 0.117 23.78 / 0.829 / 0.045 23.99 / 0.846 / 0.060 25.24 / 0.846 / 0.032 22.56 / 0.817 / 0.080 22.59 / 0.821 / 0.132

表 2: IXI 数据集 3D 体积合成定量对比

方法 T1 → T2 (PSNR / SSIM / NMSE) T1 → PD (PSNR / SSIM / NMSE) T2 → T1 (PSNR / SSIM / NMSE) T2 → PD (PSNR / SSIM / NMSE) PD → T1 (PSNR / SSIM / NMSE) PD → T2 (PSNR / SSIM / NMSE)
CycleGAN-3D 27.67 / 0.845 / 0.126 24.84 / 0.829 / 0.089 22.67 / 0.803 / 0.314 25.81 / 0.853 / 0.067 23.71 / 0.815 / 0.169 24.06 / 0.834 / 0.289
EaGAN 26.17 / 0.874 / 0.202 22.75 / 0.874 / 0.149 23.66 / 0.878 / 0.345 23.70 / 0.888 / 0.130 22.90 / 0.870 / 0.360 24.07 / 0.873 / 0.329
LDM-3D 28.79 / 0.878 / 0.107 27.74 / 0.880 / 0.044 27.66 / 0.884 / 0.112 29.09 / 0.889 / 0.034 28.33 / 0.881 / 0.070 30.31 / 0.881 / 0.070
cWDM 24.62 / 0.853 / 0.319 27.28 / 0.870 / 0.057 20.86 / 0.836 / 0.526 26.26 / 0.862 / 0.086 24.53 / 0.850 / 0.235 21.98 / 0.836 / 0.585
ALDM 29.32 / 0.874 / 0.087 27.28 / 0.875 / 0.049 27.98 / 0.883 / 0.127 28.31 / 0.885 / 0.040 27.15 / 0.876 / 0.111 29.48 / 0.877 / 0.084
MPAD (本文) 30.09 / 0.882 / 0.074 28.11 / 0.881 / 0.043 28.98 / 0.890 / 0.082 29.21 / 0.889 / 0.035 28.88 / 0.887 / 0.071 30.94 / 0.887 / 0.061

消融实验

表 3: 架构设计消融实验(MCE 维度、掩码方式、先验信息在 6 种任务上的平均表现)

消融维度 配置选择 ADNI PSNR↑ ADNI SSIM↑ ADNI NMSE↓ IXI PSNR↑ IXI SSIM↑ IXI NMSE↓
MCE 卷积架构 2D Baseline 22.82 0.813 0.091 28.18 0.879 0.090
3D 卷积 (Ours) 23.54 (+0.72) 0.832 (+0.019) 0.078 (-0.013) 29.37 (+1.19) 0.886 (+0.007) 0.061 (-0.029)
掩码策略 可学习 Token 23.22 0.803 0.133 27.37 0.869 0.102
高斯噪声 (Ours) 23.54 (+0.32) 0.832 (+0.029) 0.078 (-0.055) 29.37 (+2.00) 0.886 (+0.017) 0.061 (-0.041)
面间先验 (IPI) 无先验独立生成 22.42 0.795 0.087 27.49 0.873 0.098
带面间先验 (Ours) 23.54 (+1.12) 0.832 (+0.037) 0.078 (-0.009) 29.37 (+1.88) 0.886 (+0.013) 0.061 (-0.037)

表 4: 多平面生成与细化组合消融(\(\pi_1\): 冠状位, \(\pi_2\): 矢状位, \(\pi_3\): 轴位, \(\pi_1'\): 细化冠状位)

平面组合 ADNI PSNR↑ ADNI SSIM↑ ADNI NMSE↓ IXI PSNR↑ IXI SSIM↑ IXI NMSE↓ 说明
\(\pi_1\) 22.82 0.820 0.089 29.17 0.883 0.060 单一冠状位基线
\(\pi_1 + \pi_2\) 23.32 0.828 0.081 29.00 0.884 0.067 双正交视角聚合
\(\pi_1 + \pi_3\) 23.27 0.827 0.081 29.25 0.884 0.065 双正交视角聚合
\(\pi_2 + \pi_3\) 23.33 0.827 0.079 29.15 0.885 0.064 双正交视角聚合
\(\pi_1 + \pi_2 + \pi_3\) 23.49 0.830 0.078 29.29 0.886 0.062 三正交视角聚合
\(\pi_1 + \pi_2 + \pi_3 + \pi_1'\) (Full) 23.54 0.832 0.078 29.37 0.886 0.061 完整四步生成与回环细化

关键发现

  • 面间先验(IPI)是 3D 质量跃升的核心:引入面间先验使 ADNI 上的 PSNR 大幅提升 1.12 dB、SSIM 提升 0.037;IXI 上 PSNR 提升达 1.88 dB。这表明正交切片间的结构共享显著减轻了去噪网络在缺乏上下文下的盲目性。
  • 算力与显存断崖式下降:相比全 3D 隐式扩散(LDM-3D),MPAD 将训练计算量降低了 \(7\times\)(训练 GFLOPs 大幅节省),推理计算量降低了 \(3\times\)(从体素三维扩散变成极少步数的 2D 切片去噪),显著降低了峰值显存占用与推理耗时。
  • 平面生成顺序具备极佳鲁棒性:消融显示无论是 \(\pi_1 \to \pi_2 \to \pi_3\) 还是 \(\pi_3 \to \pi_1 \to \pi_2\),重构 PSNR 均稳定在 23.48-23.54 dB,说明先验传递是无偏且互补的渐进增益过程。

亮点与洞察

  • 2D 计算开销达成 3D 一致性的优雅解构:避免了全 3D 扩散网络的立方级爆炸,利用 3D MCE 捕获体素语义、2D 扩散执行轻量去噪、多平面正交投影提供几何对齐约束,兼具了 2D 模型的效率与 3D 模型的解剖真实感。
  • 跨正交平面的热启动去噪(\(\tau < T\)):后续正交平面不从纯高斯噪声重头采样,而是将前序平面的生成体作为先验并仅回退到中间步 \(\tau\) 加噪,仅需 2 步 DDIM 去噪即可快速收敛,兼顾了生成速度与视角一致性。
  • 统一模型支持任意一对多模态转换:不同于绝大多数一对一专项训练模型,MPAD 凭借 BioMedCLIP 文本提示与统一条件架构,单套权重即可无缝执行 T1/T2/PD 之间的任意单源到目标对比度合成。

局限与展望

  • 跨平面串行迭代带来的推理延迟:虽然总 FLOPs 和显存大幅降低,但由于需要按顺序生成 3 个正交方向并在第 1 方向细化,当前存在顺序依赖,整体端到端耗时仍受切片序列串行调度的制约(当前约为 1.8 秒/体)。
  • 极小病灶细节的微弱模糊:在极端复杂的皮层折叠与微小毛细血管区域,体素级四路均值滤波可能会平滑极高频纹理。
  • 未来改进方向:可进一步探索快速一步扩散蒸馏算法(如 Consistency Models 或 Flow Matching 蒸馏),或结合轻量级跨平面注意力机制实现正交切片的并行生成。

相关工作与启发

  • vs LDM-3D & ALDM: LDM-3D 与 ALDM 采用全 3D 卷积/注意力的去噪骨干网络,参数量大且推理显存随分辨率激增;MPAD 将去噪骨干完全限制在 2D 切片级别,训练 FLOPs 缩减 7 倍、推理 FLOPs 缩减 3 倍,且合成分辨率和指标全面超越。
  • vs Make-A-Volume & 2.5D 切片网络: 现有 2.5D 方法多局限于相邻切片堆叠输入,缺乏全脑正交视角的几何自洽性;MPAD 通过冠状/矢状/轴位循环自回归与先验传递,从机制上解决了层间断层问题。

评分

  • 新颖性: ⭐⭐⭐⭐⭐ 巧妙利用各向同性 3D 隐空间与跨正交平面自回归先验传递,将 3D 合成降维至 2D 扩散计算,思路新颖且结构极其自洽。
  • 实验充分度: ⭐⭐⭐⭐⭐ 涵盖 ADNI 与 IXI 两大权威数据集、6 种跨模态合成任务,与 5 种 3D 生成基线进行了严格的对比与全方位消融。
  • 写作质量: ⭐⭐⭐⭐⭐ 逻辑结构严谨,图表清晰易读,方法动机明确,实验分析深入透彻。
  • 价值: ⭐⭐⭐⭐⭐ 对临床 MRI 快速成像、缺失模态补齐以及高分辨率 3D 医学影像的轻量化扩散生成具有极高应用推广价值。