rumblr Work in progressWIP

● The AI Primer · Lesson 4 · Part 1: how the model works inside

Deep networks

why deep stacks fail to train, and the fixes that made them work

This lesson covers Vanishing/exploding gradients, residuals, normalization, initialization

Members · open during launch 29 min8 figures and diagrams6 interactive
How it works builds the idea from scratch. Math & code adds the formulas and the Python.

At a glance

Key takeaways

  1. Backprop multiplies one slope per layer, so gradients shrink (vanish) or grow (explode) exponentially with depth.
  2. Initialization sets weight sizes so each layer preserves signal size: Xavier for tanh/sigmoid, He (2 / fan-in) for ReLU.
  3. Residual connections add the input back, giving the gradient a path multiplied by 1; they're why very deep nets and transformers train.
  4. Normalization keeps activations in a steady range: BatchNorm across the batch (CNNs), LayerNorm/RMSNorm within each example (transformers).
  5. 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

Pass the gradient back yourself first; the table below is three settings of this chain.

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

Reading it: read right to left, the direction backprop travels. The loss hands the last layer a gradient of 1. Every hop multiplies by that layer's slope (0.25 for a sigmoid at its steepest). After ten hops the first layer receives about one millionth of the signal, so its weights barely change: it effectively stops learning while the later layers carry on.
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

0 5 10 15 20 25 30 layer (1 = next to the input) 1 0 − 3 6 1 0 − 3 0 1 0 − 2 4 1 0 − 1 8 1 0 − 1 2 1 0 − 6 1 0 0 1 0 6 1 0 1 2 1 0 1 8 1 0 2 4 gradient size ÷ size at layer 30 Gradient flow through 30 layers ReLU, He init ReLU, weights too small ReLU, weights too large sigmoid, Xavier init sigmoid, Xavier, residual

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

Reading it: the horizontal axis is the layer (1 is next to the input, 30 next to the loss). The vertical axis is the size of the gradient reaching that layer divided by its size at layer 30, on a log scale where each gridline is a factor of 10⁶. So every line starts at 1 on the right; read it from right to left, following backprop, and its height at layer 1 is the ratio in the table. The healthy ReLU + He line stays within a factor of about 5 of 1. The too-small line dives about 36 orders of magnitude. The sigmoid line dives too, but only about half as far, 18 orders: Xavier keeps the weights' own gain near 1, so what shrinks the gradient is sigmoid's slope, at most 0.25, about 4× per layer. The too-large line climbs about 22 orders. The dashed line is the same sigmoid network with skip connections, and it stays flat (see residual connections below).
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.

The figure above, with the weight size in your hands and the network running in your browser.

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

Reading it: the choice of starting weights follows the activation. Both rules have the same goal, shown in the last box: a layer should neither shrink nor grow what passes through it.

Figure 4 · Chart

0 5 10 15 20 25 30 layer (0 = the input) 1 0 − 4 0 1 0 − 3 2 1 0 − 2 4 1 0 − 1 6 1 0 − 8 1 0 0 1 0 8 1 0 1 6 1 0 2 4 typical activation size (RMS) Forward signal through 30 layers ReLU, He init ReLU, weights too small ReLU, weights too large sigmoid, Xavier init

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

Reading it: this is the forward direction: the typical size of the activations entering each layer, log scale. With He initialization the ReLU network's signal stays near 1 for all 30 layers. Too small, it fades to nothing within a few layers; too large, it grows by about 5.6× per layer. For the three ReLU lines the backward picture above mirrors this one, because the same weights scale both directions. The sigmoid line is where the mirror breaks: its signal holds steady near 0.5 for all 30 layers (sigmoid's outputs sit around 0.5 whatever comes in), yet its gradient above lost 18 orders of magnitude. The forward pass only sends values through sigmoid; the backward pass multiplies by sigmoid's slope, at most 0.25, at every layer. A healthy forward signal does not prove a healthy gradient, so check both.
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.

Both rules from the table, run on the same network, forward and backward side by side.

Chapter 4

Residual connections: an express lane for the gradient

Switch the express lanes off one at a time before reading: each one you remove costs a whole factor.

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

Reading it: the input splits. One copy goes through the block, the other goes straight around it, and the two are added. Going backwards, the gradient also splits: one part flows back through the block (and may shrink), the other flows through the "+" untouched. That untouched path is the "1" in 1 + 0.025, and it guarantees that some gradient reaches the early layers no matter how deep the stack is.
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

2.5 5.0 7.5 10.0 12.5 15.0 17.5 20.0 number of blocks 1 0 − 2 9 1 0 − 2 4 1 0 − 1 9 1 0 − 1 4 1 0 − 9 1 0 − 4 1 0 1 gradient reaching the input Skip connections keep the gradient alive plain: × 0.025 per block residual: × 1.025 per block

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

Reading it: the x-axis counts stacked blocks and the y-axis, on a log scale, is how much gradient survives the trip back to the input (1 means all of it). The plain line drops by a factor of 40 with every block, a straight plunge on a log scale, and after 10 blocks it is at 10⁻¹⁶: early layers receive essentially nothing. The residual line stays near 1 at every depth, because each block's "1 +" passes the gradient through intact. That flat line is why networks with hundreds of layers can be trained at all.

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

Change a batch-mate's number and watch whether your own grade moves: that is the whole difference.

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

Reading it: the same grid of activations can be normalized in two directions. BatchNorm takes statistics down each column, across the examples, so every example's output depends on its batch-mates. LayerNorm and RMSNorm take statistics along each row, within a single example, so an example is normalized the same way alone or in any batch.
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

Reading it: measure all layers' gradients together as one long list; if that list is longer than the limit, shrink every entry by the same factor.
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.init docs: 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.