ml.lab
Python sleeps until you run code
06 Neural networks from scratch

Lesson 6 of 6

Momentum and Adam

Plain gradient descent has no memory and one step size for every parameter. Momentum adds a memory of recent gradients, which speeds up steady directions and damps oscillations. Adam also gives every parameter its own step size. This lesson derives both, bias correction included, and has you implement them.

About 40 minutes
By the end you can
  • Explain what momentum does in a long, narrow valley, and compute a momentum run by hand.

  • Derive the Adam update as running averages of the gradient and its square, and explain why its first step is η\eta in every coordinate.

  • Derive Adam's bias correction from a geometric series.

  • Implement SGD with momentum and Adam, and compare them with plain gradient descent on a badly scaled problem.

The loop from the last lesson updates every parameter with w←w−η g\mathbf{w} \leftarrow \mathbf{w} - \eta\,\mathbf{g}, where g\mathbf{g} is the gradient of the current batch. That rule has two weaknesses you have already seen. It has no memory: each step uses only the current gradient, so a direction that points the same way step after step never builds up speed, and a direction whose gradient flips sign every step keeps zig-zagging. And it has one step size for everything: the learning rate η\eta must be small enough for the steepest direction, which leaves the gentle directions crawling.

The optimizer is the rule that turns gradients into updates, and two changes to plain gradient descent fix these weaknesses. Momentum gives the parameters a memory of recent gradients. Adam adds a separate step size for every parameter. Adam, in a variant called AdamW, is the optimizer behind most large language models.

Momentum: build up speed

Gradient descent met momentum on a quadratic valley and worked out exactly when it helps. Here is the short version, and what changes when the gradients come from mini-batches.

