All posts

World models22 min read

Can a Fly Learn the World with Causal‑JEPA?

We built a 4,096-neuron fly brain from connectome statistics and put it in a virtual dish with wind, an odour plume, a light and a hunting spider. Then we asked whether a self-supervised world model could learn how that world works from what the brain senses — and reuse that knowledge in worlds it had never seen.

Watch a world model learn. The same moment of a fresh episode in T1 (wind and a spider together, never seen in training), replayed at every saved training checkpoint. Dots show where Causal-JEPA imagines the spider over the next 0.64 s; the dashed line is where it really goes. As training proceeds the imagined spider closes in on the real one (error 73 → 12 mm with perfect entity slots, violet; 53 → 23 mm with slots read from the fly brain, green), and the map of which entity it relies on settles onto the true causal links, from 2 of 5 at the start to 4 of 5. The fly brain itself stays fixed.

The short version

  • Our first attempt, a JEPA with an explicit learned causal graph, failed in an instructive way. Its graph emptied out and its latent blocks became copies of one another. That failure pointed us toward Causal-JEPA.
  • Causal-JEPA works in FlyWorld when the world is already split into entities. Given one slot per world factor it predicts 0.64 s ahead far better than chance, performs worse when the world's mechanisms are scrambled as a real world model should, and recovers 4 of the 5 true causal links.
  • The fly's own brain holds the information needed. Slots read linearly from the brain's neurons recover almost the same causal structure as the true world factors.
  • Finding those entities without labels is the remaining challenge. Unsupervised slot attention over the brain's activity collapsed into six identical slots, and four attempted fixes failed to separate them.

01 · The questionCan a brain learn a world model it can reuse?

An animal doesn't just react, it expects. A fly that has felt wind and smelt food knows the smell comes from upwind, so when the wind turns it knows where the plume will go. That internal picture of how things work is a world model. A good one predicts what happens next. A reusable one keeps working when familiar parts are recombined — wind and a spider together, when the fly has only ever met them separately.

A world model that is reusable in this way has to capture causes, not just correlations. It has to know that the spider follows the fly, not the smell, even when the two always appear together.

Hypothesis

A self-supervised, causally structured world model trained on a connectome-constrained fly brain's experience (1) predicts well in recombined worlds it never saw, and (2) has internal dependencies that match the simulator's true causal graph.

The second prediction is the strong one. A model can transfer well just by being a good predictor, but it can't recover the true causal graph by accident.

02 · First principlesWhat the pieces are

World models and self-supervised learning

A world model maps what an agent senses now to what it expects next. It needs no labels, because the future itself is the target: the agent sees frames \(o_{1..t}\), predicts something about \(o_{t+1..t+H}\), and is corrected when the future arrives.

Predict pixels, or predict meaning?

Reconstruction models predict every raw sensor value of the future, including pure noise. A JEPA (joint-embedding predictive architecture; LeCun 2022, Assran et al. 2023) first encodes each observation into a compact latent vector and predicts future latents. The model chooses what to keep, so it can ignore noise. That freedom is also a risk: a model that chooses what to keep can discard what matters, or collapse to a constant that is trivially predictable. JEPAs prevent collapse with a slowly updated target encoder and variance/covariance penalties (VICReg).

Causality in four ideas

  • Structural causal model. Every variable has its own mechanism, \(x_i := f_i(\mathrm{pa}(x_i), \varepsilon_i)\), computed from its parents and noise.
  • Causal graph. An arrow \(j \to i\) whenever \(x_j\) is an input to \(f_i\).
  • Intervention \(do(x_j := c)\): replace one mechanism, like switching a fan off. Correlations can't predict what that does; a causal model can.
  • Sparse mechanism shift (Schölkopf et al. 2021): when the world changes, usually only a few mechanisms change. A model made of separate mechanisms can adapt by changing only those. That is what makes it reusable.

Objects first

