← Atelier
Lesson 04 · Transformers · ~35 min + code

Attention: score, normalize, mix

The trophy didn't fit in the suitcase because it was too big. To process it, a transformer has to decide which earlier tokens to read from. It decides using nothing but dot products. How?

Attention is three steps: score every key against a query, turn the scores into weights, and mix the values with those weights. Here everything lives in two dimensions, so you can see every vector and drag the query yourself.

1 · Score

Which earlier token matches the question?

Each earlier token carries a key, a vector advertising what it contains. The current token, it, emits a query, a vector describing what it's looking for. The match score is a dot product:

\[ s_j = q \cdot k_j = |q|\,|k_j| \cos\theta_j \]

Four earlier tokens with 2-D keys, and the query as the dark arrow. Drag anywhere in the left plane to move the query's tip.

Query and keys
Values and the output

Scores are divided by √d = √2 before the softmax (section 5 explains why). The output (★) is the weighted mix of the value points.

Two keys make the same angle with the query, but one is twice as long. Which scores higher?

A dot product scales with both lengths. A model can make a token "louder" by giving it a longer key, and a query can make itself more decisive by being longer (section 2). Cosine similarity would throw that information away. Some architectures normalize q and k (QK-norm) precisely to control this.

2 · Normalize

From scores to weights

Scores can be any real numbers. Softmax turns them into positive weights that sum to 1, as in lesson 1:

\[ w_j = \frac{e^{s_j/\sqrt d}}{\sum_i e^{s_i/\sqrt d}} \]
You make the query twice as long without changing its direction. What happens to the weights?

Every score doubles, so every difference between scores doubles too, and softmax depends on differences. Query length works like an inverse temperature: long queries approach an argmax, and short ones approach a uniform average. Try it in the figure above: drag the tip outward along the same direction.

3 · Mix

The output is a weighted average of values

Each token also carries a value, what it hands over if attended to. That's the right-hand plane. The output is

\[ o = \sum_j w_j\, v_j \]
Can the output ever land outside the shape formed by the value points, their convex hull?

The weights are non-negative and sum to 1, so \(o\) is a convex combination, always inside the hull. A single head can only select and blend what's already there. New directions come from the output projection, from combining several heads, and from the MLP that follows. Drag the query and watch the ★ stay inside the dashed hull.

Put the three steps together and batch over all queries at once. That's the formula in every transformer paper:

\[ \operatorname{Attention}(Q, K, V) = \operatorname{softmax}\!\Big(\frac{QK^\top}{\sqrt{d}}\Big)\, V \]

\(Q\), \(K\), \(V\) are the token representations multiplied by three learned matrices \(W_Q, W_K, W_V\). Training shapes those matrices so that the right queries find the right keys.

4 · Play it

Be the query

Each level gives you a target in value space (the ◎). Drag the query until the output ★ lands inside the ring. You control only direction and length, like a real query.

Query and keys
Values, output ★ and target ◎
5 · Why √d

Keeping softmax out of saturation

Real heads use \(d = 64\) or \(128\) dimensions, not 2.

Query and key have \(d\) independent entries, each with mean 0 and variance 1. What is the variance of \(q \cdot k\)?

\(q\cdot k = \sum_{i=1}^d q_i k_i\) is a sum of \(d\) independent terms, each with variance 1, so the variance is \(d\) and the typical size is \(\sqrt d\). Unscaled, scores grow with dimension, softmax saturates toward one-hot, and its gradient vanishes, just like the saturated sigmoid in lesson 3. Dividing by \(\sqrt d\) brings the variance back to 1 at any width.

Average largest softmax weight over 8 random keys (500 trials per point)

Random unit-variance queries and keys. Without scaling, one key takes almost all the weight at high d, before any learning has happened.

6 · Masks and the KV cache

Not reading the future, and not recomputing the past

A language model predicts each next token from the ones before it, so position \(i\) must not attend to positions after \(i\).

Where does the causal mask go?

Before. A score of \(-\infty\) becomes \(e^{-\infty} = 0\) inside the softmax, and the remaining weights still sum to 1. Zeroing after the softmax leaves rows that no longer sum to 1, because the future tokens took their share of the normalization. Zeroing values keeps the future in the denominator. Toggle the mask:

Attention weights for "the cat sat on the"

Row = the query position, column = the key it reads from. Each row sums to 1.

During generation the model produces one token at a time. The new token's query needs every earlier key and value, and those never change, since earlier tokens can't see later ones. So they're computed once and kept: the KV cache. Decoding then costs one new query per step instead of recomputing the whole sequence. The price is memory.

32 layers, model width 4096 (32 heads × 128), no grouped KV, a context of 8,192 tokens, batch size 1, fp16. How big is the KV cache?

