Course Content
Deep Learning Essentials
13 sections · 61 lessons
How do you treat vanishing and exploding gradients?
What you need to know
Each fix works by keeping the per-layer factor in the gradient product close to 1.
ReLU-family activations
ReLU's derivative is exactly 1 for positive inputs, so it does not shrink the gradient the way sigmoid's 0.25 does. GELU and LeakyReLU behave similarly. This one change made training networks with more than a handful of layers practical.
Good initialisation
If weights start too small, each layer's output shrinks and so do gradients. Too big, and they grow. He initialisation draws weights with variance 2 / fan_in (fan_in = number of inputs to the layer), which keeps the size of activations roughly constant through ReLU layers. Xavier (Glorot) initialisation does the same for tanh and sigmoid. PyTorch's default nn.Linear initialisation is a Kaiming-uniform variant, which is sensible for most cases.
Normalisation layers
Batch normalisation and layer normalisation rescale each layer's activations to a mean near 0 and a standard deviation near 1, then apply a learnable scale and shift. Activations cannot drift into the flat regions of an activation function or grow without limit, so gradients stay in a healthy range.
Residual connections (the biggest single fix)
A residual block computes y = x + F(x) instead of y = F(x). The derivative of y with respect to x is 1 + F'(x). Even if F'(x) is tiny, the gradient still has the 1 path straight back. This is why ResNets could go from about 20 layers to over 100, and why every transformer block has residual connections.
Gates for sequences
An LSTM keeps a cell state updated by addition (c_t = f × c_{t−1} + i × candidate) rather than by repeated matrix multiplication. When the forget gate f is near 1, the gradient flows back through many time steps nearly unchanged. Transformers avoid the problem another way: attention connects any two positions directly.
Gradient clipping for explosions
Clipping by norm measures the length of the whole gradient vector. If it is larger than a threshold, it scales every gradient down by the same factor, so the direction is kept but the step is capped. It goes between backward() and step():
loss.backward()torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)optimizer.step()If the total norm is 25 and max_norm is 1.0, every gradient is multiplied by 1/25. If the norm is 0.6, nothing changes. Clipping by value (capping each element separately) also exists but changes the direction, so norm clipping is preferred.
Fixes for vanishing
- ReLU or GELU in hidden layers
- He or Xavier initialisation
- Batch or layer normalisation
- Residual connections, LSTM or GRU gates
Fixes for exploding
- Gradient clipping by norm
- Lower learning rate, with warmup
- Careful initialisation
- Normalisation layers
A real-life example
A bank's team trains an LSTM on UPI transaction sequences to flag fraud: each customer's last 200 transactions go in, and the model predicts whether the next one is fraudulent. Training is stable for most of an epoch, then the loss jumps from 0.08 to NaN. Logging shows the gradient norm on that batch was about 4,000 instead of the usual 0.5 to 2, caused by a batch of merchant accounts with extreme amounts.
The fix is two lines: standardise the log of the transaction amount instead of the raw rupee value, and add clip_grad_norm_(..., max_norm=1.0). The team also logs the gradient norm every 100 steps to a dashboard, so the next spike is visible before it becomes a NaN.
Follow-up questions to expect
- "How do you choose
max_norm?" — Log the gradient norm for a few hundred steps of healthy training and set the threshold a little above its typical value; 1.0 is a common default for transformers and RNNs. - "Why did ResNet need residual connections if it already had batch norm and ReLU?" — Without them, very deep plain networks trained worse than shallower ones even on training data. The identity path made depth help instead of hurt.
- "Does mixed-precision training affect this?" — Yes. In 16-bit floats, very small gradients can round to zero. PyTorch's
GradScalermultiplies the loss by a large factor beforebackward()and divides it out before the update, so small gradients survive.