At a glance
Key takeaways
- Backprop multiplies one slope per layer, so gradients shrink (vanish) or grow (explode) exponentially with depth.
- Initialization sets weight sizes so each layer preserves signal size: Xavier for tanh/sigmoid, He (2 / fan-in) for ReLU.
- Residual connections add the input back, giving the gradient a path multiplied by 1; they're why very deep nets and transformers train.
- Normalization keeps activations in a steady range: BatchNorm across the batch (CNNs), LayerNorm/RMSNorm within each example (transformers).
- Gradient clipping caps rare spikes.
Level 2
How it works, from scratch
A deep network's gradient is a product with one factor per layer, and every fix in this lesson is a way of holding that factor near 1. This level builds the problem in a chain of ten numbers, watches it in a 30-layer network, then adds each fix and measures what it restores.
Chapter 1
The idea: a gradient is a product of slopes
Picture a game of telephone along a line of 30 people. Each person repeats
the message to the next, but everyone speaks at a quarter of the volume they
heard. By the end of the line the message is silence. If instead everyone
speaks 1.5× louder, the end of the line is a deafening roar. Training a deep
network has exactly this problem, run backwards: the learning signal (the
gradient, how much each weight should change; see primer.ml.neural_net)
starts at the output and is passed back layer by layer, and each layer
multiplies it by its own slope (derivative). Thirty multiplications by
something below 1 is almost zero: the vanishing gradient. Thirty by
something above 1 is enormous: the exploding gradient.
Worked example: a chain of ten one-number layers, each sitting at its steepest point.
| chain | slope per layer | gradient after 10 layers |
|---|---|---|
| sigmoid units, weight 1 | 0.25 | 0.25¹⁰ = 0.00000095 |
| linear units, weight 1 | 1 | 1¹⁰ = 1 |
| linear units, weight 1.5 | 1.5 | 1.5¹⁰ = 57.7 |
Figure 1 · Diagram
flowchart RL L[Loss] -- "gradient 1" --> H10[layer 10] H10 -- "× 0.25" --> H9[layer 9] H9 -- "× 0.25" --> H8[layer 8] H8 -- "× 0.25 ... " --> H2[layer 2] H2 -- "× 0.25" --> H1["layer 1<br/>receives 0.25¹⁰ ≈ 1e-6"]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the input to the first layer | ||
| the output of the last layer | ||
| the number of layers | 10 | |
| a counter over the layers | 1 to 10 | |
| "multiply together the following, for every layer" (like Σ, but multiplying) | ten factors | |
| layer 's weight | 1 | |
| the activation's slope at layer 's input | 0.25 | |
| how much the loss changes when changes | 1 at the top |
In words: "the gradient reaching the first layer is the gradient at the top times every layer's weight times every layer's slope."
With the numbers: .
Level 3: in Python
def gradient_at_input(w_l, slope, L=10):
# ∂L/∂h_L: the loss hands the top layer 1
grad = 1.0
# Π over the layers: a running product
for l in range(L):
# × w_l φ'(z_l)
grad *= w_l * slope
return grad
# sigmoid at its steepest
f"{gradient_at_input(1, 0.25):.1e}" # → '9.5e-07'
# linear, weight 1
gradient_at_input(1, 1) # → 1.0
# linear, weight 1.5
round(gradient_at_input(1.5, 1), 1) # → 57.7
chain_gradient builds that chain and backprops through it.
Why it matters this is why networks deeper than a handful of layers were considered untrainable for decades. Every fix below (better activations, careful initialization, residual connections, normalization) is a way of keeping the per-layer factor close to 1.
Chapter 2
In a real network: watching the gradient layer by layer
In a real layer, each neuron sums 64 inputs, so the multiplier per layer depends on three things together: the size of the weights, how many inputs each neuron adds up, and the activation's slope. Same telephone game, but now everyone in the line hears 64 people at once.
Worked example: a 30-layer network, 64 neurons per layer. The ratio of the gradient at the first layer to the gradient at the last:
| setup | ratio first / last |
|---|---|
| ReLU, He initialization | ≈ 4 (healthy) |
| ReLU, weights too small (std 0.01) | ≈ 10⁻³⁶ (vanished) |
| ReLU, weights too large (std 1) | ≈ 10²² (exploded) |
| sigmoid, Xavier initialization | ≈ 10⁻¹⁸ (vanished) |
Figure 2 · Chart
Gradient size at every layer, relative to layer 30: ReLU with He stays near 1, too-small weights dive 36 orders of magnitude, sigmoid dives 18, too-large weights climb 22, and sigmoid with skip connections stays flat
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the standard deviation (typical size) of the random starting weights | 0.01 (too small) | |
| "fan-in": how many inputs each neuron adds up | 64 | |
| a sum of random terms grows like , not | 8 | |
| typical | the average slope the activation passes back | ≈ 0.7 for ReLU's "half on" |
In words: "each layer multiplies the gradient by roughly the weight size, times the square root of how many inputs it sums, times the activation's typical slope."
With the numbers: too small: per layer, and . Too large: per layer, and .
Level 3: in Python
import math
n_in, typical_slope = 64, 0.7
for sigma_w in (0.01, 1.0):
# σ_w √n_in · typical φ'
gain = sigma_w * math.sqrt(n_in) * typical_slope
# per layer, then over 30 layers
print(round(gain, 3), f"{gain ** 30:.0e}") # → 0.056 3e-38 5.6 3e+22
In code: gradient_norms runs a 30-layer, 64-wide network forward and backward and returns the gradient size reaching every layer; first_to_last_gradient_ratio divides the first by the last to fill the table.
Why it matters you can't see this from the loss curve alone. A network whose early layers get no gradient still trains a little (the late layers learn), just badly. Plotting per-layer gradient norms is a standard diagnostic.
Chapter 3
Initialization: setting every amplifier's volume
Think of a chain of 30 audio amplifiers. If each is set a little too quiet, the sound fades to nothing; a little too loud, and it distorts into noise. Set each so that what comes out is exactly as loud as what went in, and the music survives the whole chain. Initialization picks the random starting weights' size so that each layer passes on a signal of the same size, forward and backward.
Worked example:
| scheme | rule for the weight standard deviation | example |
|---|---|---|
| Xavier (Glorot), for tanh/sigmoid | √(2 / (fan-in + fan-out)) | 100 in, 100 out → √(2/200) = 0.1 |
| He (Kaiming), for ReLU | √(2 / fan-in) | 50 in → √(2/50) = 0.2 |
He uses twice Xavier's variance because ReLU zeroes about half its inputs, throwing away half the signal's energy; the factor 2 puts it back.
Figure 3 · Diagram
flowchart LR
A{Activation?} -->|ReLU / GELU| He["He: std = √(2 / fan_in)"]
A -->|tanh / sigmoid / linear| X["Xavier: std = √(2 / (fan_in + fan_out))"]
He & X --> S[Signal keeps its size<br/>layer after layer]
Figure 4 · Chart
Forward signal size at every layer: ReLU with He stays near 1, too-small weights fade to 10⁻³⁸, too-large weights grow to 10²², and sigmoid holds flat near 0.5 even though its gradient vanishes
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| one neuron's weighted sum | ||
| variance: the average squared distance from the mean (standard deviation squared) | ||
| one weight | ||
| one input to the neuron (the previous layer's output) | ||
| the average of ("E" for expected value, the long-run average) | half the pre-ReLU variance | |
| fan-in | 50 | |
| "therefore" |
In words: "the variance of a neuron's sum is the number of inputs times the weight variance times the average squared input; ReLU halves that average, so to keep the variance steady the weights need variance 2 over the fan-in."
With the numbers: , so the standard deviation is .
Level 3: in Python
import math
n_in = 50
# Var(w) = 2 / n_in, for ReLU
var_w = 2 / n_in
# the variance, then the standard deviation
var_w, round(math.sqrt(var_w), 3) # → (0.04, 0.2)
# Xavier for comparison: 100 in, 100 out
round(math.sqrt(2 / (100 + 100)), 3) # → 0.1
In code: init_std returns the starting weight standard deviation for Xavier, He and two deliberately bad choices; forward_signal_rms measures the forward signal plotted above.
Why it matters every framework initializes this way by default
(PyTorch's nn.Linear uses a Kaiming-style uniform init). Custom layers or
deep stacks built without it can silently fail to train.
Chapter 4
Residual connections: an express lane for the gradient
Picture a building where messages go up by stairs, one floor at a time, and at every landing someone might mumble. Add an express lift that runs the whole height, and the message always arrives intact; each floor adds its own notes to what the lift carries. A residual (or skip) connection is that express lift: each block computes a correction and adds it to its input, instead of replacing the input.
Worked example: ten blocks, each with slope 0.025 of its own.
| per-block factor | after 10 blocks | |
|---|---|---|
| plain: h ← f(h) | 0.025 | 0.025¹⁰ ≈ 9.5 × 10⁻¹⁷ |
| residual: h ← h + f(h) | 1 + 0.025 | 1.025¹⁰ = 1.28 |
Figure 5 · Diagram
flowchart TD IN[h] --> F["block F<br/>(layers, activation)"] F --> ADD((+)) IN -- "skip: identity" --> ADD ADD --> OUT["h + F(h)"]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the signal entering block | ||
| what the block computes (its layers and activation) | slope 0.025 | |
| how the block's output changes with its input | ||
| the identity: "1" for a single number, the do-nothing matrix for vectors | 1 | |
| the block's own slope | 0.025 |
In words: "the output is the input plus a correction, so its slope is one plus the correction's slope, never just the correction's slope."
With the numbers: instead of .
Level 3: in Python
# each block's own slope
dF_dh = 0.025
plain, residual = 1.0, 1.0
for l in range(10):
# h ← F(h): the slope is just ∂F/∂h
plain *= dF_dh
# h ← h + F(h): the slope is I + ∂F/∂h
residual *= 1 + dF_dh
f"{plain:.1e}", round(residual, 2) # → ('9.5e-17', 1.28)
Figure 6 · Chart
Gradient reaching the input as blocks are stacked: without skip connections it falls 40× per block to 10⁻¹⁶ after 10 blocks, with them it stays near 1
In code: residual_chain_gradient multiplies the per-block factors from the table, with or without the skip path.
Why it matters residual connections (ResNet, 2015) made 100+ layer
networks trainable, and every transformer wraps both its attention and its
feed-forward sub-layers in one (see primer.ml.transformer). One caveat
the code shows: adding corrections forever makes the signal grow (a 30-layer
residual ReLU stack here grows its gradient about 7 million×), which is why
residuals are always paired with normalization.
Chapter 5
Normalization: grading on a curve
A teacher can "grade on a curve" in two ways. Batch normalization curves each question across the whole class: your score on question 3 is compared with everyone else's score on question 3, so your grade depends on who else sat the exam. Layer normalization curves each student across their own answers: your scores are rescaled relative to your own average, whoever else is in the room. Both re-centre and re-scale numbers into a steady range so no layer is swamped by huge or tiny values.
Worked examples:
- BatchNorm, batch [[1, 2], [3, 6]]: column means (2, 4), standard deviations (1, 2), so the output is [[−1, −1], [1, 1]]. The value 1 becomes −1.22 in the batch (1, 3, 5) but −0.93 in the batch (1, 3, 11).
- LayerNorm, one row (1, 2, 3, 4): mean 2.5, standard deviation 1.118, so the output is (−1.342, −0.447, 0.447, 1.342), whatever else is in the batch.
- RMSNorm, the same row: root-mean-square √((1+4+9+16)/4) = 2.739, so the output is (0.365, 0.730, 1.095, 1.461). No mean is subtracted.
Figure 7 · Diagram
flowchart LR
subgraph M["activations: rows = examples, columns = features"]
direction TB
r1["ex 1: a b c d"]
r2["ex 2: e f g h"]
r3["ex 3: i j k l"]
end
M -->|"down each column<br/>(across the batch)"| BN[BatchNorm]
M -->|"along each row<br/>(within one example)"| LN[LayerNorm / RMSNorm]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| one example's activations (a row) | (1, 2, 3, 4) | |
| number of features in the row | 4 | |
| the -th feature | ||
| "mu", the row's mean | 2.5 | |
| "sigma squared", the row's variance | 1.25 | |
| a tiny number so we never divide by zero | ||
| "gamma, beta", a learned scale and shift per feature, so the network can undo the normalization if it helps | 1 and 0 at the start | |
| multiply feature by feature |
In words: "LayerNorm subtracts the row's mean, divides by its standard deviation, then applies a learned scale and shift; RMSNorm skips the mean and just divides by the root-mean-square."
With the numbers: LayerNorm: . RMSNorm: .
BatchNorm is the column version of the same formula, with and computed across the batch for each feature.
In Python:
import math
x = [1, 2, 3, 4]
d, eps, gamma, beta = len(x), 1e-5, 1.0, 0.0
# the row's mean
mu = sum(x) / d
# σ², the row's variance
var = sum((x_i - mu) ** 2 for x_i in x) / d
mu, var # → (2.5, 1.25)
# LayerNorm
[round(gamma * (x_i - mu) / math.sqrt(var + eps) + beta, 3) for x_i in x] # → [-1.342, -0.447, 0.447, 1.342]
# √((1/d) Σ x_i² + ε)
rms = math.sqrt(sum(x_i ** 2 for x_i in x) / d + eps)
round(rms, 3) # → 2.739
# RMSNorm: no mean subtracted
[round(gamma * x_i / rms, 3) for x_i in x] # → [0.365, 0.73, 1.095, 1.461]
In code: batch_norm normalizes each column across the batch, layer_norm each row across its own features, and rms_norm divides each row by its root-mean-square.
Why it matters transformers use LayerNorm or RMSNorm, never BatchNorm: sequences have different lengths, batches at inference are often size 1, and an example's output must not depend on its batch-mates. Modern LLMs (Llama and others) use RMSNorm because it's cheaper and works as well, and they place it before each sub-layer ("pre-norm"), which keeps the residual path clean and trains more stably. CNNs are where BatchNorm lives on: ResNet puts it after every convolution.
Chapter 6
Gradient clipping: a circuit breaker
Even with all of the above, one unlucky batch can produce a gradient spike. A circuit breaker doesn't stop the current; it caps it. Clipping by global norm does the same to the update: if the gradient's total length exceeds a limit, it's scaled down to the limit, direction unchanged.
Worked example: in the exploding (too-large) 30-layer network above, the gradients' combined length is astronomically large. Clipped with a limit of 1, the update has length exactly 1 and points the same way.
Figure 8 · Diagram
flowchart LR
G[Gradients of all layers] --> N["‖g‖: combined length"]
N --> C{"above the limit?"}
C -->|yes| S["scale every gradient by limit / ‖g‖"]
C -->|no| K[leave unchanged]
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| every weight gradient, as one long list | the 30 layers' gradients | |
| its length: square every entry, add them up, take the square root | enormous | |
| the limit | 1 |
In words: "if the gradient is longer than the limit, rescale it to the limit."
With the numbers: in the too-large network ;
with every entry is multiplied by about and the new length
is exactly 1. (The
optimizers lesson traces (3, 4) → (0.6, 0.8); see
primer.ml.optimizers.clip_by_global_norm, which this lesson reuses.)
Level 3: in Python
import math
# a stand-in with the same enormous length
g = [4.8e46, 6.4e46]
# ‖g‖
norm = math.sqrt(sum(g_i ** 2 for g_i in g))
f"{norm:.0e}" # → '8e+46'
c = 1
# min(1, c / ‖g‖)
scale = min(1, c / norm)
f"{scale:.0e}" # → '1e-47'
# the new length
round(math.sqrt(sum((g_i * scale) ** 2 for g_i in g)), 6) # → 1.0
In code: layer_gradients collects every layer's weight gradient from a 30-layer network, and global_norm_after_clipping reports their combined length after clipping.
Why it matters clipping treats the symptom, not the cause: it makes a rare spike harmless, but a network that explodes on every step needs better initialization or normalization. Almost every large training run clips at a norm of about 1.
Test yourself
6 questions
Answer each one out loud or on paper before you open it. If you can explain it, you know it.
Question 1Why do gradients vanish in deep sigmoid networks?Think it through, then reveal
Backprop multiplies the gradient by each layer's slope, and sigmoid's slope is at most 0.25. Thirty layers can shrink it by 0.25³⁰, so early layers stop learning.
Question 2What's the difference between Xavier and He initialization?Think it through, then reveal
Both choose the weight variance so signal size is preserved. Xavier uses 2 / (fan-in + fan-out), suited to symmetric activations like tanh. He uses 2 / fan-in, doubling the variance to compensate for ReLU zeroing half its inputs.
Question 3How do residual connections fix vanishing gradients?Think it through, then reveal
The block's output is input + F(input), so its derivative is 1 + F′. The gradient always has an identity path back to early layers that isn't multiplied by small slopes.
Question 4Why do transformers use LayerNorm instead of BatchNorm?Think it through, then reveal
BatchNorm's statistics come from the batch, which breaks for variable-length sequences, tiny or single-example batches at inference, and makes an example's output depend on its batch-mates. LayerNorm normalizes each token across its own features, independent of the batch.
Question 5What is RMSNorm and why do modern LLMs use it?Think it through, then reveal
LayerNorm without the mean subtraction and shift: divide by the root-mean-square and apply a learned scale. It's cheaper and trains as well.
Question 6Gradient clipping or better initialization: which fixes exploding gradients?Think it through, then reveal
Initialization (and normalization) fix the cause, keeping per-layer gain near 1. Clipping is a safety net for occasional spikes.
Primary sources
The papers behind this lesson
Bengio, Simard & Frasconi, Learning long-term dependencies with gradient descent is difficult (IEEE Trans. Neural Networks, 1994): Proved that gradients shrink or explode exponentially through many steps, the root of the problem.
The paper ↗Glorot & Bengio, Understanding the difficulty of training deep feedforward neural networks (AISTATS 2010): Diagnosed saturation and derived Xavier initialization to keep variance steady across layers.
The paper ↗He, Zhang, Ren & Sun, Delving Deep into Rectifiers (2015): Derived the 2 / fan-in (He) initialization for ReLU networks.
The paper ↗He, Zhang, Ren & Sun, Deep Residual Learning for Image Recognition (2015): Introduced residual connections and trained networks over 100 layers deep.
Read the annotated companion →The paper ↗Ioffe & Szegedy, Batch Normalization (2015): Normalized activations across the batch, allowing much higher learning rates.
The paper ↗Ba, Kiros & Hinton, Layer Normalization (2016): Normalized within each example instead, independent of batch size; the version transformers use.
Read the annotated companion →The paper ↗Zhang & Sennrich, Root Mean Square Layer Normalization (2019): Dropped LayerNorm's mean subtraction for a cheaper normalization with the same benefit.
The paper ↗Xiong et al., On Layer Normalization in the Transformer Architecture (2020): Showed why putting the norm before each sub-layer (pre-norm) trains more stably.
The paper ↗Researcher's shelf
Further reading
- CS231n notes, Neural Networks Part 2 (initialization, batch norm): https://cs231n.github.io/neural-networks-2/
- Michael Nielsen, Why are deep neural networks hard to train?: http://neuralnetworksanddeeplearning.com/chap5.html
- Goodfellow, Bengio & Courville, Deep Learning, ch. 8 (optimization for training deep models): https://www.deeplearningbook.org/contents/optimization.html
- PyTorch
nn.initdocs: https://pytorch.org/docs/stable/nn.init.html - PyTorch
nn.LayerNorm: https://pytorch.org/docs/stable/generated/torch.nn.LayerNorm.html
About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit c8d5c21, so the two always agree: the explanation, the code that builds it and the tests that prove it.