Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

5. Attention from first principles

In this chapter

  • The problem attention solves: letting every token use information from every earlier token.
  • Attention built up in four steps: plain dot-product weights, learned queries/keys/values, scaling, causal masking.
  • Multi-head attention, and the reshapes that are easy to get subtly wrong.
  • A position-aware formulation that later serves caching, batching and sparse attention unchanged.

You will build

engine/attention.py: split_heads, merge_heads and causal_attention, the function every model in this book calls.

Time: 5-7 hours. GPU: not needed.

The problem: context

After Chapter 4, each token is a vector that describes the token in isolation. But meaning depends on context. In “The animal didn’t cross the street because it was too tired”, the vector for “it” should end up carrying information about “animal”. Predicting the next token needs the whole preceding text, not just the last word.

Before 2017, the standard answer was the recurrent neural network (RNN). It reads tokens one at a time and squeezes everything seen so far into a single fixed-size vector, its hidden state. That has two problems. Information from far back must survive many squeezes and fades. And processing is inherently sequential, so training can’t use a GPU’s parallelism across positions.

Attention takes the opposite approach: when processing token $t$, look directly at every earlier token, decide how relevant each one is, and take a weighted mixture of their information. Nothing is squeezed, and all positions can be computed at once. That one idea, from Attention Is All You Need (Vaswani et al., 2017), is the core of every model in this book. (Recurrence comes back, in a much improved form, in Chapter 28.)

Step 1: attention with no parameters at all

We’ll follow BALLM’s example: six tokens, “Your journey starts with one step”, each already embedded as a 3-dimensional vector:

x = torch.tensor([[0.43, 0.15, 0.89],   # Your
                  [0.55, 0.87, 0.66],   # journey
                  [0.57, 0.85, 0.64],   # starts
                  [0.22, 0.58, 0.33],   # with
                  [0.77, 0.25, 0.10],   # one
                  [0.05, 0.80, 0.55]])  # step

Let’s compute an enriched, context-aware vector for “journey”. The recipe has three steps.

Score every token by its similarity to “journey”. The dot product is a natural similarity measure: it’s large when two vectors point the same way.

scores = x @ x[1] = [0.9544, 1.4950, 1.4754, 0.8434, 0.7070, 1.0865]

Normalize the scores into weights that are positive and sum to 1, using softmax:

$$ w_j = \frac{e^{s_j}}{\sum_k e^{s_k}} \quad\Rightarrow\quad w = [0.1385,\ 0.2379,\ 0.2333,\ 0.1240,\ 0.1082,\ 0.1581]. $$

“journey” weighs itself highest, then “starts”, whose vector is nearly identical.

Mix: the new vector is the weighted sum of all token vectors:

$$ z_{\text{journey}} = \sum_j w_j, x_j = [0.4419,\ 0.6515,\ 0.5683]. $$

That’s attention. Doing it for every token at once is two matrix multiplications and a softmax:

$$ \text{weights} = \operatorname{softmax}_{\text{row}}(X X^{\top}), \qquad Z = \text{weights}; X . $$

Row $i$ of the [6, 6] weights matrix says how much token $i$ draws from each token $j$.

Step 2: learned queries, keys and values

Plain dot products only find vectors that are already similar. A model needs to learn what to look for, and that depends on the role a token plays. So attention gives each token three different learned views of itself, each produced by its own weight matrix:

  • Query $q_i = x_i W_q$: what token $i$ is looking for.
  • Key $k_j = x_j W_k$: what token $j$ offers, which is matched against queries.
  • Value $v_j = x_j W_v$: what token $j$ hands over if it’s chosen.

$$ \text{scores} = Q K^{\top}, \qquad Z = \operatorname{softmax}(\text{scores}), V . $$

Separating keys from values is the important design choice. A token can be found for one reason and contribute something else. For “it”, the query might look for “animate noun earlier in the sentence”; the key of “animal” advertises exactly that; and its value passes along information useful for predicting what comes after “it”. These are intuitions, not literal labels. Training simply finds whatever projections lower the loss. The mechanism makes such behavior possible.

Step 3: scale the scores

Dot products of $d$-dimensional vectors with independent, unit-variance components have a standard deviation of about $\sqrt d$. Measured on random vectors:

