Lesson 4 of 7 · 47 min
The transformer block & stacking
How attention and the FFN are wired into a block, and why the wiring is the part that determines whether a 70B model trains at all. Residual connections, LayerNorm vs RMSNorm, pre- vs post-norm (the gradient-flow argument), the 4× FFN and SwiGLU, and why stacking depth works.
The “boring” choices that decide if it trains
x = x + attn(norm(x)) then x = x + ffn(norm(x)). Read the structure carefully: the residual stream x flows straight through, untouched, and each sublayer reads a normalized copy of it and adds a delta back. The residual stream is the spine of the network — a clean additive highway from the embedding to the final logits that every interpretability and steering technique relies on.x + f(x) gives the gradient a direct route backward, so it doesn’t vanish through dozens of layers. (2) Easier optimization: each sublayer learns a residual (a correction to the running representation) rather than a full transformation from scratch, which is empirically far easier to optimize. Interview angle. The widely-misstated detail: the addition happens after the sublayer output (x + sublayer(x)), not before an activation — and in pre-norm the sublayer reads norm(x) while the unnormalized x is what’s carried forward on the residual.1import torch.nn as nn23class Block(nn.Module): # the modern (pre-norm) transformer block4 def __init__(self, dim, n_head):5 super().__init__()6 self.norm1 = nn.RMSNorm(dim) # RMSNorm, not LayerNorm7 self.attn = CausalSelfAttention(dim, n_head)8 self.norm2 = nn.RMSNorm(dim)9 self.ffn = SwiGLU(dim) # gated FFN, ~2/3*4*dim inner10 def forward(self, x):11 x = x + self.attn(self.norm1(x)) # pre-norm sublayer 1; residual add AFTER12 x = x + self.ffn(self.norm2(x)) # pre-norm sublayer 213 return x # unnormalized x is the highway forward1415# stack N of these; x flows through untouched by norm -> stable gradients at depth.
Let’s reproduce GPT-2 (124M)Andrej KarpathyPre-norm vs post-norm — the training-stability argument
LayerNorm(x + Sublayer(x)) — normalize after the residual add. Modern LLMs use pre-norm: x + Sublayer(LayerNorm(x)) — normalize the input before the sublayer. This is not a stylistic choice. Xiong et al. (2020) showed empirically that in post-norm the gradient norm at the last layer stays pinned (~1.6) regardless of model size while gradients to early layers are unstable, so post-norm requires a learning-rate warm-up to train at all. In pre-norm the residual stream is a clean unnormalized additive path, gradients are well-behaved at init, and you can remove warm-up entirely.1POST-NORM (2017) vs PRE-NORM (every modern LLM)23 x -> Sublayer -> +x -> LN x -> LN -> Sublayer -> +x4 (norm AFTER the add) (norm BEFORE the sublayer; residual stays clean)56 post-norm: gradient to early layers unstable -> NEEDS warm-up; hard past ~24 layers7 pre-norm: clean residual highway -> stable at init -> warm-up optional; scales deep89 Why it matters at 70B: pre-norm is often the difference between "trains" and "diverges."LayerNorm vs RMSNorm
LN(x) = (x − mean(x))/√(var(x)+ε)·γ + β, with learned per-dimension scale γ and shift β. Crucially it normalizes per token across features — independent of batch size and sequence length, which is exactly why transformers use it instead of BatchNorm. Interview angle. “Why LayerNorm not BatchNorm?” → BatchNorm’s statistics depend on the batch and break for variable-length sequences and small/streaming batches, and complicate distributed sync; LayerNorm is per-token, so it’s stable regardless of batch composition.β shift, keeping only the rescale: RMSNorm(x) = x/RMS(x)·γ where RMS(x) = √(mean(xᵢ²)). The bet is that re-scaling invariance is what actually stabilizes training; mean-centering is a side effect you can shed. The payoff is purely efficiency — one fewer mean computation, one fewer subtraction, one fewer learned bias vector per normalization site — which adds up across hundreds of norm layers and matters in tight inference kernels. Quality is comparable; the saving is real.1LAYERNORM vs RMSNORM (per token, across d_model features)23 LayerNorm: (x - mean(x)) / sqrt(var(x) + eps) * gamma + beta4 RMSNorm: x / sqrt(mean(x^2) + eps) * gamma (no mean subtraction, no beta)56 RMSNorm drops: 1 mean, 1 subtraction, 1 learned bias (beta) per norm site.7 Across ~2*n_layers norm sites at inference, that is a real per-token saving.8 Both are PER-TOKEN (not per-batch) -> why neither is BatchNorm.The feed-forward network — where most parameters live
FFN(x) = max(0, x·W1 + b1)·W2 + b2, a 4× expansion on the inner dimension (e.g. 768 → 3072 → 768 for GPT-2) with a nonlinearity in the middle. This is where most of a transformer’s parameters and FLOPs actually sit — attention gets the attention, but the FFN is the bulk of the compute. Intuition: attention decides what to look at; the FFN does the per-token thinking on what it found, acting like a key-value memory over learned features.FFN(x) = (Swish(x·W1) ⊙ (x·W2))·W3 — three weight matrices, a gate that lets the network modulate which features pass. It consistently improves perplexity at comparable parameter count; Llama sizes the inner dim to ≈ 2/3 · 4 · d_model (e.g. ~11008 for d_model=4096) to keep the param budget similar to a 4× ReLU FFN despite the third matrix. nanoGPT keeps plain GELU because it faithfully reproduces GPT-2. Interview angle. “What’s in a modern block that wasn’t in the 2017 one?” → pre-norm, RMSNorm, RoPE, SwiGLU, and bias-free linear layers — the crystallized 2023–2026 stack.Why stacking depth works
N identical blocks builds a hierarchy on the shared residual stream: early layers tend to resolve local/syntactic structure, middle layers compose longer-range relations, later layers assemble task-level semantics — each block reading the running representation and writing a refined delta back. Depth is a real scaling axis (GPT-2 124M: 12 layers; Llama 2 70B: 80 layers), and it only works because of the wiring above: residuals keep the gradient alive through depth, and pre-norm keeps it stable at init. Remove either and deep stacks stop training.d_model/FFN) add per-layer capacity and parallelize cleanly; more layers add representational composition but lengthen the sequential dependency and stress gradient flow. Scaling-law work guides the ratio for a given parameter budget — and a recurring finding is that for inference, depth is the costlier axis because each layer adds a sequential KV read at decode (Lesson 5), whereas width batches into the same matmul. Interview angle. “Deeper or wider for the same params?” → both add capacity, but depth composes features and raises decode-time sequential cost, while width is cheaper to serve; pick per your latency budget and follow the scaling-law ratio.Key idea
The 2023–2026 stack crystallized into a small set of defaults — pre-norm, RMSNorm, RoPE, SwiGLU, bias-free layers — not because they’re elegant, but because each one removes a training-stability or efficiency failure that bites at scale.
Common mistake
“Residual connections add the input to the output before the activation.”
x + sublayer(x)), and in pre-norm the sublayer operates on norm(x) while the unnormalized x is carried forward. The point is a clean identity path for gradients and a representation each layer only has to correct, not rebuild.Interview prep
- 01“Why residual connections?” → identity path for gradients (no vanishing) + each sublayer learns a correction, not a full mapping. Add is after the sublayer.
- 02“Pre-norm vs post-norm?” → pre-norm keeps the residual gradient stable at init (warm-up optional, scales deep); post-norm needs warm-up and is hard past ~24 layers.
- 03“Why LayerNorm not BatchNorm?” → LayerNorm is per-token across features, so it’s independent of batch size/seq length; BatchNorm’s batch stats break on variable-length sequences.
- 04“What does RMSNorm save vs LayerNorm?” → drops mean-centering and the β bias — one fewer mean, subtraction, and bias per norm site; comparable quality.
- 05“What does the FFN do and how big is it?” → per-token MLP, 4× inner expansion; it’s where most params/FLOPs live; SwiGLU is the modern gated variant.
- 06“What’s in a modern block vs the 2017 one?” → pre-norm, RMSNorm, RoPE, SwiGLU, bias-free layers — the crystallized stack.
- 07“Why does stacking depth help?” → hierarchical features on the residual stream; only viable because residuals + pre-norm keep gradients healthy through depth.
- 08“Remove residuals at depth 96 — what happens?” → gradients vanish/explode and the model fails to train; the identity path is what makes deep nets trainable.
Common mistake
The red-flag answer: treating pre-norm vs post-norm as a minor, interchangeable tweak.
Checkpoint
You scale a from-scratch decoder from 12 to 48 layers using post-norm and no warm-up; loss spikes and diverges in the first few hundred steps. The fastest principled fix?
Checkpoint
A reviewer asks why your LLM uses LayerNorm/RMSNorm rather than BatchNorm, given BatchNorm’s success in CNNs. Best answer?
Checkpoint
You’re profiling a Llama-style model and find most FLOPs are not in attention. Where are they, and what’s the modern variant of that component?
Checkpoint
An interviewer says: “post-norm sometimes reaches lower final loss, so why does everyone ship pre-norm?” Strongest response?
Checkpoint
In a pre-norm block, which tensor is carried forward on the residual stream to the next sublayer?
Could you justify pre-norm + RMSNorm, explain the residual gradient path, and name the modern block end-to-end?
Takeaways
- A block is x = x + attn(norm(x)); x = x + ffn(norm(x)) — the unnormalized residual stream is the gradient highway.
- Pre-norm keeps early-layer gradients stable at init (warm-up optional, scales to 80+ layers); post-norm needs warm-up and caps depth.
- LayerNorm/RMSNorm are per-token (not per-batch) — that’s why neither is BatchNorm; RMSNorm drops mean-centering + β for cheaper compute.
- The FFN (4× inner, SwiGLU in modern LLMs) is where most params/FLOPs live — attention mixes across tokens, the FFN thinks per token.
- Depth works only because of the wiring: residuals + pre-norm keep gradients healthy through dozens of layers.
Next: the inference lifecycle — prefill, decode, the KV cache and its linear memory growth, and sampling.
Sources
Free to read · better with Enzo
Learn it with Enzo
Save your progress, answer the checkpoints, and let Enzo quiz you on what you just read.