Skip to content

Adam and AdamW: bias correction and decoupled weight decay

ModuleM10.3 · build · Python · Pass 2 · 3 to 4 h
You buildpython/tinyllm/optim/adamw.py: Adam and AdamW, with step, zero_grad, state_dict, load_state_dict
Contractcourse/contracts/py/tinyllm/optim/adamw.pyi · checkpoint keys: formats/checkpoint.md
Testscourse/tests/M10.3/test_adamw.py (what they check: section 4); golden trajectories from torch 2.14 in course/fixtures/M10.3/adam_torch.json
Needsno code dependency. Reading: M10.2 (the Optimizer protocol and SGD with momentum), M02.2 (the exponential moving average and its bias correction), S-M10a (the Adam derivation problems)
Used byL0.5 trains every model with it from Part 0 on · L0.6 saves and restores its state in checkpoints · later L4.1, L5.5, L6.1, L7.9, C1, and L12.1 · later: L11.1, L3.6, L6.5
MilestoneMS-P2 (Pass 2 closes with every math module it teaches passing)
Optional depthKingma and Ba, “Adam: A Method for Stochastic Optimization” (2015), sections 2 and 3; Loshchilov and Hutter, “Decoupled Weight Decay Regularization” (2019), sections 2 and 3; the torch source torch/optim/adam.py, _single_tensor_adam
  • Adam keeps two exponential moving averages per coordinate, mm of the gradient and vv of its square, and moves each coordinate by η m^/(v^+ϵ)\eta\, \hat m / (\sqrt{\hat v} + \epsilon) (test_hand_example_two_steps, test_matches_torch_adamw_trajectory).
  • Bias correction divides mtm_t by 1−β1t1 - \beta_1^t and vtv_t by 1−β2t1 - \beta_2^t, which undoes the pull toward the zero they start from; the first step is then exactly η\eta per coordinate (test_first_step_moves_each_coordinate_by_lr).
  • Dividing by v^\sqrt{\hat v} makes the update invariant to the scale of the gradient, so one learning rate serves layers whose gradients differ by orders of magnitude (test_update_ignores_gradient_scale).
  • AdamW decouples weight decay: it multiplies the weights by 1−ηλ1 - \eta\lambda instead of adding λw\lambda w to the gradient, where Adam’s normalization would wash the penalty out (test_decoupled_decay_with_zero_gradient).
  • The optimizer state is part of the model: a save and load between steps must reproduce the uninterrupted run bit for bit (test_resume_is_bitwise).
Terminal window
ol start M10.3 # stubs python/tinyllm/optim/adamw.py, contract alongside
ol tests M10.3 # read the test catalog first: rung R0, you write no tests here
ol check M10.3 # exit code is the verdict
ol diff M10.3 # after passing: your code against the reference

Your training loop does not exist yet. L0.5 builds it next, and it needs an optimizer that trains every model in the course without retuning: a bigram table, a 64-unit MLP on digits, an LSTM, a Transformer, and the 10M-parameter Llama-style model of C1. Plain SGD with momentum (M10.2) needs a learning rate tuned per model, and inside one model it is too slow for layers with small gradients and unstable for layers with large ones: an embedding row that sees a rare token gets a gradient a thousand times smaller than a layer norm gain. Adam fixes this by rescaling each coordinate by its own gradient history, and AdamW is the version every modern language model trains with. The C1 training spec defaults to adamw with betas (0.9, 0.95), and the checkpoint format already reserves exp_avg and exp_avg_sq for its state, so dur.11 can kill a training worker and resume it. This module builds that optimizer and proves it equals torch.optim.AdamW step for step.

SymbolMeaningType / shape
wwone parameter tensor (the optimizer treats each coordinate the same way)float32[...], p.data
gtg_tthe gradient of the loss with respect to ww at step ttsame shape as ww, p.grad
ttthe step counter: 1 for the first updateint
η\etathe learning rate, read on every stepfloat, opt.lr
β1,β2\beta_1, \beta_2decay rates of the two moving averages, in [0,1)[0, 1)float, opt.betas
mtm_tfirst moment: moving average of ggsame shape as ww, exp_avg
vtv_tsecond moment: moving average of g2g^2 (elementwise)same shape as ww, exp_avg_sq
m^t,v^t\hat m_t, \hat v_tbias-corrected moments, mt/(1−β1t)m_t / (1 - \beta_1^t) and vt/(1−β2t)v_t / (1 - \beta_2^t)same shape as ww
ϵ\epsilona small constant that keeps the denominator away from 0float, default 10−810^{-8}
λ\lambdathe weight decay coefficientfloat, opt.weight_decay

