“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 )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-13That 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) / topA “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.131Sigmoid 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.08The 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.5This 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}")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")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.