Gradient clipping
Be able to use gradient clipping and explain when it is needed.
Prerequisites
- EVanishing and exploding gradientsrequired
Intuition
Sometimes the gradient suddenly becomes enormous — a steep cliff in the loss landscape — and a single step flings the weights so far that the model never recovers. The loss goes to NaN and the run is lost.
Gradient clipping sets a ceiling. If the norm of the gradient is larger than the threshold the whole gradient is scaled down so that the norm becomes exactly the threshold:
The direction is preserved — only the length is shortened. That is the important detail and the difference from clipping every element separately.
| The method | Does | Preserves the direction |
|---|---|---|
| Norm clipping | scales the whole gradient | yes |
| Value clipping | clips every element to [−c, c] | no |
Norm clipping (clip_grad_norm_) is nearly always the right choice.
Formal
Where it is needed:
| The model | A typical threshold | Why |
|---|---|---|
| RNNs and LSTMs | 1.0–5.0 | long chains of dependency give multiplicative growth |
| Transformers | 1.0 | the standard in nearly every recipe |
| RL (PPO) | 0.5 | rewards can vary sharply |
| CNNs | often unnecessary | the normalisation layers already keep the gradients in check |
Where in the code — the order matters:
loss.backward()
unscale if AMP is used
clip_grad_norm_(parameters, max_norm)
optimizer.step()
If you clip before backward() you clip nothing (the gradients do not exist yet). If you clip after step() it is too late. And with mixed precision the gradients have to be scaled back before the clipping, otherwise you clip against an enlarged norm.
Choose the threshold from data, not from a gut feeling. Log the gradient norm for a few hundred steps without clipping and set the threshold at roughly the 90th percentile. Then only the genuine outliers are clipped, and the normal steps are unaffected.
Clipping is a safety net, not a solution. If you have to clip at every step the threshold is too low or something else is wrong:
| The symptom | The real problem |
|---|---|
| Clipped every step | lr too high, or too tight a threshold |
| The norm grows steadily | an unstable architecture, missing normalisation |
| Sudden peaks on individual batches | bad data — outliers, mislabelling, extreme values |
| NaN despite clipping | a division by zero, log(0), or NaN already in the input |
The third row is worth investigating: log which batch caused the peak and look at it. Often it is a single broken data point, and mending the data is better than clipping harder.
Code
import torch, torch.nn as nn
# The basic pattern
for x, y in dataloader:
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad()
loss.backward()
norm = nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # returns the norm BEFORE
opt.step()
# With mixed precision: unscale first, otherwise you clip against an enlarged norm
scaler = torch.amp.GradScaler()
for x, y in dataloader:
with torch.autocast("cuda", dtype=torch.bfloat16):
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad()
scaler.scale(loss).backward()
scaler.unscale_(opt) # ← necessary
nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(opt); scaler.update()
# Choose the threshold from data: log the norms without clipping first
import numpy as np
def measure_gradient_norms(model, dataloader, opt, steps=200):
norms = []
for i, (x, y) in enumerate(dataloader):
if i >= steps:
break
loss = nn.functional.cross_entropy(model(x), y)
opt.zero_grad(); loss.backward()
norms.append(float(torch.nn.utils.clip_grad_norm_(
model.parameters(), max_norm=float("inf")))) # measure without clipping
opt.step()
n = np.array(norms)
return {"median": round(float(np.median(n)), 3),
"p90": round(float(np.percentile(n, 90)), 3),
"p99": round(float(np.percentile(n, 99)), 3),
"max": round(float(n.max()), 3),
"suggested_threshold": round(float(np.percentile(n, 90)), 2)}
# Keep track of how often it is actually clipped
clipped = total = 0
for x, y in dataloader:
...
norm = nn.utils.clip_grad_norm_(model.parameters(), 1.0)
clipped += int(float(norm) > 1.0); total += 1
print(f"clipped in {clipped / total:.1%} of the steps")
# > 50 % means the threshold is too low or the learning rate too high
The last measurement is what makes clipping a deliberate choice instead of an incantation copied in from somewhere.
Mastery means
- Uses gradient clipping the right way
- Chooses the threshold from the observed norms
- Knows when clipping is hiding another problem
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — On the difficulty of training Recurrent Neural Networks — arXiv (open access; licence per article)
- PyTorch — tutorials (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0