primer.ml.attention

Attention: how a token decides what to listen to

Run: python -m primer.ml.attention

New to the notation? primer.notation explains every symbol used here from zero. This lesson builds on the vectors and matrix multiplies of primer.ml.neural_net.

Level 1: The practitioner's guide

In one sentence. Attention is the step in which every token of a model's input scores every other token for relevance and rebuilds itself as a weighted blend of the ones that matter; because the comparison is all-pairs, it sets a model's context limit, the price of a long prompt, and the memory a conversation occupies on a GPU.

When you need it. You never switch attention on: every transformer you call already runs it in every layer. You need to understand it the day a decision turns on its cost: picking a model by its context window or by the key-value heads on its model card, deciding whether to put 200 pages in one prompt or retrieve the relevant three, sizing a GPU for a model you host, or explaining why a call with a long history is slow before the first output token appears. The tell: latency and cost that grow with the length of the conversation rather than the length of the answer. From this lesson's attention_cost: a 1,000-token prompt builds a score matrix of 1,000,000 entries per head per layer, a 2,000-token prompt 4,000,000, and a 128,000-token prompt 16.4 billion. Doubling the context quadruples that part of the work. Below a few thousand tokens you can ignore it: the parts of the model that grow linearly dominate (the lesson's crossover is at twice the model width, 8,192 tokens for a 4,096-wide model).

Your options. Some of these you choose in your prompt or API call; the rest you choose by picking a model or a serving stack. From the cheapest to the most committed:

Option What it does What it gives you What it costs Where it lives
Short prompts, retrieval for the rest Puts only the relevant text in the context Cost that stays flat as the corpus grows An index to build and a retrieval step that can miss Your code (primer.agents.rag)
Prompt caching Reuses the keys and values of a prefix that repeats across calls Cache reads at 0.1× the input price and a faster first token The prefix must be identical byte for byte; a cache write costs 1.25× The API, or the serving engine
A long context window Sends the whole document or history in one call Nothing to retrieve; exact cross-references over all of it Quadratic compute, a cache that grows with every token, weaker recall in the middle The model you pick: up to 1M tokens on current hosted models
A GQA or MQA model Shares each key-value head among several query heads A KV cache 4× to 8× smaller, so more conversations fit on one GPU A small loss of modelling capacity, decided by the model's authors The model card: num_key_value_heads
FlashAttention kernels Computes the same attention in tiles that stay in fast on-chip memory The exact result, 2× to 3× faster, and no n × n matrix in slow memory A supported GPU; nothing to tune The serving stack: PyTorch, vLLM and the rest
Sliding-window attention Lets each layer look back a fixed number of tokens A cache of fixed size per token and linear cost at any length Exact lookup beyond the window is gone; information hops layer by layer The model architecture (Mistral 7B: a 4,096-token window)
A paged KV cache Stores each request's cache in small blocks instead of one reserved strip 2× to 4× the throughput from the same GPU memory A serving engine that supports it vLLM, and most engines since
A linear-time architecture Replaces attention with a recurrence that carries a fixed-size state No n² term and a cache that does not grow Different recall behaviour and fewer mature models The model architecture (state-space models such as Mamba)

How to choose. Start from what you control.

  • A hosted API: you control the prompt. Put the stable part (system prompt, tool definitions, reference documents) first and keep it identical so it caches; put what changes last. Count tokens before you send: input that exceeds the window is rejected, not truncated.
  • Long context or retrieval: long context for one document the model must read exactly (a contract, a codebase), retrieval when the corpus outlives a single prompt or the relevant part is small. Most production systems do both: retrieve, then give the model a generous window of what came back.
  • An open model to host: read num_attention_heads, num_key_value_heads and max_position_embeddings on its card. The cache per token scales with the key-value heads, and that number, not the parameter count, decides how many conversations one GPU carries.
  • Serving it: use an engine that ships FlashAttention and paged caching rather than a loop of your own.
  • Training or fine-tuning: keep the library's score scaling and causal mask as they are; Level 2 measures what happens without them.
  • Whatever you pick, measure quality against context length on your own task, with the answer placed at the start, the middle and the end. A longer advertised window does not make a model better at using the middle of it.

