ml.lab
Python sleeps until you run code
09 From LSTMs to today's models

Lesson 2 of 5

Inside a transformer

A transformer cuts text into tokens, turns each token into a vector, sends the vectors through a stack of identical blocks of attention and small networks, and turns the last vector into a probability for every possible next token. This lesson follows one sentence through every step, with the shapes and the numbers.

About 75 minutes
By the end you can
  • Follow a sentence from characters to token ids to the matrix X\mathbf{X} that a transformer reads.

  • Explain why a transformer needs position information and how it is supplied.

  • Compute multi-head attention with its shapes, and say why a layer has several heads.

  • Write one pre-norm transformer block, count its parameters, and implement it in numpy.

  • Trace how the last position's logits become the next-token distribution, and how generation reuses cached keys and values.

Type "Transformers read tokens." into a chat model. Before any attention can happen, those 25 characters have to become numbers. After the last layer, numbers have to become a probability for every token the model knows. In between sits a stack of identical blocks, and in each block, the attention you built in the previous lesson is the one step where tokens exchange information.

This lesson assembles everything around attention, in the order the data flows: text to tokens, tokens to vectors with positions, several attention heads at once, the block that wraps them, the stack and its output, and finally the loop that generates text. Four step-by-step diagrams draw the path with real numbers, and at the end you write a whole block in numpy.

From text to tokens

A network computes with numbers, so the first job is to cut text into pieces and give each piece an integer. The choice of pieces is a trade. Single characters, as in your models of module 7 and module 8, need only a small vocabulary but make sequences long: every letter is a step. Whole words make sequences short but need an enormous vocabulary, and still cannot handle a word never seen before. Subword tokens sit in between.

Many current tokenizers learn their pieces from data with byte-pair encoding. In the byte-level form many models use, it starts from the 256 possible byte values as the first pieces. Then it counts every pair of neighboring pieces in a large sample of text, merges the most frequent pair into a new piece, and repeats, until the vocabulary reaches its target size, often tens of thousands to a few hundred thousand pieces. Common words end up as single tokens, rare words split into several, and any string at all can still be written as bytes. In ordinary English prose a token averages roughly four characters, depending on the tokenizer.

Each token's id then selects one row of an embedding table E\mathbf{E}, with one row of dd learned numbers per id. This is the lookup from module 7, turned on its side. There, a matrix times the one-hot column ek\mathbf{e}_k picked out column kk. Here tokens are rows, so the one-hot row ek⊤\mathbf{e}_k^\top times E\mathbf{E} picks out row kk of E\mathbf{E}, written E[ids] in numpy for all the ids at once. Stacking the rows in the order the tokens appear gives the matrix X\mathbf{X} of shape (T,d)(T,\allowbreak d), one row per token.

The diagram follows "Transformers read tokens." through those steps with a real tokenizer, cl100k_base (the tokenizer of GPT-3.5 and GPT-4). The strip at the top shows where each step is on the way from text to numbers, and the lime parts are what the step adds.

To recap: text becomes pieces, pieces become ids, ids select rows of E\mathbf{E}, and the rows stack into X\mathbf{X}. One thing is still missing from X\mathbf{X}: where each token sits.

Where each token sits: positions

Why does a transformer need to be told the positions? Because attention, on its own, cannot tell "dog bites man" from "man bites dog". Here is the argument in three moves, for self-attention without the causal mask.

  1. Each score qt⋅ks\mathbf{q}_t\cdot\mathbf{k}_s depends only on the two token vectors involved, not on where they sit in the sentence.
  2. Swap two tokens in the input. Their rows in Q\mathbf{Q}, K\mathbf{K} and V\mathbf{V} swap, so every score is the same number as before, moved to a new place in S\mathbf{S}.
  3. So each token computes exactly the same weighted average as before, and the output rows come out swapped the same way as the input rows. Nothing in the output records which word came first.

Run the check: the same attention on three token vectors, then on the same vectors with the first and last swapped, then again after adding a position vector to each row.

⌘+Enter runs · edit freelyPython sleeps until you run code

Without positions, the output for "bites" is identical in both sentences. With a position vector added to each row, the words move while the positions stay, so the two sentences give different results. (The causal mask reveals a little order on its own, since each position sees a different number of tokens, but most transformers are still given positions explicitly.)

