Skip to content

Learning-rate schedules (cosine, WSD, Noam) and gradient clipping

ModuleM10.4 · build · Python · Pass 2 · 2 to 3 h
You buildpython/tinyllm/optim/schedule.py: cosine_with_warmup, wsd, noam, clip_grad_norm_
Contractcourse/contracts/py/tinyllm/optim/schedule.pyi · the training spec’s schedule and grad_clip fields: formats/train-spec.schema.json
Testscourse/tests/M10.4/test_schedule.py (what they check: section 4); HF and torch goldens in course/fixtures/M10.4/schedule_hf.json
Needsno code dependency. Reading: M00.3 (sequences and the geometric sum), M10.2 (the optimizer whose lr a schedule sets)
Used byL0.5 sets opt.lr from a schedule and clips before every step · later L3.6 (clipping for recurrent networks), L5.5 (Noam), C1 (WSD) · later: L11.1, L6.1
MilestoneMS-P2 (Pass 2 closes with every math module it teaches passing)
Optional depthLoshchilov and Hutter, “SGDR: Stochastic Gradient Descent with Warm Restarts” (2017); Vaswani et al., “Attention Is All You Need” (2017), section 5.3; Hägele et al., “Scaling Laws and Compute-Optimal Training Beyond Fixed Training Durations” (2024); Pascanu, Mikolov, and Bengio, “On the difficulty of training recurrent neural networks” (2013), section 3.2
  • A schedule is a pure function from the step number to a learning rate, so a resumed run gets the same rate from the step count alone (test_schedules_hit_their_corners).
  • Warmup ramps the rate linearly from 0 while Adam’s second moment is still unreliable; cosine decay then lowers it smoothly to a floor at a fixed total (test_hand_example_cosine, test_cosine_matches_hf).
  • Warmup-stable-decay holds the peak rate and decays only at the end, so one run can be stopped and annealed at any length (test_hand_example_wsd, test_wsd_matches_hf).
  • The Noam schedule rises linearly and decays like 1/t1/\sqrt{t}, peaking at (dmodel⋅warmup)−1/2(d_{\text{model}} \cdot \text{warmup})^{-1/2} (test_noam_peak_and_inverse_sqrt_decay).
  • Global-norm clipping rescales all gradients by one factor when their joint norm exceeds a limit: it caps the step size and keeps the direction (test_hand_example_clip, test_clip_is_global_and_in_place).
Terminal window
ol start M10.4 # stubs python/tinyllm/optim/schedule.py, contract alongside
ol tests M10.4 # read the test catalog first: rung R0, you write no tests here
ol check M10.4 # exit code is the verdict
ol diff M10.4 # after passing: your code against the reference

With M10.3 your optimizer takes steps of about η\eta per coordinate, but it takes them with the same η\eta from the first step to the last. That fails at both ends. At step 1, Adam’s second moment has seen one gradient, so its normalization is noisy, and a full-size step from random initial weights can push a Transformer into a region it never recovers from (the loss jumps and stays high). At the end, a constant rate keeps the weights bouncing around the minimum at a distance proportional to η\eta, and the loss stops falling. The fix is a schedule: start small, run at the peak, end small. A second failure is a single bad batch: one gradient a thousand times larger than usual, common in recurrent networks (L3.6), which a normalized optimizer still follows for a step and momentum for many more. Gradient clipping bounds it. L0.5 builds the train loop next; it calls a schedule before every step and clips before every update, and the C1 training spec names both (schedule.name = "wsd", grad_clip = 1.0).

SymbolMeaningType / shape
ttthe step: the number of optimizer updates already taken, 0 for the firstint, step
η(t)\eta(t)the learning rate used for update number t+1t + 1float
ηmax⁡,ηmin⁡\eta_{\max}, \eta_{\min}the peak rate and the floorfloat, lr_max, lr_min
WWwarmup length in stepsint, warmup
TTthe step at which the cosine reaches the floorint, total
S,DS, DWSD’s stable and decay lengthsint, stable, decay
ppprogress through a phase, in [0,1][0, 1]float
dmodeld_{\text{model}}the Transformer’s model width (Noam)int
g(i)g^{(i)}the gradient of parameter tensor iindarray, p.grad
∥g∥\lVert g \rVertthe global norm, ∑i∑k(gk(i))2\sqrt{\sum_i \sum_k (g^{(i)}_k)^2}float
ccthe clipping limitfloat, max_norm

