Skip to content

The Optimizer protocol: SGD, momentum, Nesterov, weight decay

ModuleM10.2 · build · Python · Pass 2 · 2 to 3 h
You buildpython/tinyllm/optim/sgd.py: the Param and Optimizer protocols, SGD (step, zero_grad, state_dict, load_state_dict)
Contractcourse/contracts/py/tinyllm/optim/sgd.pyi
Testscourse/tests/M10.2/ (what they check: section 4)
NeedsM10.1 gradient descent (the tests compare, or --ref-deps)
Used byL0.5 trains the bigram · later: L2.2 the neural n-gram model, L3.6 the recurrent models; M10.3 (AdamW) and M10.4 (schedules, clipping) implement and drive the same protocol · later: L0.6, L11.1
MilestoneMS-P2 (Pass 2 gate: every math module of the pass checks green, then your autograd bigram trains)
Optional depthGoh, Why Momentum Really Works (Distill, 2017); Sutskever, Martens, Dahl, and Hinton, “On the importance of initialization and momentum in deep learning” (ICML 2013)
  • An optimizer is an object: it holds references to the parameters and the state it carries between steps, and every optimizer of the course exposes the same four methods (test_sgd_is_an_optimizer).
  • Momentum accumulates past gradients, vt=βvt−1+gtv_t = \beta v_{t-1} + g_t, which on an ill-conditioned problem turns a rate of 1−1/κ1 - 1/\kappa into roughly 1−2/κ1 - 2/\sqrt\kappa (test_hand_example, test_momentum_beats_plain_on_ill_conditioned).
  • Coupled weight decay adds λθ\lambda\theta to the gradient before the momentum buffer; for plain SGD that is the same as shrinking θ\theta by 1−ηλ1 - \eta\lambda, a coincidence Adam breaks (test_weight_decay_shrinks_toward_zero, test_matches_torch_golden).
  • Updates happen in place, parameters without a gradient are skipped, and state_dict makes a stopped run resume bit for bit (test_updates_in_place, test_state_dict_resume_bitwise).
Terminal window
ol start M10.2 # stubs sgd.py into your repo, contract alongside
ol tests M10.2 # read the test catalog first: rung R0, you write no tests here
ol check M10.2 # exit code is the verdict
ol check M10.2 --ref-deps # only if your M10.1 is not passing yet
ol diff M10.2 # after passing: your code against the reference

M10.1’s gradient_descent takes a function and returns a trajectory: fine for a two-dimensional quadratic, wrong for training. In L0.5 your bigram’s weights live inside a model, gradients arrive from your autograd one minibatch at a time, the loop runs for thousands of steps, and a run that stops (a closed laptop, a killed worker in Pass 8) must resume exactly where it was. That needs an object that holds the parameters by reference, updates them in place, keeps per-parameter state such as a momentum buffer, and can save and restore that state. The loop then reads opt.zero_grad(); loss.backward(); opt.step() whatever the optimizer is: SGD here, AdamW in M10.3, Muon in M10.6.

SymbolMeaningType / shape
θ\thetaone parameter tensor (p.data)float array
gtg_tits gradient at step tt (p.grad), or Nonesame shape
η\etalearning rate (lr)float ≥0\ge 0
β\betamomentum coefficient (momentum)float in [0,1)[0, 1)
vtv_tmomentum buffer of one parametersame shape as θ\theta
λ\lambdaweight decay coefficient (weight_decay)float ≥0\ge 0
LL, μ\mu, κ=L/μ\kappa = L/\musmoothness, strong convexity, condition number (M10.1)floats

Stochastic gradients. A training loss is an average over examples, and its gradient is an average of per-example gradients. A minibatch gives an unbiased estimate of it, so SGD is gradient descent with a noisy gradient. Everything below works per parameter tensor, with whatever gradient the backward pass produced.

The protocol. A parameter is anything with a float array data and a grad that is an array of the same shape or None. An optimizer has four methods:

MethodDoes
step()update every parameter’s data in place from its grad
zero_grad()set every grad to None
state_dict()return a deep copy of everything needed to resume
load_state_dict(sd)restore it, copying arrays in

