Skip to content

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

Conference: ECCV 2026
Paper: ECCV 2026
Area: Medical Imaging
Keywords: multi-contrast MRI synthesis, latent diffusion model, autoregressive generation, multi-plane priors, 3D anatomical consistency

TL;DR

The paper introduces the unified Multi-Plane Autoregressive Diffusion (MPAD) framework, which couples an isotropic 3D latent space with 2D masked-slice diffusion training, along with cross-orthogonal-plane autoregressive synthesis and prior propagation, synthesizing volumetrically coherent multi-contrast brain MRI scans with 2D computational efficiency.

Background & Motivation

Magnetic resonance imaging (MRI) examinations in clinical neuroimaging routinely acquire multiple contrasts—such as T1-weighted, T2-weighted, and proton density-weighted (PD-weighted) scans—each highlighting distinct physiological and pathological tissue characteristics. However, acquiring a complete set of contrasts for every patient is severely limited by acquisition time, scanner throughput, and patient discomfort. Prolonged scan durations also heighten susceptibility to motion artifacts that degrade image fidelity, a challenge that becomes especially acute for full 3D volumetric acquisitions. Consequently, synthesizing missing contrasts from acquired ones has become a vital avenue for accelerating clinical workflows and standardizing imaging protocols.

Early multi-contrast MRI synthesis approaches primarily relied on 2D slice-to-slice generative adversarial networks (GANs) or 2D diffusion models. While computationally lightweight and stable to train, treating 3D volumes as isolated 2D slices neglects through-plane anatomical continuity, leading to pronounced staircasing artifacts, jagged slice transitions, and geometric distortions when viewed along orthogonal cross-sections. To overcome this limitation, recent studies have embraced fully volumetric 3D architectures, including 3D GANs and 3D latent diffusion models (LDM-3D). However, the computational complexity and memory consumption of volumetric convolutions and 3D attention mechanisms scale cubically (\(\mathcal{O}(N^3)\)) with volume size, causing prohibitive training costs, high peak memory footprints, slow sampling speeds, and difficulties in extending a single model to one-to-many contrast synthesis.

The central challenge lies in eliminating the cubic computational scaling of 3D volumetric diffusion backbones while still rigorously enforcing global 3D anatomical consistency. Core idea: reformulate 3D volume synthesis as a 2D masked-slice latent prediction task over an isotropic 3D latent space to eliminate geometric slicing bias, and perform inference via multi-plane autoregressive generation with inter-plane prior propagation across coronal, sagittal, and axial views to achieve global volumetric coherence using purely 2D diffusion operations.

Method

Overall Architecture

MPAD comprises a two-stage training scheme and an orthogonal multi-plane autoregressive inference pipeline. In Stage 1, a 3D KL-regularized autoencoder compresses the high-dimensional volumetric MRI into an isotropic 3D latent space, enabling unbiased slicing along arbitrary anatomical orientations. In Stage 2, a 3D Multi-modal Conditioning Encoder (MCE) extracts volumetric context from the concatenated source latent and partially masked target latent guided by target modality text prompts, and injects slice-aligned conditioning into a 2D diffusion network via SPADE modulation to predict masked target-contrast latent slices.

During inference, MPAD executes plane-wise autoregressive diffusion across orthogonal views. Slices are first generated in the primary plane (coronal) via Intra-plane Autoregressive Synthesis (IAS) to establish an initial 3D anatomical structure. Next, Inter-plane Prior Inference (IPI) propagates this volume as an informative initialization prior to subsequent orthogonal planes (sagittal and axial) starting from an intermediate diffusion timestep \(\tau < T\) for rapid sampling. Finally, a refinement pass is performed on the primary plane, and the four candidate latent volumes are aggregated via voxel-wise averaging before being decoded back to the high-resolution image space.