Recovering a world's true factors from raw observations is impossible in general without extra assumptions (Locatello et al. 2019). One powerful assumption is that the world is made of objects. Slot-based encoders such as Slot Attention split a scene into a set of entity vectors, one per object. Once the world is a set of entities, “which entity affects which” becomes a question a predictor can answer.

Connectome-constrained brains

FlyWire mapped an entire adult fruit-fly brain: about 139,000 neurons and 55 million synapses (Dorkenwald et al. 2024). Lappalainen et al. (2024) showed that a network whose wiring and synapse signs are fixed from a connectome, with only its physiology learned, predicts real neural responses. We use the same constraint.

03 · What we assumedAssumptions, stated up front

AssumptionWhyWhat it could hide
Five factor blocks, known graphGives a ground-truth answer to score againstReal worlds are messier; success here is necessary, not sufficient
Surrogate connectomeKeeps the fly's regional structure at a trainable sizeIt is fly-like, not a real fly brain
Wiring and signs fixedThe defining constraint of connectome-constrained modelsReal brains also rewire
Linear read-outs measure contentSimple and comparableNon-linear content is missed, so numbers are lower bounds
Persistence is the baselineZero means no better than doing nothingNearly static factors give noisy scores

04 · The worldFlyWorld: a world with a known answer