The train loop (L0.5) does opt.lr = schedule(t, ...) and then opt.step(), for t=0,1,2,…t = 0, 1, 2, \dots Because the rate depends only on tt and the configuration, a run resumed from a checkpoint at step 5000 recomputes exactly the rate it would have used: no scheduler object has to be saved. The step is the number of updates already taken, the convention of HF’s LambdaLR schedulers, so your values equal HF’s when read the same way. Every warmup ramps linearly from 0: η(t)=ηmax⁡ t/W\eta(t) = \eta_{\max}\, t / W for t<Wt < W. Update 1 therefore runs at rate 0 and changes nothing except the optimizer’s moments, which is harmless and keeps the formulas exact at the phase boundaries.

After warmup, the rate follows half a period of a cosine from ηmax⁡\eta_{\max} down to ηmin⁡\eta_{\min}:

η(t)=ηmin⁡+(ηmax⁡−ηmin⁡) 1+cos⁡(πp)2,p=t−WT−W,W≤t≤T,\eta(t) = \eta_{\min} + (\eta_{\max} - \eta_{\min})\, \frac{1 + \cos(\pi p)}{2}, \qquad p = \frac{t - W}{T - W}, \qquad W \le t \le T,

and η(t)=ηmin⁡\eta(t) = \eta_{\min} for t≥Tt \ge T. At p=0p = 0, cos⁡0=1\cos 0 = 1 gives ηmax⁡\eta_{\max}; at p=1p = 1, cos⁡π=−1\cos \pi = -1 gives ηmin⁡\eta_{\min}; at p=1/2p = 1/2 the rate is halfway. The curve is flat at both ends (its derivative in pp is −π2(ηmax⁡−ηmin⁡)sin⁡(πp)-\frac{\pi}{2}(\eta_{\max} - \eta_{\min}) \sin(\pi p), zero at p=0p = 0 and p=1p = 1), so the peak lasts a while and the final steps change the weights very little. The weakness is TT: it must be fixed before training starts, and stopping early leaves the rate high.

WSD (warmup-stable-decay, used by MiniCPM and many recent models) splits the run into three phases: the linear warmup over WW steps, a stable phase at ηmax⁡\eta_{\max} for SS steps, and a short decay over the last DD steps, here a straight line:

η(t)=ηmax⁡−(ηmax⁡−ηmin⁡) p,p=t−W−SD,W+S≤t<W+S+D,\eta(t) = \eta_{\max} - (\eta_{\max} - \eta_{\min})\, p, \qquad p = \frac{t - W - S}{D}, \qquad W + S \le t < W + S + D,

and ηmin⁡\eta_{\min} from t=W+S+Dt = W + S + D on. The stable phase can be extended at will, and a decay can branch off any stable checkpoint, so one long run yields models at many lengths. Most of the loss improvement of a WSD run arrives during the short decay, which C1 lets you see in its loss curve. (HF’s get_wsd_schedule with a floor starts its warmup at the floor instead of 0; yours starts at 0 like the cosine, and the tests compare with HF from the end of warmup on.)

The 2017 Transformer used

η(t)=dmodel−1/2⋅min⁡ ⁣(t−1/2, t⋅W−3/2).\eta(t) = d_{\text{model}}^{-1/2} \cdot \min\!\left(t^{-1/2},\ t \cdot W^{-3/2}\right).

For t<Wt < W the second term is smaller (because t⋅W−3/2<t−1/2t \cdot W^{-3/2} < t^{-1/2} exactly when t3/2<W3/2t^{3/2} < W^{3/2}), so the rate grows linearly; for t>Wt > W it decays like 1/t1/\sqrt{t}. The two meet at t=Wt = W, where η=dmodel−1/2W−1/2=(dmodelW)−1/2\eta = d_{\text{model}}^{-1/2} W^{-1/2} = (d_{\text{model}} W)^{-1/2}, the peak. At t=0t = 0 the formula reads min⁡(∞,0)=0\min(\infty, 0) = 0, and that is the value to return: Python raises ZeroDivisionError on 0 ** -0.5. The factor dmodel−1/2d_{\text{model}}^{-1/2} ties the peak to the width: wider models get smaller rates. L5.5 trains with dmodel=512d_{\text{model}} = 512, W=4000W = 4000: peak ≈6.99×10−4\approx 6.99 \times 10^{-4}.

Treat all gradients of the model as one long vector and take its Euclidean length, the global norm ∥g∥\lVert g \rVert. If it exceeds the limit cc, scale every gradient by the same factor:

g(i)←g(i)⋅min⁡ ⁣(1, c∥g∥+10−6).g^{(i)} \leftarrow g^{(i)} \cdot \min\!\left(1,\ \frac{c}{\lVert g \rVert + 10^{-6}}\right).

