Ch. 29 · Deep Learning

Vanishing and Exploding Gradients: Causes, Init and Fixes

Why gradients vanish or explode in deep networks, measured layer by layer, and how He and Xavier init, residuals and clipping fix it.

~9 min readintermediateupdated Oct 6, 2026

“What is the vanishing gradient problem and how do you fix it?” is one of the most common deep learning interview questions, and it often arrives disguised: why did ReLU replace sigmoid, why does weight initialization matter, why do ResNets have skip connections, or “my loss is stuck at 2.30, what is wrong?”. Interviewers want a mechanism, not buzzwords: backpropagation multiplies one factor per layer, and a long product of numbers that are not close to 1 either collapses toward zero or grows without bound.

Before you start

You should know how backpropagation passes a gradient from one layer to the previous one (multiply by the activation’s derivative, then by the transposed weight matrix). If that is new, read the backpropagation walkthrough linked at the end first. The experiments use Python 3.12 or newer and only the standard library; they are slow-but-honest matrix loops on 64-unit layers. PyTorch snippets are illustrative (PyTorch 2.x) and are not needed to run anything.

The short answer

The gradient for an early layer is a product of per-layer terms: each layer contributes its weight matrix and its activation’s derivative. If those terms are consistently smaller than 1 the gradient shrinks exponentially with depth (vanishing) and early layers stop learning; if they are larger than 1 it grows exponentially (exploding) and training produces huge updates or NaN. The fixes keep the per-layer factor near 1: non-saturating activations like ReLU, variance-preserving initialization (Xavier for tanh, He for ReLU), normalization layers, residual connections, and gradient clipping as a safety net against explosions.

How it works

For a stack of layers z_l = W_l h_(l-1) and h_l = f(z_l), backpropagation computes:

