rumblr Work in progressWIP

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

Attention

how a token decides what to listen to

This lesson covers Queries, keys, values, softmax, masking, multi-head, GQA, O(n²)

Free lesson 31 min8 figures and diagrams10 interactive
How it works builds the idea from scratch. Math & code adds the formulas and the Python.

At a glance

Key takeaways

  1. Attention: each token builds a query, key and value. Query-key similarity decides how much each token listens to each other token, and the output is a weighted blend of their values.
  2. Why √d_k: dot products grow with dimension, and large scores saturate softmax and kill gradients. Scaling keeps training stable.
  3. Causal mask: future scores are set to −∞ before softmax, so each token sees only the past. That is what makes next-token training honest and generation cacheable.
  4. Long context cost: attention compares every pair of tokens, so it is O(n²), and the KV cache grows with every token.

Level 2

How it works, from scratch

Level 2 builds the mechanism from nothing, starting with a glance around a room.

The everyday picture. Imagine you're in a meeting and someone says "it's broken, can you fix it?" To know what "it" means, you glance around the room: at the laptop on the table, at the person who just walked in, at the whiteboard. You pay a lot of attention to the laptop, a little to everything else, and your understanding of "it" becomes mostly laptop.

That glance is attention. Every word in a sentence gets to look at every other word, decide how relevant each one is, and rebuild its own meaning as a mix of the relevant ones. "Bank" next to "river" becomes a riverbank; next to "loan" it becomes a lender. The word is the same, but what it paid attention to is different.

Chapter 1

A tiny worked example: what does "it" refer to?

Scroll through the whole calculation first; the lesson then decodes each step.

Take "The animal didn't cross the street because it was tired" and follow the word "it". To keep the numbers small, it compares itself with just three other words. Each comparison gives a relevance score (how the score is made comes next). Then we turn scores into shares of attention that add up to 1:

  1. Exponentiate each score (e^score), which makes every number positive and stretches the gaps between them.
  2. Divide each by the total, 7.39 + 2.72 + 1.65 = 11.76.
Word Score e^score Share of attention
animal 2.0 7.39 7.39 / 11.76 = 0.63
tired 1.0 2.72 2.72 / 11.76 = 0.23
street 0.5 1.65 1.65 / 11.76 = 0.14

Those two steps are called softmax. The new meaning of "it" is 63% "animal", 23% "tired" and 14% "street". Without being told any grammar rule, the model has worked out that "it" is the animal.

In code: worked_example_it runs these three scores through softmax and returns each word's share.

Softmax, decoded

Softmax turns any list of numbers, positive or negative, into shares that are all positive and add up to exactly 1. Think of it as a vote where louder voices get disproportionately more say.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the list of scores going in (2.0, 1.0, 0.5)
how many scores there are 3
the position we're computing a share for 1 = "animal"
the score at position
Euler's number, ≈ 2.718; is "2.718 multiplied by itself times" (works for fractions and negatives too)
"add up the following, for = 1, 2, …, "
a counter that walks over every position 1, 2, 3
the share of attention position gets 0.63

In words: "the share for item i is e raised to its score, divided by the sum of e raised to every score."

With the numbers: softmax(2.0, 1.0, 0.5)₁ = 7.39 / (7.39 + 2.72 + 1.65) = 7.39 / 11.76 = 0.63.

Level 3: in Python
import math
# animal, tired, street
z = [2.0, 1.0, 0.5]
# e^(z_i) for each score
exps = [math.exp(z_i) for z_i in z]
[round(e, 2) for e in exps]  # → [7.39, 2.72, 1.65]
# Σ_j e^(z_j)
total = sum(exps)
round(total, 2)  # → 11.76
# softmax(z)_i: each share of the total
[round(e / total, 2) for e in exps]  # → [0.63, 0.23, 0.14]