One factor for all tensors keeps the direction of the full gradient and only shortens it, so the clipped step is still a descent direction; clipping each tensor separately would change the direction. The 10−610^{-6} is torch’s guard against dividing by a zero norm. Clipping never scales up: below the limit, the gradients are untouched. The function returns the norm before clipping, which the train loop logs: a norm that keeps hitting the limit is the first sign of an unstable run. If the norm is infinite or NaN (an overflow, common under fp16 in L11.1), clipping cannot repair it, so the gradients are left alone and the caller skips the update. With Adam (M10.3), clipping matters less for the step size, since Adam normalizes per coordinate, but it still bounds what one outlier batch writes into mm and vv.

Cosine, W=2W = 2, T=6T = 6, ηmax⁡=1\eta_{\max} = 1, ηmin⁡=0.1\eta_{\min} = 0.1:

tt01234567
phasewarmupwarmupp=0p = 0p=1/4p = 1/4p=1/2p = 1/2p=3/4p = 3/4p=1p = 1after
η(t)\eta(t)00.510.1+0.9⋅1+cos⁡(π/4)2≈0.868200.1 + 0.9 \cdot \frac{1 + \cos(\pi/4)}{2} \approx 0.868200.550.1+0.9⋅1−cos⁡(π/4)2≈0.231800.1 + 0.9 \cdot \frac{1 - \cos(\pi/4)}{2} \approx 0.231800.10.1

This is the first test, test_hand_example_cosine.

WSD, W=2W = 2, S=2S = 2, D=4D = 4, ηmax⁡=1\eta_{\max} = 1, ηmin⁡=0.2\eta_{\min} = 0.2: steps 0 and 1 ramp (0, 0.5); steps 2 and 3 are stable (1, 1); the decay runs from step 4 (p=0p = 0: 1) through steps 5, 6, 7 (p=1/4,1/2,3/4p = 1/4, 1/2, 3/4: 0.8, 0.6, 0.4); step 8 and after are at the floor, 0.2 (test_hand_example_wsd).

Noam, dmodel=4d_{\text{model}} = 4, W=4W = 4, so d−1/2=1/2d^{-1/2} = 1/2 and W−3/2=1/8W^{-3/2} = 1/8: η(0)=0\eta(0) = 0; η(1)=12min⁡(1,18)=116\eta(1) = \frac12 \min(1, \frac18) = \frac1{16}; η(2)=12min⁡(0.707,14)=18\eta(2) = \frac12 \min(0.707, \frac14) = \frac18; η(4)=12min⁡(12,12)=14\eta(4) = \frac12 \min(\frac12, \frac12) = \frac14, the peak (4⋅4)−1/2(4 \cdot 4)^{-1/2}; η(16)=12⋅14=18\eta(16) = \frac12 \cdot \frac14 = \frac18, half the peak at four times the step (test_hand_example_noam).

Clipping: two tensors with gradients [3,4][3, 4] and [12][12]. ∥g∥=9+16+144=169=13\lVert g \rVert = \sqrt{9 + 16 + 144} = \sqrt{169} = 13. With c=6.5c = 6.5 the factor is 6.5/13.000001≈0.499999966.5 / 13.000001 \approx 0.49999996, so the gradients become about [1.5,2][1.5, 2] and [6][6], and the call returns 13 (test_hand_example_clip).

def cosine_with_warmup(step: int, warmup: int, total: int, lr_max: float, lr_min: float) -> float
def wsd(step: int, warmup: int, stable: int, decay: int, lr_max: float, lr_min: float) -> float
def noam(step: int, d_model: int, warmup: int) -> float
def clip_grad_norm_(params, max_norm: float) -> float # in place; returns the norm before clipping

params is any iterable of objects with a grad attribute (an ndarray or None), the same parameters M10.3 steps. The schedules raise ValueError for a negative step, a warmup longer than total, a zero decay length, or lr_min outside [0,ηmax⁡][0, \eta_{\max}]; clip_grad_norm_ raises it for max_norm <= 0. The trailing underscore follows torch: the function changes its argument in place.

