Practice / Initialisation and optimisers

Weight decay, L2 regularisation and AdamW

Ten problems on weight decay: why an L2 penalty and decoupled decay are the same update under plain SGD and what the shrink factor and half-life are, the ridge solution as the fixed point of both, how momentum amplifies an L2 decay by 1/(1 − β), how Adam normalises it away, the stationary points of AdamW and why they are not the ridge solution, the equilibrium norm and effective learning rate of a scale-invariant layer with and without decay, AdamW as an exponential moving average of updates, decay on a quadratic, and what scaling the loss does to each optimiser, with worked solutions and the mistakes that carry a λ between optimisers that do not share it.

Before you start

Weight decay is one line in every optimiser, and the line means something different in each. Under plain gradient descent, adding λ2∥θ∥2\tfrac\lambda2\|\theta\|^2 to the loss and shrinking the weights by a factor 1−ηλ1 - \eta\lambda 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 1/(1−β)1/(1 - \beta), 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 λ\lambda 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 θ∈Rd\theta \in \mathbb{R}^d, the data loss is L(θ)L(\theta), and gt=∇L(θt−1)g_t = \nabla L(\theta_{t-1}) is its gradient at step tt; η>0\eta > 0 is the learning rate and λ≥0\lambda \ge 0 the decay coefficient. Entrywise operations (∣⋅∣|\cdot|, sign⁡\operatorname{sign}, squares, division) act entry by entry as on the optimiser page.
  • L2 penalty: the optimiser is run on Lλ(θ)=L(θ)+λ2∥θ∥2L_\lambda(\theta) = L(\theta) + \tfrac\lambda2\|\theta\|^2, whose gradient is gt+λθt−1g_t + \lambda\theta_{t-1}. This is PyTorch's weight_decay= in SGD and in Adam.
  • Decoupled weight decay: the weights are multiplied by 1−ηλ1 - \eta\lambda and the optimiser's step is computed from gtg_t alone: θt=(1−ηλ)θt−1−η ut\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta\,u_t, where utu_t is the optimiser's direction (gtg_t for SGD, vtv_t for momentum, m^t/(v^t+ϵ)\hat m_t/(\sqrt{\hat v_t} + \epsilon) for Adam). With Adam this is AdamW, PyTorch's AdamW, whose decay is ηλθ\eta\lambda\theta per step with λ=0.01\lambda = 0.01 by default.
  • Momentum and Adam are as on the optimiser page: classic momentum vt=βvt−1+(gradient)v_t = \beta v_{t-1} + (\text{gradient}), θt=θt−1−ηvt\theta_t = \theta_{t-1} - \eta v_t; Adam with β1\beta_1, β2\beta_2, bias-corrected m^t\hat m_t, v^t\hat v_t, and ϵ\epsilon. That page's Problem 8 showed that a constant gradient gg gives m^t=g\hat m_t = g and v^t=g2\hat v_t = g^2, so Adam's step is −η g/(∣g∣+ϵ)-\eta\,g/(|g| + \epsilon).
  • Ridge regression in this page's units: L(θ)=12N∥Xθ−y∥2L(\theta) = \tfrac1{2N}\|X\theta - y\|^2 with gradient 1NX⊤(Xθ−y)\tfrac1NX^\top(X\theta - y), so LλL_\lambda is minimised by θ∗=(1NX⊤X+λI)−11NX⊤y\theta^* = (\tfrac1NX^\top X + \lambda I)^{-1}\tfrac1NX^\top y (the regression page, with λ\lambda scaled to match the λ2\tfrac\lambda2 here).
  • A scale-invariant weight ww is one the loss sees only through its direction, L(w)=f(w/∥w∥)L(w) = f(w/\|w\|): any weight matrix followed by batch norm or layer norm, and the embeddings of the contrastive page. GG denotes ∥∇L∥\|\nabla L\| evaluated at w^=w/∥w∥\hat w = w/\|w\|, the gradient norm at unit scale.

Builds on: Momentum, RMSProp and Adam: optimiser updates by hand, Regression gradients: linear, logistic and softmax