There are three common ways to do it:

  • Learned position vectors. A second table P\mathbf{P} with one learned row per position, added to the token rows: X=E[ids]+P[0:T]\mathbf{X} = \mathbf{E}[\text{ids}] + \mathbf{P}[0{:}T]. GPT-2 learns 1,024 of them, so it can read at most 1,024 tokens.
  • Fixed sinusoidal vectors. The original 2017 transformer added fixed patterns of sines and cosines at many different frequencies, so every position gets a distinct vector without any training.
  • Rotary embeddings. Many recent models add nothing to X\mathbf{X}. Instead they rotate each pair of coordinates in every query and key by an angle proportional to the token's position. Rotating both vectors changes their dot product only through the difference of the two angles, so position enters each score only as how far apart the two tokens are.

The next question checks the argument above.

Quick checkDog bites man

Self-attention with no causal mask and no position information reads "dog bites man", then "man bites dog": the same three token vectors, with the first and last swapped. How do the two outputs compare?

Choose one answer, then check.

Several heads at once: multi-head attention

One attention pattern per layer is limiting. Each row of the attention weights A\mathbf{A} is one set of weights that adds up to 1, so each token gathers one weighted average. But a token may need the previous word for grammar and a name twenty tokens back for meaning, two different searches. Multi-head attention runs several attentions side by side, so each can search for something different.

With hh heads and dd numbers per token:

  1. Project, as before: Q=XWQ\mathbf{Q} = \mathbf{X}\mathbf{W}_Q, K=XWK\mathbf{K} = \mathbf{X}\mathbf{W}_K and V=XWV\mathbf{V} = \mathbf{X}\mathbf{W}_V, each (T,d)(T,\allowbreak d), from three (d,d)(d,\allowbreak d) matrices.
  2. Split the columns into hh blocks of width dk=d/hd_k = d/h. Head 1 uses the first dkd_k columns of Q\mathbf{Q}, K\mathbf{K} and V\mathbf{V}, head 2 the next dkd_k, and so on. Every head still reads all dd numbers of every token: each of its columns of WQ\mathbf{W}_Q has dd entries, one per input number. Only the outputs are split.
  3. Attend in each head separately, with its own scores divided by dk\sqrt{d_k}, the causal mask and a row softmax. Head ii gives an output Oi\mathbf{O}_i of shape (T,dk)(T,\allowbreak d_k).
  4. Concatenate: place the head outputs side by side, [O1∣⋯∣Oh][\mathbf{O}_1 \mid \cdots \mid \mathbf{O}_h], of shape (T,d)(T,\allowbreak d) again.
  5. Project the result with one more learned matrix WO\mathbf{W}_O of shape (d,d)(d,\allowbreak d), which mixes what the heads found.

The cost is about the same as one head of full width: the projections are the same size, and hh score matrices of width d/hd/h take h⋅T⋅(d/h)⋅T=T2dh \cdot T \cdot (d/h) \cdot T = T^2 d multiply-adds, the same as one of width dd.

The diagram runs two heads on "the cat sat" with d=4d = 4, so each head has width 2. Head 1's columns are blue and head 2's are coral throughout.

The two patterns in the diagram were set by hand to be easy to see. In trained models, some heads do track something recognizable, such as the previous token; many do not have a tidy description.

The transformer block

Attention moves information between tokens, and that is all it does. A block wraps it with three more pieces: a small network that works on each token separately, and two devices that keep a deep stack trainable, residual connections and layer normalization. In the arrangement most current language models use, called pre-norm, one block makes four moves.

  1. Normalize each token's vector with layer normalization: subtract the vector's own mean, divide by its own standard deviation, then multiply each entry by a learned gain and add a learned bias. You wrote the first part as normalize_rows in How this lab works. It keeps the numbers at a steady scale however deep the stack. (Many recent models use a variant, RMSNorm, that skips subtracting the mean.)
  2. Attend: multi-head causal self-attention on the normalized vectors, then add the result back to the block's input. This addition is a residual connection.
  3. Normalize again.
  4. Transform each token on its own with a two-layer MLP, d→4d→dd \to 4d \to d, with a nonlinearity in between (ReLU in the original paper, GELU in GPT-2 and many later models, gated variants in many current ones), and add the result back again.