What it costs. Four currencies: compute, memory, money and recall.

  • Compute. The score-and-mix step is quadratic in tokens and is paid in full when a prompt is first read: that is the pause before the first output token. From attention_cost with a 4,096-wide model, at 8,000 tokens the quadratic part equals the linear part; at 128,000 tokens it is about 16 times larger.
  • Memory. Every token of every live conversation keeps its keys and values in the KV cache. For a Llama-3-8B-shaped model (32 layers, 8 key-value heads of width 128, 16-bit numbers) that is 131,072 bytes per token, so a 32,000-token conversation holds about 4.2 GB, and an 80 GB GPU with 16 GB of weights carries 15 such conversations. At 128,000 tokens, 32 key-value heads would need 67 GB for one request; 8 need 17 GB. The figures are primer.ml.inference's.
  • Money. Hosted APIs bill per token at a flat rate across the window: a 100,000-token prompt on a model priced at \$5 per million input tokens costs \$0.50 each time it is sent, and \$0.05 when the whole prompt is a cache read at 0.1×. A 900,000-token request costs the same per token as a 9,000-token one (Anthropic's pricing page): the quadratic compute is priced into the flat rate.
  • Recall. In the Lost in the Middle study (Liu et al., 2023), GPT-3.5-Turbo answered 75.8% of questions when the useful document was first among 20 and 53.8% when it was tenth, below its 56.1% with no documents at all.

What breaks.

  • The answer in the middle. Accuracy falls when the relevant text sits mid-prompt. Put the most important material first or last, and keep prompts as short as the task allows.
  • A cache that never hits. A timestamp or request id near the top of the prompt changes the prefix on every call, and every call pays full price. If the cache-read count in the usage report is zero across identical requests, something volatile sits before the stable part.
  • Out of memory at long context. The KV cache, not the weights, is what overflows a GPU on long requests. Prefer a GQA model, cap the context you accept, or run an engine that pages the cache.
  • A prompt that does not fit. Input beyond the window is an error; input plus the output budget beyond it stops generation early. Count tokens first and leave room for the answer.
  • Training instability from unscaled scores. Without the division by the square root of the head width, this lesson's measurement at width 512 puts 0.94 of the attention on one token on average and shrinks the training signal fivefold; the model stops learning what to attend to.

In the wild. The original transformer (Vaswani et al., 2017) ran 8 heads of width 64 on 512-wide vectors. Anthropic's context-window documentation gives current models a 1M-token window and counts everything in the request towards it: system prompt, tool definitions, tool results and the model's own thinking. Hugging Face model configs expose num_attention_heads, num_key_value_heads and max_position_embeddings; Llama 3 pairs 32 query heads with 8 key-value heads and reaches 128K tokens, and Mistral 7B uses the same 32-to-8 split plus a 4,096-token sliding window whose reach grows to about 131K tokens across its 32 layers. The GQA paper converted multi-head checkpoints with 5% of the original pretraining compute. PyTorch's scaled_dot_product_attention chooses among a FlashAttention-2 kernel, a memory-efficient kernel and a plain implementation by itself, and vLLM's PagedAttention lifted serving throughput 2 to 4 times over earlier engines, whose reserved-strip caches held real tokens in only 20% to 38% of their memory. The papers are linked at the end of the lesson.

Go deeper. Level 2 builds attention from three numbers: softmax on a pronoun's scores, queries, keys and values on a four-number example, the causal mask that makes caching possible, the square-root scaling measured at three widths, multi-head and grouped-query attention in a class you can call, and the n² curve. If you only needed to choose a model or shape a prompt, you are done.

Level 2: How it works, from scratch.

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.

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

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

$$ \text{softmax}(z)_i = \frac{e^{z_i}}{\sum_{j=1}^{n} e^{z_j}} $$

Symbols

Symbol Meaning here In the example
$z$ the list of scores going in (2.0, 1.0, 0.5)
$n$ how many scores there are 3
$i$ the position we're computing a share for 1 = "animal"
$z_i$ the score at position $i$ $z_1 = 2.0$
$e$ Euler's number, ≈ 2.718; $e^x$ is "2.718 multiplied by itself $x$ times" (works for fractions and negatives too) $e^{2.0} = 7.39$
$\sum_{j=1}^{n}$ "add up the following, for $j$ = 1, 2, …, $n$" $e^{2.0} + e^{1.0} + e^{0.5}$
$j$ a counter that walks over every position 1, 2, 3
$\text{softmax}(z)_i$ the share of attention position $i$ 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

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.

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.

Where the scores come from: queries, keys and values

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

$$ q \cdot k = \sum_{m=1}^{d_k} q_m \, k_m $$

Symbols

Symbol Meaning here In the example
$q$ the query vector: a list of $d_k$ numbers (1, 2)
$k$ the key vector: another list of $d_k$ numbers (3, 0.5)
$d_k$ how many numbers each list has (the head width) 2
$m$ a counter walking over the positions 1, 2
$q_m, k_m$ the $m$-th number in each list $q_1 = 1$, $k_1 = 3$
$\cdot$ "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

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.

flowchart LR X[Token vectors] --> Q[Query] & K[Key] & V[Value] Q --> S[Scores<br/>Q times K] K --> S S --> SC[Scale by<br/>sqrt of d_k] SC --> M[Causal mask<br/>hide future tokens] M --> SM[Softmax<br/>weights sum to 1] SM --> W[Weighted sum<br/>of values] V --> W W --> O[Context-aware vectors]

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

$$ \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V $$

Symbols

Symbol Meaning here Shape
$n$ number of tokens in the sequence
$d_k$ width of each query and key vector
$d_v$ width of each value vector
$Q$ every token's query, stacked as rows $n \times d_k$
$K$ every token's key, stacked as rows $n \times d_k$
$V$ every token's value, stacked as rows $n \times d_v$
$K^\top$ $K$ transposed: rows become columns, so each key stands upright $d_k \times n$
$QK^\top$ a matrix multiply: entry (row $i$, column $j$) is the dot product of token $i$'s query with token $j$'s key, so it is every score at once $n \times n$
$\sqrt{d_k}$ 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 $n \times n$
$(\ldots)V$ multiply the shares by the values: each output row is a share-weighted blend of value rows $n \times d_v$

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 $QK^\top/\sqrt{d_k}$ holds (2.0, 1.0, 0.5); softmax turns that row into (0.63, 0.23, 0.14); multiplying by $V$ gives 0.63·V(animal) + 0.23·V(tired) + 0.14·V(street). To see every symbol at work, take $d_k = 4$: 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 $\sqrt{4} = 2$ 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

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.

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

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.

flowchart LR subgraph S["Scores (n × n)"] direction TB r0["row 'The': ✓ · · ·"] r1["row 'cat': ✓ ✓ · ·"] r2["row 'sat': ✓ ✓ ✓ ·"] r3["row 'down': ✓ ✓ ✓ ✓"] end S --> MASK["set every · to −∞"] --> SOFT["softmax per row"] --> OUT["each row sums to 1<br/>using only the past"]

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.

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".

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 divide by √d_k?

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

$$ \operatorname{Var}(q \cdot k) = d_k \qquad\Longrightarrow\qquad \operatorname{Var}!\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = 1 $$

Symbols

Symbol Meaning here
$\operatorname{Var}(x)$ variance: the average squared distance of $x$ from its average; a measure of how widely $x$ swings. Its square root is the standard deviation, the typical size of a swing
$q \cdot k$ a raw attention score, for $q$ and $k$ whose entries are independent with average 0 and variance 1
$d_k$ the number of products added up in the dot product
$\Longrightarrow$ "which implies"
$\sqrt{d_k}$ dividing a quantity by $c$ divides its variance by $c^2$; here $c^2 = d_k$

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

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.

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.

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.

flowchart LR X["X (n × d_model)"] --> P["project with W_q, W_k, W_v"] P --> H1["head 1<br/>attention on d_model/h dims"] P --> H2["head 2"] P --> H3["..."] P --> Hh["head h"] H1 & H2 & H3 & Hh --> C["concatenate<br/>(n × d_model)"] C --> WO["mix with W_o"] --> Y["output (n × d_model)"]

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.

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.

flowchart TB subgraph MHA["MHA: 4 query heads, 4 KV heads"] q1a[Q1]-->kv1a[KV1] q2a[Q2]-->kv2a[KV2] q3a[Q3]-->kv3a[KV3] q4a[Q4]-->kv4a[KV4] end subgraph GQA["GQA: 4 query heads, 2 KV heads"] q1b[Q1]-->kv1b[KV1] q2b[Q2]-->kv1b q3b[Q3]-->kv2b[KV2] q4b[Q4]-->kv2b end subgraph MQA["MQA: 4 query heads, 1 KV head"] q1c[Q1]-->kv1c[KV1] q2c[Q2]-->kv1c q3c[Q3]-->kv1c q4c[Q4]-->kv1c end

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.

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.

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.

In 20 seconds

  • 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.
  • Why √d_k: dot products grow with dimension, and large scores saturate softmax and kill gradients. Scaling keeps training stable.
  • 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.
  • Long context cost: attention compares every pair of tokens, so it is O(n²), and the KV cache grows with every token.

Self-test questions

Explain attention to a non-engineer in 30 seconds. 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.

Now explain it to an ML engineer in two minutes. 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.

Why divide by √d_k? What breaks without it? 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.

Why is long context expensive? Name two techniques that reduce the cost. 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.

What does the causal mask buy you besides honest training? 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.

What does GQA trade away, and for what? 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.

The papers behind this lesson

  • Vaswani et al., Attention Is All You Need (2017): https://arxiv.org/abs/1706.03762. 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. Annotated companion
  • Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints (2023): https://arxiv.org/abs/2305.13245. Introduced grouped-query attention, the middle ground between multi-head and multi-query attention.
  • Dao et al., FlashAttention (2022): https://arxiv.org/abs/2205.14135. Computes exact attention in tiles sized for fast on-chip GPU memory, making long contexts practical. Annotated companion

Further reading

on GitHub
   1r"""
   2# Attention: how a token decides what to listen to
   3
   4Run: `python -m primer.ml.attention`
   5
   6New to the notation? `primer.notation` explains every symbol used here from
   7zero. This lesson builds on the vectors and matrix multiplies of
   8`primer.ml.neural_net`.
   9
  10## Level 1: The practitioner's guide
  11
  12**In one sentence.** Attention is the step in which every token of a
  13model's input scores every other token for relevance and rebuilds itself as
  14a weighted blend of the ones that matter; because the comparison is
  15all-pairs, it sets a model's context limit, the price of a long prompt, and
  16the memory a conversation occupies on a GPU.
  17
  18**When you need it.** You never switch attention on: every transformer you
  19call already runs it in every layer. You need to understand it the day a
  20decision turns on its cost: picking a model by its context window or by the
  21key-value heads on its model card, deciding whether to put 200 pages in one
  22prompt or retrieve the relevant three, sizing a GPU for a model you host, or
  23explaining why a call with a long history is slow before the first output
  24token appears. The tell: latency and cost that grow with the length of the
  25conversation rather than the length of the answer. From this lesson's
  26`attention_cost`: a 1,000-token prompt builds a score
  27matrix of 1,000,000 entries per head per layer, a 2,000-token prompt
  284,000,000, and a 128,000-token prompt 16.4 billion. Doubling the context
  29quadruples that part of the work. Below a few thousand tokens you can
  30ignore it: the parts of the model that grow linearly dominate (the lesson's
  31crossover is at twice the model width, 8,192 tokens for a 4,096-wide model).
  32
  33**Your options.** Some of these you choose in your prompt or API call; the
  34rest you choose by picking a model or a serving stack. From the cheapest to
  35the most committed:
  36
  37| Option | What it does | What it gives you | What it costs | Where it lives |
  38|---|---|---|---|---|
  39| Short prompts, retrieval for the rest | Puts only the relevant text in the context | Cost that stays flat as the corpus grows | An index to build and a retrieval step that can miss | Your code (`primer.agents.rag`) |
  40| Prompt caching | Reuses the keys and values of a prefix that repeats across calls | Cache reads at 0.1× the input price and a faster first token | The prefix must be identical byte for byte; a cache write costs 1.25× | The API, or the serving engine |
  41| A long context window | Sends the whole document or history in one call | Nothing to retrieve; exact cross-references over all of it | Quadratic compute, a cache that grows with every token, weaker recall in the middle | The model you pick: up to 1M tokens on current hosted models |
  42| A GQA or MQA model | Shares each key-value head among several query heads | A KV cache 4× to 8× smaller, so more conversations fit on one GPU | A small loss of modelling capacity, decided by the model's authors | The model card: `num_key_value_heads` |
  43| FlashAttention kernels | Computes the same attention in tiles that stay in fast on-chip memory | The exact result, 2× to 3× faster, and no n × n matrix in slow memory | A supported GPU; nothing to tune | The serving stack: PyTorch, vLLM and the rest |
  44| Sliding-window attention | Lets each layer look back a fixed number of tokens | A cache of fixed size per token and linear cost at any length | Exact lookup beyond the window is gone; information hops layer by layer | The model architecture (Mistral 7B: a 4,096-token window) |
  45| A paged KV cache | Stores each request's cache in small blocks instead of one reserved strip | 2× to 4× the throughput from the same GPU memory | A serving engine that supports it | vLLM, and most engines since |
  46| A linear-time architecture | Replaces attention with a recurrence that carries a fixed-size state | No n² term and a cache that does not grow | Different recall behaviour and fewer mature models | The model architecture (state-space models such as Mamba) |
  47
  48**How to choose.** Start from what you control.
  49
  50- A hosted API: you control the prompt. Put the stable part (system prompt,
  51  tool definitions, reference documents) first and keep it identical so it
  52  caches; put what changes last. Count tokens before you send: input that
  53  exceeds the window is rejected, not truncated.
  54- Long context or retrieval: long context for one document the model must
  55  read exactly (a contract, a codebase), retrieval when the corpus outlives
  56  a single prompt or the relevant part is small. Most production systems do
  57  both: retrieve, then give the model a generous window of what came back.
  58- An open model to host: read `num_attention_heads`, `num_key_value_heads`
  59  and `max_position_embeddings` on its card. The cache per token scales with
  60  the key-value heads, and that number, not the parameter count, decides how
  61  many conversations one GPU carries.
  62- Serving it: use an engine that ships FlashAttention and paged caching
  63  rather than a loop of your own.
  64- Training or fine-tuning: keep the library's score scaling and causal mask
  65  as they are; Level 2 measures what happens without them.
  66- Whatever you pick, measure quality against context length on your own
  67  task, with the answer placed at the start, the middle and the end. A
  68  longer advertised window does not make a model better at using the middle
  69  of it.
  70
  71**What it costs.** Four currencies: compute, memory, money and recall.
  72
  73- Compute. The score-and-mix step is quadratic in tokens and is paid in full
  74  when a prompt is first read: that is the pause before the first output
  75  token. From `attention_cost` with a 4,096-wide model, at 8,000 tokens the
  76  quadratic part equals the linear part; at 128,000 tokens it is about 16
  77  times larger.
  78- Memory. Every token of every live conversation keeps its keys and values
  79  in the KV cache. For a Llama-3-8B-shaped model (32 layers, 8 key-value
  80  heads of width 128, 16-bit numbers) that is 131,072 bytes per token, so a
  81  32,000-token conversation holds about 4.2 GB, and an 80 GB GPU with 16 GB
  82  of weights carries 15 such conversations. At 128,000 tokens, 32 key-value
  83  heads would need 67 GB for one request; 8 need 17 GB. The figures are
  84  `primer.ml.inference`'s.
  85- Money. Hosted APIs bill per token at a flat rate across the window: a
  86  100,000-token prompt on a model priced at \$5 per million input tokens
  87  costs \$0.50 each time it is sent, and \$0.05 when the whole prompt is a
  88  cache read at 0.1×. A 900,000-token request costs the same per token as a
  89  9,000-token one (Anthropic's pricing page): the quadratic compute is
  90  priced into the flat rate.
  91- Recall. In the Lost in the Middle study (Liu et al., 2023), GPT-3.5-Turbo
  92  answered 75.8% of questions when the useful document was first among 20
  93  and 53.8% when it was tenth, below its 56.1% with no documents at all.
  94
  95**What breaks.**
  96
  97- **The answer in the middle.** Accuracy falls when the relevant text sits
  98  mid-prompt. Put the most important material first or last, and keep
  99  prompts as short as the task allows.
 100- **A cache that never hits.** A timestamp or request id near the top of
 101  the prompt changes the prefix on every call, and every call pays full
 102  price. If the cache-read count in the usage report is zero across
 103  identical requests, something volatile sits before the stable part.
 104- **Out of memory at long context.** The KV cache, not the weights, is what
 105  overflows a GPU on long requests. Prefer a GQA model, cap the context you
 106  accept, or run an engine that pages the cache.
 107- **A prompt that does not fit.** Input beyond the window is an error; input
 108  plus the output budget beyond it stops generation early. Count tokens
 109  first and leave room for the answer.
 110- **Training instability from unscaled scores.** Without the division by
 111  the square root of the head width, this lesson's measurement at width 512
 112  puts 0.94 of the attention on one token on average and shrinks the
 113  training signal fivefold; the model stops learning what to attend to.
 114
 115**In the wild.** The original transformer (Vaswani et al., 2017) ran 8 heads
 116of width 64 on 512-wide vectors. Anthropic's context-window documentation gives current models a 1M-token
 117window and counts everything in the request towards it: system prompt, tool
 118definitions, tool results and the model's own thinking. Hugging Face model
 119configs expose `num_attention_heads`, `num_key_value_heads` and
 120`max_position_embeddings`; Llama 3 pairs 32 query heads with 8 key-value
 121heads and reaches 128K tokens, and Mistral 7B uses the same 32-to-8 split
 122plus a 4,096-token sliding window whose reach grows to about 131K tokens
 123across its 32 layers. The GQA paper converted multi-head checkpoints with 5%
 124of the original pretraining compute. PyTorch's
 125`scaled_dot_product_attention` chooses among a FlashAttention-2 kernel, a
 126memory-efficient kernel and a plain implementation by itself, and vLLM's
 127PagedAttention lifted serving throughput 2 to 4 times over earlier engines,
 128whose reserved-strip caches held real tokens in only 20% to 38% of their
 129memory. The papers are linked at the end of the lesson.
 130
 131**Go deeper.** Level 2 builds attention from three numbers: softmax on a
 132pronoun's scores, queries, keys and values on a four-number example, the
 133causal mask that makes caching possible, the square-root scaling measured at
 134three widths, multi-head and grouped-query attention in a class you can
 135call, and the n² curve. If you only needed to choose a model or shape a
 136prompt, you are done.
 137
 138## Level 2: How it works, from scratch
 139
 140Level 2 builds the mechanism from nothing, starting with a glance around a
 141room.
 142
 143**The everyday picture.** Imagine you're in a meeting and someone says "it's broken, can you fix it?"
 144To know what "it" means, you glance around the room: at the laptop on the
 145table, at the person who just walked in, at the whiteboard. You pay a lot of
 146attention to the laptop, a little to everything else, and your understanding
 147of "it" becomes mostly *laptop*.
 148
 149That glance is attention. Every word in a sentence gets to look at every
 150other word, decide how relevant each one is, and rebuild its own meaning as
 151a mix of the relevant ones. "Bank" next to "river" becomes a riverbank; next
 152to "loan" it becomes a lender. The word is the same, but what it paid
 153attention to is different.
 154
 155## A tiny worked example: what does "it" refer to?
 156
 157Take "The animal didn't cross the street because it was tired" and follow
 158the word "it". To keep the numbers small, it compares itself with just three
 159other words. Each comparison gives a relevance **score** (how the score is
 160made comes next). Then we turn scores into shares of attention that add up
 161to 1:
 162
 1631. Exponentiate each score (e^score), which makes every number positive and
 164   stretches the gaps between them.
 1652. Divide each by the total, 7.39 + 2.72 + 1.65 = 11.76.
 166
 167| Word   | Score | e^score | Share of attention |
 168|--------|-------|---------|--------------------|
 169| animal | 2.0   | 7.39    | 7.39 / 11.76 = **0.63** |
 170| tired  | 1.0   | 2.72    | 2.72 / 11.76 = **0.23** |
 171| street | 0.5   | 1.65    | 1.65 / 11.76 = **0.14** |
 172
 173Those two steps are called **softmax**. The new meaning of "it" is 63%
 174"animal", 23% "tired" and 14% "street". Without being told any grammar rule,
 175the model has worked out that "it" is the animal.
 176
 177**In code:** `worked_example_it` runs these three scores through softmax and returns each word's share.
 178
 179### Softmax, decoded
 180
 181Softmax turns any list of numbers, positive or negative, into shares that
 182are all positive and add up to exactly 1. Think of it as a vote where louder
 183voices get disproportionately more say.
 184
 185$$
 186\text{softmax}(z)_i = \frac{e^{z_i}}{\sum_{j=1}^{n} e^{z_j}}
 187$$
 188
 189**Symbols**
 190
 191| Symbol | Meaning here | In the example |
 192|---|---|---|
 193| $z$ | the list of scores going in | (2.0, 1.0, 0.5) |
 194| $n$ | how many scores there are | 3 |
 195| $i$ | the position we're computing a share for | 1 = "animal" |
 196| $z_i$ | the score at position $i$ | $z_1 = 2.0$ |
 197| $e$ | Euler's number, ≈ 2.718; $e^x$ is "2.718 multiplied by itself $x$ times" (works for fractions and negatives too) | $e^{2.0} = 7.39$ |
 198| $\sum_{j=1}^{n}$ | "add up the following, for $j$ = 1, 2, …, $n$" | $e^{2.0} + e^{1.0} + e^{0.5}$ |
 199| $j$ | a counter that walks over every position | 1, 2, 3 |
 200| $\text{softmax}(z)_i$ | the share of attention position $i$ gets | 0.63 |
 201
 202**In words:** "the share for item *i* is *e* raised to its score, divided by
 203the sum of *e* raised to every score."
 204
 205**With the numbers:** softmax(2.0, 1.0, 0.5)₁ = 7.39 / (7.39 + 2.72 + 1.65) =
 2067.39 / 11.76 = 0.63.
 207
 208**In Python:**
 209
 210```python
 211import math
 212# animal, tired, street
 213z = [2.0, 1.0, 0.5]
 214# e^(z_i) for each score
 215exps = [math.exp(z_i) for z_i in z]
 216[round(e, 2) for e in exps]  # → [7.39, 2.72, 1.65]
 217# Σ_j e^(z_j)
 218total = sum(exps)
 219round(total, 2)  # → 11.76
 220# softmax(z)_i: each share of the total
 221[round(e / total, 2) for e in exps]  # → [0.63, 0.23, 0.14]
 222```
 223
 224Why *e* to the power of the score, and not just the score divided by the
 225total? Three reasons you can check on the table above:
 226
 2271. **Always positive.** Scores can be negative; *e* to any power is positive,
 228   so no word ever gets a negative share.
 2292. **Order is kept.** A higher score always gets a bigger share.
 2303. **Gaps are stretched.** Adding 1 to a score multiplies its *e*-value by
 231   2.72, so the winner pulls ahead decisively. (Plain division would give
 232   "animal" 2.0 / 3.5 = 0.57, a timid 57%, and would break on negative
 233   scores.)
 234
 235One more property matters later: adding the same number to every score
 236changes nothing, because it multiplies the top and bottom of the fraction by
 237the same amount. The code uses this to subtract the largest score before
 238exponentiating, which stops *e*¹⁰⁰⁰ from overflowing to infinity. See
 239`primer.notation` for exponents and sums from scratch.
 240
 241![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](figures/primer.ml.attention.it_weights.svg)
 242
 243**Reading it:** for each word, the grey bar is its share of the raw scores
 244and the blue bar is its share of attention after softmax. Softmax exaggerates
 245differences: "animal" scores only twice as high as "tired", yet it ends up
 246with almost three times the attention, because exponentiating stretches the
 247gaps. That is how attention commits to the most relevant word while keeping
 248a little of the others.
 249
 250**In code:** `softmax` subtracts the largest score before exponentiating, and turns masked scores of −∞ into weights of exactly 0.
 251
 252## Where the scores come from: queries, keys and values
 253
 254Each word turns its vector into three new vectors by multiplying it by three
 255learned matrices:
 256
 257- a **query**: what am I looking for? ("it" is looking for a noun that can
 258  be tired.)
 259- a **key**: what do I offer to others? ("animal" offers: a noun, alive.)
 260- a **value**: what I actually hand over if someone picks me.
 261
 262A word's score for another word is the **dot product** of its query with the
 263other word's key: multiply them position by position and add.
 264
 265$$
 266q \cdot k = \sum_{m=1}^{d_k} q_m \, k_m
 267$$
 268
 269**Symbols**
 270
 271| Symbol | Meaning here | In the example |
 272|---|---|---|
 273| $q$ | the query vector: a list of $d_k$ numbers | (1, 2) |
 274| $k$ | the key vector: another list of $d_k$ numbers | (3, 0.5) |
 275| $d_k$ | how many numbers each list has (the head width) | 2 |
 276| $m$ | a counter walking over the positions | 1, 2 |
 277| $q_m, k_m$ | the $m$-th number in each list | $q_1 = 1$, $k_1 = 3$ |
 278| $\cdot$ | "dot product": multiply matching positions, then add | |
 279
 280**In words:** "multiply the first numbers together, multiply the second
 281numbers together, and so on, then add up all the products."
 282
 283**With the numbers:** (1, 2) · (3, 0.5) = 1·3 + 2·0.5 = 3 + 1 = **4**.
 284
 285**In Python:**
 286
 287```python
 288q = [1, 2]
 289k = [3, 0.5]
 290# Σ over m of q_m k_m
 291sum(q_m * k_m for q_m, k_m in zip(q, k))  # → 4.0
 292```
 293
 294Vectors that point the same way give big positive scores, vectors at right
 295angles give zero, and opposite vectors give negative scores. That is why the
 296dot product works as a relevance score. The values are then blended with
 297the softmax shares, exactly as in the table above.
 298
 299A library search is a good analogy. The query is what you type into the
 300search box, the keys are the catalogue entries, and the values are the books
 301on the shelf. You get back a *blend* of books, weighted by how well each
 302catalogue entry matched your search.
 303
 304```mermaid
 305flowchart LR
 306  X[Token vectors] --> Q[Query] & K[Key] & V[Value]
 307  Q --> S[Scores<br/>Q times K]
 308  K --> S
 309  S --> SC[Scale by<br/>sqrt of d_k]
 310  SC --> M[Causal mask<br/>hide future tokens]
 311  M --> SM[Softmax<br/>weights sum to 1]
 312  SM --> W[Weighted sum<br/>of values]
 313  V --> W
 314  W --> O[Context-aware vectors]
 315```
 316
 317**Reading it:** follow the arrows left to right. The token vectors split
 318three ways. Queries and keys meet in the Scores box: that is where relevance
 319is decided. The values take the lower path and wait, untouched, until the
 320Weighted sum box, where the relevance shares decide how much of each value
 321is mixed in. Keep the two roles apart: **Q and K decide *how much*; V carries
 322*what*.** The Scale and Mask boxes are explained in their own sections below.
 323
 324Written as one formula, with every word's query stacked into a matrix Q (and
 325likewise K and V), the whole diagram is:
 326
 327$$
 328\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V
 329$$
 330
 331**Symbols**
 332
 333| Symbol | Meaning here | Shape |
 334|---|---|---|
 335| $n$ | number of tokens in the sequence | |
 336| $d_k$ | width of each query and key vector | |
 337| $d_v$ | width of each value vector | |
 338| $Q$ | every token's query, stacked as rows | $n \times d_k$ |
 339| $K$ | every token's key, stacked as rows | $n \times d_k$ |
 340| $V$ | every token's value, stacked as rows | $n \times d_v$ |
 341| $K^\top$ | $K$ **transposed**: rows become columns, so each key stands upright | $d_k \times n$ |
 342| $QK^\top$ | a **matrix multiply**: entry (row $i$, column $j$) is the dot product of token $i$'s query with token $j$'s key, so it is every score at once | $n \times n$ |
 343| $\sqrt{d_k}$ | square root of the head width; dividing by it keeps scores from growing with width (explained below) | a single number |
 344| softmax(…) | softmax applied to each row separately, so each token's shares add to 1 | $n \times n$ |
 345| $(\ldots)V$ | multiply the shares by the values: each output row is a share-weighted blend of value rows | $n \times d_v$ |
 346
 347**In words:** "score every query against every key, shrink the scores by √d_k,
 348turn each row of scores into shares, and use the shares to blend the values."
 349
 350**With the numbers:** the "it" row of $QK^\top/\sqrt{d_k}$ holds (2.0, 1.0,
 3510.5); softmax turns that row into (0.63, 0.23, 0.14); multiplying by $V$
 352gives 0.63·V(animal) + 0.23·V(tired) + 0.14·V(street). To see every symbol
 353at work, take $d_k = 4$: the query of "it" (1, 1, 1, 1) against the keys
 354(1, 1, 1, 1), (1, 1, 0, 0) and (1, 0, 0, 0) gives raw scores (4, 2, 1), and
 355dividing by $\sqrt{4} = 2$ gives exactly that row. With toy values
 356V(animal) = (1, 0), V(tired) = (0, 1) and V(street) = (1, 1), the blend is
 357(0.63 + 0.14, 0.23 + 0.14) = (0.77, 0.37).
 358
 359**In Python:**
 360
 361```python
 362import math
 363q_it = [1, 1, 1, 1]
 364# keys: animal, tired, street
 365K = [[1, 1, 1, 1], [1, 1, 0, 0], [1, 0, 0, 0]]
 366# values, one row per word
 367V = [[1, 0], [0, 1], [1, 1]]
 368d_k = len(q_it)
 369# q Kᵀ / √d_k
 370scores = [sum(q * k for q, k in zip(q_it, k_j)) / math.sqrt(d_k) for k_j in K]
 371scores  # → [2.0, 1.0, 0.5]
 372exps = [math.exp(s) for s in scores]
 373# softmax of the row
 374weights = [e / sum(exps) for e in exps]
 375[round(w, 2) for w in weights]  # → [0.63, 0.23, 0.14]
 376# (…)V: blend the values
 377[round(sum(w * v[c] for w, v in zip(weights, V)), 2) for c in range(2)]  # → [0.77, 0.37]
 378```
 379
 380In practice this one line runs in every layer of every modern language
 381model, for every word, many times per word generated. Nearly everything
 382about a model's speed and memory use traces back to it.
 383
 384**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.
 385
 386## Shapes: the part that is easiest to get wrong
 387
 388For one sequence of `n` tokens with model width `d_model`:
 389
 390| Tensor | Shape | Meaning |
 391|---|---|---|
 392| X | (n, d_model) | input token vectors |
 393| W_q, W_k | (d_model, d_k) | learned projections |
 394| W_v | (d_model, d_v) | learned projection |
 395| Q, K | (n, d_k) | queries, keys |
 396| V | (n, d_v) | values |
 397| Q Kᵀ | **(n, n)** | every token vs. every token, hence O(n²) |
 398| weights | (n, n) | each row sums to 1 |
 399| output | (n, d_v) | one context-aware vector per token |
 400
 401## Causal masking: no peeking at the answer
 402
 403A decoder model (GPT, Claude, Llama) is trained to predict the *next* token.
 404If position i could attend to position i+1, it could simply read the answer.
 405So the scores for future positions are set to −∞ *before* softmax. Because
 406e^−∞ = 0, those weights come out as exactly zero.
 407
 408```mermaid
 409flowchart LR
 410  subgraph S["Scores (n × n)"]
 411    direction TB
 412    r0["row 'The':    ✓ · · ·"]
 413    r1["row 'cat':    ✓ ✓ · ·"]
 414    r2["row 'sat':    ✓ ✓ ✓ ·"]
 415    r3["row 'down':   ✓ ✓ ✓ ✓"]
 416  end
 417  S --> MASK["set every · to −∞"] --> SOFT["softmax per row"] --> OUT["each row sums to 1<br/>using only the past"]
 418```
 419
 420**Reading it:** each row is one token looking at the sequence. A ✓ is a
 421position it may read and a · is the future, which gets masked. The allowed
 422region is a lower triangle: the first token sees only itself and the last
 423sees everything before it. The same triangle appears as the dark region in
 424the heatmap below.
 425
 426![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](figures/primer.ml.attention.causal_heatmap.svg)
 427
 428**Reading it:** rows are the token doing the looking (the query), columns are
 429the tokens being looked at (the keys), and darker means more weight. The
 430upper-right triangle is exactly zero, so nothing reads the future. Each row
 431sums to 1, so a token's attention is a budget it spends across the past. The
 432weights come from random projections here, so the pattern itself means
 433nothing; the triangle is what matters. In a trained model, rows light up on
 434the tokens that actually help, such as "it" lighting up "animal".
 435
 436**Try it:** here is the same sentence with hand-picked queries and keys, so
 437the pattern does mean something; "it" keeps the worked example's scores for
 438"animal", "tired" and "street". Pick the row for "it", then turn the causal
 439mask on: "tired" comes later, so its weight drops to exactly 0 and "animal"
 440takes a bigger share. Drag the temperature below 1 to watch each row sharpen
 441towards one word, and above 1 to watch it flatten towards an even spread.
 442
 443<div class="viz" data-viz="attention-matrix" aria-label="Attention heatmap for the sentence, with a causal mask and a temperature"></div>
 444
 445Masking only the future is also what makes generation cheap: a token's
 446output never changes when later tokens arrive, so it can be computed once
 447and cached. See `primer.ml.inference` for the KV cache.
 448
 449**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.
 450
 451## Why divide by √d_k?
 452
 453**Everyday picture.** Roll one die and the result swings between 1 and 6.
 454Add up a hundred dice and the total swings far more, by dozens in either
 455direction. A dot product is exactly that kind of sum: one small random-ish
 456product per position. The wider the vectors, the more terms get added, and
 457the wilder the scores swing.
 458
 459**Tiny example.** With d_k = 4, a query and a key of random ±1 entries give
 460four products of +1 or −1, so the score lands between −4 and +4 and is
 461usually within ±2. With d_k = 256 it's a sum of 256 such products: usually
 462within ±16. Softmax of scores like (16, 3, −9) gives about (0.999998, …, …).
 463All the attention lands on one word, whether or not that is right.
 464
 465$$
 466\operatorname{Var}(q \cdot k) = d_k
 467\qquad\Longrightarrow\qquad
 468\operatorname{Var}\!\left(\frac{q \cdot k}{\sqrt{d_k}}\right) = 1
 469$$
 470
 471**Symbols**
 472
 473| Symbol | Meaning here |
 474|---|---|
 475| $\operatorname{Var}(x)$ | **variance**: the average squared distance of $x$ from its average; a measure of how widely $x$ swings. Its square root is the **standard deviation**, the typical size of a swing |
 476| $q \cdot k$ | a raw attention score, for $q$ and $k$ whose entries are independent with average 0 and variance 1 |
 477| $d_k$ | the number of products added up in the dot product |
 478| $\Longrightarrow$ | "which implies" |
 479| $\sqrt{d_k}$ | dividing a quantity by $c$ divides its variance by $c^2$; here $c^2 = d_k$ |
 480
 481**In words:** "the spread of a raw score grows with the head width, so we
 482divide the score by the square root of the width, which brings the spread
 483back to 1 at any width."
 484
 485**With the numbers:** at d_k = 128, the standard deviation of a raw score is
 486√128 ≈ 11.3, so scores of ±20 are routine. After dividing by √128 the
 487standard deviation is 1, so scores of ±2 are typical. At d_k = 4 you can
 488check the formula exactly: the 16 equally likely patterns of four ±1 products
 489give raw scores whose variance is 4, and dividing each by √4 = 2 brings the
 490variance to 1.
 491
 492**In Python:**
 493
 494```python
 495import itertools, math, statistics
 496d_k = 4
 497scores = [sum(products) for products in itertools.product([-1, 1], repeat=d_k)]
 498# Var(q · k) = d_k
 499statistics.pvariance(scores)  # → 4
 500# Var(q · k / √d_k) = 1
 501statistics.pvariance([s / math.sqrt(d_k) for s in scores])  # → 1.0
 502# the typical raw swing at d_k = 128
 503round(math.sqrt(128), 1)  # → 11.3
 504```
 505
 506Why "all attention on one word" is bad, beyond being a wrong answer: softmax
 507has stopped responding. Its **gradient** (the signal training uses to
 508adjust the weights) is `diag(p) − p pᵀ`, and when one share is 1 and the rest
 509are 0, every entry of that matrix is 0. The query and key weights receive no
 510signal and stop learning. This is called *saturation*.
 511`sqrt_dk_experiment` measures all of this.
 512
 513![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](figures/primer.ml.attention.sqrt_dk.svg)
 514
 515**Reading it:** the x-axis is the head width d_k, on a log scale. On the left,
 516the y-axis is the average largest attention weight: 1.0 means one-hot, all
 517attention on one token. On the right, it is the size of the softmax gradient.
 518The unscaled lines (red) climb towards 1.0 on the left and fall towards zero
 519on the right as d_k grows: the wider the head, the more attention collapses
 520onto a single token and the less signal flows back. The scaled lines (blue)
 521stay flat at every width. That flatness is the whole reason for the √d_k.
 522
 523In one sentence: *dot products grow with dimension, large scores saturate
 524softmax and kill gradients, and scaling by √d_k keeps scores in the range
 525where softmax is smooth and trainable.*
 526
 527**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.
 528
 529## Multi-head attention
 530
 531Instead of one attention with d_k = d_model, run h heads in parallel, each
 532with d_k = d_model / h and its own W_q, W_k and W_v. The total compute is the
 533same, but each head can specialise: one tracks syntax, another which noun a
 534pronoun refers to, another the previous token.
 535
 536```mermaid
 537flowchart LR
 538  X["X (n × d_model)"] --> P["project with W_q, W_k, W_v"]
 539  P --> H1["head 1<br/>attention on d_model/h dims"]
 540  P --> H2["head 2"]
 541  P --> H3["..."]
 542  P --> Hh["head h"]
 543  H1 & H2 & H3 & Hh --> C["concatenate<br/>(n × d_model)"]
 544  C --> WO["mix with W_o"] --> Y["output (n × d_model)"]
 545```
 546
 547**Reading it:** the input is projected once, and the result is sliced into h
 548narrow slabs, one per head. Each head runs the full attention recipe from the
 549first diagram on its own slab, independently and in parallel. The heads'
 550outputs are glued back side by side, and W_o lets them exchange what they
 551found. The output has the same shape as the input, which is what lets blocks
 552stack.
 553
 554**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.
 555
 556## Grouped-query attention: sharing keys and values
 557
 558With **grouped-query attention (GQA)**, several query heads share one
 559key/value head. Llama 3 8B has 32 query heads but only 8 KV heads, so it
 560stores 4× fewer keys and values per token. Multi-query attention (MQA) is the
 561extreme: one KV head for everyone.
 562
 563```mermaid
 564flowchart TB
 565  subgraph MHA["MHA: 4 query heads, 4 KV heads"]
 566    q1a[Q1]-->kv1a[KV1]
 567    q2a[Q2]-->kv2a[KV2]
 568    q3a[Q3]-->kv3a[KV3]
 569    q4a[Q4]-->kv4a[KV4]
 570  end
 571  subgraph GQA["GQA: 4 query heads, 2 KV heads"]
 572    q1b[Q1]-->kv1b[KV1]
 573    q2b[Q2]-->kv1b
 574    q3b[Q3]-->kv2b[KV2]
 575    q4b[Q4]-->kv2b
 576  end
 577  subgraph MQA["MQA: 4 query heads, 1 KV head"]
 578    q1c[Q1]-->kv1c[KV1]
 579    q2c[Q2]-->kv1c
 580    q3c[Q3]-->kv1c
 581    q4c[Q4]-->kv1c
 582  end
 583```
 584
 585**Reading it:** count the KV boxes. Every query head still asks its own
 586question, but in GQA and MQA they read from a shared set of keys and values.
 587The KV boxes are what gets stored in GPU memory for every token of every
 588conversation during generation (the KV cache). Fewer boxes means more
 589conversations fit on one GPU, at a small cost in quality. See
 590`primer.ml.inference` for the memory arithmetic.
 591
 592**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.
 593
 594## Cost: why long context is expensive
 595
 596The score matrix is n × n per head per layer, so compute and memory grow
 597with **n²**. Doubling the context roughly quadruples attention's cost.
 598
 599![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](figures/primer.ml.attention.quadratic_cost.svg)
 600
 601**Reading it:** both axes are logarithmic, so a straight line is a power law
 602and a steeper line grows faster. The projections (X·W) cost O(n·d²), a line of
 603slope 1. The score-and-mix step costs O(n²·d), a line of slope 2. The dashed
 604marker is where they cross, at n = 2·d_model. Past that point, most of the
 605work is tokens comparing themselves to other tokens, and every doubling of
 606context costs 4× there.
 607
 608The standard mitigations each attack this picture. FlashAttention produces
 609the exact same result but computes it in tiles sized for fast on-chip GPU
 610memory, so the n × n matrix is never written to slow memory. Sliding-window
 611and sparse attention let each token see only some others. GQA and MQA shrink
 612the KV cache. State-space models such as Mamba replace attention with a
 613linear-time recurrence.
 614
 615**In code:** `attention_cost` counts the quadratic score-and-mix FLOPs and the linear projection FLOPs plotted above.
 616
 617## In 20 seconds
 618
 619- **Attention:** each token builds a query, key and value. Query-key
 620  similarity decides how much each token listens to each other token, and
 621  the output is a weighted blend of their values.
 622- **Why √d_k:** dot products grow with dimension, and large scores saturate
 623  softmax and kill gradients. Scaling keeps training stable.
 624- **Causal mask:** future scores are set to −∞ before softmax, so each token
 625  sees only the past. That is what makes next-token training honest and
 626  generation cacheable.
 627- **Long context cost:** attention compares every pair of tokens, so it is
 628  O(n²), and the KV cache grows with every token.
 629
 630## Self-test questions
 631
 632**Explain attention to a non-engineer in 30 seconds.**
 633When the model reads a word, it asks which other words here help it
 634understand this one. It scores every other word for relevance, then builds
 635its understanding of the word as a mix of the relevant ones. In "the animal
 636didn't cross the street because it was tired", "it" draws mostly from
 637"animal". It does this for every word, in parallel, dozens of times over.
 638
 639**Now explain it to an ML engineer in two minutes.**
 640Project X into Q, K and V with learned matrices. Compute QKᵀ/√d_k, an n×n
 641matrix of scaled similarities; add a causal mask of −∞ above the diagonal
 642for a decoder; softmax each row; multiply by V. Run h heads in parallel on
 643d_model/h slices, concatenate them, and project with W_o. Wrap the whole
 644thing in a residual connection with pre-layer-norm and follow it with an
 645FFN. Cost is O(n²·d) in compute and O(n²) in memory for the scores, which is
 646why FlashAttention tiles it and GQA shrinks the KV cache.
 647
 648**Why divide by √d_k? What breaks without it?**
 649The variance of q·k grows linearly with d_k. Without scaling, scores spread
 650out, softmax saturates towards one-hot, its Jacobian goes to about zero, and
 651the query/key projections stop receiving gradient. Training stalls or turns
 652unstable.
 653
 654**Why is long context expensive? Name two techniques that reduce the cost.**
 655Scores are n×n per head per layer, so cost is quadratic, and the KV cache
 656grows linearly with every token held in memory. Two fixes: FlashAttention
 657(exact, memory-efficient tiling) and GQA/MQA (fewer KV heads). Others:
 658sliding-window or sparse attention, and prompt caching of stable prefixes.
 659
 660**What does the causal mask buy you besides honest training?**
 661Earlier outputs never depend on later tokens, so during generation the
 662keys and values of past tokens can be computed once and cached. Each new
 663token then costs one row of attention instead of a full recomputation.
 664
 665**What does GQA trade away, and for what?**
 666A little modelling capacity (query heads share keys and values) for a KV
 667cache that is several times smaller. That means more concurrent requests
 668and longer contexts per GPU.
 669
 670## The papers behind this lesson
 671
 672- **Vaswani et al., *Attention Is All You Need* (2017)**:
 673  https://arxiv.org/abs/1706.03762. Showed that attention alone, with no
 674  recurrence, is enough for state-of-the-art translation, and introduced
 675  scaled dot-product attention, multi-head attention and the transformer.
 676  [Annotated companion](../../papers/attention-is-all-you-need.html)
 677- **Ainslie et al., *GQA: Training Generalized Multi-Query Transformer
 678  Models from Multi-Head Checkpoints* (2023)**:
 679  https://arxiv.org/abs/2305.13245. Introduced grouped-query attention, the
 680  middle ground between multi-head and multi-query attention.
 681- **Dao et al., *FlashAttention* (2022)**: https://arxiv.org/abs/2205.14135.
 682  Computes exact attention in tiles sized for fast on-chip GPU memory,
 683  making long contexts practical.
 684  [Annotated companion](../../papers/flashattention.html)
 685
 686## Further reading
 687
 688- Vaswani et al., *Attention Is All You Need* (2017): https://arxiv.org/abs/1706.03762
 689- Jay Alammar, *The Illustrated Transformer*: https://jalammar.github.io/illustrated-transformer/
 690- Harvard NLP, *The Annotated Transformer* (the paper, line by line in code): https://nlp.seas.harvard.edu/annotated-transformer/
 691- 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
 692- 3Blue1Brown, *Attention in transformers, visually explained*: https://www.youtube.com/watch?v=eMlx5fFNoYc
 693- PyTorch `scaled_dot_product_attention`: https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html
 694- Ainslie et al., *GQA* (2023): https://arxiv.org/abs/2305.13245
 695- Dao et al., *FlashAttention* (2022): https://arxiv.org/abs/2205.14135
 696"""
 697
 698from __future__ import annotations
 699
 700import numpy as np
 701
 702from primer._show import banner, matrix, say, takeaway, table
 703
 704# ---------------------------------------------------------------------------
 705# 1. Softmax: turns any vector of scores into probabilities that sum to 1
 706# ---------------------------------------------------------------------------
 707
 708
 709def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray:
 710    """Numerically stable softmax along `axis`.
 711
 712    softmax(x)_i = exp(x_i) / sum_j exp(x_j)
 713
 714    Subtracting the max first doesn't change the result (it cancels in the
 715    ratio) but prevents overflow: exp(1000) is inf, exp(1000 - 1000) = 1.
 716    Entries equal to -inf (masked positions) come out as exactly 0.
 717    """
 718    x_max = np.max(x, axis=axis, keepdims=True)
 719    # If a whole row is -inf (fully masked), avoid inf - inf = nan.
 720    x_max = np.where(np.isfinite(x_max), x_max, 0.0)
 721    e = np.exp(x - x_max)
 722    return e / np.sum(e, axis=axis, keepdims=True)
 723
 724
 725def softmax_jacobian(p: np.ndarray) -> np.ndarray:
 726    """d softmax / d scores for one probability vector p: diag(p) - p pᵀ.
 727
 728    When p is one-hot (saturated), every entry is ~0, so no gradient
 729    reaches the scores and the query/key weights stop learning. That is the
 730    failure the sqrt(d_k) scaling prevents.
 731    """
 732    return np.diag(p) - np.outer(p, p)
 733
 734
 735# ---------------------------------------------------------------------------
 736# 2. Scaled dot-product attention
 737# ---------------------------------------------------------------------------
 738
 739
 740def causal_mask(n: int) -> np.ndarray:
 741    """Boolean (n, n) mask, True where attention is ALLOWED.
 742
 743    Row i may attend to columns 0..i (itself and the past), never the future:
 744
 745        n=4 ->  [[1 0 0 0]
 746                 [1 1 0 0]
 747                 [1 1 1 0]
 748                 [1 1 1 1]]
 749    """
 750    return np.tril(np.ones((n, n), dtype=bool))
 751
 752
 753def scaled_dot_product_attention(
 754    Q: np.ndarray,
 755    K: np.ndarray,
 756    V: np.ndarray,
 757    mask: np.ndarray | None = None,
 758    scale: bool = True,
 759) -> tuple[np.ndarray, np.ndarray]:
 760    """softmax(Q Kᵀ / sqrt(d_k)) V, the formula, one line per step.
 761
 762    Args:
 763        Q: (..., n_q, d_k) queries. Leading dims (batch, heads) broadcast.
 764        K: (..., n_k, d_k) keys.
 765        V: (..., n_k, d_v) values.
 766        mask: boolean (n_q, n_k), True = may attend. None = attend everywhere.
 767        scale: divide by sqrt(d_k). Set False only to demonstrate why you shouldn't.
 768
 769    Returns:
 770        (output, weights): output is (..., n_q, d_v); weights is
 771        (..., n_q, n_k) with rows summing to 1.
 772    """
 773    d_k = Q.shape[-1]
 774
 775    # Step 1: scores. Every query dotted with every key: (n_q, d_k) @ (d_k, n_k).
 776    # High dot product = vectors point the same way = "this key is relevant".
 777    scores = Q @ np.swapaxes(K, -1, -2)
 778
 779    # Step 2: scale so score variance is ~1 regardless of d_k.
 780    if scale:
 781        scores = scores / np.sqrt(d_k)
 782
 783    # Step 3: mask. -inf before softmax -> weight exactly 0 after.
 784    if mask is not None:
 785        scores = np.where(mask, scores, -np.inf)
 786
 787    # Step 4: softmax each row -> attention weights (each row sums to 1).
 788    weights = softmax(scores, axis=-1)
 789
 790    # Step 5: weighted sum of values: (n_q, n_k) @ (n_k, d_v).
 791    return weights @ V, weights
 792
 793
 794# ---------------------------------------------------------------------------
 795# 3. Multi-head attention (with optional grouped-query attention)
 796# ---------------------------------------------------------------------------
 797
 798
 799class MultiHeadAttention:
 800    """Multi-head self-attention with learned projections (forward pass only).
 801
 802    `n_kv_heads < n_heads` gives grouped-query attention (GQA): each K/V head
 803    is shared by `n_heads // n_kv_heads` query heads. `n_kv_heads == 1` is
 804    multi-query attention (MQA). `n_kv_heads == n_heads` is classic MHA.
 805
 806    Weights are random here (we're studying the mechanics, not training);
 807    `primer.ml.transformer` stacks this into a full block.
 808    """
 809
 810    def __init__(self, d_model: int, n_heads: int, n_kv_heads: int | None = None, seed: int = 0):
 811        n_kv_heads = n_kv_heads or n_heads
 812        assert d_model % n_heads == 0, "d_model must split evenly across heads"
 813        assert n_heads % n_kv_heads == 0, "query heads must group evenly over KV heads"
 814        self.d_model, self.n_heads, self.n_kv_heads = d_model, n_heads, n_kv_heads
 815        self.d_head = d_model // n_heads
 816        rng = np.random.default_rng(seed)
 817        # Scaled init keeps activations ~unit variance (see primer.ml.deep_nets).
 818        s = 1 / np.sqrt(d_model)
 819        self.W_q = rng.normal(0, s, (d_model, n_heads * self.d_head))
 820        self.W_k = rng.normal(0, s, (d_model, n_kv_heads * self.d_head))  # smaller when GQA
 821        self.W_v = rng.normal(0, s, (d_model, n_kv_heads * self.d_head))
 822        self.W_o = rng.normal(0, s, (n_heads * self.d_head, d_model))
 823
 824    def _split(self, x: np.ndarray, n: int) -> np.ndarray:
 825        # (seq, n*d_head) -> (n, seq, d_head): one slab per head.
 826        seq = x.shape[0]
 827        return x.reshape(seq, n, self.d_head).transpose(1, 0, 2)
 828
 829    def __call__(self, X: np.ndarray, causal: bool = True) -> tuple[np.ndarray, np.ndarray]:
 830        """X: (seq, d_model) -> (output (seq, d_model), weights (n_heads, seq, seq))."""
 831        seq = X.shape[0]
 832        Q = self._split(X @ self.W_q, self.n_heads)  # (H,   seq, d_head)
 833        K = self._split(X @ self.W_k, self.n_kv_heads)  # (Hkv, seq, d_head)
 834        V = self._split(X @ self.W_v, self.n_kv_heads)
 835
 836        # GQA: repeat each KV head for the query heads in its group.
 837        # (Real kernels avoid materializing the copies; the math is the same.)
 838        group = self.n_heads // self.n_kv_heads
 839        K = np.repeat(K, group, axis=0)
 840        V = np.repeat(V, group, axis=0)
 841
 842        mask = causal_mask(seq) if causal else None
 843        heads, weights = scaled_dot_product_attention(Q, K, V, mask)  # (H, seq, d_head)
 844
 845        # Concatenate heads back to (seq, H*d_head), then mix with W_o.
 846        concat = heads.transpose(1, 0, 2).reshape(seq, self.n_heads * self.d_head)
 847        return concat @ self.W_o, weights
 848
 849    def kv_params(self) -> int:
 850        """Parameters in the K and V projections. GQA shrinks this, and the KV cache."""
 851        return self.W_k.size + self.W_v.size
 852
 853
 854# ---------------------------------------------------------------------------
 855# 4. Worked examples
 856# ---------------------------------------------------------------------------
 857
 858IT_EXAMPLE_TOKENS = ["animal", "tired", "street"]
 859IT_EXAMPLE_SCORES = np.array([2.0, 1.0, 0.5])  # already scaled
 860
 861
 862def worked_example_it() -> dict[str, float]:
 863    """Worked example: the token "it" attending over three keys.
 864
 865    | Token  | Scaled score | e^score | Weight |
 866    |--------|--------------|---------|--------|
 867    | animal | 2.0          | 7.39    | 0.63   |
 868    | tired  | 1.0          | 2.72    | 0.23   |
 869    | street | 0.5          | 1.65    | 0.14   |
 870    """
 871    w = softmax(IT_EXAMPLE_SCORES)
 872    return dict(zip(IT_EXAMPLE_TOKENS, w.tolist()))
 873
 874
 875SENTENCE = "The animal didn't cross the street because it was tired".split()
 876
 877# Hand-picked, not learned, so the pattern means something. The query of "it"
 878# and the keys of "animal", "tired" and "street" are the worked example's, so
 879# the "it" row still scores them 2.0, 1.0 and 0.5. The other rows are chosen so
 880# each word looks for a sensible partner: "cross" for its subject, "street"
 881# for the verb it belongs to, "was" for "it". The two "the"s share a key,
 882# because without word positions identical words look identical.
 883_SENTENCE_QK = {
 884    #          query            key
 885    "The": ([0, 2, 2, 2], [0, 0, 0, -1]),
 886    "animal": ([0, -1, 0, -3], [1, 1, 1, 1]),
 887    "didn't": ([0, 0, 3, 0], [0, 0, -1, 0]),
 888    "cross": ([2, 0, 0, 2], [0, -1, 2, 0]),
 889    "the": ([3, -1, -1, 1], [0, 0, 0, -1]),
 890    "street": ([0, -1, 3, 0], [1, 0, 0, 0]),
 891    "because": ([0, 0, -3, 0], [-1, 0, 0, 0]),
 892    "it": ([1, 1, 1, 1], [0, 1, 0, -1]),
 893    "was": ([0, 3, 0, -1], [0, -1, 0, 0]),
 894    "tired": ([1, 2, 0, 1], [1, 1, 0, 0]),
 895}
 896
 897
 898def worked_example_sentence() -> tuple[np.ndarray, np.ndarray]:
 899    """The whole sentence's queries and keys, each (10, 4): one row per word of `SENTENCE`.
 900
 901    `scaled_dot_product_attention(Q, K, K)` turns them into the 10 × 10
 902    attention weights the site's interactive heatmap draws.
 903    """
 904    Q = np.array([_SENTENCE_QK[w][0] for w in SENTENCE], dtype=float)
 905    K = np.array([_SENTENCE_QK[w][1] for w in SENTENCE], dtype=float)
 906    return Q, K
 907
 908
 909def viz_data() -> dict:
 910    """The numbers the site's interactive attention heatmap starts from."""
 911    # Only Q and K: the widget recomputes softmax(QKᵀ / √d_k) itself, and
 912    # tests/test_attention.py checks it against scaled_dot_product_attention.
 913    Q, K = worked_example_sentence()
 914    return {"attention-matrix": {"tokens": SENTENCE, "Q": Q.tolist(), "K": K.tolist()}}
 915
 916
 917def sqrt_dk_experiment(d_ks=(4, 64, 512), n_keys: int = 16, trials: int = 2000, seed: int = 0) -> list[dict]:
 918    """Measure what happens to attention scores as d_k grows, with and without scaling.
 919
 920    For each d_k we sample random unit-variance q and k, then record:
 921      * std of the raw dot product (theory: sqrt(d_k))
 922      * mean max softmax weight (≈1.0 means saturated / one-hot)
 923      * mean softmax-Jacobian norm (≈0 means no gradient flows back)
 924    """
 925    rng = np.random.default_rng(seed)
 926    rows = []
 927    for d in d_ks:
 928        q = rng.standard_normal((trials, 1, d))
 929        k = rng.standard_normal((trials, n_keys, d))
 930        raw = (q @ k.transpose(0, 2, 1))[:, 0, :]  # (trials, n_keys)
 931        for scaled in (False, True):
 932            s = raw / np.sqrt(d) if scaled else raw
 933            p = softmax(s)
 934            jac = np.mean([np.linalg.norm(softmax_jacobian(pi)) for pi in p[:200]])
 935            rows.append(
 936                dict(d_k=d, scaled=scaled, score_std=float(s.std()), max_weight=float(p.max(axis=1).mean()), grad_norm=float(jac))
 937            )
 938    return rows
 939
 940
 941def attention_cost(n: int, d_model: int, n_layers: int = 1) -> dict[str, float]:
 942    """Rough FLOPs of the attention *score* and *mix* steps for n tokens.
 943
 944    QKᵀ is n·n·d multiply-adds, and weights·V is another n·n·d. Count 2 FLOPs
 945    per multiply-add. The projections (X·W) are O(n·d²), linear in n, so for
 946    long contexts the n² term dominates.
 947    """
 948    score_and_mix = 2 * (2 * n * n * d_model) * n_layers
 949    projections = 2 * (4 * n * d_model * d_model) * n_layers
 950    return dict(n=n, quadratic_flops=score_and_mix, linear_flops=projections, score_matrix_entries=n * n)
 951
 952
 953# ---------------------------------------------------------------------------
 954# 5. Figures (rendered into the HTML docs by `make figures`)
 955# ---------------------------------------------------------------------------
 956
 957
 958def figures() -> dict:
 959    """Plot this lesson's data. matplotlib is imported here, and only here,
 960    so the lesson itself needs nothing beyond NumPy."""
 961    import matplotlib
 962
 963    matplotlib.use("Agg")
 964    import matplotlib.pyplot as plt
 965
 966    SCALED, UNSCALED, MUTED = "#2563eb", "#dc2626", "#9ca3af"
 967    figs = {}
 968
 969    # --- 1. Scores -> softmax weights for "it" ------------------------------
 970    fig, ax = plt.subplots(figsize=(6, 3.2))
 971    x = np.arange(len(IT_EXAMPLE_TOKENS))
 972    weights = softmax(IT_EXAMPLE_SCORES)
 973    ax.bar(x - 0.2, IT_EXAMPLE_SCORES / IT_EXAMPLE_SCORES.sum(), 0.4, color=MUTED, label="score (share of total)")
 974    ax.bar(x + 0.2, weights, 0.4, color=SCALED, label="attention weight (softmax)")
 975    for xi, w in zip(x, weights):
 976        ax.text(xi + 0.2, w + 0.02, f"{w:.2f}", ha="center")
 977    ax.set_xticks(x, IT_EXAMPLE_TOKENS)
 978    ax.set_ylabel("share")
 979    ax.set_ylim(0, 0.78)
 980    ax.set_title('What "it" attends to: softmax sharpens the scores')
 981    ax.legend(frameon=False)
 982    figs["it_weights"] = fig
 983
 984    # --- 2. Causal attention heatmap on a sentence --------------------------
 985    tokens = "The animal didn't cross the street because it was tired".split()
 986    rng = np.random.default_rng(7)
 987    X = rng.standard_normal((len(tokens), 32))
 988    _, w = MultiHeadAttention(32, n_heads=1, seed=3)(X)
 989    fig, ax = plt.subplots(figsize=(6, 5))
 990    im = ax.imshow(w[0], cmap="Blues", vmin=0, vmax=1)
 991    ax.set_xticks(range(len(tokens)), tokens, rotation=45, ha="right")
 992    ax.set_yticks(range(len(tokens)), tokens)
 993    ax.set_xlabel("key: the token being looked at")
 994    ax.set_ylabel("query: the token doing the looking")
 995    ax.set_title("Causal attention: each row sees only the past")
 996    ax.grid(False)
 997    fig.colorbar(im, ax=ax, fraction=0.046, label="attention weight")
 998    figs["causal_heatmap"] = fig
 999
1000    # --- 3. Why sqrt(d_k): saturation and vanishing gradient ----------------
1001    dks = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024)
1002    rows = sqrt_dk_experiment(d_ks=dks, trials=600)
1003    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4))
1004    for scaled, color, label in ((False, UNSCALED, "unscaled  QKᵀ"), (True, SCALED, "scaled  QKᵀ/√d_k")):
1005        sel = [r for r in rows if r["scaled"] == scaled]
1006        a1.plot(dks, [r["max_weight"] for r in sel], "o-", color=color, label=label)
1007        a2.plot(dks, [r["grad_norm"] for r in sel], "o-", color=color, label=label)
1008    for a, title, ylabel in (
1009        (a1, "Softmax saturates (→ one-hot)", "mean largest weight"),
1010        (a2, "Gradient through softmax vanishes", "softmax Jacobian norm"),
1011    ):
1012        a.set_xscale("log", base=2)
1013        a.set_xlabel("head width d_k")
1014        a.set_ylabel(ylabel)
1015        a.set_title(title)
1016    a1.set_ylim(0, 1.05)
1017    a1.legend(frameon=False)
1018    fig.tight_layout()
1019    figs["sqrt_dk"] = fig
1020
1021    # --- 4. Quadratic vs linear cost ----------------------------------------
1022    d_model = 4096
1023    ns = np.logspace(2, 6, 60)
1024    quad = [attention_cost(int(n), d_model)["quadratic_flops"] for n in ns]
1025    lin = [attention_cost(int(n), d_model)["linear_flops"] for n in ns]
1026    fig, ax = plt.subplots(figsize=(6, 3.6))
1027    ax.loglog(ns, quad, color=UNSCALED, label="scores + mix  O(n²·d)")
1028    ax.loglog(ns, lin, color=SCALED, label="projections  O(n·d²)")
1029    ax.axvline(2 * d_model, color=MUTED, ls="--")
1030    ax.text(2 * d_model * 1.15, min(quad) * 10, "crossover\nn = 2·d_model", color="#4b5563")
1031    ax.set_xlabel("context length n (tokens)")
1032    ax.set_ylabel("FLOPs per layer")
1033    ax.set_title(f"Why long context is expensive (d_model = {d_model})")
1034    ax.legend(frameon=False)
1035    figs["quadratic_cost"] = fig
1036
1037    return figs
1038
1039
1040# ---------------------------------------------------------------------------
1041# 6. Narrated walkthrough
1042# ---------------------------------------------------------------------------
1043
1044
1045def demo() -> None:
1046    banner("1. Worked example: what does 'it' attend to?")
1047    say(
1048        """
1049        "The animal didn't cross the street because it was tired." Suppose the
1050        query for "it", dotted with three keys and scaled, gives scores 2.0,
1051        1.0 and 0.5. Softmax exponentiates and normalizes:
1052        """
1053    )
1054    w = worked_example_it()
1055    e = np.exp(IT_EXAMPLE_SCORES)
1056    table(
1057        ["token", "score", "e^score", "weight"],
1058        [(t, s, ei, wi) for t, s, ei, wi in zip(IT_EXAMPLE_TOKENS, IT_EXAMPLE_SCORES, e, w.values())],
1059        floatfmt=".2f",
1060    )
1061    say(f"Total of e^score = {e.sum():.2f}. New 'it' = 0.63·V(animal) + 0.23·V(tired) + 0.14·V(street).")
1062    takeaway("The model resolved the pronoun: most of 'it' now comes from 'animal'.")
1063
1064    banner("2. Full attention on a tiny sequence (shapes and causal mask)")
1065    rng = np.random.default_rng(0)
1066    n, d_k = 4, 8
1067    Q, K, V = (rng.standard_normal((n, d_k)) for _ in range(3))
1068    out, weights = scaled_dot_product_attention(Q, K, V, mask=causal_mask(n))
1069    say(f"Q, K, V are each ({n}, {d_k}). Scores Q·Kᵀ are ({n}, {n}): every token vs. every token.")
1070    matrix("causal attention weights (row = query token, col = key token)", weights)
1071    say(
1072        """
1073        Upper triangle is exactly zero: token 0 sees only itself, token 3 sees
1074        all four. Every row sums to 1. Output shape is (4, 8): one new
1075        context-aware vector per token.
1076        """
1077    )
1078
1079    banner("3. Why divide by sqrt(d_k)? Measure it.")
1080    rows = sqrt_dk_experiment()
1081    table(
1082        ["d_k", "scaled?", "score std", "mean max weight", "softmax grad norm"],
1083        [(r["d_k"], "yes" if r["scaled"] else "no", r["score_std"], r["max_weight"], r["grad_norm"]) for r in rows],
1084        floatfmt=".3f",
1085    )
1086    say(
1087        """
1088        Unscaled, the score std grows like sqrt(d_k) (2, 8, ~22.6). At d_k=512
1089        the softmax is essentially one-hot (max weight near 1) and the gradient
1090        norm collapses toward 0. Scaled, std stays ~1 and gradients stay
1091        healthy at every width.
1092        """
1093    )
1094    takeaway(
1095        "Dot products grow with dimension; big scores saturate softmax and kill "
1096        "gradients. Dividing by sqrt(d_k) keeps the variance at 1 so training stays stable."
1097    )
1098
1099    banner("4. Multi-head vs. grouped-query attention")
1100    X = rng.standard_normal((6, 64))
1101    mha = MultiHeadAttention(d_model=64, n_heads=8)
1102    gqa = MultiHeadAttention(d_model=64, n_heads=8, n_kv_heads=2)
1103    mqa = MultiHeadAttention(d_model=64, n_heads=8, n_kv_heads=1)
1104    for name, layer in [("MHA (8 KV heads)", mha), ("GQA (2 KV heads)", gqa), ("MQA (1 KV head)", mqa)]:
1105        y, wts = layer(X)
1106        print(f"{name:18s} out {y.shape}, weights {wts.shape}, K+V params {layer.kv_params():5d}")
1107    print()
1108    say(
1109        """
1110        All three produce the same output shape. GQA and MQA keep 8 query heads
1111        but store 4x and 8x less K/V. The KV cache (what you keep in GPU memory
1112        per token while generating) shrinks by the same factor.
1113        """
1114    )
1115
1116    banner("5. Why long context is expensive: O(n²)")
1117    table(
1118        ["tokens n", "score-matrix entries", "quadratic FLOPs", "linear FLOPs"],
1119        [
1120            (c["n"], f"{c['score_matrix_entries']:,}", f"{c['quadratic_flops']:.2e}", f"{c['linear_flops']:.2e}")
1121            for c in (attention_cost(n, 4096) for n in (1_000, 2_000, 8_000, 32_000, 128_000))
1122        ],
1123    )
1124    say(
1125        """
1126        Doubling n from 1k to 2k quadruples the score matrix. With d_model=4096
1127        the quadratic term overtakes the projections once n > 2·d_model.
1128        Mitigations: FlashAttention (tiles; never materializes n×n in slow
1129        memory), sliding-window/sparse attention, GQA for the KV cache, and
1130        prompt caching to avoid recomputing stable prefixes.
1131        """
1132    )
1133
1134
1135if __name__ == "__main__":
1136    demo()
Level 3: the code, function by function.
def softmax(x: numpy.ndarray, axis: int = -1) -> numpy.ndarray: on GitHub
710def softmax(x: np.ndarray, axis: int = -1) -> np.ndarray:
711    """Numerically stable softmax along `axis`.
712
713    softmax(x)_i = exp(x_i) / sum_j exp(x_j)
714
715    Subtracting the max first doesn't change the result (it cancels in the
716    ratio) but prevents overflow: exp(1000) is inf, exp(1000 - 1000) = 1.
717    Entries equal to -inf (masked positions) come out as exactly 0.
718    """
719    x_max = np.max(x, axis=axis, keepdims=True)
720    # If a whole row is -inf (fully masked), avoid inf - inf = nan.
721    x_max = np.where(np.isfinite(x_max), x_max, 0.0)
722    e = np.exp(x - x_max)
723    return e / np.sum(e, axis=axis, keepdims=True)

