Vanishing and exploding gradients
Be able to diagnose gradient problems and fix them with initialisation, normalisation and residuals.
Prerequisites
- DActivation functionsrequired
- DBackpropagationrequired
Intuition
The gradient reaching layer 1 of a deep network is a product of many factors. Small deviations from 1 are amplified exponentially:
| Factor per layer | After 30 layers |
|---|---|
| 0.8 | 0.001 — vanished |
| 0.95 | 0.21 |
| 1.0 | 1.0 |
| 1.05 | 4.3 |
| 1.2 | 237 — exploded |
The symptoms differ:
- Vanishing: the early layers do not learn at all; the loss levels off at a poor value; the gradient norm in layer 1 is orders of magnitude smaller than in the last layer.
- Exploding: the loss goes NaN or jumps violently; the gradient norm shoots up at a single step.
The diagnosis is simple: log the gradient norm per layer. The problem is visible immediately.
Formal
Four remedies, and what each actually does:
| Remedy | The mechanism |
|---|---|
| He/Xavier initialisation | keeps the variance constant through the layers from the start |
| ReLU/GELU instead of sigmoid/tanh | the derivative does not saturate towards 0 for large |x| |
| Normalisation (batch/layer) | keeps the activations in a sensible range layer by layer |
| Residual connections | gives — a shortcut with factor 1 |
The residual connection is the most important of them. It is the reason networks with hundreds of layers are trainable: the gradient always has a way back where it is multiplied by 1.
Gradient clipping against exploding gradients: scale the whole gradient vector down if its norm exceeds a cap. Clip on the global norm, not per parameter — otherwise the direction is distorted. c = 1.0 is the standard for transformers.
An important qualification: clipping treats the symptom. If you have to clip at every step, something else is wrong — too high an lr, poor initialisation, or a bug in the loss.
Code
import torch
def gradient_profile(model):
"""Log it per layer — this is where the problem shows up immediately."""
rows = []
for name, p in model.named_parameters():
if p.grad is not None and p.dim() > 1:
rows.append((name, float(p.grad.norm()), float(p.norm())))
for name, g, w in rows:
print(f"{name:34s} |grad| {g:.2e} |w| {w:.2e} ratio {g/max(w,1e-12):.2e}")
return rows
# In the training loop
loss.backward()
norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # the global norm
if norm > 10:
print(f"warning: gradient norm {norm:.1f} before clipping")
opt.step(); opt.zero_grad()
# A typical output with a vanishing gradient:
# blocks.0.attn.W_q |grad| 3.1e-08 ← orders of magnitude smaller
# blocks.15.attn.W_q |grad| 2.4e-04
# blocks.31.attn.W_q |grad| 1.9e-03
The ratio |grad|/|w| is often more informative than the gradient norm itself: it says how large a relative change the step amounts to.
Mastery means
- Diagnoses gradient problems with gradient norms
- Fixes them with initialisation, normalisation and residuals
- Uses gradient clipping correctly
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Deep Residual Learning for Image Recognition — arXiv (open access; licence per article)
- arXiv — On the difficulty of training Recurrent Neural Networks — arXiv (open access; licence per article)
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0