Lesson 11 / 25
Vanishing Gradients, Normalisation and Residual Connections
Why very deep plain networks stop learning.
Signals that shrink layer after layer
Backpropagation multiplies many factors together. In a deep stack of plain layers those factors are often below 1, so gradients reaching early layers become tiny (vanishing gradients) and those layers barely learn; occasionally they grow instead (exploding gradients). Fixes that made modern deep learning possible: ReLU-family activations, careful initialisation, normalisation layers (batch norm, layer norm), and residual (skip) connections that add a layer's input to its output so gradients have a direct path. Gradient clipping handles explosions.
Stable gradients, regularisation, schedules
Deep networks train well only with the right architecture tricks and regularisation.
Gradient size reaching the first layer, run
I ran this on CPU with Python 3, PyTorch 2.14.1, numpy 2.5.3 and scikit-learn 1.9.1, with fixed seeds and one thread. With 20 plain sigmoid layers the gradient at the first layer is about 3e-20, effectively zero; with 20 plain ReLU layers about 2e-11, still tiny. With 20 residual blocks (layer norm, linear, ReLU, plus a skip connection) it is about 4e-2, large enough to learn.
import torch
torch.manual_seed(0); torch.set_num_threads(1)
class Residual(torch.nn.Module):
def __init__(self):
super().__init__()
self.f = torch.nn.Sequential(torch.nn.LayerNorm(64), torch.nn.Linear(64, 64), torch.nn.ReLU())
def forward(self, x):
return x + self.f(x) # skip connection: gradient has a direct path
def first_layer_grad(blocks):
net = torch.nn.Sequential(torch.nn.Linear(64, 64), *blocks, torch.nn.Linear(64, 1))
net(torch.randn(128, 64)).pow(2).mean().backward()
return net[0].weight.grad.abs().mean().item()
plain = lambda act: [m for _ in range(20) for m in (torch.nn.Linear(64, 64), act())]
print(f"20 sigmoid layers : first-layer gradient {first_layer_grad(plain(torch.nn.Sigmoid)):.2e}")
print(f"20 ReLU layers : first-layer gradient {first_layer_grad(plain(torch.nn.ReLU)):.2e}")
print(f"20 residual ReLU blocks : first-layer gradient {first_layer_grad([Residual() for _ in range(20)]):.2e}")
Output:
20 sigmoid layers : first-layer gradient 2.70e-20 20 ReLU layers : first-layer gradient 1.85e-11 20 residual ReLU blocks : first-layer gradient 4.11e-02
Log gradient norms
Track the gradient norm per layer during training; near-zero early layers or sudden spikes reveal the problem quickly.
Quick check: How do residual connections help deep networks train?
- They freeze the first layer
- They remove all activations
- They reduce the dataset
- They give gradients a direct path around each block
Answer
They give gradients a direct path around each block — Skip connections keep gradients flowing.