Problems

  1. ·

    Show that SGD on LλL_\lambda is the decoupled update θt=(1−ηλ)θt−1−ηgt\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta g_t. With no data gradient, write θT\theta_T in terms of θ0\theta_0, and find the number of steps that halves the weights when η=0.1\eta = 0.1 and λ=10−4\lambda = 10^{-4}.

  2. ··

    Show that the L2 and decoupled SGD updates on the ridge loss share the fixed point θ∗\theta^* of Before you start, and that the fixed point does not depend on η\eta.

  3. ···

    Momentum with an L2 penalty is vt=βvt−1+gt+λθt−1v_t = \beta v_{t-1} + g_t + \lambda\theta_{t-1}, θt=θt−1−ηvt\theta_t = \theta_{t-1} - \eta v_t; decoupled is vt=βvt−1+gtv_t = \beta v_{t-1} + g_t, θt=(1−ηλ)θt−1−ηvt\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta v_t. With no data gradient, show that the L2 version satisfies θt=(1+β−ηλ)θt−1−βθt−2\theta_t = (1 + \beta - \eta\lambda)\theta_{t-1} - \beta\theta_{t-2}, find its asymptotic shrink factor per step to first order in ηλ\eta\lambda, and compare with the decoupled version.

  4. ··

    Adam with an L2 penalty feeds gt+λθt−1g_t + \lambda\theta_{t-1} to the moment estimates. For a constant gradient vector gg, and once the moment estimates have caught up with their input (so that m^t≈g+λθt−1\hat m_t \approx g + \lambda\theta_{t-1} and v^t≈(g+λθt−1)2\hat v_t \approx (g + \lambda\theta_{t-1})^2), write the step. Show that an entry with gi=0g_i = 0 shrinks by about η\eta per step whatever λ\lambda is, and that an entry with ∣gi∣≫λ∣θi∣|g_i| \gg \lambda|\theta_i| receives a decay of about ηλθi/∣gi∣\eta\lambda\theta_i/|g_i|.

  5. ···

    AdamW with ϵ>0\epsilon > 0: θt=(1−ηλ)θt−1−η m^t/(v^t+ϵ)\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta\,\hat m_t/(\sqrt{\hat v_t} + \epsilon). Show that a stationary point with gradient g=g(θ)g = g(\theta) satisfies λθi(∣gi∣+ϵ)=−gi\lambda\theta_i(|g_i| + \epsilon) = -g_i for every entry. Deduce what happens for a constant gradient gg with ∣gi∣≫ϵ|g_i| \gg \epsilon, and, for a loss whose gradient vanishes at its minimiser, that the stationary point is the ridge solution with penalty coefficient λϵ\lambda\epsilon.

  6. ···

    A scale-invariant weight: L(w)=f(w/∥w∥)L(w) = f(w/\|w\|). Show that w⊤∇L(w)=0w^\top\nabla L(w) = 0 and ∥∇L(w)∥=G/∥w∥\|\nabla L(w)\| = G/\|w\|. For decoupled SGD, wt=(1−ηλ)wt−1−ηgtw_t = (1 - \eta\lambda)w_{t-1} - \eta g_t, show that ∥wt∥2=(1−ηλ)2∥wt−1∥2+η2∥gt∥2\|w_t\|^2 = (1 - \eta\lambda)^2\|w_{t-1}\|^2 + \eta^2\|g_t\|^2, find the equilibrium norm when ∥gt∥=G/∥wt−1∥\|g_t\| = G/\|w_{t-1}\| with GG constant, and the resulting effective learning rate η/∥w∥2\eta/\|w\|^2 on the direction of ww.

  7. ··

    Same layer with no decay (λ=0\lambda = 0). Show that nt=∥wt∥2n_t = \|w_t\|^2 satisfies nt2≈nt−12+2η2G2n_t^2 \approx n_{t-1}^2 + 2\eta^2G^2, hence ∥wT∥≈(n02+2η2G2T)1/4\|w_T\| \approx (n_0^2 + 2\eta^2G^2T)^{1/4}, and describe how the effective learning rate behaves over training.

  8. ··

    Write AdamW as θt=(1−ηλ)θt−1−ηut\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta u_t with utu_t the normalised Adam direction. Unroll it to a closed form for θT\theta_T, show that the coefficients of the utu_t sum to 1/λ1/\lambda as T→∞T \to \infty, and, for a learning-rate schedule ηt\eta_t, show that the factor multiplying θ0\theta_0 after TT steps is about e−λ∑tηte^{-\lambda\sum_t\eta_t}. Evaluate it for η=10−3\eta = 10^{-3}, λ=0.1\lambda = 0.1 and T=100,000T = 100{,}000.

  9. ··

    Decoupled SGD on the quadratic L(θ)=12θ⊤Aθ−b⊤θL(\theta) = \tfrac12\theta^\top A\theta - b^\top\theta with AA symmetric positive definite. Show that θt−θ∗=(I−η(A+λI))(θt−1−θ∗)\theta_t - \theta^* = \big(I - \eta(A + \lambda I)\big)(\theta_{t-1} - \theta^*) with θ∗=(A+λI)−1b\theta^* = (A + \lambda I)^{-1}b, give the condition on η\eta for convergence, and for the optimiser page's A=(3113)A = \begin{pmatrix}3 & 1\\ 1 & 3\end{pmatrix} with λ=1\lambda = 1 find the range of η\eta, the best η\eta, its worst-case factor, and the condition number with and without the decay.

  10. ··

    The loss is multiplied by a constant c>0c > 0 (a change of units, or a different reduction over the batch), so the gradient becomes cgtcg_t. Show that SGD with an L2 penalty on cLcL with (η,λ)(\eta, \lambda) is SGD on LL with (cη,λ/c)(c\eta, \lambda/c); that AdamW with ϵ→0\epsilon \to 0 produces the same iterates for cLcL as for LL; and that Adam with an L2 penalty on cLcL with λ\lambda equals Adam with an L2 penalty on LL with λ/c\lambda/c.

Worked solutions

Problem 1