Why e to the power of the score, and not just the score divided by the total? Three reasons you can check on the table above:

  1. Always positive. Scores can be negative; e to any power is positive, so no word ever gets a negative share.
  2. Order is kept. A higher score always gets a bigger share.
  3. Gaps are stretched. Adding 1 to a score multiplies its e-value by 2.72, so the winner pulls ahead decisively. (Plain division would give "animal" 2.0 / 3.5 = 0.57, a timid 57%, and would break on negative scores.)

One more property matters later: adding the same number to every score changes nothing, because it multiplies the top and bottom of the fraction by the same amount. The code uses this to subtract the largest score before exponentiating, which stops e¹⁰⁰⁰ from overflowing to infinity. See primer.notation for exponents and sums from scratch.

Figure 1 · Drawn from the lesson's code

animal tired street 0.0 0.1 0.2 0.3 0.4 0.5 0.6 0.7 share What "it" attends to: softmax sharpens the scores 0.63 0.23 0.14 score (share of total) attention weight (softmax)

For the word it, scores 2, 1 and 0.5 become weights 0.63, 0.23 and 0.14: animal scores twice tired yet gets almost three times its attention

Reading it: for each word, the grey bar is its share of the raw scores and the blue bar is its share of attention after softmax. Softmax exaggerates differences: "animal" scores only twice as high as "tired", yet it ends up with almost three times the attention, because exponentiating stretches the gaps. That is how attention commits to the most relevant word while keeping a little of the others.

In code: softmax subtracts the largest score before exponentiating, and turns masked scores of −∞ into weights of exactly 0.

Chapter 2

Where the scores come from: queries, keys and values

Before the formulas, get a feel for it: a score is large when the query points the same way as a key.

Each word turns its vector into three new vectors by multiplying it by three learned matrices:

  • a query: what am I looking for? ("it" is looking for a noun that can be tired.)
  • a key: what do I offer to others? ("animal" offers: a noun, alive.)
  • a value: what I actually hand over if someone picks me.

A word's score for another word is the dot product of its query with the other word's key: multiply them position by position and add.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here In the example
the query vector: a list of numbers (1, 2)
the key vector: another list of numbers (3, 0.5)
how many numbers each list has (the head width) 2
a counter walking over the positions 1, 2
the -th number in each list ,
"dot product": multiply matching positions, then add

In words: "multiply the first numbers together, multiply the second numbers together, and so on, then add up all the products."

With the numbers: (1, 2) · (3, 0.5) = 1·3 + 2·0.5 = 3 + 1 = 4.

Level 3: in Python
q = [1, 2]
k = [3, 0.5]
# Σ over m of q_m k_m
sum(q_m * k_m for q_m, k_m in zip(q, k))  # → 4.0

Vectors that point the same way give big positive scores, vectors at right angles give zero, and opposite vectors give negative scores. That is why the dot product works as a relevance score. The values are then blended with the softmax shares, exactly as in the table above.

A library search is a good analogy. The query is what you type into the search box, the keys are the catalogue entries, and the values are the books on the shelf. You get back a blend of books, weighted by how well each catalogue entry matched your search.

Figure 2 · Diagram

Reading it: follow the arrows left to right. The token vectors split three ways. Queries and keys meet in the Scores box: that is where relevance is decided. The values take the lower path and wait, untouched, until the Weighted sum box, where the relevance shares decide how much of each value is mixed in. Keep the two roles apart: Q and K decide how much; V carries what. The Scale and Mask boxes are explained in their own sections below.

Written as one formula, with every word's query stacked into a matrix Q (and likewise K and V), the whole diagram is:

Level 3: the formula and its symbols

Symbols

Symbol Meaning here Shape
number of tokens in the sequence
width of each query and key vector
width of each value vector
every token's query, stacked as rows
every token's key, stacked as rows
every token's value, stacked as rows
transposed: rows become columns, so each key stands upright
a matrix multiply: entry (row , column ) is the dot product of token 's query with token 's key, so it is every score at once
square root of the head width; dividing by it keeps scores from growing with width (explained below) a single number
softmax(…) softmax applied to each row separately, so each token's shares add to 1
multiply the shares by the values: each output row is a share-weighted blend of value rows

