Decoupled Drift-Diffusion World Models for Consistent Future Frame Prediction

  1. Anton Arapin1
  2. Aditya Singh1
  3. Yixin Wang1
  4. Jun Gao1,2
  5. Bernadette Bucher1

1University of Michigan 2NVIDIA

A world model should reproduce what the actions determine, and invent only what they don't. D3WM learns where that boundary is.

D3WM teaser figure
assets/teaser.png Teaser — the drift/diffusion split beside deterministic-only and diffusion-only baselines.

Abstract

World models of visual information predict future image frames conditioned on given action sequences. This predicted scene evolution is driven by deterministic changes such as controlled motion on a robot platform, as well as stochastic changes such as rendering of previously occluded regions. Our key insight is that the spatial location of these stochastic changes is itself learnable.

We leverage this insight by jointly regressing both the deterministic next-state prediction and the spatial uncertainty map, then using the uncertainty map to guide a conditional diffusion process. We propose the Decoupled Drift-Diffusion World Model (D3WM), which uses a transformer trained with a Gaussian negative-log-likelihood objective to produce both the deterministic next-state prediction and a spatial map of heteroscedastic aleatoric uncertainty. The learned uncertainty map then provides a structured spatial gate for a conditional diffusion model: deterministic regions are anchored to the predicted state, while uncertain regions, such as disocclusions or dynamic textures, are sampled by the diffusion process. We frame this decomposition under a variance-gated stochastic differential equation that separates deterministic drift from stochastic diffusion. D3WM shows improved results over competitive baselines in two distinct types of action spaces: trajectory planning for robotic manipulation from generated images, and camera movement in long-horizon novel view synthesis.

Method

A variance-gated SDE over the observation axis

We write the transition from a start latent \(z_0\) under an action sequence \(a_{[0,t]}\) as a stochastic differential equation over the observation time axis, which names the two parts separately:

\[ \mathrm{d}z_t \;=\; \underbrace{f_\theta\!\left(t,\,z_0,\,a_{[0,t]}\right)}_{\text{drift}}\,\mathrm{d}t \;+\; \underbrace{\sigma_t \odot \mathrm{d}w_t}_{\text{diffusion}} \]
(1)

Drift — regress the integrated effect, not the rate

Rather than learning the instantaneous coefficient \(f_\theta\) and integrating it numerically, we regress its integrated effect over \([0,t]\) in a single forward pass. A deterministic transformer returns the conditional mean together with a spatial map of heteroscedastic aleatoric variance:

\[ \mu_t,\ \sigma_t^{2} \;=\; f_\theta^{\mu}\!\left(t,\,z_0,\,a_{[0,t]}\right) \]
(2)

Training by Gaussian negative log-likelihood forces \(\sigma_t^{2}\) to be large exactly where the deterministic prediction is ambiguous — disocclusions, dynamic texture:

\[ \mathcal{L}_{\text{drift}} \;=\; \tfrac{1}{2}\sum_{i,j}\left( \log \sigma_t^{2\,(i,j)} \;+\; \frac{\bigl(z_t^{(i,j)} - \mu_t^{(i,j)}\bigr)^{2}}{\sigma_t^{2\,(i,j)}} \right) \]
(3)

Diffusion — a learned sampler in place of the Wiener term

Integrating Eq. 1 by Euler–Maruyama would draw the stochastic term from an isotropic Gaussian, which fills disoccluded regions with static rather than plausible texture. We replace that term with a conditional ControlNet \(g_\theta\), conditioned on \(\mu_t\), \(\sigma_t^{2}\) and the reference frame \(z_0\) — an image-manifold-aware sampler for the same next-state distribution:

\[ \mathcal{L}_{\text{score}} \;=\; \mathbb{E}_{\tau,\,z_t,\,\epsilon\sim\mathcal{N}(0,I)} \Bigl[\; \bigl\lVert \epsilon - g_\theta\bigl(z_t^{(\tau)},\,\tau,\,\mu_t,\,\sigma_t^{2},\,z_0\bigr) \bigr\rVert_2^{2} \;\Bigr] \]
(4)

Variance-gated inference

At sampling time \(\sigma_t^{2}\) becomes a spatial gate rather than only a conditioning signal. With a diffusion-time anchoring schedule \(\lambda_\tau \in [0,1]\), define a per-pixel trust map, and at each reverse step \(\tau\) replace the diffusion model's latent with a weighted combination of it and a forward-noised version of the deterministic prediction:

\[ \gamma_t \;=\; \bigl(\mathbf{1}-\sigma_t^{2}\bigr)\odot\lambda_\tau \]
(5)
\[ \tilde{z}_t^{(\tau)} \;=\; \gamma_t \odot \tilde{\mu}_t^{(\tau)} \;+\; \bigl(\mathbf{1}-\gamma_t\bigr)\odot z_t^{(\tau)}, \qquad \tilde{\mu}_t^{(\tau)} \;=\; \sqrt{\bar{\alpha}_\tau}\,\mu_t + \sqrt{1-\bar{\alpha}_\tau}\,\xi_\tau \]
(6)

Both terms are valid noisy latents at noise level \(\tau\), so the substitution stays compatible with the next reverse step. Where \(\sigma_t^{2}\!\to\!0\) the conditional latent collapses toward \(\mu_t\) and geometric structure is preserved; where \(\sigma_t^{2}\!\to\!1\) the pixel is handed to the diffusion model and new content is synthesized. The model is deterministic where the actions decide the answer, and generative only where they do not.

