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.
- 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
- Distributed Training at Scale
- Transformer Architecture - Pre-LN, RMSNorm
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
| Technique | What it does | Why it helps |
|---|---|---|
| LR warmup | Ramp the learning rate from ~0 over the first few thousand steps | Adam'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 threshold | Caps the damage of one bad batch |
| AdamW with β₂ ≈ 0.95 | Shorter memory for the second-moment estimate than the default 0.999 | Adapts faster when gradient scale changes, reducing spikes |
| Pre-norm + RMSNorm | Normalize before each sub-layer | Keeps the residual stream well-conditioned in deep stacks |
| QK-norm | RMSNorm on queries and keys before the dot product | Bounds attention logits; prevents the attention-entropy collapse seen in very large runs |
| z-loss | Small penalty (e.g. 1e-4 × log²Z) on the softmax normalizer of the output logits | Stops output logits drifting to huge values (used in PaLM) |
| Logit soft-capping | cap · tanh(logits / cap) on attention and/or final logits | Hard ceiling on logit magnitude (Gemma 2) |
| Batch skipping and rollback | Restart from before a spike and skip the batches that triggered it | PaLM'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
- Why does a warmup-stable-decay (WSD) schedule make continual pretraining easier than cosine?
- Attention logits keep growing during a large run and training becomes unstable. Which technique targets that directly?
- What problem does μP solve?
- What does Muon do differently from AdamW for a hidden-layer weight matrix?
Exercises
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.
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
- Inspect the batches around step 41,200 for bad data (long repeated sequences, binary junk, a single dominant source).
- Check per-layer statistics: attention and output logit magnitudes, activation norms, and per-rank loss (to rule out a faulty GPU).
- Roll back to a checkpoint a few hundred steps before the spike, skip the offending data range, and resume.
- 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
- Chowdhery et al., PaLM (2022) - z-loss, batch skipping on loss spikes
- Dehghani et al., Scaling Vision Transformers to 22 Billion Parameters (2023) - QK-norm for stability
- Wortsman et al., Small-scale proxies for large-scale Transformer training instabilities (2023)
- Yang et al., Tensor Programs V: Tuning Large Neural Networks via Zero-Shot Hyperparameter Transfer (μP) (2022)
- Hu et al., MiniCPM (2024) - WSD schedule
- Jordan et al., Muon: An optimizer for hidden layers in neural networks (2024)
- Liu et al., Muon is Scalable for LLM Training (2025)
- Kimi Team, Kimi K2 (2025) - MuonClip / QK-clip
- Loshchilov & Hutter, Decoupled Weight Decay Regularization (AdamW) (2017)
Last reviewed: 2026-09