In symbols, with LN⁡\operatorname{LN} for layer normalization and each token a row:

H=X+Attn⁡(LN⁡(X)),Y=H+MLP⁡(LN⁡(H)).\mathbf{H} = \mathbf{X} + \operatorname{Attn}\big(\operatorname{LN}(\mathbf{X})\big), \qquad \mathbf{Y} = \mathbf{H} + \operatorname{MLP}\big(\operatorname{LN}(\mathbf{H})\big).

Attention is the only step where tokens exchange information. Layer norm uses each row's own mean and spread, the MLP applies the same weights to each row separately, and each residual addition adds row tt of a sublayer's output to row tt of its input. The original 2017 transformer normalized after each addition instead of before each sublayer; normalizing first turned out to train more stably in deep stacks, and most recent models do that.

Layer norm on one row. Take the row (−0.37,0.09,0.83,0.62)(-0.37,\allowbreak 0.09,\allowbreak 0.83,\allowbreak 0.62), which is the vector for "the" in the diagram below, and use a gain of 1 and a bias of 0.

  • Mean: (−0.37+0.09+0.83+0.62)/4=1.17/4=0.2925(-0.37 + 0.09 + 0.83 + 0.62)/4 = 1.17/4 = 0.2925.
  • Subtract it: (−0.6625,−0.2025,0.5375,0.3275)(-0.6625,\allowbreak -0.2025,\allowbreak 0.5375,\allowbreak 0.3275).
  • Square those: (0.4389,0.0410,0.2889,0.1073)(0.4389,\allowbreak 0.0410,\allowbreak 0.2889,\allowbreak 0.1073), whose average is 0.8761/4=0.21900.8761/4 = 0.2190. That is the variance.
  • Standard deviation: 0.2190≈0.468\sqrt{0.2190} \approx 0.468.
  • Divide: (−0.6625,−0.2025,0.5375,0.3275)/0.468≈(−1.42,−0.43,1.15,0.70)(-0.6625,\allowbreak -0.2025,\allowbreak 0.5375,\allowbreak 0.3275)/0.468 \approx (-1.42,\allowbreak -0.43,\allowbreak 1.15,\allowbreak 0.70).

The result has mean 0 and standard deviation 1, whatever the scale of the row that came in.

Counting the parameters. WQ\mathbf{W}_Q, WK\mathbf{W}_K, WV\mathbf{W}_V and WO\mathbf{W}_O are d×dd \times d each (all heads together), and the MLP has two matrices, d×4dd \times 4d and 4d×d4d \times d. That is 4d2+8d2=12 d24d^2 + 8d^2 = 12\,d^2 numbers, plus small bias and normalization vectors. Two thirds of every block is the MLP.

The diagram follows "the cat sat" through a tiny trained transformer: 2 blocks, d=4d = 4, 2 heads, an MLP of width 16 and a vocabulary of 9 words, trained on two sentences. The left side is the path, with the shape written beside every arrow; the lime stage is the one the step computes, and the panel on the right holds its numbers (rounded to 1 decimal, probabilities to 3). The text under the figure, and the walkthrough below, give the key rows to 2 decimals. The steps continue past one block to the stack and the output, which the next section explains.

To recap: a block normalizes, attends and adds, then normalizes, transforms each token and adds. The shape in equals the shape out, (T,d)(T,\allowbreak d), which is what lets blocks stack.

The identity path

Why add the input back at all? Look at a residual connection as a function: y=x+F(x)\mathbf{y} = \mathbf{x} + F(\mathbf{x}), where FF is a whole sublayer, its layer norm included. Its Jacobian, the matrix of partial derivatives from module 5, is

∂y∂x=I+JF,\frac{\partial \mathbf{y}}{\partial \mathbf{x}} = \mathbf{I} + \mathbf{J}_F ,

because the derivative of x\mathbf{x} with respect to itself is the identity matrix I\mathbf{I}, and the derivative of F(x)F(\mathbf{x}) is its Jacobian JF\mathbf{J}_F. Backpropagation through a stack multiplies one such matrix per sublayer. For two of them the product expands to