zero_grad exists because backward accumulates (M08.2): without it the second step would use the sum of two steps’ gradients. Setting None instead of zeros (PyTorch’s default) lets step skip parameters that received no gradient at all, such as a frozen layer: no update, no decay, no buffer change.

Momentum (the heavy ball). Keep a buffer per parameter:

vt=βvt−1+gt,θt+1=θt−η vt,v_t = \beta v_{t-1} + g_t, \qquad \theta_{t+1} = \theta_t - \eta\, v_t,

with v1=g1v_1 = g_1 on the first step (PyTorch’s convention: no dampening, and the learning rate multiplies the buffer). Unrolled, vt=∑k≥0βkgt−kv_t = \sum_{k \ge 0} \beta^k g_{t-k}: an exponentially weighted sum of past gradients. When gradients agree from step to step (a long shallow valley) they add up, to η/(1−β)\eta/(1 - \beta) times one gradient in the limit; when they flip sign (bouncing across a steep valley) they cancel. On a quadratic with condition number κ\kappa, choosing η=4/(L+μ)2\eta = 4/(\sqrt L + \sqrt\mu)^2 and β=(κ−1κ+1)2\beta = \left(\frac{\sqrt\kappa - 1}{\sqrt\kappa + 1}\right)^2 gives a per-step rate of κ−1κ+1≈1−2/κ\frac{\sqrt\kappa - 1}{\sqrt\kappa + 1} \approx 1 - 2/\sqrt\kappa, against 1−1/κ1 - 1/\kappa for plain gradient descent: for κ=100\kappa = 100, about 0.82 instead of 0.99.

The buffer holds gradients, not steps. η\eta multiplies vtv_t when the step is taken; it never enters the buffer. With a constant η\eta that is only bookkeeping: a buffer ut=βut−1+ηgtu_t = \beta u_{t-1} + \eta g_t with update θt+1=θt−ut\theta_{t+1} = \theta_t - u_t produces the same trajectory. It stops being the same the moment η\eta changes, and the schedules of M10.4 change it before every step. With η\eta outside, a new learning rate scales the whole next step at once; with η\eta inside, the old rate lingers in the buffer for about 1/(1−β)1/(1 - \beta) steps.

Nesterov momentum. Nesterov’s method evaluates the gradient at a look-ahead point. Rewritten in the variables PyTorch stores, it becomes one extra term:

vt=βvt−1+gt,θt+1=θt−η (gt+βvt).v_t = \beta v_{t-1} + g_t, \qquad \theta_{t+1} = \theta_t - \eta\,(g_t + \beta v_t).

It needs momentum to mean anything, so nesterov=True with β=0\beta = 0 is an error.

Weight decay, coupled. The L2 penalty λ2∥θ∥2\tfrac\lambda2 \lVert\theta\rVert^2 adds λθ\lambda\theta to the gradient. SGD applies it first, so it flows through the momentum buffer like any other gradient:

gt←gt+λθt(before momentum).g_t \leftarrow g_t + \lambda\theta_t \quad \text{(before momentum)}.

With plain SGD this is the same as shrinking the weights, θ←(1−ηλ)θ−ηg\theta \leftarrow (1 - \eta\lambda)\theta - \eta g; with a zero gradient, θt=(1−ηλ)tθ0\theta_t = (1 - \eta\lambda)^t \theta_0. With momentum the decay accumulates in the buffer, so adding λθ\lambda\theta after the buffer instead (“decoupled” decay, SGDW) is a different optimizer with a different trajectory. AdamW (M10.3) is built on exactly that difference.

In place, in the parameter’s dtype. The model, its tied weights, and the checkpoint writer hold references to each data array. p.data -= lr * g changes that array; p.data = p.data - lr * g builds a new one and leaves every other reference pointing at the old weights. In-place arithmetic also keeps a float32 parameter float32.

