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.

Three ideas: vanishing gradients, regularisation, learning-rate schedules.
Figure 4.1 — Gradients, regularisation and schedules.

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.