(I+J2)(I+J1)=I+J1+J2+J2J1.(\mathbf{I} + \mathbf{J}_2)(\mathbf{I} + \mathbf{J}_1) = \mathbf{I} + \mathbf{J}_1 + \mathbf{J}_2 + \mathbf{J}_2\mathbf{J}_1 .
Go slower: Expanding the product

Multiply out the brackets, each term of the first by each term of the second, keeping the order of the factors, since matrix products do not commute: (I+J2)(I+J1)=I I+I J1+J2 I+J2J1.(\mathbf{I} + \mathbf{J}_2)(\mathbf{I} + \mathbf{J}_1) = \mathbf{I}\,\mathbf{I} + \mathbf{I}\,\mathbf{J}_1 + \mathbf{J}_2\,\mathbf{I} + \mathbf{J}_2\mathbf{J}_1 . Multiplying by the identity changes nothing, so I I=I\mathbf{I}\,\mathbf{I} = \mathbf{I}, I J1=J1\mathbf{I}\,\mathbf{J}_1 = \mathbf{J}_1 and J2 I=J2\mathbf{J}_2\,\mathbf{I} = \mathbf{J}_2: =I+J1+J2+J2J1.= \mathbf{I} + \mathbf{J}_1 + \mathbf{J}_2 + \mathbf{J}_2\mathbf{J}_1 . With nn sublayers the same expansion gives I\mathbf{I}, plus every single Ji\mathbf{J}_i, plus every product of two, and so on up to the product of all nn.

The first term is the identity: there is always a path along which the gradient arrives unchanged, however deep the stack, and the other terms are added to it rather than multiplied into it. Compare a stack without residuals, whose Jacobian is the bare product Jn⋯J2J1\mathbf{J}_n \cdots \mathbf{J}_2\mathbf{J}_1: that is the long product which shrinks or grows geometrically, as in backpropagation through time. This is the idea you met in the LSTM's cell state: along its direct path ∂ct/∂ct−1=diag⁡(ft)\partial\mathbf{c}_t/\partial\mathbf{c}_{t-1} = \operatorname{diag}(\mathbf{f}_t), as Why gradients survive derived, and a forget gate near 1 lets gradients cross many steps. The LSTM adds through time; the residual connections add through depth.

A broad stream of light flows left to right through a row of blocks; inside each block two thin branches, one after the other, split off the stream, each passes through a narrow gate and a small machine, and each flows back into the stream a little further on, while the stream itself is never cutThe residual stream. Each sublayer reads a normalized copy of the stream and adds its result back; nothing ever replaces the stream, so information and gradients can travel the whole depth on it.

Researchers call the running sum that flows through the stack the residual stream: every sublayer reads from it and writes an addition to it.

The next two exercises check the block's parts and shapes. First, which part does what.

Quick checkWhich part moves information?

In the diagram's tiny model, the prediction at the position of "sat" changes from "on" to "by" when "cat" is replaced by "dog". Which part of a transformer block carries information from the token "cat" to the position of "sat"?

Choose one answer, then check.

Now trace every shape through one block of a real size, and count its parameters.

On paperShapes and parameters of one block

The base model of the original 2017 transformer used d=512d = 512, h=8h = 8 heads and an MLP of width 2,048. Take one pre-norm block of that size reading T=10T = 10 tokens.

  1. Write the shape (rows × columns) of each of these: X\mathbf{X}; LN⁡(X)\operatorname{LN}(\mathbf{X}); Q\mathbf{Q} for all heads together; one head's queries Qi\mathbf{Q}_i; one head's score matrix; one head's output Oi\mathbf{O}_i; the heads side by side; the attention output after WO\mathbf{W}_O; H\mathbf{H}; the MLP's hidden layer; the block output Y\mathbf{Y}.
  2. Count the numbers in the block's six weight matrices (WQ\mathbf{W}_Q, WK\mathbf{W}_K, WV\mathbf{W}_V, WO\mathbf{W}_O, W1\mathbf{W}_1, W2\mathbf{W}_2), ignoring biases and layer-norm parameters. Compare with 12d212d^2.
  3. How many numbers are in the weight matrices of a stack of 6 such blocks?
  4. After the last block and the output layer, with a vocabulary of V=32,000V = 32{,}000, what shape are the logits?

Work it on real paper: writing each step is the point. Then check your final answer here and compare your working with the walk-through.