Numerically stable softmax along axis.

softmax(x)_i = exp(x_i) / sum_j exp(x_j)

Subtracting the max first doesn't change the result (it cancels in the ratio) but prevents overflow: exp(1000) is inf, exp(1000 - 1000) = 1. Entries equal to -inf (masked positions) come out as exactly 0.

def softmax_jacobian(p: numpy.ndarray) -> numpy.ndarray: on GitHub
726def softmax_jacobian(p: np.ndarray) -> np.ndarray:
727    """d softmax / d scores for one probability vector p: diag(p) - p pᵀ.
728
729    When p is one-hot (saturated), every entry is ~0, so no gradient
730    reaches the scores and the query/key weights stop learning. That is the
731    failure the sqrt(d_k) scaling prevents.
732    """
733    return np.diag(p) - np.outer(p, p)

d softmax / d scores for one probability vector p: diag(p) - p pᵀ.

When p is one-hot (saturated), every entry is ~0, so no gradient reaches the scores and the query/key weights stop learning. That is the failure the sqrt(d_k) scaling prevents.

def causal_mask(n: int) -> numpy.ndarray: on GitHub
741def causal_mask(n: int) -> np.ndarray:
742    """Boolean (n, n) mask, True where attention is ALLOWED.
743
744    Row i may attend to columns 0..i (itself and the past), never the future:
745
746        n=4 ->  [[1 0 0 0]
747                 [1 1 0 0]
748                 [1 1 1 0]
749                 [1 1 1 1]]
750    """
751    return np.tril(np.ones((n, n), dtype=bool))