In words: "score every query against every key, shrink the scores by √d_k, turn each row of scores into shares, and use the shares to blend the values."

With the numbers: the "it" row of holds (2.0, 1.0, 0.5); softmax turns that row into (0.63, 0.23, 0.14); multiplying by gives 0.63·V(animal) + 0.23·V(tired) + 0.14·V(street). To see every symbol at work, take : the query of "it" (1, 1, 1, 1) against the keys (1, 1, 1, 1), (1, 1, 0, 0) and (1, 0, 0, 0) gives raw scores (4, 2, 1), and dividing by gives exactly that row. With toy values V(animal) = (1, 0), V(tired) = (0, 1) and V(street) = (1, 1), the blend is (0.63 + 0.14, 0.23 + 0.14) = (0.77, 0.37).

Level 3: in Python
import math
q_it = [1, 1, 1, 1]
# keys: animal, tired, street
K = [[1, 1, 1, 1], [1, 1, 0, 0], [1, 0, 0, 0]]
# values, one row per word
V = [[1, 0], [0, 1], [1, 1]]
d_k = len(q_it)
# q Kᵀ / √d_k
scores = [sum(q * k for q, k in zip(q_it, k_j)) / math.sqrt(d_k) for k_j in K]
scores  # → [2.0, 1.0, 0.5]
exps = [math.exp(s) for s in scores]
# softmax of the row
weights = [e / sum(exps) for e in exps]
[round(w, 2) for w in weights]  # → [0.63, 0.23, 0.14]
# (…)V: blend the values
[round(sum(w * v[c] for w, v in zip(weights, V)), 2) for c in range(2)]  # → [0.77, 0.37]

In practice this one line runs in every layer of every modern language model, for every word, many times per word generated. Nearly everything about a model's speed and memory use traces back to it.

In code: scaled_dot_product_attention is the whole formula, one commented step per box of the diagram, and returns both the blended values and the attention weights.

Chapter 3

Shapes: the part that is easiest to get wrong

For one sequence of n tokens with model width d_model:

Tensor Shape Meaning
X (n, d_model) input token vectors
W_q, W_k (d_model, d_k) learned projections
W_v (d_model, d_v) learned projection
Q, K (n, d_k) queries, keys
V (n, d_v) values
Q Kᵀ (n, n) every token vs. every token, hence O(n²)
weights (n, n) each row sums to 1
output (n, d_v) one context-aware vector per token

Chapter 4

Causal masking: no peeking at the answer

A decoder model (GPT, Claude, Llama) is trained to predict the next token. If position i could attend to position i+1, it could simply read the answer. So the scores for future positions are set to −∞ before softmax. Because e^−∞ = 0, those weights come out as exactly zero.

Figure 3 · Diagram

Reading it: each row is one token looking at the sequence. A ✓ is a position it may read and a · is the future, which gets masked. The allowed region is a lower triangle: the first token sees only itself and the last sees everything before it. The same triangle appears as the dark region in the heatmap below.

Figure 4 · Drawn from the lesson's code

The animal didn't cross the street because it was tired key: the token being looked at The animal didn't cross the street because it was tired query: the token doing the looking Causal attention: each row sees only the past 0.0 0.2 0.4 0.6 0.8 1.0 attention weight

Every cell above the diagonal is zero, so each of the 10 words attends only to itself and earlier words, and each row sums to 1

Reading it: rows are the token doing the looking (the query), columns are the tokens being looked at (the keys), and darker means more weight. The upper-right triangle is exactly zero, so nothing reads the future. Each row sums to 1, so a token's attention is a budget it spends across the past. The weights come from random projections here, so the pattern itself means nothing; the triangle is what matters. In a trained model, rows light up on the tokens that actually help, such as "it" lighting up "animal".