Show that SGD on LλL_\lambda is the decoupled update θt=(1−ηλ)θt−1−ηgt\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta g_t. With no data gradient, write θT\theta_T in terms of θ0\theta_0, and find the number of steps that halves the weights when η=0.1\eta = 0.1 and λ=10−4\lambda = 10^{-4}.

  1. ∇(λ2∥θ∥2)=λθ\nabla\big(\tfrac\lambda2\|\theta\|^2\big) = \lambda\theta.∇θ(θ⊤θ)=2θ\nabla_\theta(\theta^\top\theta) = 2\theta (the matrix-calculus page); the 12\tfrac12 cancels the 22.
  2. θt=θt−1−η(gt+λθt−1)=(1−ηλ)θt−1−ηgt\theta_t = \theta_{t-1} - \eta(g_t + \lambda\theta_{t-1}) = (1 - \eta\lambda)\theta_{t-1} - \eta g_t.SGD on LλL_\lambda uses the gradient of step 1 added to gtg_t; collect the θt−1\theta_{t-1} terms.
  3. With gt=0g_t = 0: θT=(1−ηλ)Tθ0\theta_T = (1 - \eta\lambda)^T\theta_0.Step 2 applied TT times is multiplication by the same scalar TT times.
  4. (1−ηλ)T=12(1 - \eta\lambda)^T = \tfrac12 gives T=ln⁡2−ln⁡(1−ηλ)≈ln⁡2ηλT = \dfrac{\ln 2}{-\ln(1 - \eta\lambda)} \approx \dfrac{\ln 2}{\eta\lambda}.Take logs; −ln⁡(1−x)=x+x22+⋯≈x-\ln(1 - x) = x + \tfrac{x^2}2 + \cdots \approx x for small xx.
  5. ηλ=10−5\eta\lambda = 10^{-5}, so T≈0.693×105≈69,300T \approx 0.693\times10^5 \approx 69{,}300 steps.ln⁡2≈0.693\ln 2 \approx 0.693; the exact value from step 4 is 69,31569{,}315.
  6. SGD with an L2 penalty is θt=(1−ηλ)θt−1−ηgt\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta g_t: the decay and the penalty are the same update; with no gradient, θT=(1−ηλ)Tθ0\theta_T = (1 - \eta\lambda)^T\theta_0, and ηλ=10−5\eta\lambda = 10^{-5} halves a weight in about 69,00069{,}000 stepsUnder plain SGD the two conventions agree exactly, step for step, with the same λ\lambda; everything after this problem is about optimisers where they do not. The shrink per step is ηλ\eta\lambda, a product: halving the learning rate halves the decay, and the characteristic time 1/(ηλ)1/(\eta\lambda) steps (here 100,000100{,}000) 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 θ∗\theta^* of Before you start, and that the fixed point does not depend on η\eta.

  1. A fixed point of θt=(1−ηλ)θt−1−ηg(θt−1)\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta g(\theta_{t-1}) satisfies θ=(1−ηλ)θ−ηg(θ)\theta = (1 - \eta\lambda)\theta - \eta g(\theta).Fixed point: the update returns the same θ\theta.
  2. ηλθ+ηg(θ)=0\eta\lambda\theta + \eta g(\theta) = 0, so g(θ)+λθ=0g(\theta) + \lambda\theta = 0.Rearrange and divide by η>0\eta > 0; η\eta has dropped out.
  3. 1NX⊤(Xθ−y)+λθ=0\tfrac1NX^\top(X\theta - y) + \lambda\theta = 0, so (1NX⊤X+λI)θ=1NX⊤y\big(\tfrac1NX^\top X + \lambda I\big)\theta = \tfrac1NX^\top y.The ridge gradient; collect the θ\theta terms.
  4. θ=θ∗=(1NX⊤X+λI)−11NX⊤y\theta = \theta^* = \big(\tfrac1NX^\top X + \lambda I\big)^{-1}\tfrac1NX^\top y.1NX⊤X+λI\tfrac1NX^\top X + \lambda I is positive definite for λ>0\lambda > 0 (the regression page), so it is invertible.
  5. The L2 update is the same map (Problem 1), so it has the same fixed point.Problem 1, step 2.
  6. Both forms have the fixed point g(θ)+λθ=0g(\theta) + \lambda\theta = 0, which is the ridge solution θ∗=(1NX⊤X+λI)−11NX⊤y\theta^* = (\tfrac1NX^\top X + \lambda I)^{-1}\tfrac1NX^\top y, for every η\etaA fixed point of SGD on LλL_\lambda is a stationary point of LλL_\lambda, whatever the learning rate; η\eta 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 vt=βvt−1+gt+λθt−1v_t = \beta v_{t-1} + g_t + \lambda\theta_{t-1}, θt=θt−1−ηvt\theta_t = \theta_{t-1} - \eta v_t; decoupled is vt=βvt−1+gtv_t = \beta v_{t-1} + g_t, θt=(1−ηλ)θt−1−ηvt\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta v_t. With no data gradient, show that the L2 version satisfies θt=(1+β−ηλ)θt−1−βθt−2\theta_t = (1 + \beta - \eta\lambda)\theta_{t-1} - \beta\theta_{t-2}, find its asymptotic shrink factor per step to first order in ηλ\eta\lambda, and compare with the decoupled version.

  1. vt=(θt−1−θt)/ηv_t = (\theta_{t-1} - \theta_t)/\eta for every tt.Rearrange θt=θt−1−ηvt\theta_t = \theta_{t-1} - \eta v_t; this holds for both versions.
  2. L2 with gt=0g_t = 0: θt=θt−1−η(βvt−1+λθt−1)=θt−1−β(θt−2−θt−1)−ηλθt−1\theta_t = \theta_{t-1} - \eta(\beta v_{t-1} + \lambda\theta_{t-1}) = \theta_{t-1} - \beta(\theta_{t-2} - \theta_{t-1}) - \eta\lambda\theta_{t-1}.Substitute the recurrence for vtv_t, then step 1 for vt−1v_{t-1}.
  3. θt=(1+β−ηλ)θt−1−βθt−2\theta_t = (1 + \beta - \eta\lambda)\theta_{t-1} - \beta\theta_{t-2}.Collect terms. A second-order linear recurrence with constant coefficients.
  4. Solutions are rtr^t with r2−(1+β−ηλ)r+β=0r^2 - (1 + \beta - \eta\lambda)r + \beta = 0.Substitute θt=rt\theta_t = r^t and divide by rt−2r^{t-2}. The general solution is a combination of the two roots' powers, and the larger root dominates.
  5. Write r=1−εr = 1 - \varepsilon: (1−ε)2−(1+β−ηλ)(1−ε)+β=ηλ−(1−β)ε+ε2−ηλε=0(1 - \varepsilon)^2 - (1 + \beta - \eta\lambda)(1 - \varepsilon) + \beta = \eta\lambda - (1 - \beta)\varepsilon + \varepsilon^2 - \eta\lambda\varepsilon = 0.Expand and cancel: the constant terms give 1−(1+β)+β=01 - (1 + \beta) + \beta = 0, the terms in ε\varepsilon give −2+(1+β)=−(1−β)-2 + (1 + \beta) = -(1 - \beta), and ηλ\eta\lambda survives from the middle product.
  6. To first order, ε=ηλ1−β\varepsilon = \dfrac{\eta\lambda}{1 - \beta}, so r≈1−ηλ1−βr \approx 1 - \dfrac{\eta\lambda}{1 - \beta}.Drop the quadratic terms ε2\varepsilon^2 and ηλε\eta\lambda\varepsilon; the next term is of order ε2/(1−β)\varepsilon^2/(1 - \beta).
  7. Decoupled with gt=0g_t = 0: vt=0v_t = 0 throughout, so θt=(1−ηλ)θt−1\theta_t = (1 - \eta\lambda)\theta_{t-1} exactly.v0=0v_0 = 0 and nothing is added to it.
  8. L2 + momentum shrinks by about 1−ηλ/(1−β)1 - \eta\lambda/(1 - \beta) per step, ten times the decoupled 1−ηλ1 - \eta\lambda at β=0.9\beta = 0.9; to match a decoupled λ\lambda, the L2 coefficient must be (1−β)λ(1 - \beta)\lambdaThe penalty gradient λθ\lambda\theta enters the momentum buffer and is applied 1/(1−β)1/(1 - \beta) times over, exactly as the optimiser page's Problem 4 found for any steady gradient. PyTorch's SGD with momentum=0.9, weight_decay=λ\lambda therefore decays weights ten times faster than the number suggests; a λ\lambda tuned there and moved to an optimiser with decoupled decay is ten times too small.