%%{init: {'flowchart': {'rankSpacing': 24, 'nodeSpacing': 28, 'padding': 6, 'wrappingWidth': 400}}}%%
flowchart TD
    A["Input: Source-contrast MRI Volume + Target Modality Text Prompt"] --> B["Stage 1: Isotropic 3D Latent Representation Learning<br/>3D KL-Autoencoder compresses volume into 32×32×32 latent space"]
    B --> C["Stage 2: 3D Multi-modal Conditioning & 2D Slice Diffusion<br/>3D MCE extracts volumetric context + SPADE modulates 2D UNet"]
    C --> D["Intra-plane Autoregressive Slice Synthesis<br/>Random-order progressive denoising enforces intra-plane continuity"]
    D --> E["Cross-orthogonal Plane Prior Propagation & Refinement<br/>Inter-plane prior initialization at timestep τ + loopback refinement"]
    E --> F["Multi-view Voxel Aggregation & 3D Decoding<br/>Voxel-wise 4-way averaging + 3D Decoder outputs final 3D volume"]

Key Designs

1. Isotropic 3D Latent Representation Learning: eliminating geometric slicing bias Directly slicing anisotropic 3D volumes in image space causes varying physical slice thicknesses and pixel aspect ratios across orthogonal planes. MPAD circumvents this by training a 3D KL-regularized autoencoder parameterized by an encoder \(\mathcal{E}\) and decoder \(\mathcal{D}\), compressing an input volume \(X \in \mathbb{R}^{H \times W \times D \times C}\) into an isotropic latent tensor \(z \in \mathbb{R}^{h \times w \times d \times c}\) (\(32 \times 32 \times 32 \times 3\)). The autoencoder optimizes a combined reconstruction, KL-divergence, and 3D adversarial objective: $\(\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})\)$ Because the resulting latent representation is strictly isotropic, slicing along axial, sagittal, or coronal directions yields identical spatial dimensions and statistical feature behaviors, establishing a geometrically unbiased foundation for a unified 2D diffusion model.

2. 3D Multi-modal Conditioning Encoder with SPADE: preserving through-plane volumetric context Training a 2D diffusion model on isolated 2D slices typically causes loss of 3D spatial awareness. To ensure the 2D denoiser perceives global anatomical context, MPAD introduces a 3D Multi-modal Conditioning Encoder (MCE). Its input combines the source-contrast latent \(z_{\text{src}}\), the partially masked target-contrast latent \(\tilde{z}_{\text{tar}}\), and a target contrast text embedding \(e_{\text{tar}}\) derived from BioMedCLIP. The 3D MCE aggregates through-plane dependencies across all slices to produce volumetric conditioning features \(c_{\text{tar}}\): $\(c_{\text{tar}} = \text{MCE}(z_{\text{src}} \oplus \tilde{z}_{\text{tar}}, e_{\text{tar}})\)$ For a sampled plane orientation \(\pi \in \{\text{axial}, \text{coronal}, \text{sagittal}\}\) and slice index \(n\), the corresponding 2D conditioning slice \(c_{\text{tar}}^{\pi, n}\) is extracted and modulates the 2D UNet intermediate features via Spatially Adaptive Normalization (SPADE): $\(\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})\)$ During training, Gaussian noise masking (masking ratios uniformly sampled between 0.7 and 1.0) replaces masked latent positions rather than learnable constant tokens. This maintains continuous activation statistics, compelling the 3D MCE to extract robust volumetric context from partially corrupted target latents.