Figure 5 · Interactive · computed from the lesson's code

Attention heatmap for the sentence

Try it: here is the same sentence with hand-picked queries and keys, so the pattern does mean something; "it" keeps the worked example's scores for "animal", "tired" and "street". Pick the row for "it", then turn the causal mask on: "tired" comes later, so its weight drops to exactly 0 and "animal" takes a bigger share. Drag the temperature below 1 to watch each row sharpen towards one word, and above 1 to watch it flatten towards an even spread.

Masking only the future is also what makes generation cheap: a token's output never changes when later tokens arrive, so it can be computed once and cached. See primer.ml.inference for the KV cache.

In code: causal_mask builds the lower triangle of allowed positions, and scaled_dot_product_attention sets every score outside it to −∞ before softmax. worked_example_sentence holds the hand-picked queries and keys the heatmap above draws.

Why masking makes generation cheap, one token at a time: once a row is written, the mask means it never changes.

Chapter 5

Why divide by √d_k?

See the problem first, then the proof.

Everyday picture Roll one die and the result swings between 1 and 6. Add up a hundred dice and the total swings far more, by dozens in either direction. A dot product is exactly that kind of sum: one small random-ish product per position. The wider the vectors, the more terms get added, and the wilder the scores swing.

Tiny example With d_k = 4, a query and a key of random ±1 entries give four products of +1 or −1, so the score lands between −4 and +4 and is usually within ±2. With d_k = 256 it's a sum of 256 such products: usually within ±16. Softmax of scores like (16, 3, −9) gives about (0.999998, …, …). All the attention lands on one word, whether or not that is right.

Level 3: the formula and its symbols

Symbols

Symbol Meaning here
variance: the average squared distance of from its average; a measure of how widely swings. Its square root is the standard deviation, the typical size of a swing
a raw attention score, for and whose entries are independent with average 0 and variance 1
the number of products added up in the dot product
"which implies"
dividing a quantity by divides its variance by ; here

In words: "the spread of a raw score grows with the head width, so we divide the score by the square root of the width, which brings the spread back to 1 at any width."

With the numbers: at d_k = 128, the standard deviation of a raw score is √128 ≈ 11.3, so scores of ±20 are routine. After dividing by √128 the standard deviation is 1, so scores of ±2 are typical. At d_k = 4 you can check the formula exactly: the 16 equally likely patterns of four ±1 products give raw scores whose variance is 4, and dividing each by √4 = 2 brings the variance to 1.

Level 3: in Python
import itertools, math, statistics
d_k = 4
scores = [sum(products) for products in itertools.product([-1, 1], repeat=d_k)]
# Var(q · k) = d_k
statistics.pvariance(scores)  # → 4
# Var(q · k / √d_k) = 1
statistics.pvariance([s / math.sqrt(d_k) for s in scores])  # → 1.0
# the typical raw swing at d_k = 128
round(math.sqrt(128), 1)  # → 11.3

Why "all attention on one word" is bad, beyond being a wrong answer: softmax has stopped responding. Its gradient (the signal training uses to adjust the weights) is diag(p) − p pᵀ, and when one share is 1 and the rest are 0, every entry of that matrix is 0. The query and key weights receive no signal and stop learning. This is called saturation. sqrt_dk_experiment measures all of this.

Figure 6 · Drawn from the lesson's code

2 2 2 4 2 6 2 8 2 1 0 head width d_k 0.0 0.2 0.4 0.6 0.8 1.0 mean largest weight Softmax saturates (→ one-hot) unscaled QKᵀ scaled QKᵀ/√d_k 2 2 2 4 2 6 2 8 2 1 0 head width d_k 0.05 0.10 0.15 0.20 0.25 0.30 softmax Jacobian norm Gradient through softmax vanishes

