Chapter 3: Transformer Architecture
Dive deep into the architecture that powers every modern LLM. Understand self-attention---the mechanism that allows tokens to attend to each other regardless of distance---multi-head attention for capturing different relationship types, feed-forward networks, layer normalization, residual connections, and how they combine into a full transformer block.
Chapter Overview
The Transformer architecture (Vaswani et al., 2017) is the foundation of all modern LLMs. Its key innovation---self-attention---allows every token to directly interact with every other token in a sequence, eliminating the information bottleneck of recurrent architectures.
A transformer block consists of two main sub-layers: a multi-head self-attention mechanism and a position-wise feed-forward network. Each sub-layer is wrapped with a residual connection and layer normalization. Stacking dozens (or hundreds) of these blocks creates the deep networks we call LLMs.
Understanding the transformer at a mathematical level is essential for: debugging model behavior, implementing efficient inference, designing architectural improvements, and understanding why certain prompting strategies work. Each component has a clear purpose, and their interplay creates a system far more powerful than the sum of its parts.
This chapter covers:
- Self-Attention: The core mechanism that enables tokens to exchange information
- Multi-Head Attention: Running multiple attention patterns in parallel
- Feed-Forward Networks: Per-token nonlinear transformations
- Layer Normalization: Stabilizing training of deep networks
- Residual Connections: Enabling gradient flow through many layers
- Full Transformer Block: How all components fit together
Chapter Roadmap
Click any topic to jump in
Self-Attention
The core mechanism — every token computes weighted attention over all other tokens via queries, keys, and values.
Multi-head attention captures diverse patterns; FFN stores and transforms knowledge
Multi-Head Attention
Parallel attention heads that each learn different relationship types — syntax, semantics, position.
Feed-Forward Networks
Per-token nonlinear transformations that store factual knowledge and apply complex feature mappings.
Normalization and skip connections make deep stacking possible
Layer Normalization
Stabilizing deep network training by normalizing activations — LayerNorm, RMSNorm, Pre-Norm vs Post-Norm.
Residual Connections
Skip connections that enable gradient flow through 100+ layers by preserving the identity path.
Full Transformer Block
Assembling all components into the Pre-Norm block that gets stacked dozens of times in modern LLMs.
A word's meaning depends on other words, often far away. In "The trophy did not fit in the suitcase because it was too big", the word "it" refers to the trophy eight words earlier, and only the final word "big" settles it. A recurrent network must carry that clue forward step by step through one hidden vector, and the signal fades over long distances. It also cannot process the positions in parallel, which makes training slow on modern hardware.
The previous chapter, Tokenization & Embeddings, turned text into a sequence of vectors with positions attached. This topic covers the operation that lets every one of those vectors look directly at every other one in a single step.
We start with queries, keys, and values, the three roles each token plays. Then we build scaled dot-product attention and explain why the scores are divided by . Next comes the causal mask that stops a decoder from reading the future. We finish with the cost: attention compares every pair of tokens, so it grows with the square of sequence length.
Definition
Self-attention is a sequence operation in which every position computes a query, a key, and a value by linear projection. It scores its query against every key with scaled dot products, turns the scores into weights with a softmax, and outputs the weighted sum of the values. The result is a context-dependent representation for each position, computed for all positions in parallel.
In this topic
Queries, Keys, and Values
Self-attention gives each token three roles, each produced by a learned matrix. The query describes what token is looking for. The key describes what token offers to be matched against. The value is the content token passes along if it is chosen. Here is the input embedding and . Separating matching (keys) from content (values) lets a token be easy to find for one reason and contribute something different. The roles are learned, not assigned, so no head is guaranteed to learn a clean grammatical role.
The projection matrices transform each token into three roles. The attention score is a bilinear form that measures compatibility between tokens and through the learned metric . This is more expressive than simple dot-product similarity because the model learns what aspects of tokens should determine attention.
In "The cat sat on the mat", how does the word "sat" use Q, K, V?
Scaled Dot-Product Attention
The attention function runs in four steps. First compute all query-key dot products at once as , an score matrix for tokens. Then divide by , where is the key dimension. Then apply a softmax to each row, so each token's weights are positive and sum to 1. Finally multiply by to mix the values. The scaling matters because a dot product of two random -dimensional vectors with unit-variance entries has variance . Unscaled, large scores push the softmax into a near one-hot regime where its gradients almost vanish, and learning stalls.
Self-attention computes . The scaling prevents dot products from growing with dimension, keeping softmax gradients useful. Without scaling, if and are random vectors with entries , then has variance . For : std , and is nearly one-hot, producing vanishing gradients for all but the top-scoring key.
Why divide by ? What happens without scaling when ?
Causal (Masked) Attention
A decoder-only model such as GPT or Llama is trained to predict each next token, so position must not see positions after it. The causal mask enforces this by adding to every score where the key position is greater than the query position , as in the formula, before the softmax. Since , those weights become exactly zero and the remaining weights still sum to 1. One forward pass then trains every position at once without leaking the answer. Forgetting the mask is a classic bug: training loss collapses toward zero while generation is useless.
The causal mask ensures for , because . This enforces the autoregressive property: cannot depend on . The mask is the only difference between encoder (bidirectional) and decoder (causal) attention — the same attention mechanism with a different mask produces fundamentally different model capabilities.
In the sequence "I love cats", which tokens can "love" attend to?
Attention Complexity
Attention's cost comes from the score matrix. Every one of the queries is compared with all keys, and each comparison is a -dimensional dot product, so time grows as , as in the formula. A naive implementation also stores the matrix, so memory grows as as well. The feedforward layers, by contrast, are linear in . At short lengths the feedforward layers dominate, but at long lengths attention takes over. FlashAttention removes the quadratic memory by computing in tiles, but the quadratic compute remains. Sparse, sliding-window, and linear attention trade exactness for lower cost.
Computing the attention matrix requires multiply-adds. Storing this matrix takes floats. For (Llama's context): the matrix has entries, requiring GB in FP32. FlashAttention avoids materializing this matrix by computing attention in tiles that fit in SRAM, reducing memory from to while maintaining exact computation.
A model processes 8K tokens with . How does doubling context to 16K affect attention compute?
Theory Exercise
Problem:
Explain why self-attention can capture long-range dependencies that RNNs cannot, using a concrete example.
Hints:
- Think about the path length between distant tokens
- Consider the vanishing gradient problem in RNNs
- Think about the attention matrix as a shortcut
Related Problems on PixelBank
One attention pattern per token is a real limit. The word "she" may need to find the noun it refers to, the verb it is the subject of, and the words right next to it, all at once. A single softmax produces one set of weights, so it must blend these needs into one average, and a strong match for one relation drowns out the others. We want several independent attention patterns computed side by side, without multiplying the cost.
The previous topic, Self-Attention Mechanism, built one attention pattern from queries, keys, and values. This topic splits that computation into many heads and then looks at how inference constraints reshaped the design.
We start with the multi-head formula, where each head works in a smaller subspace and the results are joined by an output projection. Then we look at what trained heads actually learn and how many can be removed. Next comes Grouped Query Attention, which shares keys and values across heads to shrink the inference cache. We finish with Multi-Query Attention, the extreme case with a single key-value head.
Definition
Multi-head attention runs attention operations in parallel, each with its own query, key, and value projections into a -dimensional subspace. The heads' outputs are concatenated and mixed by an output matrix . Grouped-query and multi-query variants keep query heads but share fewer key-value heads among them, which shrinks the key-value cache used during generation.
In this topic
Multi-Head Attention Formula
Multi-head attention splits the model width into heads of size . Head has its own projections and computes ordinary scaled dot-product attention inside its own subspace. The outputs, each of width , are concatenated back to width and multiplied by , as in the formula, which lets heads exchange information. Because each head is times narrower, the total cost is about that of one full-width head. Too many heads make each subspace too small to express a useful match.
Multi-head attention concatenates heads: where each head operates on dimensions. The total parameter count is , independent of the number of heads. More heads means each head works in a lower-dimensional space but captures finer-grained attention patterns.
A model has and heads. What is the dimension per head?
What Different Heads Learn
Trained heads are not interchangeable. Studies of BERT and GPT-2 find heads that attend to the previous or next token, heads that track syntactic links such as a verb and its object, heads that resolve coreference, and heads that dump most of their weight on the first token or a separator when they have nothing useful to do. Induction heads copy patterns seen earlier in the context. The specialisation is emergent and messy: many heads mix several roles. Michel et al. (2019) showed that many heads can be pruned at test time with little loss, so capacity is unevenly spread.
Probing studies show attention heads specialize: some heads have entropy (attending to a single position), while others have (uniform attention). The specialization can be quantified by the attention pattern's mutual information with linguistic annotations: . Heads with high MI for specific dependency types (subject-verb, determiner-noun) are performing implicit parsing.
In "The doctor said she would prescribe medication", what might different heads attend to for the word "she"?
Grouped Query Attention (GQA)
During generation, a model caches the keys and values of every past token so they are not recomputed. That cache grows with layers, key-value heads, head size, and context length, and it often limits batch size more than the weights do. GQA (Ainslie et al., 2023) keeps query heads but only key-value heads, as in the formula, and each group of query heads shares one key-value pair. The cache shrinks by the factor . Quality stays close to full multi-head attention. Llama 2 70B and Mistral 7B use 8 key-value heads.
GQA reduces KV cache from to bytes, a factor of reduction. For Llama 2 70B with query heads and KV heads: 8x KV cache reduction. During attention, each query head group shares keys and values: . The key insight: different query heads in a group can still attend to different positions because they have different matrices.
A model has 32 query heads and 8 KV heads. How much KV cache memory is saved vs standard MHA?
Multi-Query Attention (MQA)
Multi-Query Attention (Shazeer, 2019) is the extreme of GQA: all query heads share a single key head and a single value head, so . The cache shrinks by the full factor , and each decoding step reads far less memory, which is the real bottleneck when generating one token at a time. The cost is quality, because every head must match against the same keys, and training can be less stable. Full multi-head attention has . GQA sits between the two extremes, which is why most recent open models chose it over MQA.
MQA is the extreme case of GQA with : all query heads share a single key and value head. KV cache is reduced by x (e.g., 32x for 32-head models). The quality degradation is bounded: since the shared K and V still span the full embedding space, information is not lost — only the diversity of key-value perspectives is reduced. Empirically, MQA loses quality while GQA ( or ) loses .
Rank MHA, GQA, MQA by: (1) quality, (2) inference speed, (3) KV cache size.
Theory Exercise
Problem:
A 7B parameter model uses standard MHA with 32 heads and dimension 4096. You want to convert it to GQA with 8 KV heads for faster inference. What changes in the architecture, and how would you initialize the new KV heads?
Hints:
- Think about which weight matrices change
- Consider how to map 32 KV heads to 8
- Think about initialization strategies for the converted model
Related Problems on PixelBank
Attention moves information between tokens, but it is mostly a weighted average. Its output for each token is a mix of value vectors, and a mix of vectors cannot compute a new feature that none of them contained. Attention alone also has no obvious place to store the large body of facts a language model seems to know. Something per token, nonlinear, and large is still missing from the block.
The previous topic, Multi-Head Attention, gathered context from many positions at once. This topic covers the feed-forward network that follows it in every block. It holds about two thirds of a standard transformer's parameters.
We start with the position-wise feed-forward network, which widens each token's vector about four times, applies a nonlinearity, and projects back. Then we compare the activation functions, from ReLU to GELU. Next comes SwiGLU, the gated variant used by Llama and PaLM, and the parameter arithmetic behind its odd hidden size. We finish with the view of the feed-forward layer as a key-value memory for facts.
Definition
A transformer feed-forward network is a two-layer perceptron applied independently to each position. It expands the -dimensional token vector to a hidden width , usually about , applies a nonlinearity such as GELU or a gated SwiGLU unit, and projects back to . The same weights are shared across positions, and no information moves between tokens inside it.
In this topic
Position-wise FFN
The feed-forward network applies the same small MLP to each token separately. lifts the vector into a wider space, typically , where many hidden units can each detect a different pattern. An activation function then zeroes or dampens the units that did not fire, and writes the result back at width . The biases are often dropped in modern LLMs. Because it is per token, every position runs in parallel. Its failure mode is its size, since it dominates the parameter and compute budget.
The FFN computes where expands and contracts. The expansion to creates a higher-dimensional space where the activation function can implement more complex decision boundaries. Each of the hidden neurons computes — a learned feature detector. The FFN has parameters, which is of a transformer layer's total.
Model with and . How many parameters in one FFN layer?
Activation Functions
The activation is what makes the feed-forward network more than one big linear map. ReLU, , was the original choice. It is cheap, but its gradient is exactly zero for negative inputs, so a unit pushed negative on every input stops learning. GELU, used by BERT and GPT-2, multiplies by , the probability that a standard normal variable is below . It follows ReLU for large positive or negative inputs but curves smoothly near zero and lets small negative values through. The formula shows the tanh approximation often used in code. Smooth gradients make optimisation of deep stacks easier.
GELU approximation: . At : , . Compare ReLU: is undefined. The smooth derivative of GELU means every neuron receives a gradient signal, unlike ReLU where negative-input neurons get exactly zero gradient. This reduces the dead neuron problem from (ReLU) to (GELU).
Why did transformers switch from ReLU to GELU?
SwiGLU Activation
SwiGLU (Shazeer, 2020) replaces the single up-projection with two. One branch, , passes through SiLU, , and acts as a soft gate. The other branch, , carries values, and the two are multiplied element by element before the down-projection, as in the formula. The gate lets each hidden unit decide, from the input, how much of its value to pass. That multiplicative interaction improved quality in Shazeer's experiments at matched parameter count. It is used in PaLM, Llama, and Mistral. The cost is a third weight matrix, so the hidden width is shrunk to compensate.
SwiGLU computes where . The gating mechanism allows each dimension to be independently modulated: if gate dimension outputs , that dimension is suppressed regardless of 's output. This gives the network binary-like "switches" in addition to the continuous transformation, effectively doubling the model's expressiveness per parameter.
SwiGLU has 3 projection matrices instead of 2. How does this affect parameter count?
FFN as Knowledge Storage
Geva et al. (2021) proposed reading the feed-forward network as a key-value memory. Each row of acts as a key: its dot product with the token vector measures how well the input matches a pattern. The activation keeps the matches that fire, and each matching unit adds its row of , a value vector, to the output. So each unit maps a detected pattern to a push toward certain next tokens. Meng et al. (2022) located factual associations in mid-layer feed-forward weights and edited single facts. Facts are spread over many units, so the picture is only approximate.
The FFN implements a key-value memory: rows are keys, columns are values. For input , the activation measures how well matches key . If the match is high, value is added to the output. This interpretation explains why knowledge editing works: changing a specific fact ("Eiffel Tower built in 1889") only requires modifying a few rows of (the detector) or columns of (the stored value).
A neuron in the FFN consistently activates when the input contains context about "capital cities". What does this suggest?
Theory Exercise
Problem:
A model has 32 layers, , and uses SwiGLU with . Calculate the total FFN parameters and compare to the attention parameters (32 heads, GQA with 8 KV heads).
Hints:
- SwiGLU has 3 projection matrices per FFN layer
- GQA attention has h_Q query projections and h_KV key/value projections
- Don't forget the output projection in attention
Related Problems on PixelBank
Stack dozens of layers and the scale of the activations drifts. Each layer adds its output to the running vector, so its size can grow layer after layer, and a slight change in early weights shifts what every later layer receives. Attention scores and activations are sensitive to scale, so training becomes fragile, loss spikes appear, and the usable learning rate shrinks. We need a cheap operation that resets each token's vector to a standard scale without depending on the rest of the batch.
The previous topic, Feed-Forward Networks, completed the two sublayers in each block. This topic covers the normalisation wrapped around both of them, and why its exact placement decides whether a deep model trains at all.
We start with LayerNorm, which standardises each token across its features. Then we compare Post-Norm, from the original Transformer, with Pre-Norm, used by GPT-2 and later models. Next comes RMSNorm, the cheaper variant in Llama and Mistral. We finish with why BatchNorm, standard in CNNs, does not suit language models.
Definition
Layer normalisation rescales each token's activation vector to zero mean and unit variance across its features, then applies a learned per-feature scale and shift . The statistics are computed per token, not per batch, so the operation behaves identically in training and inference and for any batch size. RMSNorm is a variant that divides by the root mean square without subtracting the mean.
In this topic
Layer Normalization
LayerNorm works on one token vector at a time. It computes the mean and variance over the features, subtracts , and divides by , where the small prevents division by zero. The result has mean 0 and variance 1. The learned vectors and then rescale and shift each feature, so the model can undo the normalisation where that helps. Every token gets the same treatment regardless of batch size or sequence length. A caveat is that one huge outlier feature inflates and squashes the rest of the vector.
LayerNorm computes where and . The Jacobian is , which projects gradients onto the hyperplane orthogonal to and . This prevents gradients from scaling with activation magnitude, stabilizing optimization.
A token embedding has values [2.0, 4.0, 6.0, 8.0]. Apply LayerNorm (ignore and ).
Pre-Norm vs Post-Norm
The original Transformer used Post-Norm: compute , then normalise the sum. So every path from the output back to the input passes through a normalisation at each layer, and Xiong et al. (2020) showed the gradients near the output are large at initialisation. Post-Norm therefore needs a careful learning-rate warmup and becomes fragile past a few dozen layers. Pre-Norm normalises only the sublayer's input, , and leaves the residual path untouched, which gives an identity route for gradients. GPT-2 and almost every later LLM use Pre-Norm, usually with a final norm before the output head.
In Pre-Norm, the residual path is . The gradient is . The identity term guarantees a gradient of at least 1 at every layer — the "gradient highway." In Post-Norm, the normalization sits on this highway, potentially bottlenecking gradient flow. This is why Pre-Norm enables training 100+ layer models without warmup.
Why did GPT-2 and all subsequent models switch from Post-Norm to Pre-Norm?
RMSNorm
RMSNorm (Zhang and Sennrich, 2019) drops two parts of LayerNorm: the mean subtraction and the shift . It divides each token vector by its root mean square, , and multiplies by the learned scale , as in the formula. The hypothesis is that re-scaling, not re-centring, gives normalisation its benefit. With one fewer reduction and no , it is simpler and faster, and the paper reports running-time reductions of 7 to 64 percent across models with comparable quality. It is used in Llama, Mistral, and most recent open models. It is not invariant to a constant shift.
RMSNorm computes where . This skips the mean subtraction of LayerNorm, saving subtractions and 1 mean computation per token. The gradient is simpler: , with one fewer projection term. Empirically, the mean-centering step contributes to model quality, making it a free optimization to remove.
Why is removing the mean subtraction step in RMSNorm acceptable?
Why Not BatchNorm?
BatchNorm normalises each feature using statistics across the examples in a batch, which works well for image CNNs. Language models break its assumptions. Sequences vary in length and are padded, so batch statistics mix real tokens with padding. Batches in LLM training are often small per device, which makes the statistics noisy. In causal models a batch statistic lets information leak between examples and positions. BatchNorm also needs running averages at inference that differ from training statistics. LayerNorm computes everything inside one token, so it avoids all of these problems at the cost of ignoring cross-example statistics.
BatchNorm computes statistics across the batch: and . With batch size (common during inference): , making normalization undefined. LayerNorm computes statistics across the feature dimensions within a single sample, so gives robust statistics even for batch size 1. Additionally, variable sequence lengths in NLP mean batch statistics are noisy and position-dependent.
During inference, you process a single sequence (batch size 1). Why does BatchNorm fail but LayerNorm works?
Theory Exercise
Problem:
A model is training unstably---loss spikes occur every few hundred steps. The model uses Post-Norm and standard LayerNorm. Propose three changes to improve training stability.
Hints:
- Think about the normalization placement
- Consider the normalization variant
- Think about other training hyperparameters
Making a plain network deeper eventually makes it worse, even on training data. The gradient reaching early layers is a product of one factor per layer, so it can shrink toward zero or blow up, and each layer must learn to pass on everything useful from the layer below. He et al. found in 2015 that a 56-layer plain CNN had higher training error than a 20-layer one. That is an optimisation failure, not overfitting. Transformers stack up to a hundred blocks, so they need a fix.
The previous topic, Layer Normalization, kept activation scales under control inside each block. This topic covers the other half of deep trainability: the skip connection that adds each block's input to its output.
We start with the residual formula and why learning a change is easier than learning a whole mapping. Then we follow the gradient and see how the identity path keeps it alive through dozens of layers. We finish with the residual stream view, in which every layer reads from and writes to one shared vector.
Definition
A residual connection computes a layer's output as its input plus a learned update, , instead of alone. The identity shortcut lets information and gradients pass through the layer unchanged, so the sublayer only needs to learn a correction. In a transformer every attention and feed-forward sublayer is wrapped this way.
In this topic
Residual Connection Formula
A plain layer must produce its whole output, . A residual layer produces only a change, , as in the formula, so learns the difference between output and input. The input and output must have the same width, which is why every transformer sublayer keeps dimension . If a layer has nothing useful to add, it can push toward zero and pass its input through unchanged. Deep networks therefore start close to a stack of identities, which is a safe place to begin. Without normalisation, the sum of many updates can grow large.
The residual means the layer learns the residual . If the optimal function is close to identity, only needs to learn a small correction (near-zero weights). Without residuals, the layer must learn the full mapping , which requires precisely canceling the input through nonlinear transformations — much harder to optimize.
Why is learning a residual easier than learning the full transformation?
Gradient Flow
Differentiate the residual formula and each layer contributes a factor to the backward product, as in the formula. Expanding that product gives a term that is just 1, a direct path from the loss to layer that skips every sublayer, plus terms that pass through some sublayers. In a plain network each factor is only , and a product of many factors below 1 vanishes. The residual form does not guarantee a large gradient, since the terms can cancel. But the identity path makes vanishing much less likely. Pre-Norm keeps that path clean.
Through residual layers, the gradient is . Expanding the product yields terms (one for each subset of layers). The identity path (selecting at every layer) contributes directly — a gradient that bypasses all intermediate layers. This ensures the gradient norm is at least , preventing vanishing gradients.
In a 96-layer transformer, what would happen to gradients without residual connections?
Residual Stream Interpretation
Elhage et al. (2021) at Anthropic proposed reading a transformer as one shared vector per token, the residual stream, that runs from the embedding to the output. Every attention head and every feed-forward layer reads the stream through a normalisation, computes something, and adds its result back. Nothing is overwritten, only added. So the final vector is the embedding plus the sum of every sublayer's output, and the unembedding reads the prediction off that sum. This view explains why individual components can be studied separately. The limitation is that the stream has a fixed width, so many features must share dimensions.
The residual stream view interprets the representation as : the final representation is the initial embedding plus the sum of all layer contributions. Each layer reads from the stream (via its input) and writes to the stream (via its output). Layer importance can be measured by — the fraction of the final representation contributed by layer . Prunable layers are those with small contributions.
How does the residual stream view explain why some layers can be removed without much quality loss?
Theory Exercise
Problem:
A researcher proposes using multiplicative gates instead of additive residual connections: , where is a learned gate. What are the advantages and disadvantages compared to standard residuals?
Hints:
- Think about what happens when g = 1 or g = 0
- Consider the gradient flow properties
- Think about the additional parameter cost
Related Problems on PixelBank
Knowing each part of a transformer is not the same as knowing the block. The order of operations decides whether training is stable, the widths decide where the parameters go, and the attention layout decides the memory cost of generation. Given only a model's width, depth, and vocabulary, you should be able to estimate its parameter count, and given its context length, the memory it needs to serve.
The previous topic, Residual Connections, supplied the last component: the skip path that makes deep stacks trainable. This topic puts attention, the feed-forward network, normalisation, and residuals together into the block that modern LLMs repeat 32 to 80 times or more.
We start with the Pre-Norm block and trace one token through its six steps. Then we count parameters and check the count against Llama 2 7B. Next comes the key-value cache, which makes generation fast and costs memory that grows with context. We finish with the choice between a deeper model and a wider one at the same parameter budget.
Definition
A transformer block is the repeated unit of a transformer. In the Pre-Norm form it normalises its input, applies multi-head self-attention, and adds the result to the input. It then normalises that sum, applies a position-wise feed-forward network, and adds that result as well. The output has the same shape as the input, so blocks stack into a deep model.
In this topic
Pre-Norm Transformer Block
The modern block has two residual sublayers, as in the formula. First, normalise the input with RMSNorm, run multi-head attention on it, and add the result back to get . This is the only step where tokens exchange information. Second, normalise , apply the feed-forward network to each token separately, and add the result to get the output. The residual stream itself is never normalised inside the block. Stacking blocks and adding a final norm and output head gives a GPT or Llama style model. Every intermediate vector keeps width , which is what makes the residual additions possible.
The Pre-Norm block computes , then . Each sub-layer sees normalized input (stable statistics) while the residual bypasses normalization (clean gradient flow). The two sub-layers are complementary: attention aggregates across tokens ( interactions) while FFN transforms each token independently ( per token). Together, they form a complete "read from context, then process" cycle.
Trace a single token through one transformer block. Input: embedding vector .
Parameter Count Breakdown
Count per layer. Attention has four matrices for queries, keys, values, and output, so with full multi-head attention. A SwiGLU feed-forward network has three matrices, , and each of the two RMSNorms adds only scale parameters. Over layers that is about . With it becomes , as in the formula. Add the embedding table, , and an untied output head of the same size. The estimate slightly overcounts GQA models, whose key and value matrices are smaller.
For a transformer with layers, dimensions, and SwiGLU FFN (): attention per layer is (Q, K, V, O projections), FFN per layer is (gate, up, down). Total per layer: . For Llama 2 7B (): parameters plus embeddings, totaling .
Llama 2 7B: , , , . Verify the ~7B parameter count.
KV Cache for Inference
Generation produces one token per step, and each new token must attend to all earlier ones. Recomputing every earlier key and value at each step would repeat almost all the work. The key-value cache stores each token's keys and values in every layer once computed, so each step only computes the new token's query, key, and value and reads the rest from memory. The size, per the formula, is 2 (keys and values) times layers, key-value heads, head size, sequence length, and bytes per number. It grows linearly with context and batch, and at long contexts it can exceed the weights.
Without KV cache: generating tokens requires computing attention for sequences of length , totaling operations. With KV cache: each new token only computes its Q and looks up cached K, V from previous tokens: per step, total. The cache stores values per token. For Llama 2 7B at 4096 context: GB.
Llama 2 7B generates 2048 tokens. How large is the KV cache in FP16?
Scaling Depth vs Width
For a fixed parameter budget you can trade layers for width . Depth adds sequential steps of computation, which can compose features into more abstract ones, but each layer adds latency because layers run one after another. Width makes each matrix multiply larger, which GPUs execute efficiently and which tensor parallelism can split. Very deep, narrow models are harder to train, and very wide, shallow ones waste capacity. Kaplan et al. (2020) found that loss depends only weakly on the exact shape over a broad range. Llama 2 70B uses and , so its ratio is about 100.
For a fixed parameter budget : doubling depth () while halving width () preserves but changes the architecture fundamentally. Deeper models have more sequential computation (longer gradient paths but more abstraction layers). Wider models have more parallel capacity (richer per-layer representations but fewer composition steps). Empirically, the optimal depth-to-width ratio for loss minimization at a given is approximately and .
Two models with ~7B parameters each: (A) 32 layers, d=4096; (B) 64 layers, d=2896. Which is likely better?
Theory Exercise
Problem:
You are designing a 13B parameter transformer for a deployment scenario where inference latency is critical (must generate tokens in <50ms each). What architectural choices would you prioritize?
Hints:
- Think about what determines per-token latency during generation
- Consider KV cache and memory bandwidth
- Think about which components can be parallelized