3. Intra-plane Autoregression and Inter-plane Prior Inference: enforcing 3D coherence via 2D operations At inference, MPAD resolves slice-wise inconsistencies through two coordinated mechanisms: Intra-plane Autoregressive Synthesis (IAS) and Inter-plane Prior Inference (IPI). - Intra-plane Autoregressive Synthesis (IAS): Within a selected plane orientation \(\pi\), slice generation is factorized autoregressively using a random permutation order: $\(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}})\)$ In practice, slices are synthesized in parallel groups of \(g=4\). Previously completed slices are continuously fed back into the MCE input to enforce smooth slice-to-slice transitions. - Inter-plane Prior Inference (IPI): To eliminate directional slicing artifacts, generation proceeds across four passes (\(\pi_1 \to \pi_2 \to \pi_3 \to \pi_1'\)). The primary plane \(\pi_1\) (e.g., coronal) is generated from pure noise across \(T\) timesteps with 10 DDIM steps. Subsequent orthogonal planes \(\pi_2\) (sagittal) and \(\pi_3\) (axial) do not start from pure noise; instead, they extract corresponding slices from the already completed 3D volume, diffuse forward to an intermediate timestep \(\tau < T\), and perform reverse diffusion using only 2 DDIM steps. A final refinement pass on \(\pi_1'\) leverages context from all completed planes. The four candidate volumes are merged through voxel-wise averaging: $\(\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)\)$ Decoding the aggregated latent \(\hat{X}_{\text{tar}} = \mathcal{D}(\hat{z}_{\text{tar}})\) yields a structurally unified 3D volume free of slice-boundary discontinuities.

Loss & Training

The 2D diffusion backbone \(\epsilon_\theta\) minimizes the slice-wise conditional noise-prediction mean squared error: $\(\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]\)$ Slices from axial, coronal, and sagittal orientations are batched together during training, making the denoising network invariant to anatomical plane orientations. The model is trained using AdamW with a learning rate of \(4 \times 10^{-6}\) for 1000 epochs, employing a linear noise schedule from \(\beta_1 = 0.0015\) to \(\beta_T = 0.0195\) over \(T=1000\) diffusion timesteps.

Key Experimental Results

Main Results

