How Models Are Actually Trained
Warmup, cosine and WSD schedules
The learning rate is raised gently at the start and lowered at the end, and the exact shape of that curve decides whether a months-long training run works.
- 15 min read
- 3 reading levels
- Updated
Read these first
On this page 8
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
The short answer
The learning rate starts near zero, climbs to a peak, then comes back down. That shape is planned before training begins.
The analogy you have already lived
Think about riding a scooter out of a crowded lane onto a highway. You do not open the throttle at the gate. You crawl out, then build speed once the road is clear.
At the other end, you do not hit the brakes from top speed at your gate. You slow gradually over the last stretch, so you can stop exactly where you mean to.
A training run drives the same way. Slow start, fast middle, gentle stop.
Why it exists
The learning rate is how big a step the model takes each time it corrects itself. Too small and training takes forever. Too big and it overshoots and falls apart.
There is no single right value, because the right value changes during training.
At the start, the model is random. Its corrections are wild and mostly wrong. Taking huge steps based on wild corrections wrecks the model in the first few minutes. So you start tiny and ramp up. That ramp is called warmup.
At the end, the model is close to a good answer. Big steps now bounce it around the target instead of settling on it. So you shrink the steps. That is called decay or cooldown.
How it works
Three shapes are in common use.
COSINE WSD CONSTANT
/\ /-------------\ /------------------
/ \___ / \ /
/ ---___ / \ /
warm smooth fall warm flat cooldown warm, then never fall
need to know the can stop and cool down the model never
finish line upfront at any moment fully settlesCosine falls smoothly from the peak to almost nothing. It has been the default for years and works well.
It has one annoying property. You must decide the finish line before you start. Train for longer than planned and the shape is wrong.
WSD stands for warmup, stable, decay. It holds the peak flat for most of the run, then drops fast at the end.
That flat middle is the point. You can stop whenever you like, run a short cooldown, and get a finished model. You can also keep going from the same flat checkpoint. One long run becomes many possible models.
Where you have already seen this
- A train accelerating out of a station and braking into the next one.
- An oven preheating, holding temperature, then being switched off before the end.
- Every deep learning framework, where this is a "scheduler" you attach to the optimiser.
What is honestly hard here
Warmup looks like a superstition until you watch a run die without it. It is not.
Modern optimisers keep a running estimate of how noisy each parameter's corrections are. At step one they have almost no data for that estimate, so it is unreliable. Big steps taken on an unreliable estimate destroy the model before it learns anything.
Warmup buys those estimates a few hundred steps to settle. That is the whole reason.
Remember this
- Start with tiny steps and ramp up. That is warmup, and skipping it breaks big runs.
- Shrink the steps at the end so the model settles instead of bouncing.
- Cosine needs the finish line up front. WSD lets you decide later.
What to learn next
- Loss spikes and gradient clipping — what happens when the schedule is not enough.
- Finding a learning rate — how to pick the peak in the first place.
- Scaling laws and compute-optimal training — where the total step count comes from.
Developer — Code and libraries.
Setup
pip install torchThe first script needs nothing but the standard library. Both run on a CPU.
The three shapes, side by side
import math
TOTAL, WARMUP, PEAK, FLOOR = 100, 10, 1.0, 0.1
def warmup(step): # shared by all three schedules
return PEAK * (step + 1) / WARMUP
def cosine(step):
if step < WARMUP:
return warmup(step)
p = (step - WARMUP) / (TOTAL - WARMUP)
return FLOOR + 0.5 * (PEAK - FLOOR) * (1 + math.cos(math.pi * p))
def wsd(step, cooldown_frac=0.2): # warmup - stable - decay
start = int(TOTAL * (1 - cooldown_frac))
if step < WARMUP:
return warmup(step)
if step < start:
return PEAK
p = (step - start) / (TOTAL - start)
return PEAK * (1 - p) + FLOOR * p
def constant(step):
return warmup(step) if step < WARMUP else PEAK
print(f"{'step':>5} {'cosine':>8} {'WSD':>8} {'constant':>9} shape of WSD")
for step in range(0, TOTAL, 5):
c, w, k = cosine(step), wsd(step), constant(step)
print(f"{step:>5} {c:>8.3f} {w:>8.3f} {k:>9.3f} " + "#" * round(w * 30)) step cosine WSD constant shape of WSD
0 0.100 0.100 0.100 ###
5 0.600 0.600 0.600 ##################
10 1.000 1.000 1.000 ##############################
15 0.993 1.000 1.000 ##############################
20 0.973 1.000 1.000 ##############################
25 0.940 1.000 1.000 ##############################
30 0.895 1.000 1.000 ##############################
35 0.839 1.000 1.000 ##############################
40 0.775 1.000 1.000 ##############################
45 0.704 1.000 1.000 ##############################
50 0.628 1.000 1.000 ##############################
55 0.550 1.000 1.000 ##############################
60 0.472 1.000 1.000 ##############################
65 0.396 1.000 1.000 ##############################
70 0.325 1.000 1.000 ##############################
75 0.261 1.000 1.000 ##############################
80 0.205 1.000 1.000 ##############################
85 0.160 0.775 1.000 #######################
90 0.127 0.550 1.000 ################
95 0.107 0.325 1.000 ##########Note the number in the TOTAL position of cosine. It appears inside the function. Change your token budget and every learning rate in the run changes. wsd reads TOTAL only to locate the cooldown, which is why you can move that decision to the last day of training.
Warmup is not superstition — here is the measurement
import torch
import torch.nn as nn
import torch.nn.functional as F
TEXT = "the cat sat on the mat. the cat ate the rat. the rat sat on the mat. " * 40
chars = sorted(set(TEXT))
stoi = {c: i for i, c in enumerate(chars)}
V, CTX = len(chars), 16
data = torch.tensor([stoi[c] for c in TEXT])
x = torch.stack([data[i:i + CTX] for i in range(0, len(data) - CTX - 1, 3)])
y = torch.stack([data[i + 1:i + CTX + 1] for i in range(0, len(data) - CTX - 1, 3)])
def build():
torch.manual_seed(0) # identical starting weights both times
layer = nn.TransformerEncoderLayer(64, 4, 128, batch_first=True, dropout=0.0)
return nn.ModuleDict({
"emb": nn.Embedding(V, 64), "pos": nn.Embedding(CTX, 64),
"blocks": nn.TransformerEncoder(layer, 3), "head": nn.Linear(64, V),
})
def forward(m, idx):
h = m["emb"](idx) + m["pos"](torch.arange(idx.shape[1]))
mask = nn.Transformer.generate_square_subsequent_mask(idx.shape[1])
return m["head"](m["blocks"](h, mask=mask, is_causal=True))
def run(use_warmup, peak=0.02, steps=60, warm=15):
m = build()
opt = torch.optim.AdamW(m.parameters(), lr=peak)
losses = []
for s in range(steps):
lr = peak * min(1.0, (s + 1) / warm) if use_warmup else peak
for g in opt.param_groups:
g["lr"] = lr
loss = F.cross_entropy(forward(m, x).reshape(-1, V), y.reshape(-1))
opt.zero_grad()
loss.backward()
opt.step()
losses.append(loss.item())
return losses
a = run(use_warmup=False)
b = run(use_warmup=True)
print(f"peak learning rate 0.02, identical seeds and data\n")
print(f"{'step':>5} {'no warmup':>12} {'15-step warmup':>16}")
for s in range(0, 60, 5):
print(f"{s:>5} {a[s]:>12.4f} {b[s]:>16.4f}")
print(f"\nfinal loss no warmup {a[-1]:.4f} with warmup {b[-1]:.4f}")peak learning rate 0.02, identical seeds and data
step no warmup 15-step warmup
0 2.9622 2.9622
5 2.1897 1.5469
10 2.2133 0.5803
15 2.1752 0.7549
20 2.0681 0.5332
25 1.8750 0.3324
30 2.2102 0.2004
35 1.9188 0.1631
40 1.9155 0.1525
45 1.8244 0.5699
50 1.7865 0.3882
55 1.7806 0.4990
final loss no warmup 1.7556 with warmup 0.6645Written against PyTorch 2.5.1, CPU, about 40 seconds. The two runs share a seed, so the step-0 losses match exactly. Values past step 0 can differ in the last decimals on other builds; the gap between the columns is the reproducible part.
Reading that output carefully
Identical model, identical data, identical peak learning rate. The only difference is fifteen steps of ramp. Without it the run stalls near 1.78 and never recovers. With it the loss falls by an order of magnitude.
The damage happens in the first ten steps and is permanent. By step 10 the no-warmup run is already at 2.21 while the warmed-up run is at 0.58. Nothing after that closes the gap. This is what "warmup protects the early steps" means concretely.
The warmed-up column is noisy at the end — 0.15, then 0.57, then 0.39. That is the peak learning rate being too large for a nearly-converged model. It is precisely the problem that cosine or WSD decay exists to fix. Add a cooldown to this script and the bouncing stops.
The numbers real runs use
| Model | Peak LR | Warmup | Schedule |
|---|---|---|---|
| GPT-3 175B | 0.6e-4 | 375M tokens | cosine to 10% |
| Llama 2 70B | 1.5e-4 | 2000 steps | cosine to 10% |
| Llama 3 405B | 8e-5 | 8000 steps | cosine to 8e-7 |
| MiniCPM 2.4B | 0.01 | 2000 steps | WSD, 10% cooldown |
| OLMo 2 7B | 3e-4 | ~2000 steps | cosine, then linear anneal |
Two patterns hold across all of them. Warmup is a few thousand steps regardless of run length. Peak learning rate falls as model size rises.
In PyTorch, without hand-rolling it
from torch.optim.lr_scheduler import LambdaLR, SequentialLR, LinearLR, CosineAnnealingLR
opt = torch.optim.AdamW(model.parameters(), lr=3e-4)
sched = SequentialLR(
opt,
schedulers=[LinearLR(opt, start_factor=1e-8, total_iters=2000),
CosineAnnealingLR(opt, T_max=98_000, eta_min=3e-5)],
milestones=[2000],
)
# then, once per optimizer step: opt.step(); sched.step()No output block — this is a fragment needing your own model and loop. More on the built-in schedulers in learning rate schedulers.
Common mistakes
Calling sched.step() once per epoch instead of once per optimiser step. Pretraining schedules are defined in steps, not epochs. Your warmup then lasts 2000 epochs and the run never leaves warmup.
Warming up the learning rate but not the batch size. Some recipes ramp both. Mixing one recipe's warmup with another's batch size is how people accidentally reproduce nothing.
Stepping the scheduler on micro-batches. With gradient accumulation, several forward passes make one optimiser step. Step the scheduler with the optimiser, never with the forward pass.
Assuming a cosine restart is safe. Resuming a cosine run past its T_max sends the learning rate back up. If you might extend a run, use WSD.
Copying a peak learning rate across model sizes. A learning rate tuned for a 100M model will blow up a 7B one. Scale it down, or use µP — see the note below.
Try it yourself
Add a linear cooldown over the final 15 steps of run(use_warmup=True). Watch the end-of-run bouncing disappear, and the final loss drop below anything the constant run reached.
What to learn next
- Loss spikes and gradient clipping — what happens when the schedule is not enough.
- Finding a learning rate — how to pick the peak in the first place.
- Scaling laws and compute-optimal training — where the total step count comes from.
Researcher — Mathematics and papers.
Why warmup is needed, mechanically
Adam's update is
$$ \theta_{t+1} = \theta_t - \eta \cdot \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} $$
with $\hat m_t$ and $\hat v_t$ the bias-corrected first and second moment estimates of the gradient, and $\eta$ the learning rate. Because the update is normalised by $\sqrt{\hat v_t}$, its magnitude is roughly $\eta$ regardless of gradient scale.
Early in training $\hat v_t$ is estimated from a handful of samples and has enormous variance. Liu et al., 2020 (On the Variance of the Adaptive Learning Rate and Beyond, RAdam) showed the variance of the adaptive term is unbounded at $t$ small, and that linear warmup is a variance-reduction device. RAdam derives a rectification term that makes the warmup implicit; in practice explicit warmup is still preferred because it is one line.
A second, architecture-specific reason applies to post-norm transformers. Xiong et al., 2020 (On Layer Normalization in the Transformer Architecture) proved that post-LN transformers have gradients at initialisation that scale as $O(\sqrt{d}\,\ln L)$ near the output layer, and that warmup is necessary for stability. Pre-LN removes the requirement in theory, and every large model still uses warmup, because the Adam-variance argument survives the architecture change.
Cosine
$$ \eta_t = \eta_{\min} + \tfrac{1}{2}(\eta_{\max}-\eta_{\min})\left(1 + \cos!\left(\pi \frac{t - t_w}{T - t_w}\right)\right) $$
$t$ is the step, $t_w$ the warmup length, $T$ the total steps, $\eta_{\max}$ the peak and $\eta_{\min}$ the floor (typically $0.1\,\eta_{\max}$).
Introduced as SGDR (Loshchilov and Hutter, 2017) for the restarts, and adopted for LLMs without them. Chinchilla (Hoffmann et al., 2022) documented the critical constraint: the cosine cycle length must match the number of training steps. Setting $T$ to 1× the run gives the best loss; a mismatch of 10× costs a substantial fraction of the run's value. This single fact makes every cosine run a fixed-budget commitment, and it corrupts scaling-law experiments — every point on the curve needs its own full run.
WSD, and why it took over
The Warmup–Stable–Decay schedule (Hu et al., 2024, MiniCPM):
$$ \eta_t = \begin{cases} \eta_{\max}\, t/t_w & t < t_w \ \eta_{\max} & t_w \le t < T - t_d \ f!\left(\frac{t - (T - t_d)}{t_d}\right)\eta_{\max} & t \geq T - t_d \end{cases} $$
$t_d$ is the cooldown length, typically $0.1$ to $0.2\,T$, and $f$ decays from 1 to a small floor. MiniCPM used an exponential $f$; linear and $1-\sqrt{\cdot}$ variants are also used, and Hägele et al., 2024 found $1-\sqrt{\cdot}$ marginally best.
Three properties matter:
- Budget-agnostic. The stable phase is a valid checkpoint at any point. One run yields a family of models at different token counts, each finished with a short cooldown. Hägele et al. showed this reduces the compute for a scaling-law study by roughly an order of magnitude.
- Loss matches cosine. Both papers report constant-LR-plus-cooldown reaching cosine's final loss, sometimes slightly better.
- The cooldown produces a sharp, reproducible loss drop. Loss is flat and mediocre during the stable phase, then falls steeply during the decay. Watching a WSD run mid-training and concluding it has plateaued is a standard misreading.
Wen et al., 2024 (Understanding Warmup-Stable-Decay Learning Rates: A River Valley Loss Landscape Perspective) give the mechanism: the loss surface resembles a river in a steep valley. A high constant learning rate makes fast progress along the river while bouncing between the steep walls; the cooldown stops the bouncing and drops the iterate to the river bed. It also predicts, correctly, that a cooled-down checkpoint is a poor place to resume high-LR training from — resume from the stable branch instead.
The competing correction: µP
Scaling the peak learning rate down as models grow is a workaround for a parameterisation problem, not a law of nature. Yang et al., 2022 (Tensor Programs V, µTransfer) show that under Maximal Update Parametrisation, the optimal learning rate is invariant to width. You tune on a 40M-parameter proxy and transfer the hyperparameters to a 6.7B model directly. Cerebras-GPT and several frontier labs use it; it removes the most expensive hyperparameter search in pretraining.
Weight averaging as an alternative to decay
Hägele et al., 2024 and Sanyal et al., 2023 both report that stochastic weight averaging over the constant-LR phase recovers most of the cooldown's benefit with no schedule change. LAWA (latest weight averaging) keeps a rolling mean of recent checkpoints. Liu et al., 2025 (WSM) push this further, arguing checkpoint merging can replace the decay phase outright. This is an active area; treat it as promising rather than settled.
Papers
- Goyal et al., Accurate, Large Minibatch SGD, 2017 — arxiv.org/abs/1706.02677 (origin of gradual warmup)
- Loshchilov and Hutter, SGDR: Stochastic Gradient Descent with Warm Restarts, 2017 — arxiv.org/abs/1608.03983
- Liu et al., On the Variance of the Adaptive Learning Rate and Beyond (RAdam), 2020 — arxiv.org/abs/1908.03265
- Xiong et al., On Layer Normalization in the Transformer Architecture, 2020 — arxiv.org/abs/2002.04745
- Yang et al., Tensor Programs V: Tuning Large Neural Networks via Zero-Shot Hyperparameter Transfer, 2022 — arxiv.org/abs/2203.03466
- Hoffmann et al., Training Compute-Optimal Large Language Models, 2022 — arxiv.org/abs/2203.15556
- Hu et al., MiniCPM, 2024 — arxiv.org/abs/2404.06395
- Hägele et al., Scaling Laws and Compute-Optimal Training Beyond Fixed Training Durations, NeurIPS 2024 — arxiv.org/abs/2405.18392
- Wen et al., Understanding Warmup-Stable-Decay Learning Rates, 2024 — arxiv.org/abs/2410.05192
What to learn next
- Loss spikes and gradient clipping — what happens when the schedule is not enough.
- Finding a learning rate — how to pick the peak in the first place.
- Scaling laws and compute-optimal training — where the total step count comes from.