All operations on vectors are elementwise: g2g^2 squares each coordinate, v\sqrt{v} takes each square root, and m/vm / \sqrt{v} divides coordinate by coordinate.

Gradient descent moves ww against the gradient: w←w−η gtw \leftarrow w - \eta\, g_t. A minibatch gradient is noisy, and its direction changes from step to step. Momentum (M10.2) smooths it with an exponential moving average (EMA, M02.2):

mt=β1 mt−1+(1−β1) gt,m0=0.m_t = \beta_1\, m_{t-1} + (1 - \beta_1)\, g_t, \qquad m_0 = 0 .

Unrolling the recursion shows what mtm_t is: a weighted sum of all past gradients with geometrically decaying weights, mt=(1−β1)∑s=1tβ1t−sgsm_t = (1 - \beta_1) \sum_{s=1}^{t} \beta_1^{t-s} g_s. With β1=0.9\beta_1 = 0.9 the last 10 or so steps carry most of the weight. Adam keeps a second EMA, of the squared gradient:

vt=β2 vt−1+(1−β2) gt2,v0=0.v_t = \beta_2\, v_{t-1} + (1 - \beta_2)\, g_t^2, \qquad v_0 = 0 .

With β2=0.999\beta_2 = 0.999 it remembers about 1000 steps. vtv_t estimates E[g2]\mathbb{E}[g^2], the mean squared size of this coordinate’s gradient.

Adam’s update divides the first moment by the square root of the second:

wt=wt−1−η m^tv^t+ϵ.w_t = w_{t-1} - \eta\, \frac{\hat m_t}{\sqrt{\hat v_t} + \epsilon} .

v^t\sqrt{\hat v_t} has the units of the gradient, so the ratio m^t/v^t\hat m_t / \sqrt{\hat v_t} has no units: multiplying every gradient by a constant c>0c > 0 multiplies both m^t\hat m_t and v^t\sqrt{\hat v_t} by cc and leaves the update unchanged (when ϵ\epsilon is negligible). Each coordinate’s step is therefore about η\eta in size when its gradient keeps a consistent sign, and smaller when the sign flips (then ∣m^∣≪v^|\hat m| \ll \sqrt{\hat v}, because positive and negative gradients cancel in mm but not in vv). The learning rate becomes a step length in parameter space, which is why one value such as 3×10−43 \times 10^{-4} works across many models. ϵ\epsilon only matters when v^\sqrt{\hat v} is near or below it, for coordinates whose gradients are almost always zero.

Both averages start at 0, so early on they are too small. Take every gsg_s equal to the same value gg. Then mt=(1−β1) g∑s=0t−1β1s=(1−β1t) gm_t = (1 - \beta_1)\, g \sum_{s=0}^{t-1} \beta_1^{s} = (1 - \beta_1^t)\, g by the geometric sum (M00.3). The average is short of gg by exactly the factor 1−β1t1 - \beta_1^t, and dividing by it removes the bias:

m^t=mt1−β1t,v^t=vt1−β2t.\hat m_t = \frac{m_t}{1 - \beta_1^t}, \qquad \hat v_t = \frac{v_t}{1 - \beta_2^t} .

For random gradients with constant mean the same argument holds in expectation. The factor matters most for vv: with β2=0.999\beta_2 = 0.999, 1−β21=0.0011 - \beta_2^{1} = 0.001, so the uncorrected v1v_1 is 1000 times too small and the first step would be 1000≈32\sqrt{1000} \approx 32 times too large in the ratio m/vm / \sqrt{v} (partly offset by the uncorrected mm, which is 10 times too small: net 3.163.16 times too large). After a few thousand steps both factors are 1 and the correction disappears.

