Do-JEPA: From Masking to Intervention in Latent World Models
Organizations: Australian Institute for Machine Learning, Adelaide University · Responsible AI Research Centre
Abstract
Latent world models are trained to predict what happens next, so nothing in their objective separates what an action caused from what merely co-occurred with it. Object-masking models such as C-JEPA intervene on what the predictor can see; we intervene on what physically happens. From one saved simulator state we run the dynamics under an action and under a reference action , and train the model to predict the difference between the two latent futures. The resulting objective, Do-JEPA, has an effect loss, a support loss (where the action enters), a propagation loss (where its effect travels) and invariance losses (what must not change). In a synthetic system with object-aligned variables, support supervision finds the directly intervened object in 99.95% of test cases, where a sparse action mask sends the action to a nuisance slot in every case, and response-onset supervision recovers the ring-shaped propagation graph (edge AUROC 0.975 vs. 0.624). From pixels, the effect loss beats a control trained on exactly the same data: it lowers latent effect error by 28.4% on an end-to-end LeWM model and physical effect error by 13.5% when trained and tested on natural action sequences, and on three independently generated CausalWorld benchmarks it lowers responsive effect error by about 20% under physics shifts and the latent context sensitivity of predicted effects by 66%. Trained from scratch it costs factual accuracy; fine-tuning an existing model with it removes this cost. Together, these results show that intervening on the world, rather than on what the model sees, helps latent world models predict what their actions cause.
Figures & tables
| Top-1 all [obj.] | Target | Nuis. mask | Nuis. effect | Effect MSE | Pred. MSE | |
|---|---|---|---|---|---|---|
| Mask (global action) | — | — | — | |||
| Mask | — | — | — | |||
| Sparse mask | [ ] | |||||
| Sparse mask | [ ] | |||||
| Sparse mask | [ ] |
| Benefit | Cost | |||||
|---|---|---|---|---|---|---|
| OOD split | Resp. effect RMSE | Effect cosine | Latent ctx. shift | Factual RMSE | Response AUROC | Signature |
| Visual | lower [2/3] | [3/3] | lower [3/3] | higher [3/3] | [3/3] | |
| Mechanism ∗ | lower [3/3] | [2/3] | lower [3/3] | higher [3/3] | [3/3] | |
| Composed | lower [3/3] | [2/3] | lower [3/3] | higher [3/3] | [3/3] | |
| Factual RMSE | Resp. effect RMSE | Response | Latent ctx. | Factual | |||
|---|---|---|---|---|---|---|---|
| Regime | IID | OOD | IID | Physics | AUROC | shift | within |
| From scratch, | |||||||
| From scratch, | |||||||
| Fine-tuned from Aug , | |||||||
Appendix figures & tables24 assets
Supplementary material from the paper’s appendix.
Appendix
| Finding | Evidence | Where |
|---|---|---|
| Support supervision finds where the action enters | direct target found in of test cases; a sparse mask picks a nuisance slot in every case | Table 1 |
| Onset supervision recovers where the effect travels | edge AUROC vs. ; the four ring edges recovered exactly on seeds | Table 12 |
| Gates confine effects to the right objects | state Push-T edge AUROC vs. ; predicted effect on nuisances | Table 14 |
| Context twins remove a nuisance shortcut | on the two seeds not used to pick the weight: OOD error – lower, – of the IID-to-OOD gap closed, IID cost | App. E.3 |
| Intervention beats masking on C-JEPA’s own slots | all six metrics, seeds: factual error lower, AUROC , pixel context shift lower | Table 15 |
| Effect loss helps from pixels (LeWM) | latent effect error lower ( ); physical effect error lower when trained and tested on natural sequences ( ) | Fig. 3 , Table 18 |
| Symbol | Meaning | Shape | Role | When |
|---|---|---|---|---|
| physical simulator state | env.-specific | observed | build/eval | |
| action; reference action (a physical no-op unless stated) | input | train, test | ||
| context (nuisance) value and its twin | env.-specific | input | train (twin), test ( ) | |
| confounder: affects action choice and outcome | scalar | hidden | build | |
| exogenous simulator noise, shared by both branches | — | hidden | build | |
| observation (image or state vector) | env.-specific | observed | train, test |
| Loss | Question | Target | Source of the target |
|---|---|---|---|
| what happens next? | , | observed next frames | |
| what did the action change? | paired rollouts, Eq. 3 | ||
| where did it enter? | first step of the pair | ||
| where did it travel? | onset order in the pair | ||
| what must not change? | for | response set of the pair | |
| what must ignore context? | context twin |
| Information | Training | Test | Build/eval only |
|---|---|---|---|
| Observation history, action | input | input | — |
| Reference branch | input ( Aug ), target ( ) | — | — |
| Context twin | input ( Aug ), target ( ) | — | — |
| Support, onset and response labels | target (from the pair) | — | — |
| Object count and slot alignment | architecture (slot models) | architecture | synthetic metrics |
| Ground-truth graph, object identities | — | — | metrics |
| Setting | Model input | Intervention vs. reference | Shift at test | Terms used | Train/val/test | Seeds |
|---|---|---|---|---|---|---|
| Synthetic SCM | slots, | impulse at a noisy contact point vs. none | reversed action–nuisance corr. | k/ k/ k | – | |
| Temporal SCM | slots, | impulse vs. none | observation stride | k/ k/ k | ||
| State Push-T | state + nuisance slots | action vs. no-op | reversed corr. ( ) | k/ k/ k | ||
| Vision Push-T | frozen slots, | action vs. no-op | reversed pattern corr. ( ) | k/ k/ k | ||
| LeWM Push-T | RGB, , | -step macro vs. zero macro | none; mass, visual (planning) | k/ k pairs | ||
| CausalWorld | RGB, , | first joint command vs. mirrored one | visual, mechanism, composed | ( ) | k/ k/ k per split | model, data |
| Corpus | States | Resp. | Action source | |
|---|---|---|---|---|
| Enriched (E) training | episode replay | balanced toward/away proposal | ||
| Outcome-unfiltered holdout | episode replay | i.i.d. from expert action marginal | ||
| Fresh outcome-unfiltered holdout | episode replay | i.i.d. from expert action marginal | ||
| Expert contiguous | expert prefix | contiguous expert macro | ||
| Expert shuffled | expert prefix | same macro, order shuffled | ||
| Expert i.i.d. marginal | expert prefix | i.i.d. from expert marginal |
| Setting | Settings |
|---|---|
| Hard SCM | Transformer predictor (width , depth , heads), history , propagation steps ( ); epochs, batch , lr , wd ; , , mask temperature (sparse-mask rows without ) or (with , and all gated models), entropy regulariser (a cardinality term, also , is zero under the softmax mask), no-op mask penalty ; every model is trained and evaluated with one random history slot masked (masked-history loss ); edge runs add (threshold ), , gate invariance , gate , onset threshold , gate temperature . Sel.: lowest IID prediction MSE. Seeds (ablations), (gated model). |
| Easy SCM | Transformer predictor (width , depth , heads); epochs, batch , lr , wd ; , , mask temperature , entropy regulariser ; / / samples. Sel.: lowest IID prediction MSE. Seeds . |
| Temporal SCM | Width , depth , propagation steps; epochs, batch , lr , wd ; , , , . Sel.: lowest IID prediction MSE. Seeds . |
| State Push-T | Width , depth , heads; gate temperature ; epochs, batch , lr , wd ; reference-branch weight , , reward head , (threshold ), , gate invariance , gate , (context twins by shuffling nuisances within a batch; evaluated by sign flip). Sel.: IID prediction and effect error (context not used). Seeds . |
| Vision Push-T | Frozen slots; transformer predictor (width , depth , heads); epochs, batch , lr , wd ; , (soft weights, temperature ), context-twin prediction weight , history loss , (sensitivity run ). Sel.: IID validation. Seeds . |
| LeWM Push-T ( Paired vs. Aug ) | Released LeWM checkpoint (ViT-tiny/14 encoder, -D latent, history , -layer predictor); epochs, batch , lr , wd , SIGReg , reference-branch weight , , gradient clipping , validation split, expert replay in every update. Seeds (development), (confirmation). Sel.: lowest validation paired-prediction loss, the same rule for Aug and Paired (no effect or planning metric). |
| Pred. MSE | Effect MSE | Nuis. effect | Nuis. mask | |
|---|---|---|---|---|
| Mask (global action) | — | |||
| Mask | — | |||
| Sparse mask | ||||
| Sparse mask | ||||
| Sparse mask |
| Edge AUROC | True-edge gate | Off-path gate | Nuis. in-gate | Nuis. effect | Direct-effect MSE | |
|---|---|---|---|---|---|---|
| Gates without | ||||||
| (onset labels) | ||||||
| gate term of | ||||||
| same, 3 seeds |
| Direct-target | ||||
|---|---|---|---|---|
| Descendant-effect | ||||
| One-hop structural | ||||
| One-hop structural AUROC | ||||
| Onset-label | ||||
| Schedule-aware |
| Pred. MSE | Agent pos. | Nuis. effect | Edge AUROC | Nuis. in-gate | |
|---|---|---|---|---|---|
| Obs (global action) | — | — | |||
| Global action, | — | — | |||
| Action routed to agent, no gates | |||||
| Routed, gates, , gate invariance |
| Factual (px) | Resp. effect (px) | Cosine | AUROC | Latent ctx. | Phys. ctx. (px) | |
|---|---|---|---|---|---|---|
| Mask (C-JEPA-style) | ||||||
| Aug | ||||||
| Do-JEPA |
| Factual (px) | Resp. eff. (px) | Cos. | AUROC | Lat. ctx. | Phys. ctx. (px) | |
|---|---|---|---|---|---|---|
| Object-centric JEPA, no mask | ||||||
| Mask | ||||||
| Paired pred. | ||||||
| Paired pred. ctx. twins ( Aug ) | ||||||
| Paired | ||||||
| Do-JEPA ( ) |
| Held-out training-distribution pairs | Outcome-unfiltered holdout | |||||
|---|---|---|---|---|---|---|
| Seed | Effect RMSE | Cosine | Factual MSE | Effect RMSE | Cosine | Factual MSE |
| Mean change | ||||||
| Trained on | Tested on | Latent effect RMSE (%) | Physical effect RMSE (%) | Physical cosine |
|---|---|---|---|---|
| Natural | natural test ( ) | |||
| Natural | unfiltered holdout ( ) | |||
| Enriched | natural test ( ) | |||
| Enriched | unfiltered holdout ( ) |
| Held-out training-distribution pairs | Outcome-unfiltered holdout | |||||
|---|---|---|---|---|---|---|
| Seed | Effect RMSE | Cosine | Zero | Decoded true | Effect RMSE | Cosine |
| / | ||||||
| / | ||||||
| / | ||||||
| Hypothesis | Split | Resp. effect RMSE | Cosine | AUROC | Factual RMSE | Ctx. shift | Seeds | Result |
|---|---|---|---|---|---|---|---|---|
| H1: Paired vs. Aug | visual | lower | higher | — | not met | |||
| mechanism | lower | higher | — | not met | ||||
| composed | lower | higher | — | not met | ||||
| H2: Do-JEPA vs. Paired | visual | higher | higher | lower | not met | |||
| mechanism | higher | higher | lower | not met | ||||
| composed | higher | higher | lower | not met |
| Split | Factual RMSE (%) | Resp. effect RMSE (%) | Effect cosine | Response AUROC | Latent ctx. shift (%) |
|---|---|---|---|---|---|
| IID | |||||
| Visual | |||||
| Mechanism ∗ | |||||
| Composed |
| Configuration | Model | Adapted | Effect RMSE | Cosine | Planning |
|---|---|---|---|---|---|
| Base LeWM | — | — | — | — | |
| Full fine-tuning | Aug | yes | |||
| Paired | yes | ||||
| Expert data only | — | yes | |||
| Frozen representation | Aug | yes | |||
| Paired | yes |
| Property | Measured | Check |
|---|---|---|
| Prop. 1 : no-op correction | (bitwise) | pass |
| Prop. 2 : linearity | max. error | pass |
| Prop. 2 : operator bound | holds for all pairs; mean tightness | pass |
| Prop. 3 : sampled | (lower bound, not certified) | — |
| Spectral / Frobenius norm of | / | — |
| Rank of ( ) | by construction | — |
| Seed 17 | Seed 27 | Seed 37 | |
| Effect RMSE, Aug | |||
| Effect RMSE, Paired | |||
| Relative reduction | |||
| Rollout ratio, Paired | |||
| No-op and representation unchanged | yes | yes | yes |
| Planning, Aug |