The lesson in one minute
What you'll be able to explain
- 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.
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_headsandmax_position_embeddingson 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_costwith 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 builds the mechanism from nothing, starting with a glance around a room.
The everyday picture. Imagine you're in a meeting and someone says "it's broken, can you fix it?" To know what "it" means, you glance around the room: at the laptop on the table, at the person who just walked in, at the whiteboard. You pay a lot of attention to the laptop, a little to everything else, and your understanding of "it" becomes mostly laptop.
That glance is attention. Every word in a sentence gets to look at every other word, decide how relevant each one is, and rebuild its own meaning as a mix of the relevant ones. "Bank" next to "river" becomes a riverbank; next to "loan" it becomes a lender. The word is the same, but what it paid attention to is different.
Chapter 1
A tiny worked example: what does "it" refer to?
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:
- Exponentiate each score (e^score), which makes every number positive and stretches the gaps between them.
- Divide each by the total, 7.39 + 2.72 + 1.65 = 11.76.
| Word | Score | e^score | Share of attention |
|---|---|---|---|
| animal | 2.0 | 7.39 | 7.39 / 11.76 = 0.63 |
| tired | 1.0 | 2.72 | 2.72 / 11.76 = 0.23 |
| street | 0.5 | 1.65 | 1.65 / 11.76 = 0.14 |
Those two steps are called softmax. The new meaning of "it" is 63% "animal", 23% "tired" and 14% "street". Without being told any grammar rule, the model has worked out that "it" is the animal.
In code: worked_example_it runs these three scores through softmax and returns each word's share.
Softmax, decoded
Softmax turns any list of numbers, positive or negative, into shares that are all positive and add up to exactly 1. Think of it as a vote where louder voices get disproportionately more say.
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the list of scores going in | (2.0, 1.0, 0.5) | |
| how many scores there are | 3 | |
| the position we're computing a share for | 1 = "animal" | |
| the score at position | ||
| Euler's number, ≈ 2.718; is "2.718 multiplied by itself times" (works for fractions and negatives too) | ||
| "add up the following, for = 1, 2, …, " | ||
| a counter that walks over every position | 1, 2, 3 | |
| the share of attention position gets | 0.63 |
In words: "the share for item i is e raised to its score, divided by the sum of e raised to every score."
With the numbers: softmax(2.0, 1.0, 0.5)₁ = 7.39 / (7.39 + 2.72 + 1.65) = 7.39 / 11.76 = 0.63.
Level 3: in Python
import math
# animal, tired, street
z = [2.0, 1.0, 0.5]
# e^(z_i) for each score
exps = [math.exp(z_i) for z_i in z]
[round(e, 2) for e in exps] # → [7.39, 2.72, 1.65]
# Σ_j e^(z_j)
total = sum(exps)
round(total, 2) # → 11.76
# softmax(z)_i: each share of the total
[round(e / total, 2) for e in exps] # → [0.63, 0.23, 0.14]
Why e to the power of the score, and not just the score divided by the total? Three reasons you can check on the table above:
- Always positive. Scores can be negative; e to any power is positive, so no word ever gets a negative share.
- Order is kept. A higher score always gets a bigger share.
- Gaps are stretched. Adding 1 to a score multiplies its e-value by 2.72, so the winner pulls ahead decisively. (Plain division would give "animal" 2.0 / 3.5 = 0.57, a timid 57%, and would break on negative scores.)
One more property matters later: adding the same number to every score
changes nothing, because it multiplies the top and bottom of the fraction by
the same amount. The code uses this to subtract the largest score before
exponentiating, which stops e¹⁰⁰⁰ from overflowing to infinity. See
primer.notation for exponents and sums from scratch.
Figure 1 · Drawn from the lesson's code
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
In code: softmax subtracts the largest score before exponentiating, and turns masked scores of −∞ into weights of exactly 0.
Chapter 2
Where the scores come from: queries, keys and values
Each word turns its vector into three new vectors by multiplying it by three learned matrices:
- a query: what am I looking for? ("it" is looking for a noun that can be tired.)
- a key: what do I offer to others? ("animal" offers: a noun, alive.)
- a value: what I actually hand over if someone picks me.
A word's score for another word is the dot product of its query with the other word's key: multiply them position by position and add.
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| the query vector: a list of numbers | (1, 2) | |
| the key vector: another list of numbers | (3, 0.5) | |
| how many numbers each list has (the head width) | 2 | |
| a counter walking over the positions | 1, 2 | |
| the -th number in each list | , | |
| "dot product": multiply matching positions, then add |
In words: "multiply the first numbers together, multiply the second numbers together, and so on, then add up all the products."
With the numbers: (1, 2) · (3, 0.5) = 1·3 + 2·0.5 = 3 + 1 = 4.
Level 3: in Python
q = [1, 2]
k = [3, 0.5]
# Σ over m of q_m k_m
sum(q_m * k_m for q_m, k_m in zip(q, k)) # → 4.0
Vectors that point the same way give big positive scores, vectors at right angles give zero, and opposite vectors give negative scores. That is why the dot product works as a relevance score. The values are then blended with the softmax shares, exactly as in the table above.
A library search is a good analogy. The query is what you type into the search box, the keys are the catalogue entries, and the values are the books on the shelf. You get back a blend of books, weighted by how well each catalogue entry matched your search.
Figure 2 · Diagram
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]
Written as one formula, with every word's query stacked into a matrix Q (and likewise K and V), the whole diagram is:
Level 3: the formula and its symbols
Symbols
| Symbol | Meaning here | Shape |
|---|---|---|
| number of tokens in the sequence | ||
| width of each query and key vector | ||
| width of each value vector | ||
| every token's query, stacked as rows | ||
| every token's key, stacked as rows | ||
| every token's value, stacked as rows | ||
| transposed: rows become columns, so each key stands upright | ||
| a matrix multiply: entry (row , column ) is the dot product of token 's query with token 's key, so it is every score at once | ||
| square root of the head width; dividing by it keeps scores from growing with width (explained below) | a single number | |
| softmax(…) | softmax applied to each row separately, so each token's shares add to 1 | |
| multiply the shares by the values: each output row is a share-weighted blend of value rows |
In words: "score every query against every key, shrink the scores by √d_k, turn each row of scores into shares, and use the shares to blend the values."
With the numbers: the "it" row of holds (2.0, 1.0, 0.5); softmax turns that row into (0.63, 0.23, 0.14); multiplying by gives 0.63·V(animal) + 0.23·V(tired) + 0.14·V(street). To see every symbol at work, take : the query of "it" (1, 1, 1, 1) against the keys (1, 1, 1, 1), (1, 1, 0, 0) and (1, 0, 0, 0) gives raw scores (4, 2, 1), and dividing by gives exactly that row. With toy values V(animal) = (1, 0), V(tired) = (0, 1) and V(street) = (1, 1), the blend is (0.63 + 0.14, 0.23 + 0.14) = (0.77, 0.37).
Level 3: in Python
import math
q_it = [1, 1, 1, 1]
# keys: animal, tired, street
K = [[1, 1, 1, 1], [1, 1, 0, 0], [1, 0, 0, 0]]
# values, one row per word
V = [[1, 0], [0, 1], [1, 1]]
d_k = len(q_it)
# q Kᵀ / √d_k
scores = [sum(q * k for q, k in zip(q_it, k_j)) / math.sqrt(d_k) for k_j in K]
scores # → [2.0, 1.0, 0.5]
exps = [math.exp(s) for s in scores]
# softmax of the row
weights = [e / sum(exps) for e in exps]
[round(w, 2) for w in weights] # → [0.63, 0.23, 0.14]
# (…)V: blend the values
[round(sum(w * v[c] for w, v in zip(weights, V)), 2) for c in range(2)] # → [0.77, 0.37]
In practice this one line runs in every layer of every modern language model, for every word, many times per word generated. Nearly everything about a model's speed and memory use traces back to it.
In code: scaled_dot_product_attention is the whole formula, one commented step per box of the diagram, and returns both the blended values and the attention weights.
Chapter 3
Shapes: the part that is easiest to get wrong
For one sequence of n tokens with model width d_model:
| Tensor | Shape | Meaning |
|---|---|---|
| X | (n, d_model) | input token vectors |
| W_q, W_k | (d_model, d_k) | learned projections |
| W_v | (d_model, d_v) | learned projection |
| Q, K | (n, d_k) | queries, keys |
| V | (n, d_v) | values |
| Q Kᵀ | (n, n) | every token vs. every token, hence O(n²) |
| weights | (n, n) | each row sums to 1 |
| output | (n, d_v) | one context-aware vector per token |
Chapter 4
Causal masking: no peeking at the answer
A decoder model (GPT, Claude, Llama) is trained to predict the next token. If position i could attend to position i+1, it could simply read the answer. So the scores for future positions are set to −∞ before softmax. Because e^−∞ = 0, those weights come out as exactly zero.
Figure 4 · Diagram
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"]
Figure 3 · Drawn from the lesson's code
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
Figure 5 · Interactive · computed from the lesson's code
Attention heatmap for the sentence
Try it: here is the same sentence with hand-picked queries and keys, so the pattern does mean something; "it" keeps the worked example's scores for "animal", "tired" and "street". Pick the row for "it", then turn the causal mask on: "tired" comes later, so its weight drops to exactly 0 and "animal" takes a bigger share. Drag the temperature below 1 to watch each row sharpen towards one word, and above 1 to watch it flatten towards an even spread.
Masking only the future is also what makes generation cheap: a token's
output never changes when later tokens arrive, so it can be computed once
and cached. See primer.ml.inference for the KV cache.
In code: causal_mask builds the lower triangle of allowed positions, and scaled_dot_product_attention sets every score outside it to −∞ before softmax. worked_example_sentence holds the hand-picked queries and keys the heatmap above draws.
Chapter 5
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
Symbols
| Symbol | Meaning here |
|---|---|
| variance: the average squared distance of from its average; a measure of how widely swings. Its square root is the standard deviation, the typical size of a swing | |
| a raw attention score, for and whose entries are independent with average 0 and variance 1 | |
| the number of products added up in the dot product | |
| "which implies" | |
| dividing a quantity by divides its variance by ; here |
In words: "the spread of a raw score grows with the head width, so we divide the score by the square root of the width, which brings the spread back to 1 at any width."
With the numbers: at d_k = 128, the standard deviation of a raw score is √128 ≈ 11.3, so scores of ±20 are routine. After dividing by √128 the standard deviation is 1, so scores of ±2 are typical. At d_k = 4 you can check the formula exactly: the 16 equally likely patterns of four ±1 products give raw scores whose variance is 4, and dividing each by √4 = 2 brings the variance to 1.
Level 3: in Python
import itertools, math, statistics
d_k = 4
scores = [sum(products) for products in itertools.product([-1, 1], repeat=d_k)]
# Var(q · k) = d_k
statistics.pvariance(scores) # → 4
# Var(q · k / √d_k) = 1
statistics.pvariance([s / math.sqrt(d_k) for s in scores]) # → 1.0
# the typical raw swing at d_k = 128
round(math.sqrt(128), 1) # → 11.3
Why "all attention on one word" is bad, beyond being a wrong answer: softmax
has stopped responding. Its gradient (the signal training uses to
adjust the weights) is diag(p) − p pᵀ, and when one share is 1 and the rest
are 0, every entry of that matrix is 0. The query and key weights receive no
signal and stop learning. This is called saturation.
sqrt_dk_experiment measures all of this.
Figure 6 · Drawn from the lesson's code
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
In one sentence: dot products grow with dimension, large scores saturate softmax and kill gradients, and scaling by √d_k keeps scores in the range where softmax is smooth and trainable.
In code: softmax_jacobian builds the diag(p) − p pᵀ matrix, so you can watch every entry fall to 0 as one share approaches 1.
Chapter 6
Multi-head attention
Instead of one attention with d_k = d_model, run h heads in parallel, each with d_k = d_model / h and its own W_q, W_k and W_v. The total compute is the same, but each head can specialise: one tracks syntax, another which noun a pronoun refers to, another the previous token.
Figure 7 · Diagram
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)"]
In code: MultiHeadAttention holds the four projections W_q, W_k, W_v and W_o; calling it slices the projections into one slab per head, runs scaled_dot_product_attention on every head at once, and glues the results back together before W_o.
Chapter 7
Grouped-query attention: sharing keys and values
With grouped-query attention (GQA), several query heads share one key/value head. Llama 3 8B has 32 query heads but only 8 KV heads, so it stores 4× fewer keys and values per token. Multi-query attention (MQA) is the extreme: one KV head for everyone.
Figure 8 · Diagram
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
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.
Chapter 8
Cost: why long context is expensive
The score matrix is n × n per head per layer, so compute and memory grow with n². Doubling the context roughly quadruples attention's cost.
Figure 9 · Drawn from the lesson's code
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
The standard mitigations each attack this picture. FlashAttention produces the exact same result but computes it in tiles sized for fast on-chip GPU memory, so the n × n matrix is never written to slow memory. Sliding-window and sparse attention let each token see only some others. GQA and MQA shrink the KV cache. State-space models such as Mamba replace attention with a linear-time recurrence.
In code: attention_cost counts the quadratic score-and-mix FLOPs and the linear projection FLOPs plotted above.
Test yourself
6 questions
Answer each one out loud or on paper before you open it. If you can explain it, you know it.
Question 1Explain attention to a non-engineer in 30 seconds.Think it through, then reveal
When the model reads a word, it asks which other words here help it understand this one. It scores every other word for relevance, then builds its understanding of the word as a mix of the relevant ones. In "the animal didn't cross the street because it was tired", "it" draws mostly from "animal". It does this for every word, in parallel, dozens of times over.
Question 2Now explain it to an ML engineer in two minutes.Think it through, then reveal
Project X into Q, K and V with learned matrices. Compute QKᵀ/√d_k, an n×n matrix of scaled similarities; add a causal mask of −∞ above the diagonal for a decoder; softmax each row; multiply by V. Run h heads in parallel on d_model/h slices, concatenate them, and project with W_o. Wrap the whole thing in a residual connection with pre-layer-norm and follow it with an FFN. Cost is O(n²·d) in compute and O(n²) in memory for the scores, which is why FlashAttention tiles it and GQA shrinks the KV cache.
Question 3Why divide by √d_k? What breaks without it?Think it through, then reveal
The variance of q·k grows linearly with d_k. Without scaling, scores spread out, softmax saturates towards one-hot, its Jacobian goes to about zero, and the query/key projections stop receiving gradient. Training stalls or turns unstable.
Question 4Why is long context expensive? Name two techniques that reduce the cost.Think it through, then reveal
Scores are n×n per head per layer, so cost is quadratic, and the KV cache grows linearly with every token held in memory. Two fixes: FlashAttention (exact, memory-efficient tiling) and GQA/MQA (fewer KV heads). Others: sliding-window or sparse attention, and prompt caching of stable prefixes.
Question 5What does the causal mask buy you besides honest training?Think it through, then reveal
Earlier outputs never depend on later tokens, so during generation the keys and values of past tokens can be computed once and cached. Each new token then costs one row of attention instead of a full recomputation.
Question 6What does GQA trade away, and for what?Think it through, then reveal
A little modelling capacity (query heads share keys and values) for a KV cache that is several times smaller. That means more concurrent requests and longer contexts per GPU.
Primary sources
The papers behind this lesson
Showed that attention alone, with no recurrence, is enough for state-of-the-art translation, and introduced scaled dot-product attention, multi-head attention and the transformer.
Read on rumblr →The paper ↗Introduced grouped-query attention, the middle ground between multi-head and multi-query attention.
The paper ↗Computes exact attention in tiles sized for fast on-chip GPU memory, making long contexts practical.
Read the annotated companion →The paper ↗Researcher's shelf
Further reading
- Vaswani et al., Attention Is All You Need (2017): https://arxiv.org/abs/1706.03762
- Jay Alammar, The Illustrated Transformer: https://jalammar.github.io/illustrated-transformer/
- Harvard NLP, The Annotated Transformer (the paper, line by line in code): https://nlp.seas.harvard.edu/annotated-transformer/
- Andrej Karpathy, Let's build GPT: from scratch, in code, spelled out: https://www.youtube.com/watch?v=kCc8FmEb1nY and nanoGPT: https://github.com/karpathy/nanoGPT
- 3Blue1Brown, Attention in transformers, visually explained: https://www.youtube.com/watch?v=eMlx5fFNoYc
- PyTorch
scaled_dot_product_attention: https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html - Ainslie et al., GQA (2023): https://arxiv.org/abs/2305.13245
- Dao et al., FlashAttention (2022): https://arxiv.org/abs/2205.14135
About this lesson. This is the illustrated edition of a lesson from the open-source AI Primer. Its text, figures and numbers are generated from the Primer's source at commit 048aeaa, so the two always agree: the explanation, the code that builds it and the tests that prove it.