skip to content

In Adam, what changes when weight decay is applied directly to the weights instead of added to the gradient?

level: middleimportance: must knowfreq 65%

answer

  1. Which term rides the denominator
  2. The penalty is divided too
  3. Shrinkage inversely tied to gradient size
  4. Subtract decay after the adaptive step

basics

~20 s

Adding an L2 term to the gradient sends that term through the adaptive per-parameter denominator, so each weight ends up with a different effective shrinkage. Decoupled decay subtracts a fixed fraction of the weight itself, shrinking every weight at the same relative rate.

solid answer

~50 s

There are two places the pull toward zero can enter. Coupled: you form `g <- dL/dw + wd * w` and hand that to the optimizer, so the penalty feeds the moment estimates and is then divided by the adaptive denominator along with everything else, giving a per-step shrinkage of roughly `lr * wd * w / (sqrt(v) + eps)`. Decoupled: the adaptive step is computed from the loss gradient alone and the decay is subtracted separately, `w <- w - lr * m/(sqrt(v) + eps) - lr * wd * w`, which is the AdamW update. The consequence is that in the coupled form the same coefficient produces wildly different real shrinkage per parameter — inversely proportional to how large that parameter's recent gradients have been — while in the decoupled form every weight loses the same fraction of itself per step.

code

python · 13 lines
python
import math

lr, wd, eps = 1e-3, 0.01, 1e-8
w = 0.5  # same weight value in both cases

for name, v in [("large-gradient weight", 1.0), ("small-gradient weight", 1e-6)]:
    coupled = lr * wd * w / (math.sqrt(v) + eps)  # penalty rides the denominator
    decoupled = lr * wd * w                       # decay applied straight to w
    print("%-22s coupled=%.2e  decoupled=%.2e" % (name, coupled, decoupled))

# large-gradient weight  coupled=5.00e-06  decoupled=5.00e-06
# small-gradient weight  coupled=5.00e-03  decoupled=5.00e-06
# same coefficient, 1000x difference in real shrinkage once it is coupled

go deeper

for a junior

Be ready to say that decay nudges every weight toward zero a little on each step, and that the decoupled version subtracts that nudge separately instead of folding it into the gradient the optimizer sees.

for a middle

You should be able to write both updates on a whiteboard and point at the single term that moved: the penalty inside the adaptive fraction versus outside it. Explain why the denominator makes one coefficient mean different things for different parameters.

for a senior

Demonstrate that you check which form an optimizer implements before trusting a decay value inherited from someone else's recipe, and that you can tell an under-regularized run from a badly coupled one by looking at where shrinkage actually landed.

for a principal

Own the argument that regularization strength and step size should be separable knobs across a team's recipes, and be able to say what that separability buys when many models and schedules must be maintained at once.