\(2\,(\text{K and V}) \times 32\,(\text{layers}) \times 4096 \times 8192\,(\text{tokens}) \times 2\,\text{bytes} = 4.3\) GB, for one sequence. Multiply by the batch size and it often exceeds the weights. That's why decoding is limited by memory bandwidth, and why grouped-query attention shares each K/V head across several query heads. Try it:

KV cache calculator
7 · New cases, no hints

Apply it somewhere else

You shuffle the order of the input tokens and feed them to an attention layer with no positional information. What happens to the outputs?

Attention is permutation-equivariant: scores depend only on which vectors meet, not where they sit. Without positional information, "dog bites man" and "man bites dog" contain the same set of pairs. Positional encodings break the symmetry. RoPE rotates \(q\) and \(k\) by position-dependent angles, so \(q_m\cdot k_n\) depends on \(m - n\).

Every key is identical. What does the head output?

Equal scores give equal weights, \(1/n\) each, so the output is the mean of the values whatever the query. In practice this appears as heads that "attend to nothing": they dump weight on a constant position, such as the first token (an attention sink).

You double the context length from 4k to 8k tokens. The attention score computation grows by…

Every query meets every key: \(n^2\) scores per head per layer. FlashAttention doesn't change that count. It avoids writing the \(n \times n\) matrix to GPU memory by tiling and using an online softmax, which makes it much faster and memory-linear. During decoding with a KV cache, each new token costs \(O(n)\), so a whole generation still costs \(O(n^2)\).

8 · Implement, unaided

Write attention from scratch

Plain editor, hidden tests. Shapes first. ⌘/Ctrl+Enter runs.

Exercise 1: scaled dot-product attention with an optional causal mask

Exercise 2: one decoding step with a KV cache

9 · Staff-level follow-ups

Answer out loud, then check

Why divide by √d, and what breaks without it?
  • With unit-variance entries, \(\operatorname{Var}(q\cdot k) = d\). Scores grow like \(\sqrt d\), softmax saturates, and gradients through it shrink.
  • Dividing by \(\sqrt d\) restores unit variance at initialization. It's not a learned temperature, and training can still sharpen weights by growing \(|q|\) and \(|k|\).
  • Related fixes: QK-norm, and attention logit soft-capping in some large models.
Why multiple heads instead of one big head?
  • One head produces one set of weights per query: one convex mix. Several heads can attend to different things at once (syntax, coreference, position) and are concatenated and projected.
  • Same parameter count and FLOPs as one wide head, split into subspaces.
  • Many heads turn out to be redundant and can be pruned. MQA and GQA share K/V heads to cut KV-cache memory with little quality loss.
What does the MLP block do that attention can't?
  • Attention moves information between positions, with weights that are convex mixes. The MLP transforms each position independently and nonlinearly.
  • It holds most of the parameters (typically about 2/3 of a block). One useful view treats it as key-value memory for learned associations, but that is an interpretation, not a mechanism guarantee.
  • Both write into the residual stream: \(x \leftarrow x + \text{Attn}(\text{LN}(x))\), then \(x \leftarrow x + \text{MLP}(\text{LN}(x))\).
How does RoPE encode position?
  • It rotates each pair of \(q\) and \(k\) dimensions by an angle proportional to position, with a different frequency per pair. The dot product \(q_m \cdot k_n\) then depends only on \(m - n\).
  • It's relative, adds no parameters, and is applied to q and k (not v) in every layer.
  • Extending context beyond training lengths needs frequency scaling (position interpolation, NTK-aware, YaRN), because the low-frequency pairs never saw long wavelengths.
Prefill vs decode: where does the time go when serving an LLM?
  • Prefill processes the prompt in parallel: large matrix multiplies, compute-bound, which sets time-to-first-token.
  • Decode produces one token per step: it reads all weights plus the KV cache for little compute, so it's memory-bandwidth-bound and sets inter-token latency.
  • Levers: batching (continuous batching), paged KV memory (vLLM), GQA, KV quantization, speculative decoding, and separating prefill from decode onto different hardware.
What does FlashAttention change, and what doesn't it change?
  • It computes exact attention in tiles held in on-chip SRAM, with an online (running max and sum) softmax, and never writes the \(n\times n\) matrix to HBM. Memory goes from \(O(n^2)\) to \(O(n)\), and it's faster because it moves less data.
  • It doesn't change the math or the \(O(n^2)\) FLOPs. It's not an approximation.

Sources: Vaswani et al., Attention Is All You Need (2017); Su et al., RoFormer (2021); Dao et al., FlashAttention (2022); Ainslie et al., GQA (2023); Kwon et al., PagedAttention (2023).