Plateaus, sudden drops and double descent
A flat curve is not always a dead run, and a validation curve that turns upward does not always keep going up — two shapes that break the standard rules, with runnable demonstrations of each.
- 14 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
A curve can sit flat for ages and then fall off a cliff, or get worse before it gets much better. Both look like failure. Neither is.
Remember learning to ride a bicycle. Days of wobbling and putting a foot down, with nothing visibly improving. Then one afternoon you rode the length of the lane, and you never lost it again.
Nothing gradual happened on that last day. The balance was being assembled the whole time, invisibly, and then it arrived all at once.
Training curves do this. A flat stretch is called a plateau. It can mean the run is dead. It can also mean the model is halfway through assembling something.
Why this had to be invented
The standard advice is comfortable: if the curve is flat, kill the run and change something. If validation starts rising, stop, because it will keep rising.
Both rules are right most of the time. Both have famous exceptions, and the exceptions cost real money when people do not know they exist.
An epoch is one full pass over the training data. Killing a run on epoch 200 that would have solved the task on epoch 320 is a silent, unrecorded failure. So is buying a bigger model, watching it do worse, and concluding that bigger is wrong. A much bigger one might have beaten both.
You do not need to memorise the theory. You need to recognise the two shapes so you do not throw away a working run.
How it works
Shape one, the plateau and the cliff.
loss
0.69 |------------------------\
| flat for 300 epochs \
| accuracy stuck at 50% \
| (looks completely dead) \___________
0.00 | ---
+-----------------------------------------------
0 300 400 600 epochThe flat value carries information. A loss stuck near 0.69 on a two-choice problem means the model is answering "fifty-fifty" to everything. It has not started learning at all yet.
Shape two, worse then better.
error on new data
| /\
| / \ <- the bad zone, right where the model
| \___/ \ is exactly big enough to memorise
| \______
+-------------------------------
small medium very large model sizeThe bad zone sits exactly where the model has barely enough capacity to memorise its training data and nothing left over. Give it far more capacity and the problem goes away. This is called double descent — the error descends, climbs, and descends again.
Picture bending a thin metal strip through a row of nails. A strip with exactly as many bends as nails is forced into a violent zigzag. A strip with far more bends can pass through every nail while staying gentle.
A real example you have seen
Language learning. Months of vocabulary drilling with no ability to hold a conversation, then a trip where it suddenly works. The plateau was real, and so was the work happening underneath it.
Double descent shows up in the largest models in the world. Teams have watched a mid-sized model do worse than a small one on the same data, and a much larger one beat both. That surprise is the reason anyone bothered to name the shape.
Remember this
- A flat curve at 0.69 on a two-choice task means "the model is guessing", not "the model is broken".
- Before killing a plateau, give it a patience budget you decided in advance.
- Bigger can get worse before it gets better. One bad size does not prove that size is the wrong direction.
What to learn next
- Too small a model, or trained too little? — separating the two causes this lesson deliberately mixed.
- Overfit one batch — the fastest test for whether a plateau is a bug.
- Buying accuracy with size, and when to stop — how to sweep capacity without wasting a month.
Developer — Code and libraries.
Setup
pip install torch numpyVerified with torch 2.5.1 (CPU), numpy 1.26.4, Python 3.10. The first script takes about ten seconds, the second under two seconds. Everything is generated in memory.
A plateau that ends in a cliff
Parity — "is the number of negative inputs odd?" — is the classic task for this. Every input bit matters and no single bit helps, so the network learns nothing until it has assembled the whole function.
import itertools, torch, torch.nn as nn
torch.manual_seed(2)
bits = 8
X = torch.tensor(list(itertools.product([-1., 1.], repeat=bits)))
y = ((X < 0).sum(1) % 2).float().unsqueeze(1) # parity of the negative bits
net = nn.Sequential(nn.Linear(bits, 32), nn.Tanh(),
nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1))
opt = torch.optim.SGD(net.parameters(), lr=0.5, momentum=0.9)
loss_fn = nn.BCEWithLogitsLoss()
losses, accs = [], []
for _ in range(601):
opt.zero_grad(); loss = loss_fn(net(X), y); loss.backward(); opt.step()
losses.append(loss.item())
with torch.no_grad():
accs.append((((net(X) > 0).float() == y).float().mean().item()))
for e in range(0, 601, 40):
print(f"{e:5d} loss {losses[e]:.4f} acc {accs[e]:.3f}")0 loss 0.6994 acc 0.508 40 loss 0.6932 acc 0.512 80 loss 0.6931 acc 0.484 120 loss 0.6931 acc 0.508 160 loss 0.6931 acc 0.504 200 loss 0.6931 acc 0.488 240 loss 0.6928 acc 0.551 280 loss 0.6860 acc 0.547 320 loss 0.2533 acc 0.906 360 loss 0.1192 acc 0.965 400 loss 0.0263 acc 0.996 440 loss 0.0053 acc 1.000 480 loss 0.0015 acc 1.000 520 loss 0.0009 acc 1.000 560 loss 0.0007 acc 1.000 600 loss 0.0005 acc 1.000
Two hundred and eighty epochs of nothing, then perfect accuracy sixty epochs later. Killing this run at epoch 250 — which almost everyone would — throws away a model that solves the task exactly.
The walkthrough
0.6931 is ln(2), and it is the single most useful number in this section. It is the loss of a model that outputs 50/50 on a balanced two-choice problem. Seeing it means "constant output", which points at a small set of causes: the learning rate is too small, the gradient is not reaching the early layers, the inputs carry no usable signal yet, or the task genuinely needs assembling. For $k$ balanced classes the equivalent number is ln(k): 1.099 for three classes, 2.303 for ten.
Watch epochs 240 to 320. Loss moves 0.6928 → 0.6860 → 0.2533. The plateau was not perfectly flat at the end; there is a small, real decline immediately before the cliff. That tiny slope is the early-warning signal, and a smoothing setting wide enough to calm the plateau would erase it (how smoothing and log scales mislead you).
This is seed-dependent, and the honest version of that fact matters. Other seeds on this exact script escape the plateau in under 300 epochs, and some take longer. "How long is the plateau" is a random variable, not a property of the task — check it the way noisy validation curves checks everything else, by running several seeds.
Distinguish a live plateau from a dead run in about a minute. Print the gradient norm each epoch. A dead run has gradients near zero. A live plateau has small but non-zero gradients, slowly reorganising the weights (gradient flow diagnostics). Separately, confirm the model can fit — overfit one batch is the five-minute test that separates "hard task" from "broken code".
Double descent you can run in two seconds
Now the second shape. Fit a minimum-norm least-squares model on 40 noisy training rows, sweeping the number of random features through and past 40.
import numpy as np
rng = np.random.default_rng(0)
n_train, d = 40, 10
Xtr = rng.normal(size=(n_train, d)); Xte = rng.normal(size=(2000, d))
w = rng.normal(size=d)
ytr = Xtr @ w + rng.normal(0, 0.5, n_train) # noisy labels are essential here
yte = Xte @ w
def features(X, p, seed=1):
g = np.random.default_rng(seed)
W = g.normal(size=(X.shape[1], p)) / np.sqrt(X.shape[1])
return np.maximum(X @ W, 0) # random ReLU features
print("features train MSE test MSE")
for p in [5, 10, 20, 30, 38, 40, 42, 50, 80, 200, 1000, 5000]:
Ftr, Fte = features(Xtr, p), features(Xte, p)
beta = np.linalg.pinv(Ftr) @ ytr # minimum-norm least squares
tr = ((Ftr @ beta - ytr) ** 2).mean()
te = ((Fte @ beta - yte) ** 2).mean()
print(f"{p:8d} {tr:9.4f} {te:8.3f}")features train MSE test MSE
5 3.5741 4.341
10 1.2686 2.381
20 0.7949 3.572
30 0.2615 3.553
38 0.0113 8.146
40 0.0000 277.654
42 0.0000 20.597
50 0.0000 2.825
80 0.0000 0.938
200 0.0000 0.441
1000 0.0000 0.458
5000 0.0000 0.390Test error at 40 features is 277.654. At 5,000 features it is 0.390 — better than any smaller model, and 700 times better than the disaster in the middle.
Reading the double-descent table
The spike sits exactly at features = n_train = 40. That point is the interpolation threshold: the smallest model that can fit all 40 training points exactly. It has precisely enough freedom to be forced through every noisy label and none left over to stay smooth. The metal strip with one bend per nail.
Past the threshold, extra capacity is a good thing. Many different parameter settings now fit the training data perfectly, and pinv picks the one with the smallest norm — the gentlest of the available solutions. More capacity means more candidates to be gentle with.
The noisy labels are not decoration. Set rng.normal(0, 0.5, n_train) to zeros and the spike shrinks dramatically. The peak is the model contorting to fit noise it cannot know is noise.
Train MSE hits 0.0000 at 40 and never leaves. A training curve alone cannot see any of this. Every conclusion here needed held-out data.
Common mistakes
Killing every plateau on sight. Decide a patience budget before the run — "kill if no improvement for N evaluations" — and record it. That is a rule you can defend; "it looked dead" is not.
Treating 0.69 as a mysterious number. It is ln(2). Compare your flat value against ln(k) for your number of classes before doing anything else, and against the baseline you must beat right after.
Expecting epoch-wise double descent everywhere. The version you can most reliably reproduce is model-wise, as above. The epoch-wise version — validation loss rising and then falling again during a single run — is real but needs label noise and specific settings. Do not sit through 500 epochs of rising validation loss hoping for it on an ordinary problem.
Concluding "bigger is worse" from one bad size. You may have measured a single point near the interpolation threshold. If capacity is the variable under test, sweep it over an order of magnitude, not two adjacent values — buying accuracy with size covers how to run that sweep economically.
Ignoring regularisation. Ridge with a properly tuned penalty flattens most of the double-descent peak. If your curve has a bad zone, ridge regression style shrinkage is usually a cheaper fix than a much larger model.
Try it yourself
In double_descent.py, replace the noise term with np.zeros(n_train) and rerun. Watch the peak at 40 collapse. Then restore the noise and swap np.linalg.pinv(Ftr) @ ytr for a ridge solution with a small penalty, and see how much of the spike survives.
What to learn next
- Too small a model, or trained too little? — separating the two causes this lesson deliberately mixed.
- Overfit one batch — the fastest test for whether a plateau is a bug.
- Buying accuracy with size, and when to stop — how to sweep capacity without wasting a month.
Researcher — Mathematics and papers.
Plateaus: saddles, symmetry and the escape time
Flat regions in deep networks are dominated by saddle points rather than local minima. Dauphin et al. (2014), Identifying and Attacking the Saddle Point Problem in High-dimensional Non-convex Optimization (NeurIPS), argue from random matrix theory that in high dimension critical points with small loss are overwhelmingly saddles: for a random Hessian, the probability that all $n$ eigenvalues are positive falls exponentially in $n$, and the index of a critical point concentrates as a function of its loss. Gradient descent near a strict saddle escapes along the negative-curvature direction, but the escape time scales with the inverse of the negative eigenvalue's magnitude — arbitrarily long plateaus for arbitrarily flat saddles.
A second, distinct mechanism is permutation and sign symmetry. At initialisation, hidden units of a layer are exchangeable; the loss surface contains manifolds where units are functionally identical, and the network cannot make progress until symmetry breaks. Saad and Solla (1995), On-line Learning in Soft Committee Machines (Physical Review E), derived exactly this for the teacher-student setting: a long symmetric plateau in which all student units align to the same teacher direction, ending in a specialisation transition. Parity is a canonical hard case — its Fourier spectrum is concentrated on a single degree-$k$ character, so every subset of fewer than $k$ inputs is uninformative, and the statistical-query lower bounds of Kearns (1998) imply that gradient-based learners need either many queries or high precision.
Grokking (Power et al., 2022, Grokking: Generalization Beyond Overfitting on Small Algorithmic Datasets) is the sharpest documented version: training accuracy saturates early while validation stays at chance for orders of magnitude more steps, then jumps. Subsequent analyses (Nanda et al., 2023, Progress Measures for Grokking via Mechanistic Interpretability, ICLR) show the transition is not sudden internally — a Fourier-basis circuit forms gradually and the metric is what changes discontinuously. The practical lesson generalises well beyond modular arithmetic: loss is a lagging indicator of internal structure, and progress measures computed from the weights can reveal movement during an apparently dead plateau.
Double descent: the interpolation threshold
Belkin et al. (2019), Reconciling Modern Machine Learning Practice and the Classical Bias–Variance Trade-off (PNAS), named the shape and located the peak at the interpolation threshold, where model capacity first suffices to fit the training set exactly. Nakkiran et al. (2021), Deep Double Descent (ICLR 2020; JSTAT 2021), demonstrated it in modern networks and identified three axes along which it appears — model size, training epochs, and dataset size (where more data can hurt at fixed model size) — unified by an effective model complexity.
The linear case is solved. For minimum-norm least squares with $p$ features and $n$ samples, the risk diverges as $p \to n$ because the smallest singular value of the design matrix approaches zero, so $|\hat\beta| = |F^{+}y|$ blows up; the Marchenko–Pastur law gives the precise rate. Hastie, Montanari, Rosset and Tibshirani (2022), Surprises in High-Dimensional Ridgeless Least Squares Interpolation (Annals of Statistics), give exact asymptotic risk in the proportional regime $p/n \to \gamma$, and their central practical result is that optimally-tuned ridge regression removes the peak entirely and is monotone in $\gamma$. Double descent is therefore best understood as a pathology of implicit regularisation being insufficient near the threshold, not as evidence against the bias–variance decomposition. Belkin, Rakhlin and Tsybakov (2019) and Bartlett et al. (2020) (Benign Overfitting in Linear Regression, PNAS) characterise when interpolation of noise is harmless: the spectrum of the covariance must have enough small directions to absorb the noise without corrupting the signal directions.
What this changes about reading a curve
Three revisions to the default rules, each narrow.
- "Flat means dead" becomes "flat means measure the gradient". Gradient norm, weight-norm drift and a task-specific progress measure separate a live plateau from a dead one at negligible cost.
- "Validation rising means stop" keeps its default status. Epoch-wise double descent requires label noise and particular capacity regimes; treating it as a routine possibility would justify burning compute on genuinely failed runs. Keep early stopping, and note the exception when the setting matches (when validation loss rises but accuracy improves is the far more common cause of that shape).
- "Bigger got worse, so stop scaling" becomes a sampling question. Two adjacent sizes cannot distinguish a real ceiling from a threshold peak. Sweep capacity geometrically, and tune regularisation at each size, since an untuned penalty is what makes the peak visible in the first place.
What to learn next
- Too small a model, or trained too little? — separating the two causes this lesson deliberately mixed.
- Overfit one batch — the fastest test for whether a plateau is a bug.
- Buying accuracy with size, and when to stop — how to sweep capacity without wasting a month.