Quantitative evaluations are conducted on two benchmark brain MRI datasets: ADNI (737 subjects with Alzheimer's disease, 1.5T GE scanners) and IXI (577 healthy subjects, 1.5T and 3T Philips scanners), evaluating all six cross-modality translation tasks between T1, T2, and PD. Baseline models include 3D GANs (CycleGAN-3D, EaGAN) and 3D diffusion models (LDM-3D, cWDM, ALDM).

Table 1: Quantitative comparison on the ADNI dataset across full 3D volumes

Method 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 (Ours) 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

Table 2: Quantitative comparison on the IXI dataset across full 3D volumes

Method 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 (Ours) 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

Ablation Study

Table 3: Ablation study on conditional encoding and training components (averaged across all 6 translation tasks)

Ablation Category Design Choice ADNI PSNR↑ ADNI SSIM↑ ADNI NMSE↓ IXI PSNR↑ IXI SSIM↑ IXI NMSE↓
MCE Architecture 2D Conv 22.82 0.813 0.091 28.18 0.879 0.090
3D Conv (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)
Masking Strategy Learnable Token 23.22 0.803 0.133 27.37 0.869 0.102
Gaussian Noise (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)
Inter-plane Prior (IPI) Without Priors 22.42 0.795 0.087 27.49 0.873 0.098
With Priors (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)

Table 4: Multi-plane combination and refinement ablation (\(\pi_1\): coronal, \(\pi_2\): sagittal, \(\pi_3\): axial, \(\pi_1'\): refined coronal)

Planes Combination ADNI PSNR↑ ADNI SSIM↑ ADNI NMSE↓ IXI PSNR↑ IXI SSIM↑ IXI NMSE↓ Note
\(\pi_1\) 22.82 0.820 0.089 29.17 0.883 0.060 single coronal plane baseline
\(\pi_1 + \pi_2\) 23.32 0.828 0.081 29.00 0.884 0.067 2 orthogonal planes
\(\pi_1 + \pi_3\) 23.27 0.827 0.081 29.25 0.884 0.065 2 orthogonal planes
\(\pi_2 + \pi_3\) 23.33 0.827 0.079 29.15 0.885 0.064 2 orthogonal planes
\(\pi_1 + \pi_2 + \pi_3\) 23.49 0.830 0.078 29.29 0.886 0.062 3 orthogonal planes
\(\pi_1 + \pi_2 + \pi_3 + \pi_1'\) (Full) 23.54 0.832 0.078 29.37 0.886 0.061 full multi-plane with refinement

Key Findings

  • Inter-plane priors drive dramatic quality gains: Equipping multi-plane generation with inter-plane priors yields an increase of 1.12 dB PSNR and 0.037 SSIM on ADNI, and 1.88 dB PSNR on IXI over uncoordinated multi-plane sampling, verifying that cross-plane conditioning stabilizes volumetric denoising.
  • Massive efficiency and footprint reduction: Compared with fully 3D latent diffusion (LDM-3D), MPAD slashes training FLOPs by \(7\times\) and inference FLOPs by \(3\times\), while dramatically cutting peak GPU memory consumption and inference latency.
  • Invariance to plane sequence order: Ablation across different plane sequences (\(\pi_1 \to \pi_2 \to \pi_3\) vs \(\pi_3 \to \pi_1 \to \pi_2\)) achieves near-identical results (23.48-23.54 dB PSNR on ADNI), demonstrating that inter-plane context accumulation acts as an isotropic geometric constraint rather than an order-dependent bias.

Highlights & Insights

  • Decoupling 3D anatomical context from 2D diffusion compute: By allocating volumetric feature extraction to a compact 3D MCE and restricting iterative denoising to 2D latent slices, MPAD completely avoids the \(\mathcal{O}(N^3)\) computational explosion of 3D diffusion models while retaining full 3D spatial fidelity.
  • Warm-started intermediate-timestep diffusion (\(\tau < T\)): Subsequent orthogonal planes leverage the completed volume from previous passes as an informed prior, skipping the noisy initial phase and performing reverse sampling in only 2 DDIM steps from timestep \(\tau\), dramatically reducing inference latency.
  • Unified one-to-many contrast synthesis: Rather than training dedicated networks for each source-target pair, MPAD incorporates BioMedCLIP text embeddings to modulate generation, enabling a single unified model to synthesize T1, T2, or PD from any acquired contrast.

Limitations & Future Work

  • Sequential plane generation overhead: While total FLOPs and memory are reduced by multiples, sequential generation across three orthogonal planes plus a refinement pass still introduces execution latency (approximately 1.8 seconds per volume).
  • Potential texture smoothing in fine anatomical structures: Four-way voxel-level averaging can slightly smooth ultra-fine cortical folding boundaries or capillary-scale details.
  • Future directions: Integrating few-step diffusion distillation techniques (e.g., Consistency Models) or designing parallelized cross-plane attention mechanisms to synthesize orthogonal planes concurrently could further cut latency.
  • vs LDM-3D & ALDM: LDM-3D and ALDM employ full 3D convolutional and attention backbones, which incur cubic computational scaling and high memory overheads; MPAD confines the denoising network to 2D slices, cutting training FLOPs by \(7\times\) and inference FLOPs by \(3\times\) while achieving superior quantitative metrics.
  • vs Make-A-Volume & 2.5D slice models: Existing 2.5D approaches typically inject neighboring 2D slices without cross-plane geometric consistency; MPAD utilizes orthogonal multi-plane autoregression and prior propagation, effectively resolving through-plane discontinuities.

Rating

  • Novelty: ⭐⭐⭐⭐⭐ Formulating 3D MRI synthesis as coordinated 2D masked-slice diffusion over an isotropic latent space with cross-plane prior propagation is highly novel and elegant.
  • Experimental Thoroughness: ⭐⭐⭐⭐⭐ Rigorously tested across ADNI and IXI across 6 translation tasks against 5 established 3D generation baselines, supported by detailed module ablations.
  • Writing Quality: ⭐⭐⭐⭐⭐ Clear motivation, clean mathematical and architectural presentation, and thorough experimental analysis.
  • Value: ⭐⭐⭐⭐⭐ Highly impactful for accelerating clinical MRI acquisition workflows, handling missing sequences, and scaling 3D medical generative models efficiently.