Contents
Map

04 · Pretraining at Scale

Training Stability & Optimizers

View as:

Training Stability and Optimizers

A multi-month pretraining run can be lost to a single loss spike that never recovers. The techniques in this note - schedules, normalization tricks, hyperparameter transfer and newer optimizers - exist so a run that works at 100M parameters keeps working at 100B, without a full hyperparameter search at full scale.

Learning objectives 50 min
By the end of this page you will be able to:
  • Diagnose loss spikes and divergence and list the standard mitigations
  • Explain QK-norm, z-loss and logit soft-capping and the failure each one prevents
  • Compare cosine and warmup-stable-decay (WSD) schedules and explain why WSD suits continual training
  • Explain μP (hyperparameter transfer) at a conceptual level
  • Explain what Muon changes relative to AdamW, and why MuonClip was needed at trillion-parameter scale

What Instability Looks Like

flowchart LR
    S["📈 Loss spike"] --> Q{"🔍 Recovers<br/>within ~100s of steps?"}
    Q -->|Yes| OK["✅ Log it, keep going<br/>check the data batch"]
    Q -->|No| DIV["💥 Divergence"]
    DIV --> RB["⏪ Roll back to an earlier<br/>checkpoint"]
    RB --> FIX["🔧 Skip the offending batches,<br/>lower LR, add QK-norm / z-loss"]
    FIX --> RES["▶️ Resume"]

    style S fill:#e8e0d4,stroke:#c8b89a
    style DIV fill:#ddd8e4,stroke:#b8b0c8
    style OK fill:#dde4dc,stroke:#b0c4b0
    style RES fill:#d8dfe8,stroke:#b0bac8

Common causes: attention logits growing without bound, output logits drifting large, bad data (a batch of long repeated tokens or binary junk), learning rate too high for the batch size, and - at scale - hardware faults that silently corrupt values.


The Standard Stabilizers

TechniqueWhat it doesWhy it helps
LR warmupRamp the learning rate from ~0 over the first few thousand stepsAdam's variance estimates are unreliable at the start; large early updates destabilize
Gradient clipping (global norm, typically 1.0)Rescale gradients whose total norm exceeds a thresholdCaps the damage of one bad batch
AdamW with β₂ ≈ 0.95Shorter memory for the second-moment estimate than the default 0.999Adapts faster when gradient scale changes, reducing spikes
Pre-norm + RMSNormNormalize before each sub-layerKeeps the residual stream well-conditioned in deep stacks
QK-normRMSNorm on queries and keys before the dot productBounds attention logits; prevents the attention-entropy collapse seen in very large runs
z-lossSmall penalty (e.g. 1e-4 × log²Z) on the softmax normalizer of the output logitsStops output logits drifting to huge values (used in PaLM)
Logit soft-cappingcap · tanh(logits / cap) on attention and/or final logitsHard ceiling on logit magnitude (Gemma 2)
Batch skipping and rollbackRestart from before a spike and skip the batches that triggered itPaLM's practical fix when spikes didn't recover

Learning-Rate Schedules

Cosine

Warm up, then decay along a cosine curve to ~10% of the peak by the planned final step. It is robust and well understood, but the schedule is tied to the total step count: stopping early, or continuing training later, means the learning rate was wrong for most of the run.

Warmup-Stable-Decay (WSD)

Warm up, hold the peak learning rate constant for most of the run, then decay quickly (typically over the last 10-20%). Popularized by MiniCPM (2024):

  • Continual training is cheap. Keep the constant-LR checkpoint; to train longer, resume from it and decay later. With cosine you would have to restart.
  • The decay phase is where loss drops - which aligns naturally with annealing on high-quality data (see Data Curation and Mixtures).
  • One run, many budgets. Branch decays off the stable phase at several points to get scaling-law data from one run.

μP: Hyperparameters That Transfer

With standard parametrization, the best learning rate and initialization scale shift as you widen the model, so a sweep done at 100M parameters doesn't tell you the right settings at 10B. Maximal update parametrization (μP) rescales initialization and per-layer learning rates by width so the optimal hyperparameters stay (approximately) constant as width grows. The workflow becomes: sweep on a small proxy, then train the large model with the same settings. Several open model reports (e.g. Cerebras-GPT, MiniCPM) use μP or close variants; others instead fit power laws for the optimal learning rate and batch size across scales (DeepSeek LLM).