FlyWorld is a 100 mm dish simulated at 50 Hz in 12-second episodes. It is written directly as a structural causal model, so its true causal graph is known exactly. It has five factor blocks: wind (a 2D vector drifting toward a mean), odour (a plume blown downwind from a food drop, so its shape depends on the wind), predator (a jumping spider that wanders, stalks the fly within 30 mm and lunges within 7 mm), light (slow drift) and self (the fly's body, moved by its own commands and pushed by wind).

That gives 6 true edges out of 20 possible: the plume depends on the wind, the spider chases the fly, and the fly's body is affected by everything. There is also a deliberate trap. Odour and predator are correlated, because the fly goes to the food and the spider follows the fly — but the spider's rule never reads the odour.

What the fly senses: 108 channels

  • Eyes: 2 × 48 brightness columns. Ambient light sets the background of all 96, and a spider is a dark patch. No looming signal is given; the brain has to notice the patch growing.
  • Antennae: 2 × 4 odour channels.
  • Johnston's organs: 2 wind channels, measuring the wind felt by the moving fly.
  • Halteres: turning rate and speed.

Why the fly's behaviour is scripted

The fly moves by a hand-written policy: up the odour gradient, into the wind, away from looming shapes, plus a slowly drifting random wander. The spider is a state machine. We did this on purpose. This study asks whether a brain can learn a model of its world from what it senses, and that question is separate from learning to act. A fixed policy gives identical data for every model, interesting situations from the first episode, and a clean separation of perception from control. The fly's actions are still given to the world model as inputs, so it learns how actions change what happens next. Closing that loop is the natural next step.

Six test worlds

Models train in world A, where wind and spiders never appear together, and are then frozen and tested in six worlds. Two are calibration cells: T0 should score well, and T5 should score below T1.

WorldWhat changesWhat it tests
T0 nullNew random seeds onlyCalibration: should do well
T1 compositionWind and a spider togetherReusing mechanisms in a new combination
T2 parametersSpider 2× faster, wind reversed, longer plumeMechanisms as functions of their settings
T3 two predatorsTwo spidersReusing one mechanism twice
T4 antenna lesionRight antenna silencedA dead input
T5 scrambledMechanisms read the wrong parentsCalibration: a model that learned mechanisms must get worse
Six dish diagrams, one fresh episode in each test world, showing the fly's path, the spider's path and the plume
One fresh episode in each test world, as the simulator sees it from above. The brain never sees this view: it receives only the 108 sensor channels.

05 · The brainA connectome-constrained network

The brain has 4,096 neurons wired by a surrogate connectome — a graph generated to match FlyWire's published statistics, such as region sizes, connection densities and the share of inhibitory neurons. As in a real fly, most input neurons serve vision (about 1,227, against 73 for wind). Each neuron's voltage leaks toward rest and is pushed by its inputs:

$$v_i \leftarrow v_i + \frac{\Delta t}{\tau_i}\Big( -(v_i - v^{\text{rest}}_i) + g\sum_j A_{ij}\, r_j + w^{\text{in}}_i\, s_{c(i)} + b_i \Big), \qquad r_i = \min\big(\mathrm{softplus}(v_i-\theta_i),\, r_{\max}\big)$$

Here \(r_i\) is the firing rate, \(\tau_i\) a time constant, \(\theta_i\) a threshold, and \(s_{c(i)}\) the sensory channel wired to neuron \(i\) if it is an input neuron. Connections \(A_{ij}\) exist only on connectome edges, and their signs are fixed by Dale's law: a neuron is either excitatory or inhibitory, never both. Training may tune strengths, time constants and thresholds, but never which neurons connect or the sign of any synapse.

06 · First attemptGraph JEPA, and why it failed

Our first model was our own design, which we call Graph JEPA: a JEPA whose latent is split into five blocks (one hoped-for block per factor), with an explicit learned causal graph between the blocks. One shared mechanism updates each block from its parents. Edges are sampled as hard 0/1 choices and pay an L1 penalty, so only edges that help prediction should survive. Everything, the brain's readout included, was trained end to end.

It failed in three ways:

  1. The graph emptied. All 20 candidate edges started at 0.27 and decayed together to 0.02. None ever pulled away from the rest.
  2. The blocks became copies. All five ended up encoding mostly light and the fly's own motion. No block held wind, odour or the spider on its own.
  3. It wasn't learning mechanisms. It predicted as well in the scrambled world T5 (skill 0.144) as in T1 (0.146).

A series of one-change-at-a-time diagnostic runs showed why. Holding light constant made the blocks copy the next-easiest factors instead. Penalising correlation between blocks didn't separate them. We found and fixed a real wiring bug — the wind sensor reached no neuron, because haltere inputs had overwritten the wind input neurons — which put wind into the brain but changed nothing else. Most tellingly, hard-wiring every intervention to its correct block (oracle routing) still produced copies, and even supervising each block with its factor's labels produced no graph edges.

The lesson

The causal machinery sat downstream of a latent that was learned only by prediction, and five copies of the easiest signal predict themselves perfectly. With no distinct entities, there is nothing for a causal graph to connect. The representation has to be split into entities first.

07 · Causal-JEPAAn object-centric causal world model

Causal-JEPA (Nam, Le Lidec, Maes, LeCun and Balestriero, ICML 2026) learns a world model over a set of frozen entity slots. There is no explicit graph. Instead, causal structure is forced in by masking whole entities.

How it learns

Each frame is a set of slot tokens, one per entity, plus two auxiliary tokens: the fly's action and its proprioception. The model sees 5 history frames and predicts 3 future frames, with a bidirectional transformer attending over every token of every frame. Two kinds of token are hidden: all future slot tokens (actions stay visible, as a planner would supply them), and the whole history of \(|M|\) randomly chosen entities, except their first frame, which is kept as an identity anchor. A hidden token is replaced by a projection of that entity's anchor state plus a time embedding, and the loss is the squared error on every hidden token against the frozen slots themselves.

Why hiding an entity's past matters

If you can see an entity's past, the easiest way to predict its present is to extrapolate it. That shortcut is exactly what emptied Graph JEPA's graph. Hide the past, and the only way to reconstruct the entity is to use the other entities that influence it — so the model is pushed to learn interactions. With \(|M| = 0\) it is the paper's OC-JEPA baseline: object-centric, but with no causal pressure.

Three stages: where do the slots come from?

FlyWorld's brain sees no images, so we built the slots three ways, from easiest to hardest. This separates “can Causal-JEPA learn the structure?” from “can we find the entities?”

StageSlot sourceWhat it tests
A · oracleThe simulator's true factors, one slot per blockUpper bound: perfect entities
B · brain, read outEach factor read out of all 4,096 neurons by one linear map fitted on world A, then frozenIs the information in the brain, and is it usable?
C · brain, unsupervisedSlot attention over 39 groups of neurons, trained without labels, then frozenEntity discovery from neural activity alone

How we score it

Imagination skill compares the model's imagined future with the guess that nothing moves (persistence). One linear decoder, fitted once in world A, translates slots into the 17 true world-factor numbers:

$$\text{skill} = 1-\frac{\lVert f_{t+16} - M\hat s_{t+16}\rVert^2}{\lVert f_{t+16} - M s_t\rVert^2}$$

1 is perfect, 0 is no better than “the world froze”. To read a causal graph out of a model that has no explicit graph, we hide entity \(i\)'s history, then also hide entity \(j\)'s, and measure how much worse entity \(i\) is predicted:

$$D_{ij} = \frac{\mathrm{err}(i \mid \text{hide } i, j) - \mathrm{err}(i \mid \text{hide } i)}{\mathrm{err}(i \mid \text{hide } i)}$$

A large \(D_{ij}\) means the model relies on \(j\) to reconstruct \(i\). The transformer is bidirectional, so it can't tell cause from effect. The fair test is therefore whether large values fall on the 5 truly linked pairs out of 10, scored by AUC (0.5 is chance).

08 · ResultsIt works when the world is split into entities

Stage, maskingworld AT0T1T5T1 − T5Mechanisms?Pair AUCTop-5
A, |M| = 00.6440.6390.3050.253+0.052yes0.763/5
A, |M| = 10.6260.6190.2320.143+0.089yes0.804/5
A, |M| = 20.6260.6140.3080.247+0.061yes0.764/5
B, |M| = 00.2650.3590.1260.078+0.048yes0.643/5
B, |M| = 10.2510.3220.1150.035+0.081yes0.844/5
B, |M| = 20.2510.3260.1300.078+0.053yes0.763/5
C, |M| = 10.111−0.0130.0950.063+0.032voidslots identical

“Mechanisms?” asks whether the model got worse in the scrambled world T5 than in T1, by more than 0.02. A model that learned how the world works must. One that learned surface statistics won't notice.

A model must do worse in the scrambled world. Plain JEPA and every Graph JEPA variant fail, even with oracle routing and factor labels. Causal-JEPA passes with oracle slots and with slots read out of the brain. Stage C's apparent pass is meaningless, because its slots are copies.

True graph (child ← parent)

windodourpredlightself
wind·0000
odour1·000
pred00·01
light000·0
self1111·

Stage B, |M| = 1 · influence

windodourpredlightself
wind·0.460.040.030.28
odour0.41·0.050.010.15
pred0.040.06·0.030.16
light0.020.010.03·0.03
self0.050.050.170.02·

Rows are the entity being reconstructed; columns the entity whose history is also hidden; true pairs outlined. Stage B ranks 4 of the 5 true pairs on top — wind–odour, wind–self, predator–self and odour–self — missing only light–self. In stage C every value is about the same, because every slot is the same.

Three findings stand out.

  1. The scrambled-world test passes for the first time. No version of Graph JEPA ever managed it.
  2. The brain holds usable structure. Stage B's slots, read linearly out of an untrained connectome brain, recover the same interactions as the true factors.
  3. Masking adds causal pressure. \(|M| = 1\) gives the largest T1 − T5 drop and the best pair AUC in both stages, at a small cost in world-A skill.

Caveats on the good news

Direction isn't recovered, by design: the model learns which entities interact, not which causes which. Transfer is still limited — stage A falls from 0.62 in T0 to 0.23–0.31 in T1. And stage B used labels to build its slots, so it shows the information is in the brain, not that it can be found without supervision.

09 · InferenceWatching Causal-JEPA imagine

To see what the numbers mean, we ran the frozen \(|M| = 1\) models on fresh episodes they had never seen, in worlds they were never trained in. At every frame, each model takes the last 5 frames of slots and the fly's upcoming actions, imagines the next 0.64 s, and we decode that imagination back into the dish.

How to read the videos

Left, the dish from above: the plume (orange), food (green), wind (grey arrow), the fly (blue arrow) and the spider (red). Dashed black line: where the spider really goes over the next 0.64 s. Coloured dots: where each model imagines it will be, growing larger further ahead — violet is stage A (oracle slots), green is stage B (slots read from the brain). Top right: what the brain receives from its eyes. Middle and bottom right: true spider distance and sideways wind, against what each model imagined 0.64 s earlier.

T1 — wind and a spider together, a combination never seen in training. Stage A's imagined distance follows the spider's approach; across all T1 clips its spider prediction clearly beats persistence (skill 0.37, against 0.16 for stage B). Stage B's imagined spider is noisier and pulled toward a middle distance, because the brain read-out holds the spider less cleanly.
T3 — two spiders. The slots describe only the nearest spider, so the model has to track whichever one is closest. Stage A tracks the nearest spider's distance well in this episode, though training never showed two spiders.
T5 — scrambled mechanisms: here the spider chases the food, not the fly. A model that learned “spiders chase flies” should be wrong here, and that is the point. Across all T5 clips, spider prediction is worse than in T1 (0.37 → 0.27 for stage A, 0.16 → −0.02 for stage B) — the drop a model that learned how spiders behave should show.

10 · The unsolved stepFinding entities without labels

Unsupervised slot attention over the brain's activity produced six identical slots: light 0.98, self 0.63, odour 0.31, predator 0.28, wind 0.20 — the same mixture in every slot. We tried four fixes in parallel.

FixIdeaRecon. errorSpecificityOutcome
Remove common modeSubtract the shared components from every neuron group0.32−0.43No change
Linear decoderA weaker decoder, so one slot can't rebuild everything0.96−0.41Too weak to learn
3× longer training12,000 steps instead of 4,0000.18−0.43Better reconstruction, still copies
All threeCombined0.91−0.41Still copies

Specificity measures whether each slot holds its own factor and not the others. Positive means separated; negative means copies. One run looked like a success, with the linear decoder scoring 0.54 in T1. It isn't one: its T0 skill is only 0.08, its graph is at chance, and a model with identical slots can't genuinely do better in a new world than in the one it trained on.

Why slot attention fails on brain activity

Slot attention finds objects by grouping input tokens, as in an image where a spider occupies some patches and the background others. Groups of neurons don't divide up that way. Every group carries a mixture of light, self-motion, odour and the rest — mixed selectivity. With no grouping to find, the cheapest solution is for every slot to summarise everything.

11 · ConclusionsWhat we learned

Verdict

Causal-JEPA can learn a causal, partly reusable world model of FlyWorld from a fly brain's activity, if something first splits that activity into entities. With oracle or brain-read slots it passed the scrambled-world test and recovered most of the true interactions. Without labels, nothing we tried split the brain's activity into entities.

  • Entities come first. Our explicit-graph model failed because its latent never separated into entities, and no causal machinery downstream could fix that.
  • Entity masking is a workable causal signal. Hiding an entity's past forces the model to use the entities that influence it, and the dependencies it learns line up with the true graph.
  • The connectome brain carries the information. A linear read-out of its neurons gives slots good enough for causal learning.
  • Unsupervised entity discovery from neural activity is the open problem. Neurons mix many factors, so slot attention has no grouping to find. The split has to come from somewhere else.

Next steps

  1. Slot attention on raw sensory channels rather than neurons. There, entities really are localised: a spider covers some eye columns, and wind has its own channels.
  2. Close the loop. Let the brain's own outputs drive the fly, and use the world model's imagination to choose actions, so the fly learns by doing.
  3. Interventions that shape the encoder, so entity separation can be learned from experiment rather than given.