head dimension $d$std of $q\cdot k$std of $q\cdot k/\sqrt d$
21.441.02
648.081.01
102431.60.99

Large scores make softmax saturate: softmax([1, 2, 3]) is [0.09, 0.24, 0.67], but softmax([8, 16, 24]) is [0.0000, 0.0003, 0.9997]. A saturated softmax puts all its weight on one token, and its gradient is nearly zero, so training stalls. Dividing scores by $\sqrt{d}$ keeps them in a healthy range at any width. This is scaled dot-product attention:

$$ \operatorname{Attention}(Q, K, V) = \operatorname{softmax}!\left(\frac{Q K^{\top}}{\sqrt{d}}\right) V . $$

Step 4: no peeking at the future

A language model learns to predict token $t+1$ from tokens $\le t$. If position $t$ could attend to position $t+1$ during training, it would simply copy the answer. The loss would drop to near zero and the model would learn nothing useful.

The fix is a causal mask: before the softmax, set the scores for future positions to $-\infty$. Since $e^{-\infty} = 0$, they get exactly zero weight, and the remaining weights still sum to 1. For our six tokens (using the plain dot products of Step 1):

             Your  journey starts  with   one   step
Your        1.000  0.000  0.000  0.000  0.000  0.000
journey     0.368  0.632  0.000  0.000  0.000  0.000
starts      0.228  0.389  0.382  0.000  0.000  0.000
with        0.205  0.296  0.292  0.208  0.000  0.000
one         0.175  0.225  0.227  0.157  0.216  0.000
step        0.139  0.218  0.213  0.142  0.099  0.190

The first token can only attend to itself. The last row is unchanged from the unmasked version, because nothing comes after “step”.

Warning

Mask before the softmax. Masking after it (zeroing weights of future tokens) leaves their scores in the denominator, so the remaining weights no longer sum to 1, and information about the future leaks through the normalization.

The best test of a mask doesn’t compare numbers with a library. It tests the rule directly: change only the last input token, and check that every earlier output stays exactly the same. Your milestone tests do this.

Explore: causal attention weights

Each row is one query token; columns are the keys it may read. Move the query position, change the softmax temperature (a stand-in for scaling) and watch masked cells stay at zero.

Multiple heads

One attention pattern per layer is limiting. A token may need to track its subject, its previous token and the most recent comma all at once. Multi-head attention runs $H$ smaller attentions in parallel, each with its own projections into a $D_h$-dimensional space (typically $D = H \cdot D_h$), then concatenates their outputs and mixes them with an output projection $W_o$.

In code, nobody runs $H$ separate projections. One big projection produces all heads’ queries at once, and a reshape splits them:

x: [B, T, D]
q = x @ Wq.T                     [B, T, H*Dh]
q.reshape(B, T, H, Dh)           [B, T, H, Dh]   split the last axis into heads
 .transpose(1, 2)                [B, H, T, Dh]   heads become a batch axis
scores = q @ k.transpose(-2, -1) [B, H, T, T]    one T×T matrix per head
y = weights @ v                  [B, H, T, Dh]
y.transpose(1, 2)                [B, T, H, Dh]
 .reshape(B, T, H*Dh)            [B, T, D]       heads concatenated per token
out = y @ Wo.T                   [B, T, D]

Warning

The transposes are not optional. reshape(B, H, T, Dh) directly on a [B, T, H*Dh] tensor produces the right shape with the wrong data: it interleaves tokens and heads. Nothing crashes, and the model just computes garbage. Your milestone test checks that head 1 of token 2 really holds features 4-8 of token 2.

Fewer key/value heads: grouped-query attention

Modern models often use fewer key/value heads than query heads. In grouped-query attention (GQA), each KV head is shared by a group of query heads: Qwen3-0.6B has 16 query heads and 8 KV heads, and Flash-Next’s attention layers have 24 query heads and just 2 KV heads. Each query head still computes its own attention pattern, but the K and V tensors are smaller. That matters enormously for the KV cache’s memory (Chapter 16). In code, query head $h$ reads KV head $h \div (H_q/H_{kv})$. The simplest implementation repeats each KV head to match the query heads before the matmul.

Positions, not just “the last T keys”

There’s one more design decision in your causal_attention, and it pays off for the rest of the book. Instead of building a fixed triangular mask, it takes the absolute position of every query and key and applies the rule

