Hold the knob: a hands-on tour of Drive vs. Decay
This page accompanies our NeurIPS 2026 paper on why joint-embedding predictive architectures (JEPAs) collapse. You will train a JEPA in your browser, watch it collapse, and find the number that decided it before reading the theorem that says so.
Chapters 1 to 7 run the paper's linear model, solved exactly on this page, and each widget says what it computes. Chapter 8 turns to the experiments on real networks. Allow twenty minutes. The four Python cells are optional: the first Run downloads a Python runtime (about 10 MB) into your browser, and nothing is sent anywhere.
1Make it collapse
A JEPA has three parts. A context encoder reads a masked view of the input, a target encoder reads the full input, and a predictor maps the context representation onto the target one. The loss is the squared distance between prediction and target. No gradient flows through the target: its weights follow the context encoder through an exponential moving average (EMA).
There is an easy way to bring that loss down, which is to map every input to zero. Prediction and target then agree, and nothing has been learned. Press Train both. The two runs differ only in the scale of the predictor at initialisation.
Run B reaches the lower loss, and its representation shrinks to a point. Run A keeps three directions alive and grows along them. Nothing in the linear model caps that growth, so A's loss grows with its scale; what matters is the direction each run takes in its first steps, because everything after starts from there. A falling loss is therefore no evidence that a JEPA is learning, and the paper tracks the effective rank of the embeddings instead.
2Find the one number
To see what separates A from B, the paper linearises the dynamics. With a linear context encoder W, target encoder Wa and predictor Wp, gradient flow and the EMA give (Lemma 3.1)
Collapse is the point W = Wa = 0, and the question is whether training leaves it. When the data covariances and the predictor share eigenvectors (Assumption 3.2), the system splits into independent directions. Each one is a pair of numbers (w, wa) that obeys
The drive γ pulls the context weight toward the target, and the decay σ pulls it back to zero. Both come from the predictor and the data, and chapter 4 says how. Drag the sliders: the arrows show where the dynamics push each pair, and the curves start from fourteen points around the origin.
second root -2.18
Whatever EMA speed you choose, the curves leave the origin exactly when γ is larger than σ. The reason is the characteristic equation of each direction,
whose constant term changes sign at μ = 1. Above 1 one root is positive and the direction grows. Below 1 both roots are negative and the direction returns to collapse.
3Drag the clock
The EMA is often treated as a stabiliser against collapse. In the equation above, β multiplies the constant term but cannot change its sign, so the EMA decides how fast a direction escapes or dies and not whether it does. Move μ across 1 and compare three EMA speeds.
The three curves always turn together. At μ = 1.3 the slowest EMA needs about five times longer than the fastest to grow tenfold, and all three grow.
4Count the survivors
A representation has many directions, and each has its own μ. Under decoupling the ratio factorises into a data part and a predictor part (Eq. 7):
where ak and ck are the context and cross covariances along data direction k, and pℓ, qℓ come from the predictor. Take feature masking with keep probability p, a predictor αI, a target keep probability ptgt, and a unit-diagonal data covariance with spectrum sk. Then
The directions with μk > 1 can escape, and their count is a rank budget for the representation (Corollary 3.3).
Three things to try. Shrink the predictor scale and watch directions join one at a time: μ is inversely proportional to α, because the drive grows like α and the decay like α². Flatten the spectrum: when every direction carries the same variance they share one μ and cross the threshold together. Move the keep probability: it shifts strong and weak directions in opposite ways, which is why the effect of masking depends on the data.
The widget uses the closed form. The cell below does it the hard way: a full linear JEPA, trained by stochastic gradient descent on sampled data with sampled masks, with nothing decoupled by hand. It prints the growth rate the theory predicts for every direction next to the one measured in training.
With the default settings the two columns agree to the second decimal, and the number of directions that grow equals the number with μ > 1.
5Map the phase
Predictor scale and masking span a plane. Colour each point by μ of the strongest direction and the plane splits in two.
Predictor scale is the strong knob: it divides every μ by α, so it can move any direction across the line. Masking moves the boundary too, by about half a decade with the default spectrum, and the spectrum decides how far. At p = 1 the context is the whole input, every direction has the same data factor, and all the lines meet at α = ptgt. The paper measured the same plane on three tabular datasets, with a linear-probe score at every point and the μ = 1 contour in red. There too predictor scale does most of the work, and masking bends the boundary, most visibly on jannis at low context ratios:

Cell B sweeps the predictor scale and checks the rank budget of the theory against SGD runs.
6Seven tricks, one ratio
Practitioners prevent collapse with a collection of devices. The paper takes seven of them and writes each as a term in the linear model (Table 3 and Appendix J). Each acts on μ, through the drive or through the decay, by an amount the paper calls a dose. Pick one and change its dose.
The drive grows like α and the decay like α², so a larger predictor at initialisation pushes every direction toward collapse.
Read across the tabs and the devices sort themselves. Predictor scale and weight decay raise the decay. SIGReg and the VICReg hinge lower the predictor's share of it. The identity predictor takes the predictor out of μ altogether, and the EMA leaves μ untouched. These are statements about the linear model; how each device behaves in a deep network is a separate, empirical question.
Cell C lets you add a device of your own. Write its term in the dynamics and the dose you expect it to have, and the cell checks one against the other. The default is weight decay.
7Fix it with the identity
If a large predictor pushes μ down, start from a predictor that does nothing. Write it as the identity plus a correction that starts at zero, Wp = I + ΔWp. Then the predictor factor is set by the target mask alone, and μ no longer depends on how the predictor was initialised:
A transformer predictor is not a matrix, but its attention can start as one. ResidualPred adds a bias on the diagonal of the attention logits,
so that at initialisation every token attends to itself. Raise the bias and watch the grey turn into a diagonal.
Without the bias, attention at initialisation is close to uniform: every token receives nearly the same average of the others, so the differences between tokens are washed out, and the effective rank of the attention matrix is a fraction of the number of tokens. Near log n + 5 every self-weight reaches 0.99 and the matrix is the identity. The paper uses βinit = 10 everywhere. In the released code the bias is a learnable parameter, one per head, on the diagonal of the context block. It adds no loss term and no cost per step.
The identity is where the predictor starts, not where it should stay. Frozen at initialisation and denied learning, the ResidualPred predictor collapses the representation (Section 6.2).
8Beyond the linear model
Everything so far is the linear model. The results on deep networks are measurements, and the paper keeps the two apart. ResidualPred was dropped into I-JEPA on images, with the encoder and the training recipe unchanged.
Table view
| Dataset | I-JEPA | + ResidualPred |
|---|---|---|
| CIFAR-10 | 64.55 ± 1.71 | 68.61 ± 0.59 |
| CIFAR-100 | 36.83 ± 1.49 | 40.85 ± 0.66 |
| STL-10 | 71.25 ± 1.14 | 75.31 ± 0.78 |
| ImageNet-100 | 37.09 ± 0.68 | 42.73 ± 0.41 |
| ImageNet-1k pilot* | 25.22 ± 0.80 | 31.69 ± 0.22 |
On the ImageNet-1k pilot, ResidualPred also compares well with SIGReg, the regulariser of LeJEPA, and the two combine (Table 5, five seeds each):
| ImageNet-1k pilot | Linear probe (%) | Effective rank |
|---|---|---|
| I-JEPA | 25.22 ± 0.80 | 107.2 ± 7.7 |
| + SIGReg | 29.19 ± 0.54 | 185.6 ± 5.6 |
| + ResidualPred | 31.69 ± 0.22 | 188.5 ± 3.3 |
| + ResidualPred + SIGReg | 32.23 ± 0.54 | 206.7 ± 1.0 |
9Cheat sheet
- My JEPA's loss went to zero. Is that good news?
- Not by itself. The collapsed solution has zero loss too. Track the effective rank of the embeddings. See chapter 1
- Will a slower EMA prevent collapse?
- In the linear model, no. β changes how fast directions grow or die, never which ones. See chapter 3
- How large should the predictor be at initialisation?
- Small enough that μ > 1 for the directions you want to keep. μ scales like 1/α. See chapter 4
- Does more masking help?
- It depends on the spectrum of the data. In the linear model, masking moves the strong and the weak directions in opposite ways. See chapter 4
- Should I add SIGReg or VICReg?
- In the linear model both lower the predictor's share of the decay, which raises μ. On the ImageNet-1k pilot, SIGReg and ResidualPred stack. See chapter 6
- What is the cheapest change to try?
- Start the predictor's attention at the identity: one hyperparameter, β_init = 10, no extra loss term and no per-step cost. See chapter 7
Run it for real
The four cells are also a notebook: drive_vs_decay_tutorial.ipynb (numpy and matplotlib only). The experiments of the paper, tabular and image, are in the code repository.