Boolean (n, n) mask, True where attention is ALLOWED.

Row i may attend to columns 0..i (itself and the past), never the future:

n=4 ->  [[1 0 0 0]
         [1 1 0 0]
         [1 1 1 0]
         [1 1 1 1]]
def scaled_dot_product_attention( Q: numpy.ndarray, K: numpy.ndarray, V: numpy.ndarray, mask: numpy.ndarray | None = None, scale: bool = True) -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
754def scaled_dot_product_attention(
755    Q: np.ndarray,
756    K: np.ndarray,
757    V: np.ndarray,
758    mask: np.ndarray | None = None,
759    scale: bool = True,
760) -> tuple[np.ndarray, np.ndarray]:
761    """softmax(Q Kᵀ / sqrt(d_k)) V, the formula, one line per step.
762
763    Args:
764        Q: (..., n_q, d_k) queries. Leading dims (batch, heads) broadcast.
765        K: (..., n_k, d_k) keys.
766        V: (..., n_k, d_v) values.
767        mask: boolean (n_q, n_k), True = may attend. None = attend everywhere.
768        scale: divide by sqrt(d_k). Set False only to demonstrate why you shouldn't.
769
770    Returns:
771        (output, weights): output is (..., n_q, d_v); weights is
772        (..., n_q, n_k) with rows summing to 1.
773    """
774    d_k = Q.shape[-1]
775
776    # Step 1: scores. Every query dotted with every key: (n_q, d_k) @ (d_k, n_k).
777    # High dot product = vectors point the same way = "this key is relevant".
778    scores = Q @ np.swapaxes(K, -1, -2)
779
780    # Step 2: scale so score variance is ~1 regardless of d_k.
781    if scale:
782        scores = scores / np.sqrt(d_k)
783
784    # Step 3: mask. -inf before softmax -> weight exactly 0 after.
785    if mask is not None:
786        scores = np.where(mask, scores, -np.inf)
787
788    # Step 4: softmax each row -> attention weights (each row sums to 1).
789    weights = softmax(scores, axis=-1)
790
791    # Step 5: weighted sum of values: (n_q, n_k) @ (n_k, d_v).
792    return weights @ V, weights