Problem 4

Adam with an L2 penalty feeds gt+λθt−1g_t + \lambda\theta_{t-1} to the moment estimates. For a constant gradient vector gg, and once the moment estimates have caught up with their input (so that m^t≈g+λθt−1\hat m_t \approx g + \lambda\theta_{t-1} and v^t≈(g+λθt−1)2\hat v_t \approx (g + \lambda\theta_{t-1})^2), write the step. Show that an entry with gi=0g_i = 0 shrinks by about η\eta per step whatever λ\lambda is, and that an entry with ∣gi∣≫λ∣θi∣|g_i| \gg \lambda|\theta_i| receives a decay of about ηλθi/∣gi∣\eta\lambda\theta_i/|g_i|.

  1. Step =−η m^tv^t+ϵ≈−η g+λθt−1∣g+λθt−1∣+ϵ= -\eta\,\dfrac{\hat m_t}{\sqrt{\hat v_t} + \epsilon} \approx -\eta\,\dfrac{g + \lambda\theta_{t-1}}{|g + \lambda\theta_{t-1}| + \epsilon}, entry by entry.Adam's update with the converged estimates; x2=∣x∣\sqrt{x^2} = |x|.
  2. Entry with gi=0g_i = 0: stepi≈−η λθi∣λθi∣+ϵ≈−ηsign⁡(θi)_i \approx -\eta\,\dfrac{\lambda\theta_i}{|\lambda\theta_i| + \epsilon} \approx -\eta\operatorname{sign}(\theta_i).λ\lambda cancels between numerator and denominator when λ∣θi∣≫ϵ\lambda|\theta_i| \gg \epsilon.
  3. Entry with ∣gi∣≫λ∣θi∣|g_i| \gg \lambda|\theta_i|: gi+λθi∣gi+λθi∣≈gi+λθi∣gi∣\dfrac{g_i + \lambda\theta_i}{|g_i + \lambda\theta_i|} \approx \dfrac{g_i + \lambda\theta_i}{|g_i|}, so stepi≈−ηgi∣gi∣−ηλθi∣gi∣_i \approx -\eta\dfrac{g_i}{|g_i|} - \eta\dfrac{\lambda\theta_i}{|g_i|}.The denominator is dominated by ∣gi∣|g_i| (and ϵ\epsilon is negligible next to it); split the numerator. The first term is the plain Adam step.
  4. Adam + L2 step ≈−η (g+λθ)/(∣g+λθ∣+ϵ)\approx -\eta\,(g + \lambda\theta)/(|g + \lambda\theta| + \epsilon): a weight with no data gradient moves by η\eta per step towards 00 regardless of λ\lambda, while a weight with a large gradient is decayed by only ηλθi/∣gi∣\eta\lambda\theta_i/|g_i| per stepThe decay is normalised by the same v^t\sqrt{\hat v_t} 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 λ\lambda 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 λθ\lambda\theta.

Problem 5

AdamW with ϵ>0\epsilon > 0: θt=(1−ηλ)θt−1−η m^t/(v^t+ϵ)\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta\,\hat m_t/(\sqrt{\hat v_t} + \epsilon). Show that a stationary point with gradient g=g(θ)g = g(\theta) satisfies λθi(∣gi∣+ϵ)=−gi\lambda\theta_i(|g_i| + \epsilon) = -g_i for every entry. Deduce what happens for a constant gradient gg with ∣gi∣≫ϵ|g_i| \gg \epsilon, and, for a loss whose gradient vanishes at its minimiser, that the stationary point is the ridge solution with penalty coefficient λϵ\lambda\epsilon.

  1. At a stationary point the gradient is a constant gg, so m^t=g\hat m_t = g and v^t=g2\hat v_t = g^2 once the averages settle.The optimiser page, Problem 8: constant input gives exact bias-corrected moments.
  2. θ=(1−ηλ)θ−η g∣g∣+ϵ\theta = (1 - \eta\lambda)\theta - \eta\,\dfrac{g}{|g| + \epsilon}, so λθi=−gi∣gi∣+ϵ\lambda\theta_i = -\dfrac{g_i}{|g_i| + \epsilon}.Fixed-point condition; cancel η\eta.
  3. λθi(∣gi∣+ϵ)=−gi\lambda\theta_i(|g_i| + \epsilon) = -g_i.Multiply out.
  4. Constant gg with ∣gi∣≫ϵ|g_i| \gg \epsilon: λθi≈−sign⁡(gi)\lambda\theta_i \approx -\operatorname{sign}(g_i), so ∣θi∣≈1/λ|\theta_i| \approx 1/\lambda.gi/(∣gi∣+ϵ)≈sign⁡(gi)g_i/(|g_i| + \epsilon) \approx \operatorname{sign}(g_i). The size of the gradient has no effect on where the weight settles.
  5. Gradient that vanishes at the minimiser: near the stationary point ∣gi∣|g_i| is small, and if ∣gi∣≪ϵ|g_i| \ll \epsilon then λθiϵ≈−gi\lambda\theta_i\epsilon \approx -g_i, that is g(θ)+λϵ θ≈0g(\theta) + \lambda\epsilon\,\theta \approx 0.Step 3 with ∣gi∣+ϵ≈ϵ|g_i| + \epsilon \approx \epsilon. Consistency: it requires ∣gi∣=λϵ∣θi∣≪ϵ|g_i| = \lambda\epsilon|\theta_i| \ll \epsilon, that is λ∣θi∣≪1\lambda|\theta_i| \ll 1, which holds for the usual λ=0.01\lambda = 0.01 and weights of order 11.
  6. On the ridge loss, g(θ)+λϵθ=0g(\theta) + \lambda\epsilon\theta = 0 is Problem 2's fixed-point equation with λϵ\lambda\epsilon in place of λ\lambda.Compare with Problem 2, step 2.
  7. AdamW's stationary points satisfy λθi(∣gi∣+ϵ)=−gi\lambda\theta_i(|g_i| + \epsilon) = -g_i; under a persistent gradient the weight settles at ∣θi∣≈1/λ|\theta_i| \approx 1/\lambda whatever ∣gi∣|g_i|; where the gradient vanishes at the optimum, the stationary point is the ridge solution with penalty λϵ\lambda\epsilon, essentially unregularised at ϵ=10−8\epsilon = 10^{-8}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 1/λ1/\lambda, 100100 at the default λ\lambda) but not the solution of a problem it can solve exactly. Adam with the L2 penalty, by contrast, is a method for LλL_\lambda and does converge to the ridge solution with coefficient λ\lambda, at the cost of Problem 4's normalised decay on the way there.