Weight-matrix parameters in one block, then in all 6 blocks

One entry per box, top to bottom. 0.25, -2, 3/4 and sqrt(2) all work. Enter moves to the next empty box and checks once all are filled.

Then build the block. The exercise gives you the reference attention from the previous lesson and the diagram's tiny model, so your block can reproduce step 6 of the diagram exactly.

Code itOne transformer block

Build one pre-norm transformer block in numpy, rows as tokens. The setup gives you attention(Q, K, V, mask) and causal_mask(T), reference versions of what you wrote in the attention lesson, and TINY_X and TINY_BLOCK, the input and the weights of block 1 in the lesson's diagram "Inside a transformer".

  1. layer_norm(X, gain, bias, eps=1e-5): normalize every row to mean 0 and standard deviation 1, dividing the variance by dd (as np.var does) and adding eps inside the square root, then multiply by gain and add bias, both of shape (d,)(d,\allowbreak ).
  2. multi_head_attention(X, Wq, Wk, Wv, Wo, n_heads): causal multi-head self-attention. Head hh uses columns h⋅dhh \cdot d_h to (h+1)⋅dh−1(h + 1) \cdot d_h - 1 of Q\mathbf{Q}, K\mathbf{K} and V\mathbf{V}, with dh=d/nheadsd_h = d/n_{\text{heads}}. Return the output (T,d)(T,\allowbreak d) and the weights (nheads,T,T)(n_{\text{heads}},\allowbreak T,\allowbreak T).
  3. mlp(X, W1, b1, W2, b2): ReLU⁡(XW1+b1) W2+b2\operatorname{ReLU}(\mathbf{X}\mathbf{W}_1 + \mathbf{b}_1)\,\mathbf{W}_2 + \mathbf{b}_2, row by row.
  4. transformer_block(X, p, n_heads): H=X+MHA⁡(LN⁡1(X))\mathbf{H} = \mathbf{X} + \operatorname{MHA}(\operatorname{LN}_1(\mathbf{X})), then Y=H+MLP⁡(LN⁡2(H))\mathbf{Y} = \mathbf{H} + \operatorname{MLP}(\operatorname{LN}_2(\mathbf{H})), with the weights in the dictionary p (keys ln1_g, ln1_b, Wq, Wk, Wv, Wo, ln2_g, ln2_b, W1, b1, W2, b2).

The tests check each piece, reproduce step 6 of the diagram for "the cat sat", check that the block returns X\mathbf{X} unchanged when both sublayers output zero (the identity path), and check that later tokens never change earlier outputs.

⌘+Enter runsPython sleeps until you run code
Write your code where the starter says raise NotImplementedError, then press Run tests. Each check says what it expects.

Stack, then predict

One block lets every token gather information from earlier tokens once. Stacking blocks lets later blocks build on what earlier ones gathered: a token can collect something from its neighbor in one block and pass it on in the next. Because a block's output has the shape of its input, (T,d)(T,\allowbreak d), blocks stack without any glue. The diagram's tiny model has L=2L = 2 blocks; GPT-2's smallest model has 12 with d=768d = 768.

After the last block come two more steps, which are steps 8 and 9 of the diagram above.

  1. A final layer norm, then a linear layer Wout\mathbf{W}_{\text{out}} of shape (d,V)(d,\allowbreak V), where VV is the vocabulary size. It turns every row into VV logits, one score per token in the vocabulary: (T,d)(d,V)=(T,V)(T,\allowbreak d)(d,\allowbreak V) = (T,\allowbreak V). Some models reuse the embedding table as this layer (GPT-2 does), scoring each token by the dot product of its embedding with the final vector.
  2. A softmax turns a row of logits into probabilities that add up to 1, the softmax of module 6. Row tt is the model's prediction of the token after position tt.

In training, every row is used: the loss is the average cross-entropy of the true next token over all TT positions, the objective of module 7, and the causal mask makes that honest. To generate, only the last row matters. In the diagram, the last row after "the cat sat" gives "on" 0.840. Given "the dog sat" instead, the same model gives "by" 0.890, as in its training sentence "the dog sat by the door .": attention carried the second word's identity to the last position.