At t=1t = 1 the corrections give m^1=g1\hat m_1 = g_1 and v^1=g12\hat v_1 = g_1^2, so the first update is η g1/(∣g1∣+ϵ)\eta\, g_1 / (|g_1| + \epsilon): exactly η\eta times the sign of g1g_1 when ϵ=0\epsilon = 0. torch folds the corrections into the step for speed, and you implement the same arrangement, because it decides where ϵ\epsilon goes:

wt=wt−1−η1−β1t⋅mtvt/1−β2t+ϵ.w_t = w_{t-1} - \frac{\eta}{1 - \beta_1^t}\cdot \frac{m_t}{\sqrt{v_t} / \sqrt{1 - \beta_2^t} + \epsilon} .

Weight decay shrinks the weights toward 0 a little on every step, which regularizes the model. For SGD it is the same as adding the penalty λ2∥w∥2\frac{\lambda}{2}\lVert w \rVert^2 to the loss, whose gradient is λw\lambda w: the step w−η(g+λw)=(1−ηλ) w−ηgw - \eta(g + \lambda w) = (1 - \eta\lambda)\, w - \eta g decays ww by the factor 1−ηλ1 - \eta\lambda.

For Adam the two are not the same. Adam with L2 (torch.optim.Adam(weight_decay=...)) adds λw\lambda w to gg before both moments see it, and then the normalization divides it by v^\sqrt{\hat v}: a coordinate with large gradients gets almost no decay, and with g=0g = 0 the penalty itself becomes a normalized step of size η\eta, whatever the weight’s size. AdamW applies the decay directly to the weights, before the Adam update, and leaves the gradient alone:

w←(1−ηλ) w,then the Adam step on gt.w \leftarrow (1 - \eta\lambda)\, w, \quad \text{then the Adam step on } g_t .

Every coordinate shrinks by the same factor, as intended. The decay is scaled by η\eta, so a schedule (M10.4) that lowers η\eta lowers the decay with it. Parameters whose gradient is None (the loss did not reach them this step) are skipped entirely: no decay, no moment update.

The optimizer state is tt, mm, and vv for every parameter, plus the hyperparameters. Training that stops and resumes (dur.11 kills workers on purpose) must continue as if it never stopped, so state_dict returns copies of everything, and load_state_dict restores all of it, including tt: with tt reset to 0 the bias correction would restart and the next step would be far too large. formats/checkpoint.md writes the moments as <name>.exp_avg and <name>.exp_avg_sq, in the parameter’s dtype.

One scalar parameter x0=1x_0 = 1, gradients g1=2g_1 = 2 then g2=−1g_2 = -1 (given, not computed from a loss), η=0.1\eta = 0.1, β1=0.9\beta_1 = 0.9, β2=0.999\beta_2 = 0.999, ϵ=0\epsilon = 0, no weight decay.

Stepmtm_tvtv_t1−β1t1 - \beta_1^t1−β2t1 - \beta_2^tm^t\hat m_tv^t\hat v_tupdate ηm^/v^\eta \hat m / \sqrt{\hat v}xtx_t
10.1⋅2=0.20.1 \cdot 2 = 0.20.001⋅4=0.0040.001 \cdot 4 = 0.0040.10.10.0010.00122440.1⋅2/2=0.10.1 \cdot 2 / 2 = 0.10.90.9
20.9⋅0.2+0.1⋅(−1)=0.080.9 \cdot 0.2 + 0.1 \cdot (-1) = 0.080.999⋅0.004+0.001⋅1=0.0049960.999 \cdot 0.004 + 0.001 \cdot 1 = 0.0049960.190.190.0019990.0019998/19≈0.421058/19 \approx 0.421054996/1999≈2.499254996/1999 \approx 2.499250.1⋅0.42105/1.58090≈0.0266340.1 \cdot 0.42105 / 1.58090 \approx 0.026634≈0.873366\approx 0.873366

Two things to see. Step 1 moves by exactly η=0.1\eta = 0.1: that is bias correction at work (without it, 0.1⋅0.2/0.004≈0.3160.1 \cdot 0.2 / \sqrt{0.004} \approx 0.316). Step 2 has a negative gradient, yet xx keeps decreasing: mm still points the old way, and the step is smaller because v^\hat v remembers the large first gradient. This is the first test, test_hand_example_two_steps.