As head width grows from 2 to 1024, unscaled attention's top weight climbs from 0.29 to 0.96 and its gradient shrinks fivefold; scaled stays flat

Reading it: the x-axis is the head width d_k, on a log scale. On the left, the y-axis is the average largest attention weight: 1.0 means one-hot, all attention on one token. On the right, it is the size of the softmax gradient. The unscaled lines (red) climb towards 1.0 on the left and fall towards zero on the right as d_k grows: the wider the head, the more attention collapses onto a single token and the less signal flows back. The scaled lines (blue) stay flat at every width. That flatness is the whole reason for the √d_k.

In one sentence: dot products grow with dimension, large scores saturate softmax and kill gradients, and scaling by √d_k keeps scores in the range where softmax is smooth and trainable.

In code: softmax_jacobian builds the diag(p) − p pᵀ matrix, so you can watch every entry fall to 0 as one share approaches 1.

Chapter 6

Multi-head attention

Instead of one attention with d_k = d_model, run h heads in parallel, each with d_k = d_model / h and its own W_q, W_k and W_v. The total compute is the same, but each head can specialise: one tracks syntax, another which noun a pronoun refers to, another the previous token.

Figure 7 · Diagram

Reading it: the input is projected once, and the result is sliced into h narrow slabs, one per head. Each head runs the full attention recipe from the first diagram on its own slab, independently and in parallel. The heads' outputs are glued back side by side, and W_o lets them exchange what they found. The output has the same shape as the input, which is what lets blocks stack.

In code: MultiHeadAttention holds the four projections W_q, W_k, W_v and W_o; calling it slices the projections into one slab per head, runs scaled_dot_product_attention on every head at once, and glues the results back together before W_o.

Chapter 7

Grouped-query attention: sharing keys and values

With grouped-query attention (GQA), several query heads share one key/value head. Llama 3 8B has 32 query heads but only 8 KV heads, so it stores 4× fewer keys and values per token. Multi-query attention (MQA) is the extreme: one KV head for everyone.

Figure 8 · Diagram

Reading it: count the KV boxes. Every query head still asks its own question, but in GQA and MQA they read from a shared set of keys and values. The KV boxes are what gets stored in GPU memory for every token of every conversation during generation (the KV cache). Fewer boxes means more conversations fit on one GPU, at a small cost in quality. See primer.ml.inference for the memory arithmetic.

In code: MultiHeadAttention takes a number of KV heads: fewer than the query heads is GQA, one is MQA. MultiHeadAttention.kv_params counts the key and value weights that shrink.

Try it on a real model's shape: how many conversations fit on one GPU?

Chapter 8

Cost: why long context is expensive

The score matrix is n × n per head per layer, so compute and memory grow with n². Doubling the context roughly quadruples attention's cost.

Figure 9 · Drawn from the lesson's code

1 0 2 1 0 3 1 0 4 1 0 5 1 0 6 context length n (tokens) 1 0 8 1 0 9 1 0 1 0 1 0 1 1 1 0 1 2 1 0 1 3 1 0 1 4 1 0 1 5 1 0 1 6 FLOPs per layer Why long context is expensive (d_model = 4096) crossover n = 2·d_model scores + mix O(n²·d) projections O(n·d²)

On log axes, the n-squared attention cost overtakes the linear projection cost at 8,192 tokens (twice d_model of 4,096) and then pulls away

Reading it: both axes are logarithmic, so a straight line is a power law and a steeper line grows faster. The projections (X·W) cost O(n·d²), a line of slope 1. The score-and-mix step costs O(n²·d), a line of slope 2. The dashed marker is where they cross, at n = 2·d_model. Past that point, most of the work is tokens comparing themselves to other tokens, and every doubling of context costs 4× there.

The standard mitigations each attack this picture. FlashAttention produces the exact same result but computes it in tiles sized for fast on-chip GPU memory, so the n × n matrix is never written to slow memory. Sliding-window and sparse attention let each token see only some others. GQA and MQA shrink the KV cache. State-space models such as Mamba replace attention with a linear-time recurrence.