TestKINDChecksWhy it matters downstream
test_hand_example_cosineunitsection 3’s cosine row, steps 0 to 7you and the tests agree on the step convention
test_hand_example_wsdunitsection 3’s WSD values, steps 0 to 9the C1 schedule
test_hand_example_noamunitsection 3’s Noam valuesthe L5.5 schedule
test_hand_example_clipunitnorm 13 clipped to 6.5the train loop’s clip
test_cosine_matches_hfgoldenHF’s cosine schedules (with and without a floor and warmup) at every stepcomparability with HF-trained baselines
test_wsd_matches_hfgoldenHF’s linear WSD, including steps past the endsame
test_clip_matches_torchgoldentorch.nn.utils.clip_grad_norm_ on four gradient sets, including None and float32same clip as every reference run
test_schedules_hit_their_cornerspropertyover random configurations: peak exactly at the end of warmup, never above it, never rising after it, never below the floor, the floor from the end onevery phase-boundary off-by-one
test_cosine_after_total_stays_at_floorboundarysteps past totalruns that train a little longer than planned
test_noam_peak_and_inverse_sqrt_decaypropertypeak (dW)−1/2(d W)^{-1/2} at t=Wt = W; ηt\eta\sqrt{t} constant after it; linear beforethe 2017 recipe
test_noam_step_zero_is_zeroboundarynoam(0, ...) == 0.0 with no exceptionthe first update
test_clip_below_max_is_a_no_opboundarysmall gradients unchanged bit for bitclipping only shrinks
test_clip_is_global_and_in_placepropertyone factor for all tensors; the norm after is max_norm; the arrays are the same objectsthe direction is kept; the optimizer sees the change
test_clip_nonfinite_norm_leaves_gradsboundaryan inf gradient returns inf and changes nothingL11.1 skips overflowed steps
test_clip_skips_missing_gradsboundarygrad = None is skipped and stays None; no grads gives 0unused parameters
test_rejects_bad_argumentsboundaryimpossible configurations raise ValueErrorerrors at the first step, not a rising “decay”
PitfallSymptomCaught by
decaying to 0 and ignoring lr_minthe last part of the run does nothing; WSD and cosine runs end at different floorstest_hand_example_cosine, test_cosine_matches_hf (mutant s01)
measuring cosine progress over total instead of total - warmupthe floor is never reached; the rate at total is above lr_mintest_cosine_matches_hf (mutant s02)
warmup from step + 1the rate reaches the peak one step early and every value differs from HFtest_hand_example_cosine (mutant s03)
letting the cosine run past totalthe rate climbs back toward lr_max after the planned endtest_cosine_after_total_stays_at_floor (mutant s04)
counting WSD’s stable phase from step 0the decay starts warmup steps earlytest_hand_example_wsd (mutant s05)
evaluating Noam’s formula at step 0ZeroDivisionError on the first updatetest_noam_step_zero_is_zero (mutant s08)
clipping each tensor by its own normthe update direction changes; small layers are never clipped while large ones always aretest_clip_is_global_and_in_place (mutant s11)
scaling small gradients up to max_normevery step has the same length; early training divergestest_clip_below_max_is_a_no_op (mutant s12)
clipping an infinite normmax_norm / inf = 0 zeroes the update and the overflow goes unnoticedtest_clip_nonfinite_norm_leaves_grads (mutant s15)

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

DirectionModuleHow it uses this
BackM00.3sequences and the geometric sum; the warmup is an arithmetic sequence (reading)
BackM10.2the optimizer protocol; a schedule writes opt.lr (reading)
ForwardL0.5train_step(..., clip=...) clips, then the loop sets opt.lr from the spec’s schedule and steps
ForwardM10.3AdamW reads opt.lr on every step, and its decay scales with it
ForwardL3.6clipping keeps truncated backpropagation through time stable
ForwardL5.5the Noam schedule for the 2017 Transformer
ForwardC1WSD with decay_frac and grad_clip = 1.0 from formats/train-spec.schema.json

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

Your pieceProduction equivalentWhat it addsWhere to look
schedule functionstorch.optim.lr_scheduler (LambdaLR, SequentialLR, CosineAnnealingWarmRestarts)stateful schedulers that compose phases and save their own statetorch/optim/lr_scheduler.py
wsdHF get_wsd_schedulecosine and 1-sqrt decay shapes, warmup shapestransformers/optimization.py
fixed schedulesschedule-free AdamWaverages the iterates instead of decaying the rate, so no total is neededDefazio et al., “The Road Less Scheduled” (2024)
clip_grad_norm_torch foreach norms, FSDP clip_grad_norm_one fused kernel per dtype; a norm reduced across data-parallel shardstorch/nn/utils/clip_grad.py; torch/distributed/fsdp
skip on a non-finite normAMP GradScalerlowers the loss scale and skips the step when gradients overflow; your L11.1torch/amp/grad_scaler.py