Skip to content
AI-grafen
EUniversityDeep learning· about 60 min· evolving, reviewed regularly· verified 2026-09-20· EN

Gradient clipping

Be able to use gradient clipping and explain when it is needed.

Prerequisites

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:

g←g⋅min⁡(1,τ∥g∥)g \leftarrow g \cdot \min\left(1, \frac{\tau}{\|g\|}\right)

The direction is preserved — only the length is shortened. That is the important detail and the difference from clipping every element separately.

The methodDoesPreserves the direction
Norm clippingscales the whole gradientyes
Value clippingclips every element to [−c, c]no

Norm clipping (clip_grad_norm_) is nearly always the right choice.

Formal

Where it is needed:

The modelA typical thresholdWhy
RNNs and LSTMs1.0–5.0long chains of dependency give multiplicative growth
Transformers1.0the standard in nearly every recipe
RL (PPO)0.5rewards can vary sharply
CNNsoften unnecessarythe 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 symptomThe real problem
Clipped every steplr too high, or too tight a threshold
The norm grows steadilyan unstable architecture, missing normalisation
Sudden peaks on individual batchesbad data — outliers, mislabelling, extreme values
NaN despite clippinga 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

All the sources and licences