← Drive vs. Decay (NeurIPS 2026)

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.

t = 10.0
small predictor (α = 0.8)
large predictor (α = 2)
Size of the representation
1e-41e-21e01e2starttraining time
JEPA loss
1e-81e-51e-21e11e4training time
A: small predictor (α = 0.8)B: large predictor (α = 2)
A: directions learned
3 of 6
B: directions learned
0 of 6
The paper's linear model, solved exactly: 6 data directions, context keep probability 0.7, EMA β = 1, predictor held fixed. Cloud shape is exact; its overall size is drawn on a log scale.

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)

W˙=−Wp⊤WpWΣ11+Wp⊤M¯2WaΣX1
W˙a=β(W−Wa)

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

w˙=−σw+γwa,w˙a=β(w−wa)

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.

μ = γ / σ
1.40
leading root λ
0.183
verdict
learns

second root -2.18

One mode of the paper's linear model (Lemma 3.1), solved exactly from 14 starting points. Arrows: the direction the gradient flow and the EMA push the pair (w, wₐ).

Whatever EMA speed you choose, the curves leave the origin exactly when γ is larger than σ. The reason is the characteristic equation of each direction,

λ2+(β+σ)λ+βσ(1−μ)=0,μ=γσ

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.

1e-41e-21e01e21e4startβ = 0.2β = 1β = 5training timesize of the direction
β = 0.2: time to ×10
47.9
β = 1: time to ×10
16.4
β = 5: time to ×10
9.6
One mode of the linear model, decay σ = 1 and drive γ = μ, started at w = wₐ = 1, at three EMA speeds. Exact solution.

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.

Proven
In the linear model with spectral decoupling, a direction leaves collapse if and only if μ > 1, and β changes the rate but not the sign. Near collapse the criterion still holds when the predictor is trained jointly with the encoder (Lemma 3.4).
Measured
How far real covariances are from decoupled is measured on the paper's datasets (Appendix B). Without decoupling, the appendix gives two certificates, one that forces collapse and one that forces escape.

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):

μkℓ=(ckak)·(qℓpℓ)

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

μk=ptgtskα(psk+1−p)

The directions with μk > 1 can escape, and their count is a rank budget for the representation (Corollary 3.3).

0.31312345678910μ = 1
direction (strongest data direction first)
1e-61e-31e01e3starttraining timesize of each direction
μ > 1: growsμ < 1: decays
predicted rank budget (μ > 1)
3 of 10
measured in the simulation
3 of 10
The paper's linear model with 10 data directions (power-law spectrum, unit-diagonal covariance), predictor αI, EMA β = 1, solved exactly. Predicted = directions with μ > 1; measured = directions still growing at the end of the simulation.

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.

Python · A. The full model, trained by SGD
.py

With the default settings the two columns agree to the second decimal, and the number of directions that grow equals the number with μ > 1.

Proven
Corollary 3.3: the growing part of the encoder has rank at most the number of super-critical directions.
Measured
Across 810 linear Tabular-JEPA runs, the rank budget tracks the measured effective rank with a correlation of 0.91 on helena and aloi, and 0.62 on jannis, whose top spectrum is near-degenerate.

5Map the phase

Predictor scale and masking span a plane. Colour each point by μ of the strongest direction and the plane splits in two.

α, p
hover the map
μ of the top direction
–
rank budget
–
The paper's linear model: μ of the strongest direction over predictor scale and context masking, for the spectrum you set. Thick line: μ = 1. Thin lines: where the 2nd to 5th directions cross μ = 1, so the rank budget is constant between two lines.

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:

Three heatmaps (ALOI, HELENA, JANNIS) of normalised linear-probe score over log predictor scale and context ratio, green on the left, red on the right, with a red mu = 1 contour near log alpha between -1.5 and -0.5.
measuredFigure 4 of the paper: sweeps over predictor scale and context ratio on three tabular datasets. Colour is the normalised linear-probe score; the thick red curve is μ = 1.

Cell B sweeps the predictor scale and checks the rank budget of the theory against SGD runs.

Python · B. Sweep the predictor scale
.py

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.

0.31312345678910μ = 1
direction (strongest data direction first)
μ(α) = μ(1) / α
directions that escape
3 of 10

The drive grows like α and the decay like α², so a larger predictor at initialisation pushes every direction toward collapse.

The paper's linear model (10 directions, base predictor scale 1, context keep probability 0.7), with each mechanism written as its derived effect on the decay σ or the drive γ (Table 3 and Appendix J). The paper checks every derived count against 152 simulations of the linear model; all 152 match.

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.

Python · C. Your own mechanism
.py

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:

μk=ptgtskpsk+1−p

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,

Attention(Q,K,V)=softmax(QK⊤d+βinit·I)V

so that at initialisation every token attends to itself. Raise the bias and watch the grey turn into a diagonal.

row: query tokencolumn: attended token
saturation bound log n + 5 = 7.5
smallest self-attention Aᵢᵢ
0.057
effective rank of A
3.3 of 12
One attention head at initialisation: 12 context tokens, queries and keys at the scale of a standard initialisation (logit standard deviation about 0.16), softmax(QKᵀ/√d + β_init·I). The threshold log n + 5 is the paper's saturation bound (Appendix D); the paper uses β_init = 10.

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).

Python · D. ResidualPred's attention in numpy
.py

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.

I-JEPA, standard predictor+ ResidualPred
20%40%60%80%CIFAR-10+4.1CIFAR-100+4.0STL-10+4.1ImageNet-100+5.6ImageNet-1k pilot*+6.5
Table view
DatasetI-JEPA+ ResidualPred
CIFAR-1064.55 ± 1.7168.61 ± 0.59
CIFAR-10036.83 ± 1.4940.85 ± 0.66
STL-1071.25 ± 1.1475.31 ± 0.78
ImageNet-10037.09 ± 0.6842.73 ± 0.41
ImageNet-1k pilot*25.22 ± 0.8031.69 ± 0.22
Linear-probe accuracy of I-JEPA with the standard predictor and with ResidualPred, mean over seeds with ± one standard deviation (paper, Table 4; ImageNet-1k from Table 5). *ImageNet-1k: a pilot at a matched 30-epoch budget, ViT-S/16, five seeds per arm; tripling the schedule to 90 epochs attenuates the margin.

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 pilotLinear probe (%)Effective rank
I-JEPA25.22 ± 0.80107.2 ± 7.7
+ SIGReg29.19 ± 0.54185.6 ± 5.6
+ ResidualPred31.69 ± 0.22188.5 ± 3.3
+ ResidualPred + SIGReg32.23 ± 0.54206.7 ± 1.0
Proven
In linear networks with spectral decoupling: the μ > 1 criterion, the rank budget, and the dose of each of the seven mechanisms. Near collapse the criterion also holds with the predictor trained jointly.
Measured
The phase boundary across more than 800 Tabular-JEPA configurations, and the gains of ResidualPred on images. Whether μ is what carries those gains in deep networks is not established: at initialisation a whole-block estimator does not separate the two predictors, and the first link of the mechanism is measured after a few predictor-only steps, in 9 of 9 matched comparisons.

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.