State, and resuming bit for bit. SGD’s state is one momentum buffer per parameter, keyed by the parameter’s position in the list. state_dict copies the buffers and the hyperparameters (PyTorch’s layout: {"state": {i: {"momentum_buffer": ...}}, "param_groups": [{...}]}); a copy, because a checkpoint must not change when training continues. Since every operation is deterministic, 10 steps, save, load into a fresh optimizer, and 10 more steps give exactly the bits of 20 uninterrupted steps.

Minimize f(x)=x2/2f(x) = x^2/2, whose gradient is xx, from x0=1x_0 = 1 with η=0.1\eta = 0.1.

Momentum β=0.9\beta = 0.9:

ttgt=xt−1g_t = x_{t-1}vt=0.9 vt−1+gtv_t = 0.9\,v_{t-1} + g_txt=xt−1−0.1 vtx_t = x_{t-1} - 0.1\,v_t
1110.9
20.90.9+0.9=1.80.9 + 0.9 = 1.80.9−0.18=0.720.9 - 0.18 = 0.72
30.721.62+0.72=2.341.62 + 0.72 = 2.340.72−0.234=0.4860.72 - 0.234 = 0.486

Plain gradient descent would give 0.9,0.81,0.7290.9, 0.81, 0.729: the buffer has already doubled the effective step.

Nesterov:

ttgtg_tvtv_tstep direction gt+0.9 vtg_t + 0.9\,v_txtx_t
1111.90.81
20.810.9+0.81=1.710.9 + 0.81 = 1.710.81+1.539=2.3490.81 + 1.539 = 2.3490.81−0.2349=0.57510.81 - 0.2349 = 0.5751

Weight decay λ=0.1\lambda = 0.1, no momentum: the gradient becomes x+0.1x=1.1x + 0.1x = 1.1, so x1=1−0.11=0.89x_1 = 1 - 0.11 = 0.89, which is (1−ηλ)⋅1−η⋅1=0.99−0.1(1 - \eta\lambda)\cdot 1 - \eta \cdot 1 = 0.99 - 0.1.

These numbers are the first test case in section 4, test_hand_example.

Changing the learning rate, momentum β=0.9\beta = 0.9: take the first momentum step above with η=0.1\eta = 0.1 (v1=1v_1 = 1, x1=0.9x_1 = 0.9), then set η=0.01\eta = 0.01. The buffer is v2=0.9+0.9=1.8v_2 = 0.9 + 0.9 = 1.8 as before, and the step is 0.01⋅1.8=0.0180.01 \cdot 1.8 = 0.018, so x2=0.882x_2 = 0.882. A buffer that had absorbed η\eta would hold u1=0.1u_1 = 0.1, then u2=0.09+0.009=0.099u_2 = 0.09 + 0.009 = 0.099, and land on x2=0.801x_2 = 0.801: ten times the step the schedule asked for. This is test_lr_change_takes_effect_on_the_next_step.

python/tinyllm/optim/sgd.py
@runtime_checkable
class Param(Protocol): data: NDArray; grad: NDArray | None
@runtime_checkable
class Optimizer(Protocol):
def step(self) -> None; def zero_grad(self) -> None
def state_dict(self) -> dict; def load_state_dict(self, sd: dict) -> None
class SGD:
def __init__(self, params, lr: float, momentum: float = 0.0, nesterov: bool = False, weight_decay: float = 0.0)

SGD follows torch.optim.SGD with dampening 0, so the contract’s update order is PyTorch’s: weight decay, then the buffer, then Nesterov, then the in-place update. ValueError for a negative lr, momentum, or weight_decay, and for Nesterov without momentum. load_state_dict refuses a state for a different number or shape of parameters.

