Practice / Initialisation and optimisers
Momentum, RMSProp and Adam: optimiser updates by hand
Before you start
An optimiser turns a sequence of gradients into a sequence of parameter updates, and every popular one is a short recurrence that can be run by hand. Running it by hand is how you find out what the step size really is. These ten problems start with plain gradient descent on a quadratic, where the learning-rate limit comes from the Hessian's largest eigenvalue, then unroll momentum into a weighted sum of past gradients, compare its two common forms and Nesterov's variant, follow AdaGrad and RMSProp under a constant gradient, derive Adam's bias correction, take two Adam steps with numbers, and separate an L2 penalty from AdamW's weight decay. The five mistakes at the end each change the effective step size by a constant factor: a stability bound from the wrong eigenvalue, a momentum step taken to be the plain step, a learning rate carried between momentum conventions, an Adam step without bias correction, and an L2 penalty read as weight decay.
- The parameters are and the loss is . is the starting point and the parameters after step . The gradient used at step is , taken at the parameters before the step. is the learning rate.
- Gradient descent: .
- Momentum (the "classic" or summed form): , , , with . This is PyTorch's
torch.optim.SGDwithmomentum=and the defaultdampening=0. The EMA form is , , . - Nesterov momentum, in the form PyTorch uses with
nesterov=True: as above, and . - AdaGrad: , , . RMSProp: , , , with .
- Adam: , , , the bias-corrected estimates and , and . PyTorch's defaults are , , , . Here is to the power ; in Adam is its second-moment estimate, not momentum's , and each problem says which one it uses.
- In the adaptive methods, , , and the division act entry by entry, and is a small constant that keeps the denominator away from .
- A quadratic loss is with symmetric positive definite, so and the Hessian is . Its eigenvalues are , and is the minimiser.
Builds on: Regression gradients: linear, logistic and softmax
Problems
- ·
Gradient descent on the quadratic . Let be the error after step . Show that and hence write in terms of .
- ··
Show that gradient descent on the quadratic converges to from every starting point exactly when . For , find that range, and find the that makes the worst-case error factor per step as small as possible.
- ···
Momentum, classic form, with any gradients treated as given. Show that and that .
- ·
Classic momentum with a constant gradient, for every . Find and the size of the step as . How large is it for compared with plain gradient descent at the same ?
- ··
The EMA form of momentum is , , with . Show that for the classic fed the same gradients, and find the for which the classic form produces exactly the same parameters as the EMA form with learning rate .
- ··
One parameter, , , , . Compute and with classic momentum and with Nesterov momentum in the form .
- ··
One parameter with a constant gradient , and . Find the step size of AdaGrad and of RMSProp as functions of . What happens to each as , and what is RMSProp's first step for ?
- ··
Adam. Show that , whose weights sum to , so that a constant gradient gives and exactly. Then show that, whatever the gradients, Adam's first step is , entry by entry.
- ···
One parameter, , , Adam with , , and . Compute and , rounding to four decimal places at the end.
- ···
Two parameters, , first gradient of the data loss , , , . Compute for (a) Adam with an L2 penalty, which replaces by before the moment updates (PyTorch's
Adamwithweight_decay=), and (b) AdamW, which sets with the moments built from alone (PyTorch'sAdamW). Write the decay part of the step in each.
Answers
- Converges from every start iff ; for this , , and gives the smallest worst-case factor, per step
- and
- , and the step tends to : ten times the plain step when
- ; the classic form with reproduces the EMA form exactly
- Classic: , ; Nesterov: ,
- AdaGrad: ; RMSProp: , with a first step of for
- For a constant gradient, and ; for any gradients,
- and
- (a) Adam + L2: , with the decay divided by ; (b) AdamW: , decay part
Worked solutions
Problem 1
Gradient descent on the quadratic . Let be the error after step . Show that and hence write in terms of .
- , so . is where the gradient vanishes; writing this way lets the gradient be expressed through the error.
- .The gradient of the quadratic is ; substitute step 1 and factor out .
- .Subtract from both sides of the update and use step 2.
- Step 3 applied times. On a quadratic, gradient descent is a fixed linear map applied to the error over and over, so whether it converges depends only on the eigenvalues of (Problem 2).
Problem 2
Show that gradient descent on the quadratic converges to from every starting point exactly when . For , find that range, and find the that makes the worst-case error factor per step as small as possible.
- with orthogonal and , so .A symmetric matrix has an orthonormal basis of eigenvectors, and , so and are diagonal in the same basis.
- With , .Problem 1 in the eigenvector basis: each coordinate is multiplied by its own factor at every step, independently of the others.
- for every iff for every , iff for every . because is orthogonal. A start along eigenvector has only , and exactly when .
- Every , so the condition is .The largest eigenvalue gives the tightest upper bound; the others are then satisfied automatically.
- gives and ..
- The worst factor is ; it is smallest where , at , giving .As grows, falls while rises once ; the maximum of the two is least where they cross. In general the crossing is at , with factor .
- Converges from every start iff ; for this , , and gives the smallest worst-case factor, per stepThe steep direction () caps the learning rate and the flat direction () then sets the speed; the ratio , the condition number, is what momentum and the adaptive methods below try to work around.
Problem 3
Momentum, classic form, with any gradients treated as given. Show that and that .
- .; this matches the claimed sum for , which has the single term .
- If , then .Multiplying by raises every exponent by one, and the new gradient enters with weight . Induction on completes the first claim.
- .Each step subtracts , so after steps the parameters have moved by the sum of all of them.
- .Swap the order of summation: the pairs with can be listed by first, with running from to .
- .A geometric series with terms; so the denominator is not zero.
- and Step 2 and steps 3 to 5. An old gradient's total effect on approaches times the gradient: every gradient is eventually applied times over, spread across the following steps.
Problem 4
Classic momentum with a constant gradient, for every . Find and the size of the step as . How large is it for compared with plain gradient descent at the same ?
- .Problem 3 with every ; substitute .
- .A geometric series with terms.
- , so ..
- , and the step tends to : ten times the plain step when On a long stretch where the gradient barely changes, momentum moves times as far per step as gradient descent with the same learning rate. Where the gradient flips sign from step to step, the terms of step 1 alternate and partly cancel instead, which is what damps oscillation across a narrow valley.
Problem 5
The EMA form of momentum is , , with . Show that for the classic fed the same gradients, and find the for which the classic form produces exactly the same parameters as the EMA form with learning rate .
- .Both start at zero.
- If , then .Substitute and factor out ; the bracket is the classic recurrence. Induction gives the claim for every , as long as both are fed the same .
- , so the two updates agree when .With the same starting point and the same steps, the two runs pass through the same parameters, so their gradients are the same too and step 2 keeps applying.
- ; the classic form with reproduces the EMA form exactlyWith , an EMA-form learning rate of is a classic-form learning rate of . Problem 4's factor lives in the classic form; the EMA form has already divided it out, since its weights sum to .
Problem 6
One parameter, , , , . Compute and with classic momentum and with Nesterov momentum in the form .
- ., evaluated before the step.
- Classic, step 1: , , ..
- Classic, step 2: , , ., then .
- Nesterov, step 1: , , , .The buffer is updated exactly as in the classic form; only the direction used for the step changes.
- Nesterov, step 2: , , , ., then the same two lines as step 4.
- Classic: , ; Nesterov: , Nesterov's direction is : the newest gradient counts times and the older history is damped by one more factor of , so the step leans towards where momentum is about to carry the parameters. Under a constant gradient both forms tend to the same step, , because ; they differ when the gradient changes.
Problem 7
One parameter with a constant gradient , and . Find the step size of AdaGrad and of RMSProp as functions of . What happens to each as , and what is RMSProp's first step for ?
- AdaGrad: , so . adds the squared gradient at every step and never forgets.
- AdaGrad step: .The gradient's size cancels; only the count of steps is left.
- RMSProp: .Problem 3's unrolling with for and the extra factor , then the geometric sum .
- RMSProp step: ..
- At with : ; as the step falls to ., and is smallest at .
- AdaGrad: ; RMSProp: , with a first step of for AdaGrad's sum grows without limit, so its steps keep shrinking even when the gradient does not; RMSProp's moving average forgets, so its step settles at . Neither depends on : multiplying the loss by a constant leaves both updates unchanged when . RMSProp's oversized early steps come from biasing towards zero, which is the bias Adam corrects (Problem 8).
Problem 8
Adam. Show that , whose weights sum to , so that a constant gradient gives and exactly. Then show that, whatever the gradients, Adam's first step is , entry by entry.
- .Problem 5, step 2 shows the EMA recurrence is times the classic one, and Problem 3 unrolls the classic one.
- .Geometric series with terms. The weights fall short of by , the weight that would have carried.
- With : , so .Step 2 times . Dividing by the sum of the weights turns a sum weighted towards zero into a true weighted average.
- The same steps for , with and : and . is the same recurrence applied to .
- At for any : and , so and , .One gradient is a constant sequence of length one, so step 3 applies; entry by entry.
- For a constant gradient, and ; for any gradients, Each entry of the first step is times the sign of its gradient (slightly less when is comparable to ), whatever the gradient's size. Adam's learning rate is therefore close to the actual distance each parameter moves early in training, which is why does not need retuning when the loss is rescaled.
Problem 9
One parameter, , , Adam with , , and . Compute and , rounding to four decimal places at the end.
- , , . at ; both moments start at .
- , , step , so . and ; the first step is times the sign, as Problem 8 predicts.
- , , . and ; then the two recurrences.
- , . and .
- Step ..
- and The second step is , again almost exactly , although the gradient fell by : early on is a ratio of two averages of the same few gradients and stays near . Gradient descent with would have moved and then .
Problem 10
Two parameters, , first gradient of the data loss , , , . Compute for (a) Adam with an L2 penalty, which replaces by before the moment updates (PyTorch's Adam with weight_decay=), and (b) AdamW, which sets with the moments built from alone (PyTorch's AdamW). Write the decay part of the step in each.
- (a) The gradient Adam sees is .The penalty has gradient , added to the data gradient before anything else happens.
- (a) .Problem 8: the first step is times the sign of the gradient it is given; and are both to six decimal places.
- (a) In general contains a weighted average of the past terms, and that share of the step is divided entry by entry by , exactly like the data gradient. is linear in the gradients it is fed, so it splits into a data part and a penalty part; the division by applies to both, with built from the squared total gradient.
- (b) , , so the Adam part is .Problem 8 with the data gradient alone; the first entry's gradient is exactly , so its Adam step is .
- (b) .: each weight shrinks by the same fraction before the Adam step.
- (a) Adam + L2: , with the decay divided by ; (b) AdamW: , decay part With L2, the weight that has no data gradient moved , a hundred times the that "weight decay" suggests, while for the weight with a large gradient the penalty barely changed the step. Dividing by makes the decay strongest on weights with small gradient history and weakest on those with large ones. AdamW keeps the decay a fixed fraction of every weight, which is why its (PyTorch's default is ) behaves like weight decay in plain SGD.
Where this goes wrong
1. Learning-rate limit from the smallest eigenvalue
The minimum lies along the flat directions, and it is natural to tune the step to the direction that has furthest to go.
- The error along eigenvector is multiplied by at each stepRight so far: Problem 2, step 2.
- “The slow direction is the bottleneck, so the learning rate should be set by .”The reasoning that causes the mistake: the flat direction sets how fast gradient descent can go, but the steep direction sets whether it converges at all.
- For with eigenvalues and , any converges, so take At the steep direction's factor is : the error there flips sign and grows by every step. The limit is (Problem 2); the flat direction's slowness has to be fixed by momentum or preconditioning, not by a larger .
2. Momentum step taken to equal the plain step ηg
Adding momentum to a working SGD setup looks like adding smoothing, not like changing the learning rate.
- , Right so far: classic momentum, PyTorch's
SGDwithmomentum=. - “Momentum averages the recent gradients, so the step is still about , only smoother.”The analogy that causes the mistake: reading the classic form as if it were the EMA form, whose weights sum to at most .
- Switching from plain SGD at to momentum at keeps the step size at about The classic buffer sums the gradients, so on a steady gradient the step tends to , ten times larger (Problem 4). A learning rate that was near the stability limit of Problem 2 is now well past it; to keep the step the same, multiply by .
3. Same learning rate for the EMA and summed momentum forms
Papers and libraries write momentum both ways, and is called the learning rate in both.
- Right so far: Problem 5, step 2.
- “Both forms are momentum with the same , so a learning rate from one carries over to the other.”The shortcut that causes the mistake: treating the factor in the EMA recurrence as a detail of the bookkeeping rather than a scale on every step.
- An EMA-form run with , is reproduced by the classic form with The classic form needs (Problem 5). With every step is ten times too large; going the other way, from classic to EMA with the number unchanged, every step is ten times too small and training looks merely slow.
4. First Adam step computed without the bias correction
The bias correction looks like a small fix-up that matters only for a few steps, so a hand-written Adam often leaves it out.
- , Right so far: the first moment updates from zero.
- “The correction only matters for small ; use and directly.”The shortcut that causes the mistake: assuming the two biases are similar and cancel in the ratio .
- The biases do not cancel: is shrunk by , by . The uncorrected first step is times instead of (Problem 8), and the factor by which every uncorrected step is too large rises to about near and is still about at , so the early steps are oversized just when the parameters are furthest from sensible.
5. Adam's L2 penalty taken as a decay of ηλθ
In plain SGD an L2 penalty and weight decay are the same thing: the penalty's gradient times shrinks each weight by .
- The gradient Adam sees is Right so far: Problem 10, step 1, PyTorch's
Adamwithweight_decay=. - “Adding to the gradient is weight decay, so each step shrinks the weights by as in SGD.”The analogy that causes the mistake: SGD multiplies the whole gradient by the same ; Adam divides it entry by entry by .
- In Problem 10, Adam with
weight_decay=0.01moves the weight with zero data gradient by It moves by , a full Adam step, because the penalty gradient is normalised by its own size (Problem 10, step 2). Weights with small gradients are decayed far more than and weights with large gradients far less; AdamW applies separately, outside the normalisation.
Print this set: momentum-rmsprop-and-adam.pdf (problems, answers, and worked solutions on separate pages).