### The two places a decay term can enter Weight decay is a small pull of every weight toward zero on each optimizer step. There are two structurally different places to inject that pull, and once the optimizer rescales the gradient per parameter, the two stop being the same operation. **Form 1 — an L2 term added to the gradient.** Before the optimizer runs, you form the sum `g <- dL/dw + wd * w`. The optimizer never learns that part of `g` came from a penalty rather than from the data. It builds its running first- and second-moment estimates from this summed vector and then takes its usual step, `w <- w - lr * m / (sqrt(v) + eps)`. **Form 2 — decoupled decay.** The adaptive step is computed from the loss gradient alone, and a separate term is subtracted afterwards: `w <- w - lr * m / (sqrt(v) + eps) - lr * wd * w` That second line is AdamW. Written side by side on a whiteboard the diff is one term moving from inside the fraction to outside it, and that is the whole idea. ### Why the denominator changes the meaning of the coefficient The quantity `sqrt(v)` in the denominator is, roughly, the root-mean-square magnitude of that parameter's recent gradients. Dividing by it is what makes the optimizer adaptive: a parameter with tiny gradients still gets a step of usable size, a parameter with huge gradients gets its step damped. In form 1 the injected term `wd * w` rides through that same division. The shrinkage actually applied per step is approximately `lr * wd * w / (sqrt(v) + eps)` which is *inversely proportional to how big that parameter's gradients have been*. In form 2 the shrinkage is `lr * wd * w` for every parameter, so each weight is multiplied by roughly `(1 - lr * wd)` per step regardless of its gradient history. There is a second, subtler effect in form 1: the penalty contaminates the moment estimates themselves. The running averages now describe gradient-plus-penalty, so the direction of the adaptive step is altered, and the penalty's influence is smeared across many steps by the averaging rather than applied cleanly on the step where it was added. ### What this looks like in a real network Take one network with two very differently scaled layers — an early block whose gradients have RMS around 1e-1 and a late block whose gradients have RMS around 1e-4. Under form 1, one single coefficient produces shrinkage that differs by a factor of about a thousand between them: the high-gradient layer is barely regularized while the quiet layer is squeezed hard. Nothing in the configuration says this; it falls out of the denominator. Under form 2 the same coefficient means the same relative shrinkage everywhere, and the amount of regularization a layer receives is no longer an accident of its gradient scale. This is why practitioners describe the coupled form as regularizing exactly the wrong parameters: the weights doing the most work, which receive the largest gradients, are the ones the penalty touches least. ### Reading the coefficient The number labelled "weight decay" means two different things in the two forms. In the coupled form it is the coefficient of a term that the optimizer will divide by a per-parameter estimate you never see. In the decoupled form it is, directly, the fraction of itself each weight loses per step (times the learning rate). A value tuned under one form carries no guarantee under the other, and the mapping between them is not a single constant — it varies per parameter and drifts over training as gradient magnitudes change. ### Common confusions worth clearing up - *Thinking decoupled decay is still minimizing an explicit L2 objective.* It is not. Under form 2 there is no term in the loss whose gradient the optimizer follows; the shrinkage is a modification of the update rule itself. - *Thinking the difference is just a rescaling of the learning rate.* A rescale would affect all parameters by the same factor. The coupled denominator affects each parameter by a different factor. - *Thinking the epsilon in the denominator makes this negligible.* Epsilon is there to prevent division by zero; it is far smaller than typical values of `sqrt(v)` and does not flatten the per-parameter spread. ### The short version to say out loud Coupled means the penalty goes through the adaptive rescaling and every parameter gets a different real decay. Decoupled means the decay is applied straight to the weights after the adaptive step, so the coefficient means one thing across the whole network. That is the change AdamW makes.

  • Which parameters end up under-regularized when the L2 term is routed through the adaptive denominator?
    The ones with the largest recent gradient magnitudes. Their `sqrt(v)` is large, so the injected `wd * w` term is divided down hardest and the shrinkage they actually receive is smallest. The parameters doing the most work get the least regularization, which is the opposite of what you usually want.
  • Does the decoupled decay term depend on the loss at all?
    No. It is a function only of the current weight value and the coefficient, so it is applied on every step the parameter is updated, including steps where the loss gradient for that parameter is essentially zero. That independence is exactly what makes its effect predictable across parameters.
  • Does decoupling change what objective the optimizer is minimizing?
    It changes the update rule rather than adding an explicit term to the objective. With an adaptive step there is no simple loss whose gradient reproduces a uniform per-step shrinkage, so decoupled decay is best described as a directly specified pull toward zero applied alongside the gradient step, not as a penalty the optimizer is descending.

Think of the adaptive denominator as a per-parameter volume knob. Route the decay signal through the knob and every weight gets a different amount of it; plumb it around the knob and every weight loses the same fraction.

saying these in an interview costs you the question

  • Claims the two forms are interchangeable in any optimizer
  • Says decoupling is just a rescaled learning rate
  • Assumes the same coefficient gives the same shrinkage in both forms
  • Believes the adaptive denominator leaves the penalty untouched
  • Thinks epsilon in the denominator makes the difference negligible

context