The lesson in one minute
What you'll be able to explain
- 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 1
The practitioner's guide
In one sentence
Backpropagation multiplies one factor per layer, so in a deep stack the learning signal shrinks to nothing or grows without bound unless every layer is built to pass it on at about its original size, and the four standard fixes (a good activation, matched initialization, residual connections, normalization) plus a safety net (gradient clipping) are what make every modern architecture trainable.
When you need it
You need this when you read a model's config and
meet rms_norm_eps, initializer_range or layer_norm_eps, when a
network you built stops improving while its loss curve looks merely slow,
when a training run turns NaN in its first steps, or when a paper says
"pre-norm" and you have to decide whether it matters. The tell: a model
whose late layers learn while its early layers stay at their random start,
which no loss curve shows and a plot of gradient size per layer shows at
once. You don't need it to fine-tune a published transformer: the fixes
are baked into its architecture, and your job is to leave them alone. One
number from this lesson says why they are there: in a 30-layer network of
ReLU units, weights drawn a little too small shrink the gradient reaching
the first layer by 36 orders of magnitude, a little too large grow it by
22, and the right size keeps it within a factor of about 4.
Your options
The fixes, from the ones a framework applies for you to the ones that shape an architecture:
| Option | What it does | What it guarantees | What it costs | Where it lives |
|---|---|---|---|---|
| An activation with slope near 1 (ReLU exactly; GELU and SiLU for large positive inputs) | Passes the gradient through unshrunk for positive inputs, where sigmoid passes at most 0.25 (GELU and SiLU pass 0.5 at zero, rising toward 1) | No 0.25-per-layer decay: ten sigmoid layers lose a factor of a million, ten ReLU layers lose nothing | ReLU units pushed negative pass nothing and can die; GELU and SiLU leak a little instead | hidden_act in a config |
| Initialization matched to the activation (Xavier, He) | Sets the starting weights' size so each layer passes signal on at the same size, forward and backward | The healthy line in this lesson's figure: a factor of about 4 over 30 layers, against 10⁻³⁶ or 10²² | Nothing at runtime; a per-layer rule you must apply to custom layers | The framework's default init; initializer_range in a config |
| Residual connections | Each block adds its correction to its input instead of replacing it | Some gradient always reaches the early layers: 1.28 after ten blocks against 10⁻¹⁶ without | The signal grows as corrections pile up (7 million times over 30 layers here) unless normalized; block input and output must share a shape | The architecture: every transformer block |
| Normalization (BatchNorm, LayerNorm, RMSNorm) | Re-centres and rescales activations, across the batch or within each example | Activations in a steady range at every depth; with residuals, a stable stack of any depth | A mean and a variance per layer per step (RMSNorm drops the mean); BatchNorm ties each example to its batch-mates | The architecture: layer_norm_eps, rms_norm_eps |
| Gradient clipping | Rescales the whole update when its length exceeds a limit | A rare spike cannot wreck the run: a gradient of length 8 × 10⁴⁶ becomes length 1, same direction | One norm per step; a network that explodes every step is hidden, not fixed | The training loop: max_grad_norm |
How to choose
Start from whether you are reading an architecture or building one.
- Fine-tuning a published model: read the config and change nothing. A
Llama model is built with RMSNorm placed before each sub-layer
(pre-norm), residuals in every block and SiLU-based activations, and its
released config records the numbers, such as Llama 2's
rms_norm_epsof 10⁻⁵; the trained weights assume every one of them. - Building a network more than a few layers deep: ReLU or GELU, the
framework's default initialization (PyTorch's
nn.Linearscales its starting weights by the fan-in), a residual path around every block, and a normalization layer beside it. - Sequences, or inference one example at a time: LayerNorm or RMSNorm, never BatchNorm, because an example's output must not depend on who else is in the batch. Convolutional networks with large batches: BatchNorm, which ResNet places after every convolution.
- A run that spikes: clip at 1.0, the limit GPT-3 and Llama 2 trained with. If the clip fires on every step, the fault is initialization or normalization, and clipping is masking it.
- A network that trains slowly for no visible reason: plot the gradient norm per layer. Vanishing shows up as a slope of many orders of magnitude from the last layer to the first.
- Whatever you pick, the goal is one number: a per-layer factor near 1 in both directions. Check the forward signal and the backward gradient separately, because a healthy one does not prove a healthy other.
What it costs
Initialization is free. Residual connections cost one addition per value, nothing beside the block's matrix multiplies, but fix the shape of every block's output to its input. A normalization layer costs a mean and a variance per row per layer, which is why RMSNorm, dropping the mean, is the cheaper choice modern language models make. Clipping costs one norm over all parameters per step. What they buy is depth itself: ResNet trained networks over 100 layers deep with residuals, a 7-billion-parameter Llama config stacks 32 blocks, and Llama 3's largest model is a dense transformer with 405 billion parameters, none of which could be trained if the per-layer factor drifted from 1. Depth is also what you pay for at inference: every layer runs on every token.
What breaks
- Early layers never learn. The gradient vanished on the way back: sigmoid or tanh stacked deep (18 orders of magnitude lost over 30 layers even with Xavier initialization), or weights initialized too small.
- NaN in the first steps. Weights too large (22 orders of magnitude of growth), or residual blocks stacked without normalization.
- A healthy forward pass with a dead backward pass. The sigmoid network's signal holds steady near 0.5 through all 30 layers while its gradient collapses, because the forward pass sends values through the activation and the backward pass multiplies by its slope. Check both.
- BatchNorm where the batch is not a population. The value 1 becomes −1.22 in one batch and −0.93 in another; at batch size 1, or with variable-length sequences, the statistics are meaningless. Use LayerNorm.
- A custom layer that silently fails to train. It skipped the initialization rule the framework applies to its own layers.
- Clipping that fires every step. Not a spike: an explosion. Fix the cause.
- Post-norm instability. Placing the norm after the residual add trains less stably than before it (Xiong et al., 2020); pre-norm is what Llama and most recent models use.
In the wild
Llama 2's paper describes its blocks as pre-normalization
with RMSNorm, the SwiGLU activation and rotary position embeddings, and
Hugging Face's LlamaConfig exposes the settings (initializer_range
0.02, num_hidden_layers 32, hidden_act silu, and an rms_norm_eps that
defaults to 1e-6 while Meta's Llama 2 code and released config use 1e-5).
PyTorch's nn.LayerNorm takes the shape to normalize over with
eps=1e-05 and a learned per-element scale and shift; its nn.Linear
initializes from a uniform range set by the fan-in; and
torch.nn.utils.clip_grad_norm_ clips by the norm over all parameters
together, which Hugging Face's TrainingArguments calls with a default
max_grad_norm of 1.0, the same limit Llama 2 trained with. The fixes are
He et al. (ResNet and He initialization, 2015), Glorot and Bengio (Xavier,
2010), Ioffe and Szegedy (BatchNorm, 2015), Ba, Kiros and Hinton
(LayerNorm, 2016) and Zhang and Sennrich (RMSNorm, 2019), with the problem
itself diagnosed by Bengio, Simard and Frasconi (1994); all are linked at
the end of the lesson.
Go deeper
Level 2 multiplies the slopes of a ten-layer chain by hand, watches the gradient at every layer of a 30-layer network under four initializations, derives the Xavier and He rules from one variance equation, shows the "1 +" that residual connections add, normalizes one row three ways with the numbers shown, and clips an exploding gradient of length 8 × 10⁴⁶ down to 1. If you only needed to read a config, you are done.
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 · Drawn from the lesson's code
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 4 · 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 3 · Drawn from the lesson's code
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 6 · 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 5 · Drawn from the lesson's code
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 048aeaa, so the two always agree: the explanation, the code that builds it and the tests that prove it.