Where the parameters live. Add up a whole model: LL blocks of 12d212d^2, plus the embedding table (V×dV \times d), plus the output layer (d×Vd \times V, unless it is shared), plus learned positions if the model has them. Take d=4096d = 4096, 32 blocks and a vocabulary of 32,000:

  • One block: 12×40962=12×16,777,216=201,326,59212 \times 4096^2 = 12 \times 16{,}777{,}216 = 201{,}326{,}592.
  • All 32 blocks: 32×201,326,592=6,442,450,94432 \times 201{,}326{,}592 = 6{,}442{,}450{,}944, about 6.4 billion.
  • Embedding table and output layer: 2×32,000×4,096=262,144,0002 \times 32{,}000 \times 4{,}096 = 262{,}144{,}000.
  • Total: about 6.7 billion, the size sold as "7B".

About 96% of it sits in the blocks, and two thirds of each block is its MLP. Layers and parameter counts counted a published 7B design entry by entry and found 6,738,415,616. The difference, about 34 million, is its gated MLP (three 4096×110084096 \times 11008 matrices, a little more than 8d28d^2) and its normalization weights.

Small models look different, as the next exercise shows.

Work it outCounting GPT-2's smallest model

GPT-2's smallest model has d=768d = 768, 12 blocks, a vocabulary of 50,257 tokens and 1,024 learned position vectors, and its output layer reuses the embedding table (there is no separate output matrix). Estimate its parameter count as 12d212d^2 per block, plus the embedding table and the position table. Give the answer in millions, to one decimal place.

millions of parameters

Type a number: 0.25, -2, 3/4 and sqrt(2) all work. Enter checks.

Generating text, and the KV cache

A trained model gives a probability for every possible next token. To write text it has to choose one, append it, and ask again. Generating text introduced greedy choice, sampling and temperature for your character models, with the diagram below; a transformer generates the same way. As a reminder, the diagram runs the loop with a tiny character model continuing the prompt "th", with made-up logits and temperature τ=0.8\tau = 0.8.

The KV cache. Taken literally, the loop runs the whole transformer again on the whole text for every new token. Most of that work would repeat itself, because of one fact about the causal mask:

  1. In the first block, the key and value of token ss come from token ss's own row of X\mathbf{X} only (its embedding plus its position).
  2. The block's output row ss depends only on rows 11 to ss of its input: the causal mask stops attention from reading later rows, and layer norm, the MLP and the residual additions work row by row.
  3. So the input to the second block, row ss, depends only on tokens 11 to ss, and the same holds in every later block.
  4. So appending a new token changes nothing in the rows of earlier tokens, in any layer. Their keys and values stay exactly the same.

An inference server therefore computes each token's keys and values once, in every layer, and keeps them in memory: the KV cache. For each new token it computes only that token's own query, key and value, compares the query with the cached keys, and averages the cached values: one row of work per layer instead of the whole matrix.

The cache has a size. Per token it holds one key vector and one value vector per layer, each of dd numbers counting all heads together: 2×d×L2 \times d \times L numbers. For d=4096d = 4096 and 32 layers that is 2×4096×32=262,1442 \times 4096 \times 32 = 262{,}144 numbers per token. At 2 bytes per number, a context of 32,768 tokens needs 262,144×32,768×2=17,179,869,184262{,}144 \times 32{,}768 \times 2 = 17{,}179{,}869{,}184 bytes, about 17.2 GB, more than the 12.9 GB taken by the 6.4 billion block weights at the same 2 bytes each. The cache grows linearly with the context, where a recurrent network carries a state of fixed size.

The next exercise sizes a cache for a smaller model.

Work it outSize a KV cache

A smaller model has 24 layers and d=2048d = 2048 (all heads together), keeps a key and a value for every head, and stores its cache in 16-bit numbers (2 bytes each). How large is its KV cache for a context of 8,192 tokens? Answer in gigabytes (11 GB =109= 10^9 bytes), to 2 decimal places.

GB

Type a number: 0.25, -2, 3/4 and sqrt(2) all work. Enter checks.

Finally, put the whole path into words.

In your own wordsFrom a sentence to the next token

Explain to a colleague, step by step, how a transformer language model turns the text "Transformers read tokens." into a probability for the next token. Give the shapes where they help, and say which part does what.

A few sentences first: 0 of 60 characters.

Saved in this browser as you type.

Next: Pretraining and fine-tuning