Practice / Initialisation and optimisers
Weight decay, L2 regularisation and AdamW
Before you start
Weight decay is one line in every optimiser, and the line means something different in each. Under plain gradient descent, adding to the loss and shrinking the weights by a factor each step are the same thing; under momentum the shrink is amplified; under Adam it is divided away entry by entry; and AdamW, which applies the shrink outside the adaptive step, no longer minimises any penalised loss at all. These ten problems work each case out: the shrink factor and its half-life, the ridge solution as a fixed point, the amplification by , the normalised decay inside Adam, AdamW's stationary points, what decay does to a weight matrix that is followed by a normalisation layer (where the loss cannot see the weight's norm at all, and decay turns into a learning rate), AdamW as a moving average of its own updates, the geometry of decay on a quadratic, and how each optimiser responds to the loss being multiplied by a constant. The five mistakes are each a carried from one optimiser to another: between SGD and momentum, between a loss and a normalised layer, between Adam and the ridge solution, between two learning rates, and between two loss scales.
- The parameters are , the data loss is , and is its gradient at step ; is the learning rate and the decay coefficient. Entrywise operations (, , squares, division) act entry by entry as on the optimiser page.
- L2 penalty: the optimiser is run on , whose gradient is . This is PyTorch's
weight_decay=inSGDand inAdam. - Decoupled weight decay: the weights are multiplied by and the optimiser's step is computed from alone: , where is the optimiser's direction ( for SGD, for momentum, for Adam). With Adam this is AdamW, PyTorch's
AdamW, whose decay is per step with by default. - Momentum and Adam are as on the optimiser page: classic momentum , ; Adam with , , bias-corrected , , and . That page's Problem 8 showed that a constant gradient gives and , so Adam's step is .
- Ridge regression in this page's units: with gradient , so is minimised by (the regression page, with scaled to match the here).
- A scale-invariant weight is one the loss sees only through its direction, : any weight matrix followed by batch norm or layer norm, and the embeddings of the contrastive page. denotes evaluated at , the gradient norm at unit scale.
Builds on: Momentum, RMSProp and Adam: optimiser updates by hand, Regression gradients: linear, logistic and softmax
Problems
- ·
Show that SGD on is the decoupled update . With no data gradient, write in terms of , and find the number of steps that halves the weights when and .
- ··
Show that the L2 and decoupled SGD updates on the ridge loss share the fixed point of Before you start, and that the fixed point does not depend on .
- ···
Momentum with an L2 penalty is , ; decoupled is , . With no data gradient, show that the L2 version satisfies , find its asymptotic shrink factor per step to first order in , and compare with the decoupled version.
- ··
Adam with an L2 penalty feeds to the moment estimates. For a constant gradient vector , and once the moment estimates have caught up with their input (so that and ), write the step. Show that an entry with shrinks by about per step whatever is, and that an entry with receives a decay of about .
- ···
AdamW with : . Show that a stationary point with gradient satisfies for every entry. Deduce what happens for a constant gradient with , and, for a loss whose gradient vanishes at its minimiser, that the stationary point is the ridge solution with penalty coefficient .
- ···
A scale-invariant weight: . Show that and . For decoupled SGD, , show that , find the equilibrium norm when with constant, and the resulting effective learning rate on the direction of .
- ··
Same layer with no decay (). Show that satisfies , hence , and describe how the effective learning rate behaves over training.
- ··
Write AdamW as with the normalised Adam direction. Unroll it to a closed form for , show that the coefficients of the sum to as , and, for a learning-rate schedule , show that the factor multiplying after steps is about . Evaluate it for , and .
- ··
Decoupled SGD on the quadratic with symmetric positive definite. Show that with , give the condition on for convergence, and for the optimiser page's with find the range of , the best , its worst-case factor, and the condition number with and without the decay.
- ··
The loss is multiplied by a constant (a change of units, or a different reduction over the batch), so the gradient becomes . Show that SGD with an L2 penalty on with is SGD on with ; that AdamW with produces the same iterates for as for ; and that Adam with an L2 penalty on with equals Adam with an L2 penalty on with .
Answers
- SGD with an L2 penalty is : the decay and the penalty are the same update; with no gradient, , and halves a weight in about steps
- Both forms have the fixed point , which is the ridge solution , for every
- L2 + momentum shrinks by about per step, ten times the decoupled at ; to match a decoupled , the L2 coefficient must be
- Adam + L2 step : a weight with no data gradient moves by per step towards regardless of , while a weight with a large gradient is decayed by only per step
- AdamW's stationary points satisfy ; under a persistent gradient the weight settles at whatever ; where the gradient vanishes at the optimum, the stationary point is the ridge solution with penalty , essentially unregularised at
- and ; ; equilibrium and effective learning rate
- : the norm grows like and the effective learning rate on the direction falls like , independent of
- ; the update weights sum to , so is times an EMA of the Adam directions with time constant steps; is multiplied by about , which is in the example
- , ; converges iff ; for this and : , best with factor (was ), condition number (was )
- SGD + L2: ; AdamW is invariant to (up to ); Adam + L2: at the same
Worked solutions
Problem 1
Show that SGD on is the decoupled update . With no data gradient, write in terms of , and find the number of steps that halves the weights when and .
- . (the matrix-calculus page); the cancels the .
- .SGD on uses the gradient of step 1 added to ; collect the terms.
- With : .Step 2 applied times is multiplication by the same scalar times.
- gives .Take logs; for small .
- , so steps.; the exact value from step 4 is .
- SGD with an L2 penalty is : the decay and the penalty are the same update; with no gradient, , and halves a weight in about stepsUnder plain SGD the two conventions agree exactly, step for step, with the same ; everything after this problem is about optimisers where they do not. The shrink per step is , a product: halving the learning rate halves the decay, and the characteristic time steps (here ) is the horizon over which a weight forgets its initial value (Problem 8).
Problem 2
Show that the L2 and decoupled SGD updates on the ridge loss share the fixed point of Before you start, and that the fixed point does not depend on .
- A fixed point of satisfies .Fixed point: the update returns the same .
- , so .Rearrange and divide by ; has dropped out.
- , so .The ridge gradient; collect the terms.
- . is positive definite for (the regression page), so it is invertible.
- The L2 update is the same map (Problem 1), so it has the same fixed point.Problem 1, step 2.
- Both forms have the fixed point , which is the ridge solution , for every A fixed point of SGD on is a stationary point of , whatever the learning rate; decides how fast and whether the iteration gets there (Problem 9), not where it goes. This is the property AdamW gives up (Problem 5).
Problem 3
Momentum with an L2 penalty is , ; decoupled is , . With no data gradient, show that the L2 version satisfies , find its asymptotic shrink factor per step to first order in , and compare with the decoupled version.
- for every .Rearrange ; this holds for both versions.
- L2 with : .Substitute the recurrence for , then step 1 for .
- .Collect terms. A second-order linear recurrence with constant coefficients.
- Solutions are with .Substitute and divide by . The general solution is a combination of the two roots' powers, and the larger root dominates.
- Write : .Expand and cancel: the constant terms give , the terms in give , and survives from the middle product.
- To first order, , so .Drop the quadratic terms and ; the next term is of order .
- Decoupled with : throughout, so exactly. and nothing is added to it.
- L2 + momentum shrinks by about per step, ten times the decoupled at ; to match a decoupled , the L2 coefficient must be The penalty gradient enters the momentum buffer and is applied times over, exactly as the optimiser page's Problem 4 found for any steady gradient. PyTorch's
SGDwithmomentum=0.9, weight_decay=therefore decays weights ten times faster than the number suggests; a tuned there and moved to an optimiser with decoupled decay is ten times too small.
Problem 4
Adam with an L2 penalty feeds to the moment estimates. For a constant gradient vector , and once the moment estimates have caught up with their input (so that and ), write the step. Show that an entry with shrinks by about per step whatever is, and that an entry with receives a decay of about .
- Step , entry by entry.Adam's update with the converged estimates; .
- Entry with : step. cancels between numerator and denominator when .
- Entry with : , so step.The denominator is dominated by (and is negligible next to it); split the numerator. The first term is the plain Adam step.
- Adam + L2 step : a weight with no data gradient moves by per step towards regardless of , while a weight with a large gradient is decayed by only per stepThe decay is normalised by the same as the data gradient, so it is strongest exactly where the gradient history is smallest and weakest where it is largest: the opposite of a uniform shrink, and a that is invisible on busy weights and irrelevant on idle ones. The optimiser page's Problem 10 is the first step of this; here it is the steady state, approximate because the estimates trail the slowly moving .
Problem 5
AdamW with : . Show that a stationary point with gradient satisfies for every entry. Deduce what happens for a constant gradient with , and, for a loss whose gradient vanishes at its minimiser, that the stationary point is the ridge solution with penalty coefficient .
- At a stationary point the gradient is a constant , so and once the averages settle.The optimiser page, Problem 8: constant input gives exact bias-corrected moments.
- , so .Fixed-point condition; cancel .
- .Multiply out.
- Constant with : , so .. The size of the gradient has no effect on where the weight settles.
- Gradient that vanishes at the minimiser: near the stationary point is small, and if then , that is .Step 3 with . Consistency: it requires , that is , which holds for the usual and weights of order .
- On the ridge loss, is Problem 2's fixed-point equation with in place of .Compare with Problem 2, step 2.
- AdamW's stationary points satisfy ; under a persistent gradient the weight settles at whatever ; where the gradient vanishes at the optimum, the stationary point is the ridge solution with penalty , essentially unregularised at AdamW is not gradient descent on any penalised loss: its decay regularises the trajectory (a weight needs a persistent gradient to stay large, and cannot exceed , at the default ) but not the solution of a problem it can solve exactly. Adam with the L2 penalty, by contrast, is a method for and does converge to the ridge solution with coefficient , at the cost of Problem 4's normalised decay on the way there.
Problem 6
A scale-invariant weight: . Show that and . For decoupled SGD, , show that , find the equilibrium norm when with constant, and the resulting effective learning rate on the direction of .
- for , so .; chain rule along the ray, as on the contrastive page.
- .Chain rule through with the Jacobian (the Jacobians page).
- At the same formula gives , so and .Step 2 at and again at : the projection is the same, only the differs. by definition.
- .Expand the squared norm; the cross term vanishes because is orthogonal to (step 1). Pythagoras: the decay shrinks along itself and the gradient step moves it sideways.
- With at equilibrium and : , so .Set and multiply through by .
- , so and .; divide and take the fourth root.
- .Substitute step 6.
- and ; ; equilibrium and effective learning rate For a weight the loss cannot see the size of, decay does not regularise the function at all: it sets the norm, and through the norm the angular step per iteration. That step is : it depends on and only through their product, and only as a square root, so dividing by at the end of a schedule slows the rotation of a normalised layer by only , and the same decay with no normalisation layer would be doing something else entirely.
Problem 7
Same layer with no decay (). Show that satisfies , hence , and describe how the effective learning rate behaves over training.
- .Problem 6, step 4 with and .
- .Square step 1; the last term is smaller than the middle one by the factor , which is tiny once is of order .
- , so .Sum step 2 over steps.
- The effective learning rate for large .Substitute step 3 and drop .
- : the norm grows like and the effective learning rate on the direction falls like , independent of With no decay, a normalised layer applies its own learning-rate schedule, because every gradient step lengthens (step 1: the orthogonal step can only add to the norm) and a longer turns more slowly. Weight decay (Problem 6) is what stops this: it pulls the norm back to an equilibrium and holds the angular step at instead of letting it decay to zero. This is also why removing decay from such a layer makes training appear to stall late on rather than diverge.
Problem 8
Write AdamW as with the normalised Adam direction. Unroll it to a closed form for , show that the coefficients of the sum to as , and, for a learning-rate schedule , show that the factor multiplying after steps is about . Evaluate it for , and .
- , .Apply the update twice; each earlier term picks up one more factor of .
- .Induction on : multiplying by raises every exponent by one and the new update enters with exponent .
- .Geometric series with ratio ; the cancels.
- So with weights , an exponential moving average with time constant steps.The weights of step 3 normalised to sum to ; they halve every steps (Problem 1).
- With a schedule, 's factor is , and . for small ; the error is of order .
- , so the factor is .Constant : .
- ; the update weights sum to , so is times an EMA of the Adam directions with time constant steps; is multiplied by about , which is in the exampleEach Adam direction has entries of size about , so cannot exceed (Problem 5 again) and a weight reflects only the last steps of updates: steps in the example, against a run of . The decay in PyTorch's
AdamWis multiplied by the scheduled , so the EMA horizon stretches as the learning rate decays and the total forgetting is set by the area under the schedule, not by alone.
Problem 9
Decoupled SGD on the quadratic with symmetric positive definite. Show that with , give the condition on for convergence, and for the optimiser page's with find the range of , the best , its worst-case factor, and the condition number with and without the decay.
- .The gradient of the quadratic is ; collect the terms.
- satisfies ., so .
- Subtracting: .Step 1 minus step 2; the terms cancel.
- has eigenvalues with the eigenvectors of , so convergence from every start needs for all : .The optimiser page's Problem 2 with in place of . Adding shifts every eigenvalue by and leaves the eigenvectors alone.
- For the given , eigenvalues and become and : , best , worst-case factor .The optimiser page's formulas and applied to and .
- Condition number , against without decay.The ratio of the largest to the smallest eigenvalue of the matrix the iteration sees.
- , ; converges iff ; for this and : , best with factor (was ), condition number (was )Decay moves the minimiser from to , which is the bias it pays for, and makes the problem better conditioned, which is why it speeds up convergence along the flat directions most: an eigenvalue is effectively replaced by . The stability limit tightens slightly, from to .
Problem 10
The loss is multiplied by a constant (a change of units, or a different reduction over the batch), so the gradient becomes . Show that SGD with an L2 penalty on with is SGD on with ; that AdamW with produces the same iterates for as for ; and that Adam with an L2 penalty on with equals Adam with an L2 penalty on with .
- SGD + L2 on : .Factor out of the bracket.
- That is SGD + L2 on with learning rate and coefficient .Compare with Problem 1, step 2. The regularisation relative to the data has weakened by .
- AdamW on : and , so is unchanged.Both moment estimates are built from ; the first is linear in it and the second quadratic, and . With the ratio cancels exactly.
- The decay does not involve the gradient, so AdamW's iterates are identical.The decoupled form touches only and .
- Adam + L2 on feeds to the moments; by step 3 the factor cancels in the normalised step.Factor out inside the input to the moment estimates, then apply the scale invariance.
- SGD + L2: ; AdamW is invariant to (up to ); Adam + L2: at the same Two of the three change meaning when the loss is rescaled, and only AdamW's is a property of the optimiser alone. Switching a loss from a sum over the batch to a mean is ; under Adam with
weight_decayit multiplies the effective regularisation by , while the Adam step itself does not change, which is a quiet way to make a that worked at one batch size useless at another.
Where this goes wrong
1. Carrying λ from SGD to SGD with momentum
Adding momentum=0.9 to an SGD run looks like changing the direction, not the regularisation.
- SGD with
weight_decay=shrinks weights by per stepRight so far: Problem 1. - “Momentum averages the gradients; the decay is a separate term and stays .”The shortcut that causes the mistake: in PyTorch's
SGDthe penalty gradient is added to the gradient before the momentum buffer, so it is accumulated like everything else. - With momentum the weights still shrink by per stepThey shrink by about (Problem 3), ten times more at . The run trains, with weights ten times more strongly decayed than intended; the that reproduces the old behaviour is , and a decoupled implementation would have needed no change.
2. Decay on a normalised layer taken to shrink the layer's output
Weight decay is introduced as shrinking the weights, and a smaller weight matrix sounds like a smaller, more regular function.
- for a weight matrix followed by a normalisation layerRight so far: the layer's output is unchanged by the scale of .
- “Decay pulls down, so the layer's output gets smaller and the model is regularised.”The assumption that causes the mistake: the output is invariant to (step 1), so it cannot get smaller.
- Smaller means smaller activations and a simpler functionNothing downstream changes with . What decay sets on such a layer is the equilibrium norm and through it the effective learning rate on the direction of (Problem 6); raising there makes the layer learn faster, and removing decay makes its learning rate decay like (Problem 7). Decay on the parameters that the normalisation does not absorb, the gains and biases and the final layer, is where a function-space effect lives.
3. Expecting AdamW to converge to the ridge solution
AdamW is described as "Adam with weight decay", and weight decay is described as L2 regularisation.
- The ridge loss has minimiser Right so far: Problem 2.
- “AdamW minimises , so it converges to .”The assumption that causes the mistake: the decoupled decay is not the gradient of a term in the loss once the data gradient is normalised by , so there is no penalised loss whose stationary points AdamW finds.
- AdamW with
weight_decay=converges to Its stationary points satisfy (Problem 5): on the ridge loss, where the gradient vanishes at the optimum, that is the ridge solution with coefficient , about at the defaults, so AdamW converges to the unregularised least-squares solution. Adam with the L2 penalty does converge to . The decay in AdamW bounds and forgets the trajectory (Problems 5 and 8); it does not pick the solution.
4. Halving the learning rate with the decay held fixed
The decay coefficient is called and the learning rate , and tuning one is not supposed to touch the other.
- AdamW: Right so far.
- “The decay is , independent of the learning rate.”The slip that causes the mistake: the shrink per step is the product (Problem 1), and in PyTorch's
AdamWthe decay is multiplied by the scheduled learning rate. - Halving leaves the weight decay unchangedIt halves the decay per step and doubles the EMA time constant (Problem 8): the weights remember twice as many updates and is forgotten half as fast, and on a normalised layer the effective learning rate drops by rather than (Problem 6). A learning-rate sweep with fixed is also a sweep over the decay; to vary one alone, hold fixed.
5. Rescaling the loss under Adam with an L2 penalty
Adam does not care about the scale of the loss, and the penalty is just another term in it.
- Adam's step is unchanged when the gradient is multiplied by Right so far: Problem 10, step 3.
- “So multiplying the loss by (switching from a mean to a sum over the batch) changes nothing under Adam.”The assumption that causes the mistake: the penalty gradient is added to the data gradient before normalisation, and it was not multiplied by .
- Adam with
weight_decay=gives the same iterates for as for It gives the iterates for with (Problem 10, step 5): a sum over a batch of instead of a mean divides the effective regularisation by , while the loss curve, in its new units, looks the same. AdamW is the version for which the claim holds, which is one reason its transfers between setups and Adam's does not.
Print this set: weight-decay-l2-and-adamw.pdf (problems, answers, and worked solutions on separate pages).