softmax(Q Kᵀ / sqrt(d_k)) V, the formula, one line per step.

Arguments:

  • Q: (..., n_q, d_k) queries. Leading dims (batch, heads) broadcast.
  • K: (..., n_k, d_k) keys.
  • V: (..., n_k, d_v) values.
  • mask: boolean (n_q, n_k), True = may attend. None = attend everywhere.
  • scale: divide by sqrt(d_k). Set False only to demonstrate why you shouldn't.

Returns:

(output, weights): output is (..., n_q, d_v); weights is (..., n_q, n_k) with rows summing to 1.

class MultiHeadAttention: on GitHub
800class MultiHeadAttention:
801    """Multi-head self-attention with learned projections (forward pass only).
802
803    `n_kv_heads < n_heads` gives grouped-query attention (GQA): each K/V head
804    is shared by `n_heads // n_kv_heads` query heads. `n_kv_heads == 1` is
805    multi-query attention (MQA). `n_kv_heads == n_heads` is classic MHA.
806
807    Weights are random here (we're studying the mechanics, not training);
808    `primer.ml.transformer` stacks this into a full block.
809    """
810
811    def __init__(self, d_model: int, n_heads: int, n_kv_heads: int | None = None, seed: int = 0):
812        n_kv_heads = n_kv_heads or n_heads
813        assert d_model % n_heads == 0, "d_model must split evenly across heads"
814        assert n_heads % n_kv_heads == 0, "query heads must group evenly over KV heads"
815        self.d_model, self.n_heads, self.n_kv_heads = d_model, n_heads, n_kv_heads
816        self.d_head = d_model // n_heads
817        rng = np.random.default_rng(seed)
818        # Scaled init keeps activations ~unit variance (see primer.ml.deep_nets).
819        s = 1 / np.sqrt(d_model)
820        self.W_q = rng.normal(0, s, (d_model, n_heads * self.d_head))
821        self.W_k = rng.normal(0, s, (d_model, n_kv_heads * self.d_head))  # smaller when GQA
822        self.W_v = rng.normal(0, s, (d_model, n_kv_heads * self.d_head))
823        self.W_o = rng.normal(0, s, (n_heads * self.d_head, d_model))
824
825    def _split(self, x: np.ndarray, n: int) -> np.ndarray:
826        # (seq, n*d_head) -> (n, seq, d_head): one slab per head.
827        seq = x.shape[0]
828        return x.reshape(seq, n, self.d_head).transpose(1, 0, 2)
829
830    def __call__(self, X: np.ndarray, causal: bool = True) -> tuple[np.ndarray, np.ndarray]:
831        """X: (seq, d_model) -> (output (seq, d_model), weights (n_heads, seq, seq))."""
832        seq = X.shape[0]
833        Q = self._split(X @ self.W_q, self.n_heads)  # (H,   seq, d_head)
834        K = self._split(X @ self.W_k, self.n_kv_heads)  # (Hkv, seq, d_head)
835        V = self._split(X @ self.W_v, self.n_kv_heads)
836
837        # GQA: repeat each KV head for the query heads in its group.
838        # (Real kernels avoid materializing the copies; the math is the same.)
839        group = self.n_heads // self.n_kv_heads
840        K = np.repeat(K, group, axis=0)
841        V = np.repeat(V, group, axis=0)
842
843        mask = causal_mask(seq) if causal else None
844        heads, weights = scaled_dot_product_attention(Q, K, V, mask)  # (H, seq, d_head)
845
846        # Concatenate heads back to (seq, H*d_head), then mix with W_o.
847        concat = heads.transpose(1, 0, 2).reshape(seq, self.n_heads * self.d_head)
848        return concat @ self.W_o, weights
849
850    def kv_params(self) -> int:
851        """Parameters in the K and V projections. GQA shrinks this, and the KV cache."""
852        return self.W_k.size + self.W_v.size

