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.

On this page 8
  1. The short answer
  2. The analogy you have already lived
  3. Why it exists
  4. How it works
  5. Where you have already seen this
  6. What is honestly hard here
  7. Remember this
  8. What to learn next

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 settles

Cosine 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

Developer — Code and libraries.

Setup

bash
pip install torch

The first script needs nothing but the standard library. Both run on a CPU.

The three shapes, side by side

schedules.py
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))
Output
 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

warmup_matters.py
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}")
Output
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.6645

Written 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

ModelPeak LRWarmupSchedule
GPT-3 175B0.6e-4375M tokenscosine to 10%
Llama 2 70B1.5e-42000 stepscosine to 10%
Llama 3 405B8e-58000 stepscosine to 8e-7
MiniCPM 2.4B0.012000 stepsWSD, 10% cooldown
OLMo 2 7B3e-4~2000 stepscosine, 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

python
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

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:

  1. 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.
  2. Loss matches cosine. Both papers report constant-LR-plus-cooldown reaching cosine's final loss, sometimes slightly better.
  3. 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

What to learn next