Gradient clipping is the standard, direct fix for exploding gradients: before applying a weight update, rescale the gradient so its magnitude never exceeds a chosen threshold — capping the damage a single bad update can do, without otherwise changing the optimizer.
Formula — Clipping by Norm
\(\|\mathbf{g}\|\) is the gradient's L2 norm (see Vector Norms), computed across all parameters together (not per-parameter). \(\tau\) is the clipping threshold — a chosen hyperparameter, commonly something like 1.0 or 5.0. If the gradient's norm exceeds \(\tau\), every component is scaled down proportionally so the resulting norm becomes exactly \(\tau\) — crucially, the gradient's direction is preserved; only its magnitude is capped.
Numerical Example
A gradient vector \(\mathbf{g}=[3, 4]\) has norm \(\|\mathbf{g}\|=\sqrt{9+16}=5\). With threshold \(\tau=2\):
Check: \(\|[1.2,1.6]\| = \sqrt{1.44+2.56}=\sqrt{4}=2\) — exactly the threshold, direction unchanged (still pointing the same way as \([3,4]\), just shorter).
Code
import torch
import torch.nn as nn
model = nn.Sequential(nn.Linear(10, 10), nn.Linear(10, 1))
x = torch.randn(1, 10)
loss = model(x).sum()
loss.backward()
# Clip the combined gradient norm across ALL parameters to at most 1.0
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer = torch.optim.Adam(model.parameters())
optimizer.step() # the update now uses the clipped gradient
The clipping call happens after loss.backward() (gradients must exist first) and before optimizer.step() (the optimizer must use the already-clipped values) — this ordering is essential.
Where It's Standard Practice
- RNN/LSTM training — where exploding gradients (see Exploding Gradient Problem) are especially common due to the same weight matrix being applied repeatedly across time steps.
- Transformer training — gradient clipping is routinely applied as a standard stabilization measure, often alongside learning rate warmup (see Warmup Learning Rate), for very large-scale models.
Clipping by Value — A Simpler Alternative
A less common variant simply clips each individual gradient component to a fixed range (e.g. \([-1, 1]\)) independently, rather than rescaling the whole gradient vector by its norm. This is simpler but doesn't preserve the gradient's overall direction the way norm-based clipping does, and is generally used less often in practice.
Common Mistakes
- Calling the clipping function before
loss.backward()— gradients don't exist yet at that point, so there's nothing to clip. - Setting the clipping threshold \(\tau\) so low that it interferes with normal, healthy training even when gradients aren't actually exploding — this can slow convergence unnecessarily; \(\tau\) should be chosen based on the gradient norms observed during genuinely stable training, not set arbitrarily low as a precaution.
Interview Relevance
Q: "How does gradient clipping prevent exploding gradients without changing what the optimizer or model architecture does?" It intercepts the gradient exactly between backpropagation and the optimizer's update step, rescaling it (by its L2 norm) if it exceeds a threshold — capping the update's magnitude while preserving its direction. Neither the model's forward pass, backpropagation's gradient computation, nor the optimizer's update rule needs any modification; clipping is a clean, self-contained safeguard inserted at one specific point in the training loop.
Key Takeaways — Backpropagation
- Backpropagation exists because computing gradients any other way (like finite differences) is computationally infeasible at the scale of real networks.
- It applies the chain rule systematically, layer by layer, reusing each layer's local derivative rather than recomputing shared work — the forward pass caches what the backward pass needs, and the backward pass computes each layer's error signal recursively.
- Gradient calculation (the outer product formula) and weight updates (handled by a separate, swappable optimizer) are the final two steps, cleanly decoupled from backpropagation itself.
- Vanishing and exploding gradients are the two failure modes of the same underlying multiplicative chain-rule structure — one shrinking, one growing exponentially with depth — with gradient clipping as the standard direct fix for the exploding case.
Next: Training Deep Networks zooms out from the mechanics of a single gradient computation to the full training loop — epochs, checkpointing, early stopping, and the underfitting/overfitting tradeoff that determines whether a trained model actually generalizes.
Practice Question
A gradient vector has norm 8.0, with threshold \(\tau=2.0\). By what factor is every component of the gradient scaled down during clipping?