Skip to content
Kudos AI

The Step, and the Edge of Stability

What a gradient step is minimising, why a minibatch gradient is the full gradient plus noise rather than a different gradient, and the exact step size above which descent stops descending.

IntermediateModule 125 min · 100 XP
A loss surface with its two curvature directions drawn as arrows, a gradient step contracting the error along each one by a factor 1 − ηλ, and the same run at 0.99 and 1.01 times 2/L - one settling, the other flung off the frame.

Every model on this site is fitted the same way. Linear regression, logistic regression, a neural network, a transformer: each one writes down a number that says how wrong it currently is, asks which way that number goes down, and moves. The model changes; the loop does not.

What does change, and what decides whether the loop takes forty steps or blows up in three, is the geometry around it. This lesson is about the smallest piece of that geometry: one step.

What the loop is minimising

Training minimises empirical risk - the average loss over the data you have:

f(w)  =  1n∑i=1nℓ ⁣(w;xi,yi).f(w) \;=\; \frac{1}{n}\sum_{i=1}^{n} \ell\!\left(w; x_i, y_i\right).

The examples throughout this path use a least-squares problem with n=400n = 400 points and two parameters, because everything about it can be computed exactly and compared against what the theory claims. Its loss is f(w)=1n∥Xw−y∥2f(w) = \frac{1}{n}\lVert Xw - y\rVert^2, its minimiser is w⋆=(1.973613,−2.947388)w^\star = (1.973613, -2.947388), and the loss there is f⋆=0.233943f^\star = 0.233943. Knowing the answer in advance is the point: it turns every claim below into a measurement.

A minibatch gradient is the full gradient plus noise

Computing ∇f(w)\nabla f(w) means touching all nn examples. A minibatch of size BB touches BB of them and averages:

gB(w)  =  1B∑i∈B∇ℓ(w;xi,yi).g_B(w) \;=\; \frac{1}{B}\sum_{i \in \mathcal{B}} \nabla \ell(w; x_i, y_i).

Because B\mathcal{B} is a uniform sample, E[gB(w)]=∇f(w)\mathbb{E}[g_B(w)] = \nabla f(w) exactly. This is worth stating carefully, because the alternative belief - that a small batch points somewhere systematically different - leads to the wrong conclusions about batch size.

At the origin w=(0,0)w = (0, 0) the full gradient is (−3.416495,0.179597)(-3.416495, 0.179597). Averaging 20,000 independent minibatches at that same point gives:

batch size BBmean errortypical size of the noise
80.0321.4382
320.0120.7052
1280.0020.3027

The first column shrinks toward zero because it is Monte-Carlo residue in the average, not bias. The second column is the real story: sixteen times the batch buys 1.4382/0.3027=4.751.4382 / 0.3027 = 4.75 times the precision, close to the 16=4\sqrt{16} = 4 that averaging independent draws predicts. The 19% excess over 4 is not noise: these batches are drawn without replacement from only n=400n = 400 points, and the finite-population factor (n−B)/(n−1)(n - B)/(n - 1) shrinks a batch of 128 far more than a batch of 8. That predicts 4392/272=4.804\sqrt{392/272} = 4.80, and drawing with replacement, which really is independent, gives 3.89 instead. That exchange rate is why training runs use small batches and many steps rather than the reverse. The compute buys precision at the square root, and it buys steps linearly.

The factor (1 − ηλ)

Now the step itself. Near a minimum the loss looks like a quadratic bowl, and the shape of that bowl is the Hessian HH. Write the error as et=wt−w⋆e_t = w_t - w^\star; a gradient step with rate η\eta gives

et+1  =  (I−ηH) et.e_{t+1} \;=\; (I - \eta H)\, e_t .

Decompose ete_t along the eigenvectors of HH. Each component is simply multiplied by 1−ηλ1 - \eta\lambda, where λ\lambda is that direction's curvature, and the directions do not interact at all. One step is therefore not one motion; it is as many independent contractions as there are curvature directions, each running at its own rate.