Problem 6

A scale-invariant weight: L(w)=f(w/∥w∥)L(w) = f(w/\|w\|). Show that w⊤∇L(w)=0w^\top\nabla L(w) = 0 and ∥∇L(w)∥=G/∥w∥\|\nabla L(w)\| = G/\|w\|. For decoupled SGD, wt=(1−ηλ)wt−1−ηgtw_t = (1 - \eta\lambda)w_{t-1} - \eta g_t, show that ∥wt∥2=(1−ηλ)2∥wt−1∥2+η2∥gt∥2\|w_t\|^2 = (1 - \eta\lambda)^2\|w_{t-1}\|^2 + \eta^2\|g_t\|^2, find the equilibrium norm when ∥gt∥=G/∥wt−1∥\|g_t\| = G/\|w_{t-1}\| with GG constant, and the resulting effective learning rate η/∥w∥2\eta/\|w\|^2 on the direction of ww.

  1. L(αw)=L(w)L(\alpha w) = L(w) for α>0\alpha > 0, so ddαL(αw)∣α=1=w⊤∇L(w)=0\tfrac{d}{d\alpha}L(\alpha w)\big|_{\alpha=1} = w^\top\nabla L(w) = 0.αw/∥αw∥=w/∥w∥\alpha w/\|\alpha w\| = w/\|w\|; chain rule along the ray, as on the contrastive page.
  2. ∇L(w)=1∥w∥(I−w^w^⊤)∇f(w^)\nabla L(w) = \dfrac{1}{\|w\|}(I - \hat w\hat w^\top)\nabla f(\hat w).Chain rule through w^=w/∥w∥\hat w = w/\|w\| with the Jacobian 1∥w∥(I−w^w^⊤)\tfrac1{\|w\|}(I - \hat w\hat w^\top) (the Jacobians page).
  3. At ∥w∥=1\|w\| = 1 the same formula gives ∇L(w^)=(I−w^w^⊤)∇f(w^)\nabla L(\hat w) = (I - \hat w\hat w^\top)\nabla f(\hat w), so ∇L(w)=∇L(w^)/∥w∥\nabla L(w) = \nabla L(\hat w)/\|w\| and ∥∇L(w)∥=G/∥w∥\|\nabla L(w)\| = G/\|w\|.Step 2 at w=w^w = \hat w and again at ww: the projection is the same, only the 1/∥w∥1/\|w\| differs. G=∥∇L(w^)∥G = \|\nabla L(\hat w)\| by definition.
  4. ∥wt∥2=(1−ηλ)2∥wt−1∥2−2η(1−ηλ) wt−1⊤gt+η2∥gt∥2=(1−ηλ)2∥wt−1∥2+η2∥gt∥2\|w_t\|^2 = (1 - \eta\lambda)^2\|w_{t-1}\|^2 - 2\eta(1 - \eta\lambda)\,w_{t-1}^\top g_t + \eta^2\|g_t\|^2 = (1 - \eta\lambda)^2\|w_{t-1}\|^2 + \eta^2\|g_t\|^2.Expand the squared norm; the cross term vanishes because gt=∇L(wt−1)g_t = \nabla L(w_{t-1}) is orthogonal to wt−1w_{t-1} (step 1). Pythagoras: the decay shrinks ww along itself and the gradient step moves it sideways.
  5. With n=∥w∥2n = \|w\|^2 at equilibrium and ∥gt∥2=G2/n\|g_t\|^2 = G^2/n: n=(1−ηλ)2n+η2G2/nn = (1 - \eta\lambda)^2n + \eta^2G^2/n, so n2(1−(1−ηλ)2)=η2G2n^2\big(1 - (1 - \eta\lambda)^2\big) = \eta^2G^2.Set ∥wt∥2=∥wt−1∥2=n\|w_t\|^2 = \|w_{t-1}\|^2 = n and multiply through by nn.
  6. 1−(1−ηλ)2=2ηλ−η2λ2≈2ηλ1 - (1 - \eta\lambda)^2 = 2\eta\lambda - \eta^2\lambda^2 \approx 2\eta\lambda, so n2≈ηG22λn^2 \approx \dfrac{\eta G^2}{2\lambda} and ∥w∥≈(ηG22λ)1/4\|w\| \approx \Big(\dfrac{\eta G^2}{2\lambda}\Big)^{1/4}.ηλ≪1\eta\lambda \ll 1; divide and take the fourth root.
  7. η∥w∥2=ηn≈η2ληG2=2ηλG\dfrac{\eta}{\|w\|^2} = \dfrac{\eta}{n} \approx \eta\sqrt{\dfrac{2\lambda}{\eta G^2}} = \dfrac{\sqrt{2\eta\lambda}}{G}.Substitute step 6.
  8. w⊤∇L=0w^\top\nabla L = 0 and ∥∇L(w)∥=G/∥w∥\|\nabla L(w)\| = G/\|w\|; ∥wt∥2=(1−ηλ)2∥wt−1∥2+η2∥gt∥2\|w_t\|^2 = (1 - \eta\lambda)^2\|w_{t-1}\|^2 + \eta^2\|g_t\|^2; equilibrium ∥w∥≈(ηG2/2λ)1/4\|w\| \approx (\eta G^2/2\lambda)^{1/4} and effective learning rate η/∥w∥2≈2ηλ/G\eta/\|w\|^2 \approx \sqrt{2\eta\lambda}/GFor 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 η∥g∥/∥w∥=ηG/∥w∥2\eta\|g\|/\|w\| = \eta G/\|w\|^2 per iteration. That step is 2ηλ/G\sqrt{2\eta\lambda}/G: it depends on η\eta and λ\lambda only through their product, and only as a square root, so dividing η\eta by 1010 at the end of a schedule slows the rotation of a normalised layer by only 10\sqrt{10}, and the same decay with no normalisation layer would be doing something else entirely.

