rumblr Work in progressWIP

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

Attention

how a token decides what to listen to

You'll be able to explain Queries, keys, values, softmax, masking, multi-head, GQA, O(n²)

Free lesson 31 min8 figures and diagrams10 interactive
Guide is what to use and when. How it works builds it from scratch. Math & code adds the formulas and the Python.

The lesson in one minute

What you'll be able to explain

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

Level 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 builds the mechanism from nothing, starting with a glance around a room.

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

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

Chapter 1

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

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

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

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

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

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

Softmax, decoded

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

Level 3: the formula and its symbols

Symbols

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

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

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

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

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

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

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

Figure 1 · Drawn from the lesson's code

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

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

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

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

Chapter 2

Where the scores come from: queries, keys and values

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

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

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

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

Level 3: the formula and its symbols

Symbols

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

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

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

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

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

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

Figure 2 · Diagram

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

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

Level 3: the formula and its symbols

Symbols

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

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

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

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

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

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

Chapter 3

Shapes: the part that is easiest to get wrong

For one sequence of n tokens with model width d_model:

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

Chapter 4

Causal masking: no peeking at the answer

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

Figure 4 · Diagram

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

Figure 3 · Drawn from the lesson's code

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

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

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

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

Attention heatmap for the sentence

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

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

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

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

Chapter 5

Why divide by √d_k?

See the problem first, then the proof.

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

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

Level 3: the formula and its symbols

Symbols

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

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

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

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

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

Figure 6 · Drawn from the lesson's code

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

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

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

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

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

Chapter 6

Multi-head attention

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

Figure 7 · Diagram

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

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

Chapter 7

Grouped-query attention: sharing keys and values

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

Figure 8 · Diagram

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

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

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

Chapter 8

Cost: why long context is expensive

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

Figure 9 · Drawn from the lesson's code

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

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

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

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

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

Test yourself

6 questions

Answer each one out loud or on paper before you open it. If you can explain it, you know it.

Question 1Explain attention to a non-engineer in 30 seconds.Think it through, then reveal

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

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

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

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

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

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

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

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

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

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

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

Primary sources

The papers behind this lesson

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

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

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

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

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

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

Read the annotated companion →The paper ↗

Researcher's shelf

Further reading

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

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