dL/dh_(l-1) = W_l^T ( f'(z_l) * dL/dh_l )
Text

Unroll that across L layers and the gradient at the bottom is the top gradient multiplied by L matrices of the form W^T diag(f'(z)). Two things set the size of each factor.

The activation derivative. The sigmoid’s derivative is at most 0.25 (at z = 0) and close to zero once the unit saturates. Even in the best case, 10 sigmoid layers multiply the gradient by 0.25**10, about one in a million. Tanh peaks at 1 but also saturates. ReLU’s derivative is exactly 1 for positive inputs, which is why it trains deep networks so much better, though it is 0 for negative inputs.

The weight scale. If each weight has variance s**2 and a layer has n inputs, each pre-activation sums n terms, so its variance is about n * s**2 times the mean square of the inputs. Pick s too small and signals shrink at every layer, forwards and backwards; too large and they grow. Initialization schemes solve for the s that keeps the variance constant:

Scheme Weight variance Designed for
Xavier (Glorot) 2 / (fan_in + fan_out) tanh, sigmoid, linear
He (Kaiming) 2 / fan_in ReLU and its variants
LeCun 1 / fan_in SELU, and a common default elsewhere

He initialization has the extra factor of 2 because ReLU zeroes half of its inputs on average, which halves the signal’s mean square. Doubling the weight variance compensates.

Step-by-step walkthrough

Step 1: Multiply sigmoid derivatives by hand

import math

def sigmoid(z):
    return 1 / (1 + math.exp(-z))

print("max sigmoid'(z) =", sigmoid(0) * (1 - sigmoid(0)))   # max sigmoid'(z) = 0.25
for depth in (5, 10, 20):
    print(depth, f"{0.25 ** depth:.2e}")
# 5 9.77e-04
# 10 9.54e-07
# 20 9.09e-13
python

That is the optimistic bound, with every unit at the steepest point of the curve. Real networks do worse.

Step 2: Build a probe that measures a deep network

The probe pushes one input through depth dense layers, then sends a random gradient back down and reports two numbers: the root-mean-square of the final activations and how much the gradient grew or shrank on its way to the input.

import math, random

ACTS = {
    "sigmoid": (sigmoid, lambda z: sigmoid(z) * (1 - sigmoid(z))),
    "tanh": (math.tanh, lambda z: 1 - math.tanh(z) ** 2),
    "relu": (lambda z: max(0.0, z), lambda z: 1.0 if z > 0 else 0.0),
}

def rms(v):
    return math.sqrt(sum(x * x for x in v) / len(v))

def probe(act_name, weight_std, depth=20, width=64, seed=0):
    f, df = ACTS[act_name]
    rng = random.Random(seed)
    Ws = [[[rng.gauss(0, weight_std) for _ in range(width)] for _ in range(width)]
          for _ in range(depth)]
    h = [rng.gauss(0, 1) for _ in range(width)]
    zs = []
    for W in Ws:                                   # forward
        z = [sum(w * x for w, x in zip(row, h)) for row in W]
        zs.append(z)
        h = [f(v) for v in z]
    g = [rng.gauss(0, 1) for _ in range(width)]    # dL/dh at the top
    top = rms(g)
    for W, z in zip(reversed(Ws), reversed(zs)):   # backward: g = W^T (g * f'(z))
        gz = [gi * df(zi) for gi, zi in zip(g, z)]
        g = [sum(W[i][j] * gz[i] for i in range(width)) for j in range(width)]
    return rms(h), rms(g) / top
python

A “grad ratio” of 1 means the bottom layer receives a gradient as large as the top one. That is the target.

Step 3: Compare initializations on a 20-layer network

n = 64
for act, std, label in [
    ("sigmoid", 0.01, "small N(0, 0.01)"),
    ("sigmoid", math.sqrt(1 / n), "Xavier"),
    ("relu", 1.0, "N(0, 1)"),
    ("relu", 0.01, "small N(0, 0.01)"),
    ("relu", math.sqrt(2 / n), "He"),
    ("tanh", math.sqrt(1 / n), "Xavier"),
]:
    out, ratio = probe(act, std)
    print(f"{act:7s} {label:17s} output rms {out:9.3g}   grad ratio {ratio:9.3g}")
# sigmoid small N(0, 0.01)  output rms     0.501   grad ratio  7.37e-35
# sigmoid Xavier            output rms     0.519   grad ratio   2.2e-13
# relu    N(0, 1)           output rms  1.58e+15   grad ratio  6.44e+14
# relu    small N(0, 0.01)  output rms  1.58e-25   grad ratio  6.44e-26
# relu    He                output rms      1.41   grad ratio     0.572
# tanh    Xavier            output rms     0.155   grad ratio     0.131
python

Sigmoid vanishes even with a sensible init, because its derivative caps every factor at 0.25. ReLU with unit-variance weights explodes by fifteen orders of magnitude, and the same ReLU network with tiny weights vanishes by twenty-five: the activation is identical, only the scale changed. He initialization keeps both the forward signal and the gradient within a small factor of 1 after 20 layers. Tanh with Xavier is usable but still drifts downward.

Step 4: Add a residual path

A residual block computes h = h + f(W h). Its backward pass becomes g = g + W^T (g * f'(z)): the incoming gradient is copied straight through, and the layer’s contribution is added on top. Change two lines of the probe and rerun with the small init that failed before:

def probe_residual(act_name, weight_std, depth=20, width=64, seed=0):
    f, df = ACTS[act_name]
    rng = random.Random(seed)
    Ws = [[[rng.gauss(0, weight_std) for _ in range(width)] for _ in range(width)]
          for _ in range(depth)]
    h = [rng.gauss(0, 1) for _ in range(width)]
    zs = []
    for W in Ws:
        z = [sum(w * x for w, x in zip(row, h)) for row in W]
        zs.append(z)
        h = [hi + f(v) for hi, v in zip(h, z)]                          # skip connection
    g = [rng.gauss(0, 1) for _ in range(width)]
    top = rms(g)
    for W, z in zip(reversed(Ws), reversed(zs)):
        gz = [gi * df(zi) for gi, zi in zip(g, z)]
        g = [g[j] + sum(W[i][j] * gz[i] for i in range(width)) for j in range(width)]
    return rms(h), rms(g) / top

for act in ("sigmoid", "tanh"):
    print(f"{act:7s} plain {probe(act, 0.01)[1]:8.3g}   residual {probe_residual(act, 0.01)[1]:6.3g}")
# sigmoid plain 7.37e-35   residual      1
# tanh    plain 8.11e-23   residual   1.08
python

The identity path guarantees a gradient highway regardless of what the layers do. This is the core idea behind ResNets and the residual stream in every transformer.

Step 5: Clip exploding gradients by their global norm

Initialization fixes the starting point, but gradients can still spike later, especially in recurrent networks. Clipping rescales the whole gradient vector when its norm exceeds a threshold, which preserves its direction:

def clip_by_global_norm(grads, max_norm):
    total = math.sqrt(sum(g * g for g in grads))
    scale = min(1.0, max_norm / (total + 1e-6))
    return [g * scale for g in grads], total

for grads in ([30.0, -40.0], [0.3, -0.4]):
    clipped, norm = clip_by_global_norm(grads, 5.0)
    print([round(g, 4) for g in clipped], norm)
# [3.0, -4.0] 50.0
# [0.3, -0.4] 0.5
python

This mirrors what torch.nn.utils.clip_grad_norm_ does. It returns the pre-clipping norm, and logging that number is one of the cheapest training diagnostics you can add.

Worked scenario

A 30-layer ReLU MLP for a 10-class problem trains for an hour and the loss never moves from 2.3026. The code initializes every Linear layer with normal_(std=0.02) because a transformer codebase used that value.

The number itself is the clue: ln(10) = 2.3026 is the cross-entropy of a uniform prediction over 10 classes, so the logits are essentially zero. With width 256, each ReLU layer scales the signal by about sqrt(256 * 0.02**2 / 2) = 0.226, and 0.226**30 is about 4e-20. The output carries no information about the input, and the gradient reaching early layers has shrunk by the same factor on the way back.

Logging per-layer gradient norms confirms it. The illustrative PyTorch version:

# Illustrative PyTorch 2.x: print each layer's gradient norm after loss.backward()
for name, p in model.named_parameters():
    if p.grad is not None:
        print(f"{name:30s} {p.grad.norm().item():.2e}")
python

The fix is He initialization (torch.nn.init.kaiming_normal_(w, nonlinearity="relu"), std sqrt(2/256) = 0.088 here) or simply keeping PyTorch’s default Linear init, which scales by 1/sqrt(fan_in). For a network this deep, adding residual connections and normalization makes training robust to the exact scale. The transformer got away with 0.02 because its residual stream and LayerNorm keep the signal alive; copying the number without the architecture broke it.

Common mistake

  • “ReLU solves vanishing gradients.” It removes saturation for positive inputs, but bad weight scales still vanish or explode (the table above), and units stuck at negative inputs pass no gradient at all.
  • “Gradient clipping fixes vanishing gradients.” Clipping only shrinks large gradients. It does nothing for small ones.
  • “Vanishing gradients mean the model has converged.” A stalled loss far above what a simple baseline achieves is a symptom, not convergence. Check gradient norms per layer.
  • “Batch norm makes initialization irrelevant.” It makes networks far less sensitive, but the init still matters for layers it does not cover and for the residual branches.
  • Confusing exploding gradients with a high learning rate. Both produce divergence. Gradient norms that grow before the loss explodes point to the former.

Verify the behavior

Assert the properties instead of eyeballing them. This runs in a few seconds with python3:

def test_he_init_keeps_relu_signal_alive():
    out, ratio = probe("relu", math.sqrt(2 / 64))
    assert 0.1 < out < 10 and 0.1 < ratio < 10

def test_scale_alone_decides_vanish_or_explode():
    assert probe("relu", 0.01)[1] < 1e-10
    assert probe("relu", 1.0)[1] > 1e10

def test_clipping_caps_the_norm_and_keeps_direction():
    clipped, _ = clip_by_global_norm([30.0, -40.0], 5.0)
    assert abs(math.hypot(*clipped) - 5.0) < 1e-5
    assert abs(clipped[0] / clipped[1] - 30.0 / -40.0) < 1e-12

test_he_init_keeps_relu_signal_alive(); test_scale_alone_decides_vanish_or_explode()
test_clipping_caps_the_norm_and_keeps_direction()
print("ok")
python

Follow-up questions

How do LSTMs address vanishing gradients in recurrent networks? The cell state is updated additively and scaled by a learned forget gate, so the gradient along it is not forced through a squashing nonlinearity at every time step. It is the same idea as a residual connection, applied across time.

Clip by norm or by value? Clipping by global norm rescales every component equally and keeps the update direction. Clipping by value caps each component separately and can change the direction. Norm clipping is the usual default.

Why does mixed-precision training need loss scaling? Float16 cannot represent very small numbers, so tiny gradients underflow to zero. Multiplying the loss by a scale factor before backward() and dividing the gradients afterwards keeps them representable. bfloat16 has float32’s exponent range and usually does not need it.

What does PyTorch use by default for nn.Linear? A Kaiming-uniform variant that works out to a uniform distribution with bound 1/sqrt(fan_in) for the weights.

Interview exercise

A colleague builds a 50-layer tanh network, width 512, and initializes every weight from N(0, 1). They expect vanishing gradients because tanh saturates. Predict what actually happens to the forward activations and to the gradients, and propose the right initialization.

Answer and reasoning

Each pre-activation sums 512 terms with unit-variance weights, so its standard deviation is roughly sqrt(512), about 22.6. Almost every tanh unit is saturated at plus or minus 1, so the forward signal is a nearly binary pattern that carries little information. The surprise is the backward pass: most derivatives are near zero, but the weight matrices are enormous, and the product does not reliably vanish. Running the probe on a 20-layer, 128-wide version of this network gives a gradient ratio around 1e7, so gradients explode while activations saturate. This is the “chaotic” regime of a badly scaled network. The fix is Xavier initialization, standard deviation sqrt(1/512) (about 0.044) for equal fan-in and fan-out, which keeps pre-activations in tanh’s linear region. A strong answer also says: measure activation statistics and per-layer gradient norms on the first batch instead of reasoning from the activation function alone.

Continue learning

More in Deep Learning

esc