With weight decay λ=0.1\lambda = 0.1. AdamW first multiplies xx by 1−ηλ=0.991 - \eta\lambda = 0.99: step 1 gives 0.99−0.1=0.890.99 - 0.1 = 0.89; step 2 gives 0.89⋅0.99−0.026634≈0.8544660.89 \cdot 0.99 - 0.026634 \approx 0.854466. Adam with L2 instead feeds g1+λx0=2.1g_1 + \lambda x_0 = 2.1 and g2+λx1=−1+0.09=−0.91g_2 + \lambda x_1 = -1 + 0.09 = -0.91 to the moments: step 1 still moves by exactly 0.10.1 (the first step is η\eta times a sign, whatever the gradient), to 0.90.9; step 2 lands at ≈0.868123\approx 0.868123. The penalty changed x2x_2 by 0.00520.0052 under L2 and by 0.01890.0189 under AdamW (test_hand_example_adamw_and_adam_l2_differ).

class Adam:
lr: float; betas: tuple[float, float]; eps: float; weight_decay: float
def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0.0) -> None
def step(self) -> None # one update of every p whose grad is not None, in place
def zero_grad(self) -> None # every p.grad = None
def state_dict(self) -> dict # {"step", "lr", "betas", "eps", "weight_decay", "exp_avg": [...], "exp_avg_sq": [...]}
def load_state_dict(self, sd: dict) -> None
class AdamW(Adam): # decoupled decay; weight_decay defaults to 0.01
def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=0.01) -> None

A parameter is any object with data (a float ndarray) and grad (an ndarray or None); the tests use a two-attribute class, and L0.5 passes autograd tensors (L0.1). The step counter tt is one global count of step() calls, as formats/checkpoint.md stores one step. torch keeps a counter per parameter instead; the two agree whenever every parameter has a gradient on every step, which is how the course trains. A schedule changes the learning rate by assigning opt.lr before step().

TestKINDChecksWhy it matters downstream
test_hand_example_two_stepsunitsection 3, both steps, to float64 precisionyou and the test agree on the update before any code
test_hand_example_adamw_and_adam_l2_differunitsection 3 with λ=0.1\lambda = 0.1 for both kinds of decaythe one difference between the two classes
test_matches_torch_adamw_trajectorygolden20 steps of torch AdamW on two tensors, and the final momentsevery training run assumes torch’s update
test_matches_torch_adam_l2_trajectorygolden20 steps of torch Adam with weight_decay = 0.1L2 coupling feeds both moments
test_lr_set_between_steps_is_usedgoldenlr changes on every step (betas (0.9, 0.95))M10.4 schedules write opt.lr
test_eps_is_added_after_bias_correctiongoldengradients of size 10−610^{-6} with ϵ=10−6\epsilon = 10^{-6}where ϵ\epsilon goes is visible for rare features
test_float32_parameters_stay_float32goldentorch’s float32 trajectory; moments stay float32model parameters are float32 (L0.4)
test_first_step_moves_each_coordinate_by_lrpropertywith ϵ=0\epsilon = 0 the first step is η⋅sign(g)\eta \cdot \mathrm{sign}(g) at every scalebias correction, stated as a law
test_update_ignores_gradient_scalepropertygradients times 10−310^{-3} or 250250 give the same trajectoryone lr for every layer
test_decoupled_decay_with_zero_gradientunitg=0g = 0: AdamW gives 3.83.8, Adam with L2 gives 3.93.9decoupling, in one line
test_parameter_without_grad_is_untouchedboundarygrad = None: no decay, no moment change, tt still advancesunused embeddings and frozen branches
test_update_is_in_placeunitp.data keeps its identity, a view writes through to its bufferthe model and L11 flat buffers hold the arrays
test_zero_grad_sets_noneunitevery grad becomes NoneL0.1 accumulates into grad
test_resume_is_bitwiseproperty5 steps, save, load into a fresh optimizer, 5 steps equals 10 steps bit for bitdur.11 kill and resume
test_state_dict_is_a_snapshotunitlater steps do not change a saved dictL0.6 may write it after training continues
test_state_dict_layoutunitkeys, shapes, hyperparameters restored, mismatches rejectedformats/checkpoint.md
test_rejects_bad_hyperparametersboundarynegative lr, eps, or decay, a beta outside [0,1)[0, 1), no parameters, integer dataerrors at construction, not NaN at step 9000
PitfallSymptomCaught by
no bias correction, or correcting only one momentthe first steps are about 3 times too large (or too small); early loss spikestest_first_step_moves_each_coordinate_by_lr, test_hand_example_two_steps (mutants s01, s03)
counting tt from 0 inside the correction, or using t+1t + 11−β0=01 - \beta^0 = 0 divides by zero, or every step is mis-scaledtest_hand_example_two_steps (mutant s02)
AdamW adding λw\lambda w to the gradientAdamW behaves like Adam with L2: decay is normalized away for large-gradient weightstest_decoupled_decay_with_zero_gradient, test_hand_example_adamw_and_adam_l2_differ (mutant s05)
decaying by 1−λ1 - \lambda instead of 1−ηλ1 - \eta\lambdawith λ=0.1\lambda = 0.1 the weights lose 10% per step and collapse toward 0test_matches_torch_adamw_trajectory (mutant s06)
ϵ\epsilon inside the square root, or before the correctionrare-feature coordinates take steps of the wrong size; only visible when gradients are near ϵ\epsilontest_eps_is_added_after_bias_correction (mutants s08, s14)
caching the learning rate at constructionthe schedule (M10.4) is silently ignoredtest_lr_set_between_steps_is_used (mutant s13)
p.data = p.data - ... instead of p.data -= ...the optimizer updates a copy; the model keeps its old weightstest_update_is_in_place (mutant s16)
returning live arrays from state_dict, or not restoring tta resumed run differs from the uninterrupted one; bias correction restarts after every resumetest_state_dict_is_a_snapshot, test_resume_is_bitwise (mutants s19, s22)
decaying parameters whose grad is Noneembeddings of tokens absent from the batch shrink anywaytest_parameter_without_grad_is_untouched (mutant s07)