Problem 7

Same layer with no decay (λ=0\lambda = 0). Show that nt=∥wt∥2n_t = \|w_t\|^2 satisfies nt2≈nt−12+2η2G2n_t^2 \approx n_{t-1}^2 + 2\eta^2G^2, hence ∥wT∥≈(n02+2η2G2T)1/4\|w_T\| \approx (n_0^2 + 2\eta^2G^2T)^{1/4}, and describe how the effective learning rate behaves over training.

  1. nt=nt−1+η2G2/nt−1n_t = n_{t-1} + \eta^2G^2/n_{t-1}.Problem 6, step 4 with λ=0\lambda = 0 and ∥gt∥2=G2/nt−1\|g_t\|^2 = G^2/n_{t-1}.
  2. nt2=nt−12+2η2G2+η4G4/nt−12≈nt−12+2η2G2n_t^2 = n_{t-1}^2 + 2\eta^2G^2 + \eta^4G^4/n_{t-1}^2 \approx n_{t-1}^2 + 2\eta^2G^2.Square step 1; the last term is smaller than the middle one by the factor η2G2/(2nt−12)\eta^2G^2/(2n_{t-1}^2), which is tiny once nn is of order 11.
  3. nT2≈n02+2η2G2Tn_T^2 \approx n_0^2 + 2\eta^2G^2T, so ∥wT∥=nT1/2≈(n02+2η2G2T)1/4\|w_T\| = n_T^{1/2} \approx (n_0^2 + 2\eta^2G^2T)^{1/4}.Sum step 2 over TT steps.
  4. The effective learning rate η/nT≈η/2η2G2T=1/(G2T)\eta/n_T \approx \eta/\sqrt{2\eta^2G^2T} = 1/(G\sqrt{2T}) for large TT.Substitute step 3 and drop n02n_0^2.
  5. ∥wT∥≈(n02+2η2G2T)1/4\|w_T\| \approx (n_0^2 + 2\eta^2G^2T)^{1/4}: the norm grows like T1/4T^{1/4} and the effective learning rate on the direction falls like 1/T1/\sqrt T, independent of η\etaWith no decay, a normalised layer applies its own 1/T1/\sqrt T learning-rate schedule, because every gradient step lengthens ww (step 1: the orthogonal step can only add to the norm) and a longer ww 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 2ηλ/G\sqrt{2\eta\lambda}/G 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 θt=(1−ηλ)θt−1−ηut\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta u_t with utu_t the normalised Adam direction. Unroll it to a closed form for θT\theta_T, show that the coefficients of the utu_t sum to 1/λ1/\lambda as T→∞T \to \infty, and, for a learning-rate schedule ηt\eta_t, show that the factor multiplying θ0\theta_0 after TT steps is about e−λ∑tηte^{-\lambda\sum_t\eta_t}. Evaluate it for η=10−3\eta = 10^{-3}, λ=0.1\lambda = 0.1 and T=100,000T = 100{,}000.

  1. θ1=(1−ηλ)θ0−ηu1\theta_1 = (1 - \eta\lambda)\theta_0 - \eta u_1, θ2=(1−ηλ)2θ0−η(1−ηλ)u1−ηu2\theta_2 = (1 - \eta\lambda)^2\theta_0 - \eta(1 - \eta\lambda)u_1 - \eta u_2.Apply the update twice; each earlier term picks up one more factor of 1−ηλ1 - \eta\lambda.
  2. θT=(1−ηλ)Tθ0−η∑t=1T(1−ηλ)T−tut\theta_T = (1 - \eta\lambda)^T\theta_0 - \eta\sum_{t=1}^{T}(1 - \eta\lambda)^{T-t}u_t.Induction on TT: multiplying by 1−ηλ1 - \eta\lambda raises every exponent by one and the new update enters with exponent 00.
  3. η∑t=1T(1−ηλ)T−t=η 1−(1−ηλ)Tηλ→1λ\eta\sum_{t=1}^{T}(1 - \eta\lambda)^{T-t} = \eta\,\dfrac{1 - (1 - \eta\lambda)^T}{\eta\lambda} \to \dfrac1\lambda.Geometric series with ratio 1−ηλ<11 - \eta\lambda < 1; the η\eta cancels.
  4. So θT≈−1λ⋅(weighted average of the ut)\theta_T \approx -\dfrac1\lambda\cdot\big(\text{weighted average of the } u_t\big) with weights ηλ(1−ηλ)T−t\eta\lambda(1 - \eta\lambda)^{T-t}, an exponential moving average with time constant 1/(ηλ)1/(\eta\lambda) steps.The weights of step 3 normalised to sum to 11; they halve every ln⁡2/(ηλ)\ln 2/(\eta\lambda) steps (Problem 1).
  5. With a schedule, θ0\theta_0's factor is ∏t(1−ηtλ)\prod_t(1 - \eta_t\lambda), and ln⁡∏t(1−ηtλ)=∑tln⁡(1−ηtλ)≈−λ∑tηt\ln\prod_t(1 - \eta_t\lambda) = \sum_t\ln(1 - \eta_t\lambda) \approx -\lambda\sum_t\eta_t.ln⁡(1−x)≈−x\ln(1 - x) \approx -x for small xx; the error is of order λ2∑tηt2\lambda^2\sum_t\eta_t^2.
  6. λ∑tηt=0.1×10−3×105=10\lambda\sum_t\eta_t = 0.1\times10^{-3}\times10^5 = 10, so the factor is e−10≈4.5×10−5e^{-10} \approx 4.5\times10^{-5}.Constant η\eta: ∑tηt=ηT\sum_t\eta_t = \eta T.
  7. θT=(1−ηλ)Tθ0−η∑t(1−ηλ)T−tut\theta_T = (1 - \eta\lambda)^T\theta_0 - \eta\sum_t(1 - \eta\lambda)^{T-t}u_t; the update weights sum to 1/λ1/\lambda, so θ\theta is 1/λ1/\lambda times an EMA of the Adam directions with time constant 1/(ηλ)1/(\eta\lambda) steps; θ0\theta_0 is multiplied by about e−λ∑tηte^{-\lambda\sum_t\eta_t}, which is e−10e^{-10} in the exampleEach Adam direction has entries of size about 11, so ∣θi∣|\theta_i| cannot exceed 1/λ1/\lambda (Problem 5 again) and a weight reflects only the last 1/(ηλ)1/(\eta\lambda) steps of updates: 10,00010{,}000 steps in the example, against a run of 100,000100{,}000. The decay in PyTorch's AdamW is multiplied by the scheduled ηt\eta_t, so the EMA horizon stretches as the learning rate decays and the total forgetting is set by the area under the schedule, not by λ\lambda alone.

