Reading Training Curves

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.

On this page 5
  1. Why this had to be invented
  2. How it works
  3. A real example you have seen
  4. Remember this
  5. What to learn next

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  epoch

The 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 size

The 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

Developer — Code and libraries.

Setup

bash
pip install torch numpy

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

plateau.py
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}")
Output
    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.

double_descent.py
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}")
Output
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.390

Test 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

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