In code: attention_cost counts the quadratic score-and-mix FLOPs and the linear projection FLOPs plotted above.

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 1Explain attention to a non-engineer in 30 seconds.Think it through, then reveal

When the model reads a word, it asks which other words here help it understand this one. It scores every other word for relevance, then builds its understanding of the word as a mix of the relevant ones. In "the animal didn't cross the street because it was tired", "it" draws mostly from "animal". It does this for every word, in parallel, dozens of times over.

Question 2Now explain it to an ML engineer in two minutes.Think it through, then reveal

Project X into Q, K and V with learned matrices. Compute QKᵀ/√d_k, an n×n matrix of scaled similarities; add a causal mask of −∞ above the diagonal for a decoder; softmax each row; multiply by V. Run h heads in parallel on d_model/h slices, concatenate them, and project with W_o. Wrap the whole thing in a residual connection with pre-layer-norm and follow it with an FFN. Cost is O(n²·d) in compute and O(n²) in memory for the scores, which is why FlashAttention tiles it and GQA shrinks the KV cache.

Question 3Why divide by √d_k? What breaks without it?Think it through, then reveal

The variance of q·k grows linearly with d_k. Without scaling, scores spread out, softmax saturates towards one-hot, its Jacobian goes to about zero, and the query/key projections stop receiving gradient. Training stalls or turns unstable.

Question 4Why is long context expensive? Name two techniques that reduce the cost.Think it through, then reveal

Scores are n×n per head per layer, so cost is quadratic, and the KV cache grows linearly with every token held in memory. Two fixes: FlashAttention (exact, memory-efficient tiling) and GQA/MQA (fewer KV heads). Others: sliding-window or sparse attention, and prompt caching of stable prefixes.

Question 5What does the causal mask buy you besides honest training?Think it through, then reveal

Earlier outputs never depend on later tokens, so during generation the keys and values of past tokens can be computed once and cached. Each new token then costs one row of attention instead of a full recomputation.

Question 6What does GQA trade away, and for what?Think it through, then reveal

A little modelling capacity (query heads share keys and values) for a KV cache that is several times smaller. That means more concurrent requests and longer contexts per GPU.

Primary sources

The papers behind this lesson

Vaswani et al., Attention Is All You Need (2017)

Showed that attention alone, with no recurrence, is enough for state-of-the-art translation, and introduced scaled dot-product attention, multi-head attention and the transformer.

Read on rumblr →The paper ↗
Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints (2023)

Introduced grouped-query attention, the middle ground between multi-head and multi-query attention.

The paper ↗
Dao et al., FlashAttention (2022)

Computes exact attention in tiles sized for fast on-chip GPU memory, making long contexts practical.

Read the annotated companion →The paper ↗

Researcher's shelf

Further reading

  • Vaswani et al., Attention Is All You Need (2017): https://arxiv.org/abs/1706.03762
  • Jay Alammar, The Illustrated Transformer: https://jalammar.github.io/illustrated-transformer/
  • Harvard NLP, The Annotated Transformer (the paper, line by line in code): https://nlp.seas.harvard.edu/annotated-transformer/
  • Andrej Karpathy, Let's build GPT: from scratch, in code, spelled out: https://www.youtube.com/watch?v=kCc8FmEb1nY and nanoGPT: https://github.com/karpathy/nanoGPT
  • 3Blue1Brown, Attention in transformers, visually explained: https://www.youtube.com/watch?v=eMlx5fFNoYc
  • PyTorch scaled_dot_product_attention: https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
  • Ainslie et al., GQA (2023): https://arxiv.org/abs/2305.13245
  • Dao et al., FlashAttention (2022): https://arxiv.org/abs/2205.14135

About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit c8d5c21, so the two always agree: the explanation, the code that builds it and the tests that prove it.