Multi-head self-attention with learned projections (forward pass only).

n_kv_heads < n_heads gives grouped-query attention (GQA): each K/V head is shared by n_heads // n_kv_heads query heads. n_kv_heads == 1 is multi-query attention (MQA). n_kv_heads == n_heads is classic MHA.

Weights are random here (we're studying the mechanics, not training); primer.ml.transformer stacks this into a full block.

MultiHeadAttention( d_model: int, n_heads: int, n_kv_heads: int | None = None, seed: int = 0) on GitHub
811    def __init__(self, d_model: int, n_heads: int, n_kv_heads: int | None = None, seed: int = 0):
812        n_kv_heads = n_kv_heads or n_heads
813        assert d_model % n_heads == 0, "d_model must split evenly across heads"
814        assert n_heads % n_kv_heads == 0, "query heads must group evenly over KV heads"
815        self.d_model, self.n_heads, self.n_kv_heads = d_model, n_heads, n_kv_heads
816        self.d_head = d_model // n_heads
817        rng = np.random.default_rng(seed)
818        # Scaled init keeps activations ~unit variance (see primer.ml.deep_nets).
819        s = 1 / np.sqrt(d_model)
820        self.W_q = rng.normal(0, s, (d_model, n_heads * self.d_head))
821        self.W_k = rng.normal(0, s, (d_model, n_kv_heads * self.d_head))  # smaller when GQA
822        self.W_v = rng.normal(0, s, (d_model, n_kv_heads * self.d_head))
823        self.W_o = rng.normal(0, s, (n_heads * self.d_head, d_model))
W_q
W_k
W_v
W_o
def kv_params(self) -> int: on GitHub
850    def kv_params(self) -> int:
851        """Parameters in the K and V projections. GQA shrinks this, and the KV cache."""
852        return self.W_k.size + self.W_v.size

Parameters in the K and V projections. GQA shrinks this, and the KV cache.

IT_EXAMPLE_TOKENS = ['animal', 'tired', 'street']
IT_EXAMPLE_SCORES = array([2. , 1. , 0.5])
def worked_example_it() -> dict[str, float]: on GitHub
863def worked_example_it() -> dict[str, float]:
864    """Worked example: the token "it" attending over three keys.
865
866    | Token  | Scaled score | e^score | Weight |
867    |--------|--------------|---------|--------|
868    | animal | 2.0          | 7.39    | 0.63   |
869    | tired  | 1.0          | 2.72    | 0.23   |
870    | street | 0.5          | 1.65    | 0.14   |
871    """
872    w = softmax(IT_EXAMPLE_SCORES)
873    return dict(zip(IT_EXAMPLE_TOKENS, w.tolist()))

Worked example: the token "it" attending over three keys.

Token Scaled score e^score Weight
animal 2.0 7.39 0.63
tired 1.0 2.72 0.23
street 0.5 1.65 0.14
SENTENCE = ['The', 'animal', "didn't", 'cross', 'the', 'street', 'because', 'it', 'was', 'tired']
def worked_example_sentence() -> tuple[numpy.ndarray, numpy.ndarray]: on GitHub
899def worked_example_sentence() -> tuple[np.ndarray, np.ndarray]:
900    """The whole sentence's queries and keys, each (10, 4): one row per word of `SENTENCE`.
901
902    `scaled_dot_product_attention(Q, K, K)` turns them into the 10 × 10
903    attention weights the site's interactive heatmap draws.
904    """
905    Q = np.array([_SENTENCE_QK[w][0] for w in SENTENCE], dtype=float)
906    K = np.array([_SENTENCE_QK[w][1] for w in SENTENCE], dtype=float)
907    return Q, K

The whole sentence's queries and keys, each (10, 4): one row per word of SENTENCE.

scaled_dot_product_attention(Q, K, K) turns them into the 10 × 10 attention weights the site's interactive heatmap draws.

def viz_data() -> dict: on GitHub
910def viz_data() -> dict:
911    """The numbers the site's interactive attention heatmap starts from."""
912    # Only Q and K: the widget recomputes softmax(QKᵀ / √d_k) itself, and
913    # tests/test_attention.py checks it against scaled_dot_product_attention.
914    Q, K = worked_example_sentence()
915    return {"attention-matrix": {"tokens": SENTENCE, "Q": Q.tolist(), "K": K.tolist()}}

The numbers the site's interactive attention heatmap starts from.

def sqrt_dk_experiment( d_ks=(4, 64, 512), n_keys: int = 16, trials: int = 2000, seed: int = 0) -> list[dict]: on GitHub
918def sqrt_dk_experiment(d_ks=(4, 64, 512), n_keys: int = 16, trials: int = 2000, seed: int = 0) -> list[dict]:
919    """Measure what happens to attention scores as d_k grows, with and without scaling.
920
921    For each d_k we sample random unit-variance q and k, then record:
922      * std of the raw dot product (theory: sqrt(d_k))
923      * mean max softmax weight (≈1.0 means saturated / one-hot)
924      * mean softmax-Jacobian norm (≈0 means no gradient flows back)
925    """
926    rng = np.random.default_rng(seed)
927    rows = []
928    for d in d_ks:
929        q = rng.standard_normal((trials, 1, d))
930        k = rng.standard_normal((trials, n_keys, d))
931        raw = (q @ k.transpose(0, 2, 1))[:, 0, :]  # (trials, n_keys)
932        for scaled in (False, True):
933            s = raw / np.sqrt(d) if scaled else raw
934            p = softmax(s)
935            jac = np.mean([np.linalg.norm(softmax_jacobian(pi)) for pi in p[:200]])
936            rows.append(
937                dict(d_k=d, scaled=scaled, score_std=float(s.std()), max_weight=float(p.max(axis=1).mean()), grad_norm=float(jac))
938            )
939    return rows

Measure what happens to attention scores as d_k grows, with and without scaling.

For each d_k we sample random unit-variance q and k, then record:

  • std of the raw dot product (theory: sqrt(d_k))
  • mean max softmax weight (≈1.0 means saturated / one-hot)
  • mean softmax-Jacobian norm (≈0 means no gradient flows back)
def attention_cost(n: int, d_model: int, n_layers: int = 1) -> dict[str, float]: on GitHub
942def attention_cost(n: int, d_model: int, n_layers: int = 1) -> dict[str, float]:
943    """Rough FLOPs of the attention *score* and *mix* steps for n tokens.
944
945    QKᵀ is n·n·d multiply-adds, and weights·V is another n·n·d. Count 2 FLOPs
946    per multiply-add. The projections (X·W) are O(n·d²), linear in n, so for
947    long contexts the n² term dominates.
948    """
949    score_and_mix = 2 * (2 * n * n * d_model) * n_layers
950    projections = 2 * (4 * n * d_model * d_model) * n_layers
951    return dict(n=n, quadratic_flops=score_and_mix, linear_flops=projections, score_matrix_entries=n * n)

Rough FLOPs of the attention score and mix steps for n tokens.

QKᵀ is n·n·d multiply-adds, and weights·V is another n·n·d. Count 2 FLOPs per multiply-add. The projections (X·W) are O(n·d²), linear in n, so for long contexts the n² term dominates.

def figures() -> dict: on GitHub
 959def figures() -> dict:
 960    """Plot this lesson's data. matplotlib is imported here, and only here,
 961    so the lesson itself needs nothing beyond NumPy."""
 962    import matplotlib
 963
 964    matplotlib.use("Agg")
 965    import matplotlib.pyplot as plt
 966
 967    SCALED, UNSCALED, MUTED = "#2563eb", "#dc2626", "#9ca3af"
 968    figs = {}
 969
 970    # --- 1. Scores -> softmax weights for "it" ------------------------------
 971    fig, ax = plt.subplots(figsize=(6, 3.2))
 972    x = np.arange(len(IT_EXAMPLE_TOKENS))
 973    weights = softmax(IT_EXAMPLE_SCORES)
 974    ax.bar(x - 0.2, IT_EXAMPLE_SCORES / IT_EXAMPLE_SCORES.sum(), 0.4, color=MUTED, label="score (share of total)")
 975    ax.bar(x + 0.2, weights, 0.4, color=SCALED, label="attention weight (softmax)")
 976    for xi, w in zip(x, weights):
 977        ax.text(xi + 0.2, w + 0.02, f"{w:.2f}", ha="center")
 978    ax.set_xticks(x, IT_EXAMPLE_TOKENS)
 979    ax.set_ylabel("share")
 980    ax.set_ylim(0, 0.78)
 981    ax.set_title('What "it" attends to: softmax sharpens the scores')
 982    ax.legend(frameon=False)
 983    figs["it_weights"] = fig
 984
 985    # --- 2. Causal attention heatmap on a sentence --------------------------
 986    tokens = "The animal didn't cross the street because it was tired".split()
 987    rng = np.random.default_rng(7)
 988    X = rng.standard_normal((len(tokens), 32))
 989    _, w = MultiHeadAttention(32, n_heads=1, seed=3)(X)
 990    fig, ax = plt.subplots(figsize=(6, 5))
 991    im = ax.imshow(w[0], cmap="Blues", vmin=0, vmax=1)
 992    ax.set_xticks(range(len(tokens)), tokens, rotation=45, ha="right")
 993    ax.set_yticks(range(len(tokens)), tokens)
 994    ax.set_xlabel("key: the token being looked at")
 995    ax.set_ylabel("query: the token doing the looking")
 996    ax.set_title("Causal attention: each row sees only the past")
 997    ax.grid(False)
 998    fig.colorbar(im, ax=ax, fraction=0.046, label="attention weight")
 999    figs["causal_heatmap"] = fig
1000
1001    # --- 3. Why sqrt(d_k): saturation and vanishing gradient ----------------
1002    dks = (2, 4, 8, 16, 32, 64, 128, 256, 512, 1024)
1003    rows = sqrt_dk_experiment(d_ks=dks, trials=600)
1004    fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4))
1005    for scaled, color, label in ((False, UNSCALED, "unscaled  QKᵀ"), (True, SCALED, "scaled  QKᵀ/√d_k")):
1006        sel = [r for r in rows if r["scaled"] == scaled]
1007        a1.plot(dks, [r["max_weight"] for r in sel], "o-", color=color, label=label)
1008        a2.plot(dks, [r["grad_norm"] for r in sel], "o-", color=color, label=label)
1009    for a, title, ylabel in (
1010        (a1, "Softmax saturates (→ one-hot)", "mean largest weight"),
1011        (a2, "Gradient through softmax vanishes", "softmax Jacobian norm"),
1012    ):
1013        a.set_xscale("log", base=2)
1014        a.set_xlabel("head width d_k")
1015        a.set_ylabel(ylabel)
1016        a.set_title(title)
1017    a1.set_ylim(0, 1.05)
1018    a1.legend(frameon=False)
1019    fig.tight_layout()
1020    figs["sqrt_dk"] = fig
1021
1022    # --- 4. Quadratic vs linear cost ----------------------------------------
1023    d_model = 4096
1024    ns = np.logspace(2, 6, 60)
1025    quad = [attention_cost(int(n), d_model)["quadratic_flops"] for n in ns]
1026    lin = [attention_cost(int(n), d_model)["linear_flops"] for n in ns]
1027    fig, ax = plt.subplots(figsize=(6, 3.6))
1028    ax.loglog(ns, quad, color=UNSCALED, label="scores + mix  O(n²·d)")
1029    ax.loglog(ns, lin, color=SCALED, label="projections  O(n·d²)")
1030    ax.axvline(2 * d_model, color=MUTED, ls="--")
1031    ax.text(2 * d_model * 1.15, min(quad) * 10, "crossover\nn = 2·d_model", color="#4b5563")
1032    ax.set_xlabel("context length n (tokens)")
1033    ax.set_ylabel("FLOPs per layer")
1034    ax.set_title(f"Why long context is expensive (d_model = {d_model})")
1035    ax.legend(frameon=False)
1036    figs["quadratic_cost"] = fig
1037
1038    return figs

Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.