Problem 9

Decoupled SGD on the quadratic L(θ)=12θ⊤Aθ−b⊤θL(\theta) = \tfrac12\theta^\top A\theta - b^\top\theta with AA symmetric positive definite. Show that θt−θ∗=(I−η(A+λI))(θt−1−θ∗)\theta_t - \theta^* = \big(I - \eta(A + \lambda I)\big)(\theta_{t-1} - \theta^*) with θ∗=(A+λI)−1b\theta^* = (A + \lambda I)^{-1}b, give the condition on η\eta for convergence, and for the optimiser page's A=(3113)A = \begin{pmatrix}3 & 1\\ 1 & 3\end{pmatrix} with λ=1\lambda = 1 find the range of η\eta, the best η\eta, its worst-case factor, and the condition number with and without the decay.

  1. θt=(1−ηλ)θt−1−η(Aθt−1−b)=(I−η(A+λI))θt−1+ηb\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta(A\theta_{t-1} - b) = \big(I - \eta(A + \lambda I)\big)\theta_{t-1} + \eta b.The gradient of the quadratic is Aθ−bA\theta - b; collect the θt−1\theta_{t-1} terms.
  2. θ∗=(A+λI)−1b\theta^* = (A + \lambda I)^{-1}b satisfies θ∗=(I−η(A+λI))θ∗+ηb\theta^* = \big(I - \eta(A + \lambda I)\big)\theta^* + \eta b.(A+λI)θ∗=b(A + \lambda I)\theta^* = b, so η(A+λI)θ∗=ηb\eta(A + \lambda I)\theta^* = \eta b.
  3. Subtracting: θt−θ∗=(I−η(A+λI))(θt−1−θ∗)\theta_t - \theta^* = \big(I - \eta(A + \lambda I)\big)(\theta_{t-1} - \theta^*).Step 1 minus step 2; the ηb\eta b terms cancel.
  4. A+λIA + \lambda I has eigenvalues λi+λ\lambda_i + \lambda with the eigenvectors of AA, so convergence from every start needs ∣1−η(λi+λ)∣<1|1 - \eta(\lambda_i + \lambda)| < 1 for all ii: 0<η<2λmax⁡+λ0 < \eta < \dfrac{2}{\lambda_{\max} + \lambda}.The optimiser page's Problem 2 with A+λIA + \lambda I in place of AA. Adding λI\lambda I shifts every eigenvalue by λ\lambda and leaves the eigenvectors alone.
  5. For the given AA, eigenvalues 22 and 44 become 33 and 55: 0<η<250 < \eta < \tfrac25, best η=23+5=14\eta = \dfrac{2}{3 + 5} = \tfrac14, worst-case factor 5−35+3=14\dfrac{5 - 3}{5 + 3} = \tfrac14.The optimiser page's formulas 2/(λmin⁡+λmax⁡)2/(\lambda_{\min} + \lambda_{\max}) and (λmax⁡−λmin⁡)/(λmax⁡+λmin⁡)(\lambda_{\max} - \lambda_{\min})/(\lambda_{\max} + \lambda_{\min}) applied to 33 and 55.
  6. Condition number λmax⁡+λλmin⁡+λ=53\dfrac{\lambda_{\max} + \lambda}{\lambda_{\min} + \lambda} = \tfrac53, against 42=2\tfrac42 = 2 without decay.The ratio of the largest to the smallest eigenvalue of the matrix the iteration sees.
  7. θt−θ∗=(I−η(A+λI))(θt−1−θ∗)\theta_t - \theta^* = (I - \eta(A + \lambda I))(\theta_{t-1} - \theta^*), θ∗=(A+λI)−1b\theta^* = (A + \lambda I)^{-1}b; converges iff 0<η<2/(λmax⁡+λ)0 < \eta < 2/(\lambda_{\max} + \lambda); for this AA and λ=1\lambda = 1: η<25\eta < \tfrac25, best η=14\eta = \tfrac14 with factor 14\tfrac14 (was 13\tfrac13), condition number 53\tfrac53 (was 22)Decay moves the minimiser from A−1bA^{-1}b to (A+λI)−1b(A + \lambda I)^{-1}b, 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 λi≪λ\lambda_i \ll \lambda is effectively replaced by λ\lambda. The stability limit tightens slightly, from 2/λmax⁡2/\lambda_{\max} to 2/(λmax⁡+λ)2/(\lambda_{\max} + \lambda).

Problem 10