The component shrinks when ∣1−ηλ∣<1\lvert 1 - \eta\lambda \rvert < 1, which is exactly η<2/λ\eta < 2/\lambda. Every direction must shrink, so the binding constraint is the sharpest one:

η  <  2L,L=λmax⁡(H).\eta \;<\; \frac{2}{L}, \qquad L = \lambda_{\max}(H).

The cliff, measured

On this problem L=1.760627L = 1.760627, so the threshold is 2/L=1.1359592/L = 1.135959. Two runs from the same start, two hundred steps each:

step sizeas a multiple of 2/Lloss after 200 steps
1.1246000.990.236592
1.1473191.0123585.65

A two per cent change in the step separates a converged run from one five orders of magnitude away. This is not a gentle degradation with a grey zone in the middle; it is a sign change in 1−ηL1 - \eta L, and once that factor is below −1-1 the sharpest direction is amplified by a constant factor every step forever.

The practical shape of this is familiar to anyone who has watched a loss curve go to NaN within a few dozen steps after a learning-rate change that looked harmless. Nothing was unstable and then became unstable; the run crossed a threshold that was there from the beginning.

Interactive: the step, and the edge of stability

η = 0.6816, and the cliff is at 2/L = 1.1360.

w*sharp direction →↑ flat direction
Contraction per step
0.949667
Steps for six digits
> 260
Cliff at 2/L
1.1360
Try:

Six digits are out of reach in 260 steps at this setting. The condition number is 23.8410 - a mild elongation, on a problem with two parameters - and it alone sets the rate. Raise β and the zig-zag across the valley cancels while the crawl along it accumulates: the dependence drops from κ to √κ.

Stable is not the same as fast

Knowing LL does not mean setting η=2/L\eta = 2/L. Stability and speed are separate questions, and they have different answers.

The error contracts by max⁡(∣1−ηL∣,∣1−ημ∣)\max\big(\lvert 1 - \eta L\rvert, \lvert 1 - \eta\mu \rvert\big) each step, where μ=λmin⁡(H)\mu = \lambda_{\min}(H) is the flattest direction. Raising η\eta speeds up the flat direction and slows the sharp one, so the best step is where the two costs meet:

η⋆  =  2L+μ.\eta^\star \;=\; \frac{2}{L + \mu}.

Here μ=0.073849\mu = 0.073849, giving η⋆=1.090230\eta^\star = 1.090230 - noticeably below the divergence threshold. Every step size between η⋆\eta^\star and 2/L2/L is slower and closer to the cliff.

That formula is worth checking rather than trusting. Running the same 200 steps at each of 2000 step sizes and taking the one that ends closest to w⋆w^\star gives an empirical optimum of 1.088582, against the derived 1.090230 - agreement to 1.6×10−31.6 \times 10^{-3}. The small gap is not the grid: over a finite 200 steps the best step sits just below the asymptotic optimum, and closes on it as the run lengthens. A derivation and a grid search are three lines apart, and the grid catches the sign slips that algebra hides.

What this does not yet explain

Two facts sit uncomfortably beside each other. The best step on this problem is around 1.09, and μ\mu is 0.074 - so the flattest direction moves by only ημ≈8%\eta\mu \approx 8\% of its remaining error per step even at the optimal rate. That ratio, not the step size, is what makes training slow, and it has a name and a fix.

The name is the condition number, and the fix is the next lesson.

References & further reading

  • Ian Goodfellow, Yoshua Bengio, Aaron Courville, Deep Learning, MIT Press (Adaptive Computation and Machine Learning), 2016source ↗
  • Stephen Boyd, Lieven Vandenberghe, Convex Optimization, Cambridge University Press, 2004source ↗

Copyrighted works are cited for reference only and are not hosted here; please consult the publisher for access.

Unlock the full path

This first lesson is free. Enrol to take the mastery quiz, earn XP, and unlock every module, with more interactive, runnable examples throughout.