In a long, narrow valley, where the curvature is large across the valley and small along it, the learning rate has to be small enough for the steep direction, so progress along the gentle direction crawls. Momentum (Polyak's "heavy ball", 1964) gives the parameters a velocity v\mathbf{v}. Module 5 wrote the momentum coefficient as β\beta; here it is μ\mu (mu), because Adam below needs β1\beta_1 and β2\beta_2 for something else. Typically μ=0.9\mu = 0.9, and each step does two things, starting from v=0\mathbf{v} = \mathbf{0}:

v←μ v+g,w←w−η v.\mathbf{v} \leftarrow \mu\,\mathbf{v} + \mathbf{g}, \qquad \mathbf{w} \leftarrow \mathbf{w} - \eta\,\mathbf{v}.

The first line decays the old velocity by the factor μ\mu and adds the current gradient g\mathbf{g}. The second moves the weights along the velocity instead of along the gradient. Unrolled, the velocity is a running sum of past gradients in which each older gradient counts μ\mu times less. Three consequences:

  • Steady directions speed up. If the gradient is the same g\mathbf{g} every step, the velocity after tt steps is g(1+μ+μ2+⋯+μt−1)\mathbf{g}(1 + \mu + \mu^2 + \cdots + \mu^{t-1}). That geometric series approaches 11−μ\frac{1}{1 - \mu}, so the velocity approaches g/(1−μ)\mathbf{g}/(1 - \mu). With μ=0.9\mu = 0.9 that is 1/0.1=101/0.1 = 10 times the plain step. Along the floor of the valley, where the gradient points the same way step after step, momentum accelerates.
  • Oscillating directions cancel. Across the valley the gradient flips sign every step, so consecutive contributions to the velocity mostly cancel, and the zig-zag is damped. Module 5 showed the price: with coefficient μ\mu, no direction can shrink faster than a factor μ\sqrt{\mu} per step, so on a valley that plain descent already handles well, heavy momentum is slower.
  • Mini-batch noise averages out. The velocity sums about the last 1/(1−μ)1/(1 - \mu) gradients (10 of them for μ=0.9\mu = 0.9). Each mini-batch gradient is the true gradient plus noise; the true parts add up step after step, while noise from different batches points in different directions and partly cancels. So the velocity is a steadier guide than any single batch gradient.

A worked example, one parameter, μ=0.9\mu = 0.9, η=0.1\eta = 0.1, and a steady gradient g=2g = 2:

  1. v1=0.9(0)+2=2v_1 = 0.9(0) + 2 = 2, and the parameter moves 0.1×2=0.20.1 \times 2 = 0.2.
  2. v2=0.9(2)+2=1.8+2=3.8v_2 = 0.9(2) + 2 = 1.8 + 2 = 3.8, and it moves 0.1×3.8=0.380.1 \times 3.8 = 0.38.
  3. v3=0.9(3.8)+2=3.42+2=5.42v_3 = 0.9(3.8) + 2 = 3.42 + 2 = 5.42, and it moves 0.1×5.42=0.5420.1 \times 5.42 = 0.542.

After three steps it has moved 0.2+0.38+0.542=1.1220.2 + 0.38 + 0.542 = 1.122, against 3×0.2=0.63 \times 0.2 = 0.6 for plain gradient descent, and the step size keeps growing toward ηg/(1−μ)=0.1×2/0.1=2\eta g/(1 - \mu) = 0.1 \times 2/0.1 = 2.

The widget answers one question: when does momentum help, and when does it hurt? It shows the narrow valley L=12x2+5y2L = \frac{1}{2}x^2 + 5y^2 from module 5 as a contour map: each closed curve joins points of equal loss, the lime cross marks the minimum at the center, and the blue dot is the start at (−2.8,1.6)(-2.8,\allowbreak 1.6). The curvature is 10 across the valley (the yy direction) and 1 along it (the xx direction), so plain descent is stable only for η<2/10=0.2\eta < 2/10 = 0.2. Play runs gradient descent from the start and draws its path; the chart under the map plots how far the loss is above its minimum, as log⁡10(L−L∗)\log_{10}(L - L^*) with L∗=0L^* = 0 the lowest possible loss, step by step (a straight falling line means the gap shrinks by the same factor every step), and a live note gives each direction's shrink factor per step and, at the end, how many steps the run took. The widget calls the momentum coefficient β\beta, as module 5 did; it is the μ\mu above.

Now run momentum by hand with different numbers.

Work it outMomentum builds up speed

SGD with momentum uses μ=0.5\mu = 0.5 and η=0.1\eta = 0.1, starts from velocity v=0v = 0, and sees the same gradient g=4g = 4 at every step. By how much has the parameter decreased after three steps? (Enter a positive number.)

distance moved

Type a number: 0.25, -2, 3/4 and sqrt(2) all work. Enter checks.

To recap: momentum keeps a decaying sum of past gradients and steps along it. Directions where the gradient keeps its sign speed up by up to 1/(1−μ)1/(1 - \mu), directions where it flips cancel out, and batch noise averages away.

Adam: a step size for every parameter

Momentum still uses one learning rate for every parameter. But gradients in a network differ in size by orders of magnitude from one parameter to another: the output layer's biases may see gradients a thousand times larger than an early layer's weights. One learning rate is then too large for some parameters and too small for others.

Adam (Kingma and Ba; the paper appeared in 2014 and was published at ICLR in 2015), which mini-batches and linear regression previewed, keeps two running averages for every parameter, and divides one by the square root of the other. Here is the full algorithm. With tt counting steps from 1, β1\beta_1 and β2\beta_2 two decay rates between 0 and 1, and every operation entry by entry (so gt2\mathbf{g}_t^2 squares each entry of the gradient gt\mathbf{g}_t at step tt):

mt=β1 mt−1+(1−β1) gtaverage of the gradientvt=β2 vt−1+(1−β2) gt2average of its squarem^t=mt1−β1t,v^t=vt1−β2tbias correctionwt=wt−1−η m^tv^t+ϵthe step\begin{aligned} \mathbf{m}_t &= \beta_1\,\mathbf{m}_{t-1} + (1 - \beta_1)\,\mathbf{g}_t && \text{average of the gradient} \\ \mathbf{v}_t &= \beta_2\,\mathbf{v}_{t-1} + (1 - \beta_2)\,\mathbf{g}_t^2 && \text{average of its square} \\ \hat{\mathbf{m}}_t &= \frac{\mathbf{m}_t}{1 - \beta_1^t}, \qquad \hat{\mathbf{v}}_t = \frac{\mathbf{v}_t}{1 - \beta_2^t} && \text{bias correction} \\ \mathbf{w}_t &= \mathbf{w}_{t-1} - \eta\,\frac{\hat{\mathbf{m}}_t}{\sqrt{\hat{\mathbf{v}}_t} + \epsilon} && \text{the step} \end{aligned}

with m0=v0=0\mathbf{m}_0 = \mathbf{v}_0 = \mathbf{0}. The defaults are β1=0.9\beta_1 = 0.9, β2=0.999\beta_2 = 0.999 and ϵ=10−8\epsilon = 10^{-8}, a tiny number that only prevents division by zero.

Why it works. Take the two averages one at a time.

  • m\mathbf{m} is momentum's velocity in another form. Call the momentum velocity u\mathbf{u} for a moment, since v\mathbf{v} now means the average square, and give it the same decay: ut=β1ut−1+gt\mathbf{u}_t = \beta_1\mathbf{u}_{t-1} + \mathbf{g}_t. It adds each gradient at full weight, while m\mathbf{m} adds it at weight 1−β11 - \beta_1. Both start at 0\mathbf{0}, and multiplying the update of u\mathbf{u} by 1−β11 - \beta_1 gives exactly the update of m\mathbf{m}, so mt=(1−β1) ut\mathbf{m}_t = (1 - \beta_1)\,\mathbf{u}_t at every step: an average of recent gradients instead of a sum.
  • v^\sqrt{\hat{\mathbf{v}}} is the typical size of recent gradients, their root-mean-square over roughly the last 1/(1−β2)=1,0001/(1 - \beta_2) = 1{,}000 steps.
  • Their ratio is a pure number, roughly between −1-1 and 11, whatever the scale of the gradient: near ±1\pm 1 when recent gradients agree in sign, and near 0 when they are noise around zero.

So every parameter moves by roughly η\eta per step when its gradient is consistent, and less when it is not, whether its raw gradients are 10−610^{-6} or 10310^{3}. This is why one Adam learning rate, such as 10−310^{-3} or 3×10−43 \times 10^{-4}, is a reasonable starting point for wildly different models.

Why the bias correction. Both averages start at zero, so for the first several steps they are too small. With the defaults, after one step m1=(1−0.9) g1=0.1 g1\mathbf{m}_1 = (1 - 0.9)\,\mathbf{g}_1 = 0.1\,\mathbf{g}_1 and v1=(1−0.999) g12=0.001 g12\mathbf{v}_1 = (1 - 0.999)\,\mathbf{g}_1^2 = 0.001\,\mathbf{g}_1^2. Dividing by 1−β11=0.11 - \beta_1^1 = 0.1 and 1−β21=0.0011 - \beta_2^1 = 0.001 undoes exactly this: m^1=g1\hat{\mathbf{m}}_1 = \mathbf{g}_1 and v^1=g12\hat{\mathbf{v}}_1 = \mathbf{g}_1^2. The first step is then m^1/v^1=g1/∣g1∣\hat{\mathbf{m}}_1/\sqrt{\hat{\mathbf{v}}_1} = \mathbf{g}_1/|\mathbf{g}_1|, entry by entry: every coordinate moves by exactly η\eta (up to ϵ\epsilon), in the direction opposite to its gradient. Without the correction, the first step would be 0.1/0.001≈0.1/0.0316≈3.160.1/\sqrt{0.001} \approx 0.1/0.0316 \approx 3.16 times larger than that, because v\mathbf{v} starts even further below its true value than m\mathbf{m} does.

Go slower: Where 1−βt1 - \beta^t comes from

Step 1, unroll the average. Start from m0=0\mathbf{m}_0 = \mathbf{0} and apply the update tt times: m1=(1−β1)g1,m2=β1(1−β1)g1+(1−β1)g2,…,mt=(1−β1)∑k=1tβ1 t−k gk.\mathbf{m}_1 = (1 - \beta_1)\mathbf{g}_1,\allowbreak \quad \mathbf{m}_2 = \beta_1(1 - \beta_1)\mathbf{g}_1 + (1 - \beta_1)\mathbf{g}_2,\allowbreak \quad\ldots,\allowbreak \quad \mathbf{m}_t = (1 - \beta_1)\sum_{k=1}^{t}\beta_1^{\,t-k}\,\mathbf{g}_k .

Step 2, suppose the gradient were constant, gk=g\mathbf{g}_k = \mathbf{g} for every kk. A good average should then equal g\mathbf{g}. Instead, pulling g\mathbf{g} out of the sum and writing i=t−ki = t - k, mt=(1−β1) g∑i=0t−1β1 i.\mathbf{m}_t = (1 - \beta_1)\,\mathbf{g}\sum_{i=0}^{t-1}\beta_1^{\,i} .

Step 3, sum the geometric series. ∑i=0t−1β1 i=1−β1t1−β1\sum_{i=0}^{t-1}\beta_1^{\,i} = \frac{1 - \beta_1^t}{1 - \beta_1}. Multiply by (1−β1)(1 - \beta_1): mt=(1−β1t) g\mathbf{m}_t = (1 - \beta_1^t)\,\mathbf{g}. The weights on the gradients add up to 1−β1t1 - \beta_1^t instead of 1, because the missing weight went to the zero we started from.

Step 4, correct. Dividing by 1−β1t1 - \beta_1^t makes the weights add up to 1: m^t=g\hat{\mathbf{m}}_t = \mathbf{g} exactly when the gradient is constant. The same argument with g2\mathbf{g}^2 gives v^t=vt/(1−β2t)\hat{\mathbf{v}}_t = \mathbf{v}_t/(1 - \beta_2^t). As tt grows, βt→0\beta^t \to 0 and the correction fades away: with β1=0.9\beta_1 = 0.9 it is under 1% after 44 steps (0.944≈0.00970.9^{44} \approx 0.0097), and with β2=0.999\beta_2 = 0.999 after about 4,600.

Now run two Adam steps by hand, where the two parameters' gradients differ by a factor of 20.

On paperTwo Adam steps by hand

Run Adam by hand for two steps on a model with two parameters, starting from w=(1,1)\mathbf{w} = (1,\allowbreak 1), with η=0.01\eta = 0.01, β1=0.9\beta_1 = 0.9, β2=0.999\beta_2 = 0.999 and ϵ\epsilon ignored. The gradients are g1=(0.2,−4)\mathbf{g}_1 = (0.2,\allowbreak -4) at step 1 and g2=(0.2,2)\mathbf{g}_2 = (0.2,\allowbreak 2) at step 2.

  1. Compute m1\mathbf{m}_1, v1\mathbf{v}_1, the corrected m^1\hat{\mathbf{m}}_1 and v^1\hat{\mathbf{v}}_1, and w1\mathbf{w}_1. How far did each coordinate move?
  2. Compute m2\mathbf{m}_2, v2\mathbf{v}_2, the corrected averages (use 1−0.92=0.191 - 0.9^2 = 0.19 and 1−0.9992=0.0019991 - 0.999^2 = 0.001999) and w2\mathbf{w}_2.
  3. At step 1 the two gradients differ in size by a factor of 20. Why did both coordinates move by the same amount? And why did the second coordinate move so little at step 2?

To check part 2 before reading the steps, enter w2\mathbf{w}_2 as a column, the first coordinate at the top, to 4 decimal places.

Work it on real paper: writing each step is the point. Then check your final answer here and compare your working with the walk-through.

w after two steps (4 decimal places)

One entry per box, top to bottom. 0.25, -2, 3/4 and sqrt(2) all work. Enter moves to the next empty box and checks once all are filled.

To recap: Adam's step for each parameter is about η\eta times (average gradient) divided by (typical gradient size). It is close to ±η\pm\eta when recent gradients agree, close to 0 when they conflict, and independent of the gradient's overall scale.

Racing the optimizers

Now implement both optimizers. The code ends by racing plain SGD, momentum and Adam on a badly scaled bowl, L(w)=12(w12+100 w22)L(\mathbf{w}) = \frac{1}{2}(w_1^2 + 100\,w_2^2). Its curvature is 1 along w1w_1 and 100 along w2w_2, so the stable learning rate for plain gradient descent is below 2/100=0.022/100 = 0.02, the same trap as the widget's valley, ten times worse.

Here is what the race does. With η=0.019\eta = 0.019, plain descent multiplies the steep coordinate by 1−0.019×100=−0.91 - 0.019 \times 100 = -0.9 each step, so it flips across the valley every step, while the gentle coordinate shrinks by only 1−0.019=0.9811 - 0.019 = 0.981 per step: after 100 steps w1w_1 is still 0.981100≈0.150.981^{100} \approx 0.15. Momentum at the same learning rate swings across the valley too, but its speed builds along the floor, so much that it overshoots the minimum along the floor, to w1≈−0.28w_1 \approx -0.28 at step 23. It swings back past the minimum with ever smaller swings (w1≈+0.08w_1 \approx +0.08 at step 47, −0.02-0.02 at step 71) and settles. The picture below shows the two kinds of path on a valley like this one, seen from above.

A long narrow valley seen from above as nested ellipses around one lowest point, each ellipse ten times as long as it is wide. Two paths leave the same starting point, up on one wall and some way along the valley. The thin path zig-zags from wall to wall on every step, its zig-zags dying out quickly, then crawls along the floor and ends well short of the lowest point. The heavy ball's path swings from wall to wall more slowly, with swings that shrink steadily, races along the floor, rolls past the lowest point by about a quarter of the distance it came, swings back with a smaller overshoot, and settles on the lowest pointOn a badly scaled bowl, plain gradient descent zig-zags across the valley and creeps along it. A heavy ball keeps its speed along the floor, where the gradient keeps pointing the same way, overshoots, and settles at the bottom.

Code itMomentum and Adam

Implement the two optimizers from the lesson. Each takes a list of parameter arrays when it is created, keeps its own state for each one, and has step(grads), which updates every parameter in place from a list of gradients in the same order.

  1. SGDMomentum(params, lr, momentum=0.9): v←μv+g\mathbf{v} \leftarrow \mu\mathbf{v} + \mathbf{g}, then w←w−ηv\mathbf{w} \leftarrow \mathbf{w} - \eta\mathbf{v}.
  2. Adam(params, lr=1e-3, beta1=0.9, beta2=0.999, eps=1e-8), with bias correction.

The last lines race plain SGD, momentum and Adam on the badly scaled bowl L=12(w12+100 w22)L = \frac{1}{2}(w_1^2 + 100\,w_2^2), starting from (1,1)(1,\allowbreak 1). Plain SGD's learning rate is just under its stability limit 2/1002/100.

⌘+Enter runsPython sleeps until you run code
Write your code where the starter says raise NotImplementedError, then press Run tests. Each check says what it expects.
Next: the module checkpoint