| Forward | L11.1 | Registered call site uses this module. | | Forward | L3.6 | Registered call site uses this module. | | Forward | L6.5 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackM10.2the Optimizer protocol and SGD with momentum; Adam’s mm is momentum’s EMA form (reading)
BackM02.2the EMA as a geometric series and its bias correction, here applied per coordinate (reading)
BackS-M10athe Adam derivation problems check section 2.3 by hand (reading)
ForwardL0.5train_step calls opt.zero_grad(), the backward pass, then opt.step() on every model in Part 0
ForwardL0.6the checkpoint writer stores state_dict() as <name>.exp_avg, <name>.exp_avg_sq, and the step, and resume calls load_state_dict
ForwardM10.4schedules set opt.lr between steps; clipping runs before step()
ForwardL4.1trains the encoder-decoder with teacher forcing
ForwardL5.5the 2017 Transformer with the Noam schedule
ForwardL6.1pretraining objectives
ForwardL7.9 and C1the modern decoder and the TinyStories capstone, betas (0.9,0.95)(0.9, 0.95), λ=0.1\lambda = 0.1, checkpoints with exp_avg and exp_avg_sq
ForwardL12.1supervised fine-tuning

If you skip this module, ol check L0.5 stops with BLOCKED ... needs M10.3: build it, or pass --ref-deps.

Your pieceProduction equivalentWhat it addsWhere to look
AdamW.step loop over parameterstorch.optim.AdamW foreach and fused pathsone kernel launch for all parameters at once; the fused CUDA kernel also handles AMP grad scalingtorch/optim/adam.py, _multi_tensor_adam, _fused_adam
one global step countertorch per-parameter state["step"]correct bias correction for parameters that start receiving gradients latetorch/optim/adam.py
float32 moments8-bit optimizers (bitsandbytes)block-wise quantized mm and vv: optimizer memory drops by 4 timesDettmers et al., “8-bit Optimizers via Block-wise Quantization” (2022)
state on one processZeRO stage 1 (DeepSpeed, FSDP)each data-parallel rank owns a shard of mm and vv; your optional L11.3Rajbhandari et al., “ZeRO” (2020)
AdamW for every matrixMuonorthogonalized momentum for 2-D weights, AdamW for the rest; your optional M10.6Jordan et al., “Muon” (2024)