For roughly two decades, deep networks had a reputation for being nearly impossible to train. The main culprit had a precise mathematical cause, identified by Sepp Hochreiter in his 1991 diploma thesis and by Bengio, Simard and Frasconi in 1994: gradients vanish or explode as they propagate through many layers or time steps. Understanding this problem explains the design of nearly every modern architecture.
The mathematics#
By the chain rule, the gradient of the loss with respect to an early layer's activations is a product of Jacobians:
The norm of a product of $L$ matrices behaves roughly like the product of their typical scaling factors. If each Jacobian shrinks vectors by a factor $\alpha < 1$, the gradient shrinks like $\alpha^L$ โ vanishing. If each stretches by $\alpha > 1$, it grows like $\alpha^L$ โ exploding.
Symptoms#
Vanishing gradients:
- early layers barely change; their weights stay near initialisation;
- training loss plateaus early;
- in RNNs, the model cannot learn long-range dependencies.
Exploding gradients:
- loss spikes or suddenly becomes NaN or inf;
- weights grow very large;
- training is unstable and sensitive to the learning rate.
Diagnosis#
Log the gradient norm per layer during training. A healthy network has gradient norms of similar order of magnitude across layers.
import torch
import torch.nn as nn
def layer_grad_norms(depth, act):
torch.manual_seed(0)
layers = []
for _ in range(depth):
layers += [nn.Linear(64, 64), act()]
net = nn.Sequential(*layers, nn.Linear(64, 1))
x, y = torch.randn(256, 64), torch.randn(256, 1)
nn.functional.mse_loss(net(x), y).backward()
linears = [m for m in net if isinstance(m, nn.Linear)]
return [f"{m.weight.grad.norm().item():.1e}" for m in linears[::max(1, depth // 5)]]
print("sigmoid, 30 layers:", layer_grad_norms(30, nn.Sigmoid))
print("relu, 30 layers:", layer_grad_norms(30, nn.ReLU))With sigmoid, early-layer gradient norms are many orders of magnitude smaller than late-layer norms.
The cures#
1. Better activations#
ReLU's derivative is exactly 1 for active units, unlike sigmoid's maximum of 0.25. This single change made much deeper networks trainable.
2. Careful initialisation#
Xavier and He initialisation set weight variances so the Jacobians neither shrink nor stretch signals on average (previous lecture).
3. Normalisation layers#
Batch normalisation and layer normalisation keep activations in a well-scaled range at every layer, preventing drift in scale that compounds with depth.
4. Residual (skip) connections#
With $\mathbf{h}_l = \mathbf{h}_{l-1} + F(\mathbf{h}_{l-1})$, the Jacobian becomes
The identity term gives gradients a direct highway back to early layers. This is why ResNets train with hundreds of layers and why every transformer uses residual connections.
5. Gating#
LSTMs and GRUs control information flow with multiplicative gates and an additive cell state, allowing gradients to flow across many time steps (the "constant error carousel").
6. Gradient clipping (for explosions)#
Rescale the gradient when its norm exceeds a threshold $c$:
Clipping by global norm preserves the gradient's direction while bounding its size. It is standard for RNNs and transformers (often $c = 1.0$).
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()7. Smaller learning rates, warm-up and mixed-precision care#
Explosions are often triggered by an overly aggressive learning rate early in training. Warm-up and appropriate loss scaling (for float16) reduce the risk.
A summary table#
| Problem | Main causes | Fixes |
|---|---|---|
| Vanishing | Saturating activations, small weights, long chains | ReLU/GELU, He init, residuals, normalisation, LSTM/GRU gating |
| Exploding | Large weights, recurrent loops, high learning rate | Clipping, proper init, normalisation, warm-up, lower LR |