The loss is multiplied by a constant c>0c > 0 (a change of units, or a different reduction over the batch), so the gradient becomes cgtcg_t. Show that SGD with an L2 penalty on cLcL with (η,λ)(\eta, \lambda) is SGD on LL with (cη,λ/c)(c\eta, \lambda/c); that AdamW with ϵ→0\epsilon \to 0 produces the same iterates for cLcL as for LL; and that Adam with an L2 penalty on cLcL with λ\lambda equals Adam with an L2 penalty on LL with λ/c\lambda/c.

  1. SGD + L2 on cLcL: θt=θt−1−η(cgt+λθt−1)=θt−1−cη(gt+λcθt−1)\theta_t = \theta_{t-1} - \eta(cg_t + \lambda\theta_{t-1}) = \theta_{t-1} - c\eta\big(g_t + \tfrac\lambda c\theta_{t-1}\big).Factor cc out of the bracket.
  2. That is SGD + L2 on LL with learning rate cηc\eta and coefficient λ/c\lambda/c.Compare with Problem 1, step 2. The regularisation relative to the data has weakened by cc.
  3. AdamW on cLcL: m^t→cm^t\hat m_t \to c\hat m_t and v^t→c2v^t\hat v_t \to c^2\hat v_t, so m^t/v^t\hat m_t/\sqrt{\hat v_t} is unchanged.Both moment estimates are built from cgtcg_t; the first is linear in it and the second quadratic, and c2v^t=cv^t\sqrt{c^2\hat v_t} = c\sqrt{\hat v_t}. With ϵ→0\epsilon \to 0 the ratio cancels cc exactly.
  4. The decay (1−ηλ)θt−1(1 - \eta\lambda)\theta_{t-1} does not involve the gradient, so AdamW's iterates are identical.The decoupled form touches only θ\theta and ηλ\eta\lambda.
  5. Adam + L2 on cLcL feeds cgt+λθt−1=c(gt+λcθt−1)cg_t + \lambda\theta_{t-1} = c\big(g_t + \tfrac\lambda c\theta_{t-1}\big) to the moments; by step 3 the factor cc cancels in the normalised step.Factor out cc inside the input to the moment estimates, then apply the scale invariance.
  6. SGD + L2: (cL,η,λ)≡(L,cη,λ/c)(cL, \eta, \lambda) \equiv (L, c\eta, \lambda/c); AdamW is invariant to cc (up to ϵ\epsilon); Adam + L2: (cL,λ)≡(L,λ/c)(cL, \lambda) \equiv (L, \lambda/c) at the same η\etaTwo of the three change meaning when the loss is rescaled, and only AdamW's λ\lambda is a property of the optimiser alone. Switching a loss from a sum over the batch to a mean is c=1/Bc = 1/B; under Adam with weight_decay it multiplies the effective regularisation by BB, while the Adam step itself does not change, which is a quiet way to make a λ\lambda 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.

  1. SGD with weight_decay=λ\lambda shrinks weights by ηλ\eta\lambda per stepRight so far: Problem 1.
  2. “Momentum averages the gradients; the decay is a separate term and stays ηλ\eta\lambda.”The shortcut that causes the mistake: in PyTorch's SGD the penalty gradient λθ\lambda\theta is added to the gradient before the momentum buffer, so it is accumulated like everything else.
  3. With momentum β\beta the weights still shrink by ηλ\eta\lambda per stepThey shrink by about ηλ/(1−β)\eta\lambda/(1 - \beta) (Problem 3), ten times more at β=0.9\beta = 0.9. The run trains, with weights ten times more strongly decayed than intended; the λ\lambda that reproduces the old behaviour is (1−β)λ(1 - \beta)\lambda, 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.

  1. L(w)=f(w/∥w∥)L(w) = f(w/\|w\|) for a weight matrix followed by a normalisation layerRight so far: the layer's output is unchanged by the scale of ww.
  2. “Decay pulls ∥w∥\|w\| down, so the layer's output gets smaller and the model is regularised.”The assumption that causes the mistake: the output is invariant to ∥w∥\|w\| (step 1), so it cannot get smaller.
  3. Smaller ∥w∥\|w\| means smaller activations and a simpler functionNothing downstream changes with ∥w∥\|w\|. What decay sets on such a layer is the equilibrium norm (ηG2/2λ)1/4(\eta G^2/2\lambda)^{1/4} and through it the effective learning rate 2ηλ/G\sqrt{2\eta\lambda}/G on the direction of ww (Problem 6); raising λ\lambda there makes the layer learn faster, and removing decay makes its learning rate decay like 1/T1/\sqrt T (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.

  1. The ridge loss LλL_\lambda has minimiser θ∗=(1NX⊤X+λI)−11NX⊤y\theta^* = (\tfrac1NX^\top X + \lambda I)^{-1}\tfrac1NX^\top yRight so far: Problem 2.
  2. “AdamW minimises L+λ2∥θ∥2L + \tfrac\lambda2\|\theta\|^2, so it converges to θ∗\theta^*.”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 v^t\sqrt{\hat v_t}, so there is no penalised loss whose stationary points AdamW finds.
  3. AdamW with weight_decay=λ\lambda converges to θ∗\theta^*Its stationary points satisfy λθi(∣gi∣+ϵ)=−gi\lambda\theta_i(|g_i| + \epsilon) = -g_i (Problem 5): on the ridge loss, where the gradient vanishes at the optimum, that is the ridge solution with coefficient λϵ\lambda\epsilon, about 10−1010^{-10} at the defaults, so AdamW converges to the unregularised least-squares solution. Adam with the L2 penalty does converge to θ∗\theta^*. 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 λ\lambda and the learning rate η\eta, and tuning one is not supposed to touch the other.

  1. AdamW: θt=(1−ηλ)θt−1−ηut\theta_t = (1 - \eta\lambda)\theta_{t-1} - \eta u_tRight so far.
  2. “The decay is λ\lambda, independent of the learning rate.”The slip that causes the mistake: the shrink per step is the product ηλ\eta\lambda (Problem 1), and in PyTorch's AdamW the decay is multiplied by the scheduled learning rate.
  3. Halving η\eta leaves the weight decay unchangedIt halves the decay per step and doubles the EMA time constant 1/(ηλ)1/(\eta\lambda) (Problem 8): the weights remember twice as many updates and θ0\theta_0 is forgotten half as fast, and on a normalised layer the effective learning rate 2ηλ/G\sqrt{2\eta\lambda}/G drops by 2\sqrt2 rather than 22 (Problem 6). A learning-rate sweep with λ\lambda fixed is also a sweep over the decay; to vary one alone, hold ηλ\eta\lambda 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.

  1. Adam's step m^t/(v^t+ϵ)\hat m_t/(\sqrt{\hat v_t} + \epsilon) is unchanged when the gradient is multiplied by ccRight so far: Problem 10, step 3.
  2. “So multiplying the loss by cc (switching from a mean to a sum over the batch) changes nothing under Adam.”The assumption that causes the mistake: the penalty gradient λθ\lambda\theta is added to the data gradient before normalisation, and it was not multiplied by cc.
  3. Adam with weight_decay=λ\lambda gives the same iterates for cLcL as for LLIt gives the iterates for LL with λ/c\lambda/c (Problem 10, step 5): a sum over a batch of 256256 instead of a mean divides the effective regularisation by 256256, 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 λ\lambda 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).