def demo() -> None: on GitHub
1046def demo() -> None:
1047    banner("1. Worked example: what does 'it' attend to?")
1048    say(
1049        """
1050        "The animal didn't cross the street because it was tired." Suppose the
1051        query for "it", dotted with three keys and scaled, gives scores 2.0,
1052        1.0 and 0.5. Softmax exponentiates and normalizes:
1053        """
1054    )
1055    w = worked_example_it()
1056    e = np.exp(IT_EXAMPLE_SCORES)
1057    table(
1058        ["token", "score", "e^score", "weight"],
1059        [(t, s, ei, wi) for t, s, ei, wi in zip(IT_EXAMPLE_TOKENS, IT_EXAMPLE_SCORES, e, w.values())],
1060        floatfmt=".2f",
1061    )
1062    say(f"Total of e^score = {e.sum():.2f}. New 'it' = 0.63·V(animal) + 0.23·V(tired) + 0.14·V(street).")
1063    takeaway("The model resolved the pronoun: most of 'it' now comes from 'animal'.")
1064
1065    banner("2. Full attention on a tiny sequence (shapes and causal mask)")
1066    rng = np.random.default_rng(0)
1067    n, d_k = 4, 8
1068    Q, K, V = (rng.standard_normal((n, d_k)) for _ in range(3))
1069    out, weights = scaled_dot_product_attention(Q, K, V, mask=causal_mask(n))
1070    say(f"Q, K, V are each ({n}, {d_k}). Scores Q·Kᵀ are ({n}, {n}): every token vs. every token.")
1071    matrix("causal attention weights (row = query token, col = key token)", weights)
1072    say(
1073        """
1074        Upper triangle is exactly zero: token 0 sees only itself, token 3 sees
1075        all four. Every row sums to 1. Output shape is (4, 8): one new
1076        context-aware vector per token.
1077        """
1078    )
1079
1080    banner("3. Why divide by sqrt(d_k)? Measure it.")
1081    rows = sqrt_dk_experiment()
1082    table(
1083        ["d_k", "scaled?", "score std", "mean max weight", "softmax grad norm"],
1084        [(r["d_k"], "yes" if r["scaled"] else "no", r["score_std"], r["max_weight"], r["grad_norm"]) for r in rows],
1085        floatfmt=".3f",
1086    )
1087    say(
1088        """
1089        Unscaled, the score std grows like sqrt(d_k) (2, 8, ~22.6). At d_k=512
1090        the softmax is essentially one-hot (max weight near 1) and the gradient
1091        norm collapses toward 0. Scaled, std stays ~1 and gradients stay
1092        healthy at every width.
1093        """
1094    )
1095    takeaway(
1096        "Dot products grow with dimension; big scores saturate softmax and kill "
1097        "gradients. Dividing by sqrt(d_k) keeps the variance at 1 so training stays stable."
1098    )
1099
1100    banner("4. Multi-head vs. grouped-query attention")
1101    X = rng.standard_normal((6, 64))
1102    mha = MultiHeadAttention(d_model=64, n_heads=8)
1103    gqa = MultiHeadAttention(d_model=64, n_heads=8, n_kv_heads=2)
1104    mqa = MultiHeadAttention(d_model=64, n_heads=8, n_kv_heads=1)
1105    for name, layer in [("MHA (8 KV heads)", mha), ("GQA (2 KV heads)", gqa), ("MQA (1 KV head)", mqa)]:
1106        y, wts = layer(X)
1107        print(f"{name:18s} out {y.shape}, weights {wts.shape}, K+V params {layer.kv_params():5d}")
1108    print()
1109    say(
1110        """
1111        All three produce the same output shape. GQA and MQA keep 8 query heads
1112        but store 4x and 8x less K/V. The KV cache (what you keep in GPU memory
1113        per token while generating) shrinks by the same factor.
1114        """
1115    )
1116
1117    banner("5. Why long context is expensive: O(n²)")
1118    table(
1119        ["tokens n", "score-matrix entries", "quadratic FLOPs", "linear FLOPs"],
1120        [
1121            (c["n"], f"{c['score_matrix_entries']:,}", f"{c['quadratic_flops']:.2e}", f"{c['linear_flops']:.2e}")
1122            for c in (attention_cost(n, 4096) for n in (1_000, 2_000, 8_000, 32_000, 128_000))
1123        ],
1124    )
1125    say(
1126        """
1127        Doubling n from 1k to 2k quadruples the score matrix. With d_model=4096
1128        the quadratic term overtakes the projections once n > 2·d_model.
1129        Mitigations: FlashAttention (tiles; never materializes n×n in slow
1130        memory), sliding-window/sparse attention, GQA for the KV cache, and
1131        prompt caching to avoid recomputing stable prefixes.
1132        """
1133    )