TestKINDChecksWhy it matters downstream
test_hand_exampleunitthe section 3 tables for momentum, Nesterov, and decayyou and the test agree on the update order
test_lr_change_takes_effect_on_the_next_stepunitη\eta from 0.1 to 0.01 after one momentum step gives x2=0.882x_2 = 0.882 and v2=1.8v_2 = 1.8schedules (M10.4) set opt.lr before every step
test_matches_torch_goldengolden20-step torch.optim.SGD trajectories, five settings, two parametersrecipes from PyTorch work on your optimizer
test_plain_sgd_equals_gradient_descentdifferentialbit for bit equal to M10.1’s gradient_descentthe protocol wraps the same update
test_updates_in_placeunitp.data is the same array after step; float32 stays float32the model’s references see the update
test_none_grad_is_skipped_and_zero_grad_clearsboundarya None grad means no move, no decay, no buffer; zero_grad sets Nonefrozen and unused parameters
test_state_dict_resume_bitwiseproperty10 + save + load + 10 steps equals 20 steps exactlyresumable training (L0.5, C1)
test_state_dict_is_a_copyunitsaved buffers do not change when training continues; loading copiescheckpoints are snapshots
test_rejects_bad_hyperparametersboundarynegative values and Nesterov without momentumPyTorch’s rules
test_momentum_beats_plain_on_ill_conditionedpropertyκ=100\kappa = 100: plain leaves a gap above 0.02, tuned heavy ball below 10−810^{-8}why momentum exists
test_weight_decay_shrinks_toward_zeropropertyzero gradient: θt=(1−ηλ)tθ0\theta_t = (1 - \eta\lambda)^t \theta_0decay as shrinkage, before M10.3
test_buffers_are_per_parameterunittwo parameters with opposite gradients keep separate buffersno velocity leaks between tensors
test_sgd_is_an_optimizerunitisinstance against both runtime-checkable protocolsloops accept any optimizer
PitfallSymptomCaught by
1. adding weight decay after the momentum buffera different optimizer (SGDW); matches PyTorch only without momentumtest_matches_torch_golden (mutant s04)
2. one momentum buffer for every parametervelocity leaks between tensors, or shapes fail to broadcasttest_buffers_are_per_parameter (mutant s10)
3. p.data = p.data - lr * gthe model keeps training on stale weights it still referencestest_updates_in_place (mutant s05)
4. a state_dict of live buffers, or a load that drops thema checkpoint that changes after saving; a resumed run that diverges from the originaltest_state_dict_is_a_copy (mutant s08), test_state_dict_resume_bitwise (mutant s09)
5. momentum as an average, v=βv+(1−β)gv = \beta v + (1 - \beta) gsteps 1−β1 - \beta times smaller than PyTorch’stest_hand_example (mutant s01)
6. Nesterov ignored or with its terms swappedplain momentum where look-ahead was asked fortest_hand_example (mutants s02, s03)
7. zeros instead of None, or decaying parameters without a gradientfrozen layers shrink every steptest_none_grad_is_skipped_and_zero_grad_clears (mutants s06, s07)
8. the sign of the decay termweights pushed away from zerotest_weight_decay_shrinks_toward_zero (mutant s12)
9. folding η\eta into the momentum bufferidentical while η\eta is constant; under a schedule the old rate lingers for about 1/(1−β)1/(1-\beta) stepstest_lr_change_takes_effect_on_the_next_step (mutant s11)

| Forward | L0.6 | Registered call site uses this module. | | Forward | L11.1 | Registered call site uses this module. |

DirectionModuleHow it uses this
BackM10.1the gradient descent step this generalizes; plain SGD reproduces it bit for bit
ForwardL0.5the bigram’s training loop: zero_grad, backward, step, and state_dict in the checkpoint
ForwardL2.2the neural n-gram model trains with SGD and momentum
ForwardL3.6the recurrent language models train through the same protocol
ForwardM10.3AdamW implements the same four methods, with decoupled decay
ForwardM10.4schedules set lr between steps; clipping rescales grad before step

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

Your pieceProduction equivalentWhat it addsWhere to look
SGD.stepPyTorch torch.optim.SGDparameter groups with their own hyperparameters, dampening, foreach and fused multi-tensor kernelstorch/optim/sgd.py (_single_tensor_sgd)
the Optimizer protocoloptaxoptimizers as composable pure gradient transformations (trace for momentum, add_decayed_weights, scale)optax/_src/alias.py (sgd)
state_dictPyTorch distributed checkpointingsharded optimizer state saved and resharded across rankstorch/distributed/checkpoint/