How Models Are Actually Trained
Loss spikes and gradient clipping
Long training runs suddenly jump to a terrible loss for no visible reason, and the standard defence is a hard cap on how large a single correction may be.
- 14 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
Sometimes a training run suddenly gets much worse for one step. Gradient clipping puts a hard ceiling on how big any single correction can be.
The analogy you have already lived
Most Indian homes have a voltage stabiliser sitting under the fridge. The grid voltage wanders. Now and then it surges.
The fridge does not need the extra voltage. It needs to survive the surge. The stabiliser caps whatever comes in at a safe level and passes it on.
Gradient clipping is that stabiliser, wired into training. A surge arrives, the cap holds, the model survives.
Why it exists
Training makes a correction after every batch of data. Most corrections are small and sensible.
Every so often, one is enormous. When that huge correction is applied, the model lurches somewhere terrible. The loss — the score of how wrong the model is — jumps upward on the chart. That jump is a loss spike.
A small spike recovers on its own in a few hundred steps. A large one can destroy weeks of work.
At the scale of real pretraining, this is not a rare curiosity. Google reported around twenty spikes in one large run.
How it works
without clipping with clipping
correction size correction size
| |
| X <- surge | _ <- capped here
| |
| . . . . . . . | . . . . . . .
+----------------- steps +----------------- steps
the model lurches the model takes a normal-sized
somewhere terrible step in the surge's directionTwo details matter, and both are easy to miss.
The direction is kept, only the size is cut. Clipping does not throw the correction away. It says "go that way, but no further than this".
The cap is on the whole model at once, not on each number separately. All the corrections are treated as one big arrow, and that arrow is shortened.
The uncomfortable truth about why spikes happen
You would expect the cause to be one bad piece of text in the data. Researchers checked this, carefully.
They took the exact batches that caused a spike and fed them to an earlier copy of the same model. No spike.
So it is not the data alone. It is a particular batch meeting a particular state of the model. That combination is what breaks. Nobody can predict it in advance.
Where you have already seen this
- A voltage stabiliser protecting a fridge or an air conditioner.
- A fuse that blows before the wiring burns.
- Speed limiters on trucks, which cap the top speed without changing direction.
Remember this
- A loss spike is a sudden jump to a much worse score, mid-training.
- Gradient clipping caps the size of any single correction, keeping its direction.
- Spikes come from data and model state together, so no data cleaning removes them all.
What to learn next
- Reading a training run while it happens — the dashboard that catches this early.
- Debugging NaN loss — the version of this problem that kills the run outright.
- Mixed precision training — why bf16 removes a whole class of spikes.
Developer — Code and libraries.
Setup
pip install torchRuns on a CPU in about fifteen seconds.
What clipping does, and what it measures
import statistics
import torch
import torch.nn as nn
import torch.nn.functional as F
# ---------- part 1: what clipping actually does to a gradient ----------
g = torch.tensor([3.0, 4.0]) # norm 5
print("gradient ", g.tolist(), " norm", g.norm().item())
for c in (10.0, 1.0):
clipped = g * min(1.0, c / g.norm())
print(f"clipped at {c:>5}", [round(v, 3) for v in clipped.tolist()],
" norm", round(clipped.norm().item(), 3),
" same direction:", torch.allclose(clipped / clipped.norm(), g / g.norm()))
# ---------- part 2: gradient norms during a run with one bad batch ----------
V, CTX, POISON = 24, 12, 25
PATTERN = torch.arange(V).repeat(60)[:720].reshape(-1, CTX)
def run(clip, seed, lr=0.02, steps=45):
torch.manual_seed(seed)
layer = nn.TransformerEncoderLayer(32, 4, 64, batch_first=True, dropout=0.0)
m = nn.ModuleDict({"emb": nn.Embedding(V, 32),
"blocks": nn.TransformerEncoder(layer, 2),
"head": nn.Linear(32, V)})
opt = torch.optim.AdamW(m.parameters(), lr=lr)
gen = torch.Generator().manual_seed(seed + 999)
losses, norms = [], []
for step in range(steps):
if step == POISON: # one batch of pure noise
rows = torch.randint(0, V, (8, CTX + 1), generator=gen)
else:
i = torch.randint(0, len(PATTERN), (8,), generator=gen)
rows = torch.stack([torch.cat([PATTERN[j], PATTERN[j][:1]]) for j in i])
x, y = rows[:, :-1], rows[:, 1:]
mask = nn.Transformer.generate_square_subsequent_mask(x.shape[1])
logits = m["head"](m["blocks"](m["emb"](x), mask=mask, is_causal=True))
loss = F.cross_entropy(logits.reshape(-1, V), y.reshape(-1))
opt.zero_grad()
loss.backward()
# clip_grad_norm_ RETURNS the norm measured BEFORE clipping - log it, always
norms.append(nn.utils.clip_grad_norm_(m.parameters(), clip).item())
opt.step()
losses.append(loss.item())
return losses, norms
print("\ngradient norms in a healthy run, and at the bad batch (no clipping, seed 0)")
losses, norms = run(clip=1e9, seed=0)
healthy = norms[15:POISON]
print(f" median norm over steps 15-24 : {statistics.median(healthy):.2f}")
print(f" largest norm over steps 15-24: {max(healthy):.2f}")
print(f" norm at the bad batch (25) : {norms[POISON]:.2f}")
print(f" norm one step later (26) : {norms[POISON + 1]:.2f}")
print("\nmean loss over the 10 steps AFTER the bad batch, across 5 seeds")
for clip in (1e9, 1.0, 0.25):
per_seed = [statistics.mean(run(clip, s)[0][POISON + 1:POISON + 11]) for s in range(5)]
label = "no clipping" if clip > 1e8 else f"clip at {clip}"
print(f" {label:<14} {statistics.mean(per_seed):.3f} per seed: "
+ " ".join(f"{v:.2f}" for v in per_seed))gradient [3.0, 4.0] norm 5.0 clipped at 10.0 [3.0, 4.0] norm 5.0 same direction: True clipped at 1.0 [0.6, 0.8] norm 1.0 same direction: True gradient norms in a healthy run, and at the bad batch (no clipping, seed 0) median norm over steps 15-24 : 0.01 largest norm over steps 15-24: 0.03 norm at the bad batch (25) : 3.44 norm one step later (26) : 0.46 mean loss over the 10 steps AFTER the bad batch, across 5 seeds no clipping 0.094 per seed: 0.10 0.06 0.08 0.11 0.13 clip at 1.0 0.025 per seed: 0.03 0.01 0.03 0.03 0.02 clip at 0.25 0.004 per seed: 0.00 0.00 0.01 0.00 0.00
Written against PyTorch 2.5.1 on CPU. Seeds are fixed, so this reproduces on this machine; the last decimal of a loss can move on a different build. The five-seed spread is printed deliberately — a one-seed comparison of a stability fix proves very little.
Reading that output
Part 1 is the whole definition. A norm-5 gradient clipped at 10 is untouched. Clipped at 1 it becomes [0.6, 0.8] — one fifth the length, exactly the same direction. Clipping is a rescale, not a truncation, and it does nothing at all on a normal step.
The bad batch produced a gradient roughly 300 times the median. Median 0.01, spike 3.44. That ratio is the reason a fixed clip threshold works: healthy gradients are far below it, so the cap only ever engages on genuine outliers.
clip_grad_norm_ returns the pre-clip norm. That return value is the single most useful number in your logs. A rising trend in it is an early warning; a sudden jump is the spike itself, visible before the loss curve shows anything.
Clipping cut post-spike damage roughly fourfold, and the tighter clip more. Consistently across all five seeds. It does not prevent the spike — loss still jumped at step 25 — it limits how far the model is thrown.
The full stabilisation checklist
Clipping is one line of a longer defence.
# 1. clip, and log the pre-clip norm
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
logger.log({"grad_norm": norm.item()})
# 2. skip the step entirely if the loss is non-finite
if not torch.isfinite(loss):
optimizer.zero_grad(set_to_none=True)
continueBeyond that, in rough order of how much they help large runs:
- Use bf16, not fp16. bf16 has the same exponent range as fp32, so activations do not overflow to
inf. fp16 needs a loss scaler and is a common source of spikes. See mixed precision training. - Pre-norm, not post-norm. Layer normalisation inside the residual branch, not after it.
- QK normalisation. Normalise queries and keys before the attention dot product. This bounds attention logits and removed most spikes in published ablations.
- A z-loss. A small auxiliary penalty keeping the softmax normaliser near 1, stopping output logits from drifting large.
- Checkpoint often enough to rewind. Every large lab's actual recovery procedure is: roll back a few hundred steps, skip the batches, resume.
Common mistakes
Clipping after optimizer.step(). The order is backward(), then clip, then step(). Clip afterwards and it does nothing.
Clipping with gradient scaling still applied. Under fp16 GradScaler, gradients are inflated. Call scaler.unscale_(optimizer) before clipping, or you are clipping the wrong numbers.
Clipping each parameter separately. clip_grad_value_ caps each element, which changes the update direction. clip_grad_norm_ caps the whole vector and preserves it. Use the norm version.
Setting the threshold from a single run's peak. Pick it from the median healthy norm — one order of magnitude above it is the usual choice, and 1.0 is the near-universal default for LLMs.
Treating a spike as a data bug and stopping to clean data. Sometimes it is. Frequently it is not, and you will burn a week finding nothing. Check the gradient norm history first.
Not logging the pre-clip norm. Without it, a spike is invisible until the loss moves, which is one step too late.
Try it yourself
Change POISON to 5, so the bad batch lands while the model is still near random. Compare the gradient-norm outlier ratio. Early spikes and late spikes behave very differently, and that difference is the subject of the researcher block.
What to learn next
- Reading a training run while it happens — the dashboard that catches this early.
- Debugging NaN loss — the version of this problem that kills the run outright.
- Mixed precision training — why bf16 removes a whole class of spikes.
Researcher — Mathematics and papers.
Clipping, formally
Global-norm clipping (Pascanu et al., 2013, On the difficulty of training recurrent neural networks) rescales the concatenated gradient $g$ of all parameters:
$$ \tilde g = g \cdot \min!\left(1, \frac{\tau}{\lVert g \rVert_2}\right) $$
$\tau$ is the threshold and $\lVert g \rVert_2$ the global L2 norm. The map is the identity inside the ball of radius $\tau$ and a projection onto its boundary outside. Direction is exactly preserved.
This is not a benign modification of gradient descent. Clipped SGD does not converge to a stationary point of the original objective in general; it converges under a relaxed smoothness condition. Zhang et al., 2020 (Why Gradient Clipping Accelerates Training) introduced $(L_0, L_1)$-smoothness, where the local smoothness constant grows with the gradient norm:
$$ \lVert \nabla^2 f(x) \rVert \leq L_0 + L_1 \lVert \nabla f(x) \rVert $$
Under this condition — which they verified empirically on language models — clipped gradient descent has a strictly better convergence rate than any fixed-step-size method. Clipping is not damage control bolted onto SGD; it is the correct algorithm for this loss surface.
What actually causes spikes
Four mechanisms are documented, and they are distinct.
1. Data × state interaction. PaLM (Chowdhery et al., 2022) saw roughly 20 spikes in the 540B run and none in the 8B or 62B runs at the same data order. The decisive ablation: replaying the offending batches from an earlier checkpoint produced no spike. The conclusion in the paper is explicit — spikes require a specific batch and a specific parameter state. Their fix was operational: rewind ~100 steps, skip 200–500 batches, resume.
2. Adam's update decorrelating from the descent direction. Molybog et al., 2023 (A Theory on Adam Instability in Large-Scale Machine Learning) studied models from 7B to 546B parameters and describe the failure state precisely: Adam enters a regime where the update vector has a large norm and is essentially uncorrelated with the descent direction. The mechanism runs through the ratio $\hat m_t / (\sqrt{\hat v_t} + \epsilon)$. For a parameter whose recent gradients have been near zero, $\hat v_t$ decays until $\epsilon$ dominates the denominator, and the resulting step size no longer reflects the gradient that produced it. They tie the severity to large batch sizes, which is exactly the regime of LLM pretraining. The practical levers this implicates are $\epsilon$, $\beta_2$, and batch size — all of which are tuned in practice, and none of which is a complete fix.
3. Attention-logit growth. Wortsman et al., 2023 (Small-scale proxies for large-scale Transformer training instabilities) reproduced large-model instabilities at small scale by raising the learning rate, then showed two fixes generalise: qk-layernorm (normalise Q and K before the dot product) removes attention-logit divergence, and a z-loss $\lambda \log^2 Z$ on the output softmax normaliser $Z$ removes output-logit divergence. PaLM used $\lambda = 10^{-4}$.
4. Embedding-layer gradient scale. Takase et al., 2024 (Spike No More) argue spikes originate in large embedding gradients and derive initialisation conditions bounding the update norm. GLM-130B (Zeng et al., 2023) shipped embedding gradient shrink, scaling the embedding gradient by $\alpha \approx 0.1$ via a straight-through trick, for the same reason.
Adaptive alternatives to a fixed threshold
A fixed $\tau$ has a defect: as training proceeds the healthy gradient norm falls, so a threshold set at step 1000 becomes progressively less protective.
- AutoClip (Seetharaman et al., 2020) sets $\tau$ to a running percentile of observed norms.
- ZClip (2025) models the norm's history with an exponential moving average and standard deviation, clipping on a z-score rather than an absolute value. Reported to reduce spike frequency in LLM pretraining relative to fixed clipping.
Both share a failure mode worth knowing: if the norm rises slowly, an adaptive threshold rises with it and stops protecting anything.
The monitoring signal
Log $\lVert g \rVert_2$ every step, and watch its distribution rather than the loss. In a healthy run, $\log \lVert g \rVert$ is close to stationary with a slow downward trend and light tails. The observable precursors of a spike, in order of how early they appear:
- Rising maximum attention logit.
- Rising output-logit magnitude (equivalently, falling output entropy).
- Gradient-norm outliers.
- The loss itself.
By the time (4) is visible on a dashboard the step has already been applied. This is why the pre-clip norm belongs in your logging, and why serious runs also log per-layer norms — a spike localised to one layer points at a very different cause from a global one.
Papers
- Pascanu et al., On the difficulty of training recurrent neural networks, 2013 — arxiv.org/abs/1211.5063
- Zhang et al., Why Gradient Clipping Accelerates Training, ICLR 2020 — arxiv.org/abs/1905.11881
- Seetharaman et al., AutoClip, 2020 — arxiv.org/abs/2007.14469
- Chowdhery et al., PaLM, 2022 — arxiv.org/abs/2204.02311 (section 5.1 documents the spikes)
- Zeng et al., GLM-130B, 2023 — arxiv.org/abs/2210.02414
- Molybog et al., A Theory on Adam Instability in Large-Scale Machine Learning, 2023 — arxiv.org/abs/2304.09871
- Wortsman et al., Small-scale proxies for large-scale Transformer training instabilities, 2023 — arxiv.org/abs/2309.14322
- Takase et al., Spike No More, 2024 — arxiv.org/abs/2312.16903
- Kumar et al., ZClip: Adaptive Spike Mitigation for LLM Pre-Training, 2025 — arxiv.org/abs/2504.02507
What to learn next
- Reading a training run while it happens — the dashboard that catches this early.
- Debugging NaN loss — the version of this problem that kills the run outright.
- Mixed precision training — why bf16 removes a whole class of spikes.