Optimizers Beyond AdamW

Concept

AdamW is still the default: per-parameter adaptive step sizes from first and second moment estimates, with decoupled weight decay. Its cost is two FP32 state tensors per parameter - 8 of the ~16 bytes/parameter of mixed-precision training memory.

Muon (Jordan et al., 2024) treats each 2D weight matrix as a matrix rather than a bag of numbers. It takes the momentum-averaged gradient and orthogonalizes it (approximately, with a few Newton-Schulz iterations), so the update pushes equally in all directions of the matrix instead of being dominated by a few large singular directions. Embeddings, the output head and 1D parameters (norm gains) still use AdamW.

  • Moonshot AI (Liu et al., 2025) scaled Muon to large LLMs with weight decay and per-matrix update-scale matching and reported roughly 2× compute efficiency relative to AdamW at compute-optimal training.
  • Muon needs only one momentum buffer per parameter (less optimizer memory than Adam's two).
  • MuonClip (Kimi K2): at trillion-parameter scale Muon produced exploding attention logits. Kimi K2 added QK-clip - rescale the query and key projection weights whenever the maximum attention logit exceeds a threshold - and reported pretraining on 15.5T tokens with no loss spikes.
flowchart LR
    G["📉 Gradient of W<br/>(a matrix)"] --> M["🌀 Momentum"]
    M --> A1["AdamW: per-element<br/>m / (√v + ε)"]
    M --> A2["Muon: orthogonalize the matrix<br/>(Newton-Schulz ≈ U·Vᵀ)"]
    A1 --> U["🔧 Update W"]
    A2 --> U

    style A1 fill:#d8dfe8,stroke:#b0bac8
    style A2 fill:#dde4dc,stroke:#b0c4b0

Check Yourself

Check yourself
0 / 4 answered
  1. Why does a warmup-stable-decay (WSD) schedule make continual pretraining easier than cosine?
  2. Attention logits keep growing during a large run and training becomes unstable. Which technique targets that directly?
  3. What problem does μP solve?
  4. What does Muon do differently from AdamW for a hidden-layer weight matrix?

Exercises

Exercise - WSD vs cosine in the GPT lab

Implement WSD in the GPT From Scratch lab (hold the peak LR until 80% of steps, then decay linearly). Train the same model with cosine and WSD for the same steps and compare final validation loss. Then continue the WSD run for 50% more steps from its 80% checkpoint, and explain why you can't do the same cleanly with cosine.

Exercise - Diagnose a spike

At step 41,200 of a 7B run, loss jumps from 2.10 to 3.40, the global gradient norm jumps 20×, and loss has not recovered after 500 steps. List, in order, what you would check and what you would do.

Solution
  1. Inspect the batches around step 41,200 for bad data (long repeated sequences, binary junk, a single dominant source).
  2. Check per-layer statistics: attention and output logit magnitudes, activation norms, and per-rank loss (to rule out a faulty GPU).
  3. Roll back to a checkpoint a few hundred steps before the spike, skip the offending data range, and resume.
  4. If it recurs: lower the peak LR or extend warmup; add QK-norm or z-loss if logits were growing; confirm gradient clipping is active; check for hardware errors on the ranks involved.

Study Notes

Must-know:

  • Stabilizers: warmup, grad-norm clipping, β₂ ≈ 0.95, pre-norm RMSNorm, QK-norm (attention logits), z-loss (output logits), soft-capping, skip-and-rollback
  • WSD: constant LR then short decay - enables continual training and branching decays; the decay phase pairs with annealing data
  • μP: parametrize so small-model hyperparameters transfer to large widths
  • Muon: orthogonalized momentum updates for 2D matrices; ~2× compute efficiency reported at scale; MuonClip (QK-clip) fixed exploding attention logits in Kimi K2

References

Last reviewed: 2026-09

⚡AI-assisted content - always verify, always explore multiple perspectives·