D3WM architecture
assets/architecture.png Architecture — the drift network fθ runs once; the diffusion loop runs the variance-gated refinement of Eq. 6.
The deterministic prediction network runs a single forward pass per query. Only the diffusion network iterates.

Results — Real PushT

Predicted frames good enough to plan on

We predict 16 frames ahead from a fixed and a wrist camera, then hand the predicted frame to a pre-trained PushT Diffusion Policy and let it plan. The question is not how the image looks — it is whether a policy can act on it.

  • 3.5× lower trajectory error than iVideoGPT: 12.86 ADE against 45.13.
  • Gripper and object land where the actions put them, instead of the hallucinated grippers and deformed T that iVideoGPT produces.
Trajectory error on Real PushTTable 1
Trajectories generated from predicted frames. Ground truth is the floor: the same policy run twice on the real frame.
MethodADE ↓FDE ↓
Ground truth frame1.5521.319
iVideoGPT45.13465.228
D3WM12.86417.550
Real PushT rollout 1
assets/pusht_1.gif
Real PushT rollout 2
assets/pusht_2.gif
Real PushT rollout 3
assets/pusht_3.gif
Real PushT rollout 4
assets/pusht_4.gif
Predicted rollouts on Real PushT. Drop in four animated GIFs, or swap each block for a <video autoplay loop muted playsinline> element.

Results — BridgeData V2

Predict, rank, and act

Tabletop manipulation from a fixed camera, against IRASim and WorldGym — measured on the three jobs a world model actually has, not just on how well it reconstructs a frame.

  • Best open-loop prediction at every horizon: μt leads PSNR by +1.4 to +1.6 dB over IRASim, with the full model second.
  • Best action verifier. Ranking candidate action chunks, D3WM closes 38–40% of the random-to-oracle gap against IRASim's 23–26%, at half the compute per candidate.
  • 4.9 ms per candidate for the deterministic pass — 204 candidates a second against IRASim's 0.2, the only variant cheap enough to screen at control rate.
  • Driving a frozen Octo policy, the stochastic component is worth 1.063 ADE against 1.445 for the deterministic prediction alone. IRASim still leads here, since it diffuses a whole trajectory where we predict one frame.
Open-loop predictionTable 2
Columns are the prediction gap g in frames. Both reference metrics reward a hedged prediction, which is why the two tables below exist.
MethodPSNR ↑LPIPS ↓
12341234
copy-frame (control)24.4821.4820.0419.19.051.084.105.121
IRASim23.3422.3721.6421.12.101.116.127.136
WorldGym23.8021.1019.7018.86.059.093.118.137
D3WM μt24.7323.5423.0922.76.086.099.107.112
D3WM24.2922.8022.3221.91.104.118.124.130
Candidate verificationTable 3
Candidate action chunks are ranked by predicted similarity to the true t+4 frame; the selected chunk is scored against the demonstrated action. Random and oracle bracket every row.
Verifierms/candK@1sMSE @ 41632
random——.040.038.038
oracle——.006.003.003
IRASim54780.2.031.029.030
WorldGym10940.9.044.056.063
D3WM μt4.9204.054.043.038
D3WM26430.4.027.024.022
Policy consistencyTable 4
A frozen Octo policy acts on the predicted future and on the real one. A perfect simulator scores 0.
World modelADE ↓FDE ↓
copy-frame (control)1.5771.994
WorldGym1.4481.783
D3WM μt1.4451.649
D3WM1.0631.193
IRASim0.8461.035
Qualitative BridgeData V2 comparison
assets/bridge_qualitative.png Qualitative BridgeData V2 — the hedged μt beside the sharp full-model sample, against IRASim and WorldGym.

Results — Novel view synthesis

Extrapolating 12 frames from two views

Camera movement as the action: two input views on the DL3DV benchmark, extrapolating 3 to 12 frames, then the same model transferred to RealEstate10K with no retraining.

  • Beats DepthSplat on PSNR at every horizon, and on all three metrics by 9 frames out — the gap widens with distance as splatting's black holes expand and ours are filled.
  • +3.1 dB PSNR and 39% lower LPIPS than NWM on RealEstate10K, a dataset we never trained on.
DL3DV, in-distributionTable 5
Two-view extrapolation against DepthSplat, from the same two reference frames.
Extrap.MethodPSNR ↑SSIM ↑LPIPS ↓
3 fr.DepthSplat20.870.7610.200
D3WM21.480.6690.222
6 fr.DepthSplat17.390.6430.298
D3WM20.120.6090.283
9 fr.DepthSplat15.400.5570.371
D3WM19.190.5630.338
12 fr.DepthSplat14.140.4960.421
D3WM18.350.5250.382
RealEstate10K, zero-shotTable 6
Out-of-distribution transfer against NWM. Trained on DL3DV only.
Extrap.MethodPSNR ↑SSIM ↑LPIPS ↓
3 fr.NWM18.450.6320.233
D3WM21.550.7900.142
6 fr.NWM15.860.5540.315
D3WM19.210.6190.189
9 fr.NWM14.450.5090.375
D3WM17.630.5560.237
12 fr.NWM13.570.4800.422
D3WM16.780.5220.266
Qualitative NVS comparison
assets/nvs_qualitative.png Qualitative NVS — reference, deterministic μt, variance map, D3WM, DepthSplat, ground truth.
The variance map column is the part to look at: it locates the disocclusions before anything is generated.

Citation

BibTeX

@article{
}