$$ \text{key } j \text{ is visible to query } i \iff \text{pos}(k_j) \le \text{pos}(q_i). $$

During ordinary training, queries and keys are both positions $0..T-1$, and this reproduces the triangle. But the same rule also handles, unchanged:

  • cached decoding (Chapter 16): 1 new query at position 57 against 58 cached keys;
  • chunked prefill: queries 40-47 against keys 0-47, a rectangular mask that a square triangle gets wrong;
  • batches of different lengths (Chapters 19 and 24), where each row has its own positions;
  • sparse attention (Chapter 29), via the extra allowed mask.
def causal_attention(q, k, v, query_positions=None, key_positions=None, allowed=None, scale=None):
    """Scaled dot-product attention with a causal rule.  (Your engine: Chapter 5)

    query_positions: [T] or [B, T]; defaults to the last T positions of the S keys.
    key_positions:   [S] or [B, S]; defaults to 0..S-1.
    allowed: optional extra boolean mask broadcastable to [B, Hq, T, S] (sliding windows,
             sparse selections). True means "may attend".
    Grouped-query attention: Hq must be a multiple of Hkv; each KV head serves Hq/Hkv query heads.
    """
    batch, q_heads, t, d = q.shape
    kv_heads, s = k.shape[1], k.shape[2]
    if q_heads % kv_heads:
        raise ValueError(f"Query heads ({q_heads}) must be a multiple of KV heads ({kv_heads})")
    if kv_heads != q_heads:
        k = k.repeat_interleave(q_heads // kv_heads, dim=1)
        v = v.repeat_interleave(q_heads // kv_heads, dim=1)
    scale = 1.0 / math.sqrt(d) if scale is None else scale
    scores = (q.float() @ k.float().transpose(-2, -1)) * scale          # [B, Hq, T, S]
    qp = _positions(query_positions, batch, t, s - t, q.device)          # [B, T]
    kp = _positions(key_positions, batch, s, 0, q.device)                # [B, S]
    mask = kp[:, None, None, :] <= qp[:, None, :, None]                  # [B, 1, T, S]
    if allowed is not None:
        mask = mask & allowed
    scores = scores.masked_fill(~mask, float("-inf"))
    weights = torch.softmax(scores, dim=-1)
    return (weights @ v.float()).to(q.dtype)
// q [T, Hq*D] for positions S-T..S-1; k, v [S, Hkv*D]. Query head h reads KV head h/(Hq/Hkv).
inline std::vector<float> causal_attention(const std::vector<float>& q, const std::vector<float>& k, const std::vector<float>& v,
                                           size_t T, size_t S, size_t hq, size_t hkv, size_t d) {
    std::vector<float> out(T * hq * d, 0.f), scores(S);
    float scale = 1.f / std::sqrt(float(d));
    for (size_t i = 0; i < T; ++i)
        for (size_t h = 0; h < hq; ++h) {
            size_t kvh = h / (hq / hkv), visible = S - T + i + 1;
            const float* qv = &q[(i * hq + h) * d];
            for (size_t j = 0; j < visible; ++j) {
                const float* kv = &k[(j * hkv + kvh) * d];
                scores[j] = std::inner_product(qv, qv + d, kv, 0.f) * scale;
            }
            softmax(scores.data(), visible);
            for (size_t j = 0; j < visible; ++j)
                for (size_t e = 0; e < d; ++e) out[(i * hq + h) * d + e] += scores[j] * v[(j * hkv + kvh) * d + e];
        }
    return out;
}
#![allow(unused)]
fn main() {
/// q: [T, Hq*D] for the new tokens at positions start..start+T; k, v: [S, Hkv*D] for every
/// cached position 0..S. Each query reads only keys at positions <= its own (causality).
/// Query head h reads KV head h / (Hq/Hkv): grouped-query attention.
pub fn causal_attention(q: &[f32], k: &[f32], v: &[f32], t: usize, s: usize,
                        q_heads: usize, kv_heads: usize, d: usize) -> Vec<f32> {
    let start = s - t;
    let group = q_heads / kv_heads;
    let scale = 1.0 / (d as f32).sqrt();
    let mut out = vec![0.0; t * q_heads * d];
    let mut scores = vec![0.0; s];
    for i in 0..t {
        let visible = start + i + 1; // keys 0..=start+i
        for h in 0..q_heads {
            let kvh = h / group;
            let qv = &q[(i * q_heads + h) * d..][..d];
            for j in 0..visible {
                let kv = &k[(j * kv_heads + kvh) * d..][..d];
                scores[j] = qv.iter().zip(kv).map(|(a, b)| a * b).sum::<f32>() * scale;
            }
            softmax(&mut scores[..visible]);
            let o = &mut out[(i * q_heads + h) * d..][..d];
            for j in 0..visible {
                let vv = &v[(j * kv_heads + kvh) * d..][..d];
                for (oo, vvv) in o.iter_mut().zip(vv) {
                    *oo += scores[j] * vvv;
                }
            }
        }
    }
    out
}
}

The C++ and Rust versions process one sequence with explicit loops and no batch axis. Read them to see the computation without tensor broadcasting. They implement the same rule: query $i$ of $T$ new tokens sees keys $0 \ldots S-T+i$.

What attention costs

For a sequence of $T$ tokens with model width $D$, the score matrix has $T^2$ entries per head, and computing scores and outputs costs about $4T^2D$ FLOPs per layer. At $T = 8{,}192$ and FP32, one head’s score matrix alone is 256 MiB. This quadratic growth is attention’s central problem at long context. FlashAttention (Chapter 15) removes the memory cost without changing the result. Linear attention (Chapter 28) and sparse attention (Chapter 29) change the computation itself.

Build it

Engine milestone 5: causal attention. In engine/attention.py, implement:

  1. split_heads(x, heads): [B, T, H*Dh] → [B, H, T, Dh].
  2. merge_heads(x): the inverse.
  3. causal_attention(q, k, v, query_positions=None, key_positions=None, allowed=None, scale=None): grouped-query attention with the position rule above. The _positions helper that fills in defaults is provided.
pytest tests/test_ch05_attention.py
python run.py attention --impl engine

The tests compare with PyTorch’s scaled_dot_product_attention, check this chapter’s worked example, perturb a future token to prove causality, and exercise GQA and rectangular queries.

Tip

Compute scores in FP32 (q.float() @ k.float().transpose(-2, -1)) and cast the result back to q.dtype. Build the mask by broadcasting positions: key_pos[:, None, None, :] <= query_pos[:, None, :, None] has shape [B, 1, T, S] and broadcasts over heads. Use masked_fill(~mask, float("-inf")) before torch.softmax.

Stretch exercises

  1. ★ Recompute this chapter’s six-token causal weight table with your function (use x as q, k and v with a single head, and scale=1.0). Where: experiments/ch05.py (create it), importing engine.attention.causal_attention.
  2. ★★ Write a test that catches the “reshape without transpose” bug: build q so that each head’s features are distinguishable and assert that merge_heads(split_heads(x)) equals x, but that a wrong reshape doesn’t. Where: create tests/test_ch05_stretch.py, importing engine.attention.
  3. ★★ Implement a sliding-window variant using the allowed argument: each query sees only the last $w$ keys. (You’ll meet this again in Chapter 29.) Where: add a window-mask helper to engine/attention.py and pass its result to causal_attention(allowed=...).
  4. ★★★ Measure the runtime and peak memory of causal_attention for $T$ = 512, 1,024, 2,048 and 4,096. Confirm the quadratic growth and estimate when it stops fitting in your device’s memory. Where: experiments/ch05.py (create it), importing engine.attention.causal_attention.

Check your understanding

  1. Why does attention normalize over keys (each row) rather than over queries (each column)?
  2. Why mask before the softmax?
  3. Which inputs can affect the output at position 0 of a causal attention layer?
  4. What does dividing by $\sqrt{d}$ fix, and what goes wrong without it?
  5. With 16 query heads and 8 KV heads, which KV head does query head 11 read?

Going deeper

  • BALLM Chapter 3 (pp. 50-91): attention from simplified weights to causal multi-head attention. This chapter’s worked example comes from §3.3.
  • PMPP §20.2-20.3 (pp. 482-488): multi-head attention as matrix operations, and a first CUDA implementation.
  • Vaswani et al., Attention Is All You Need (2017); Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints (2023).
  • Jay Alammar, The Illustrated Transformer: the classic visual walkthrough.