PidokuInfra

Attention from First Principles

Foundations Intermediate 2h Difficulty 4/5 Topic 08 of 12

Prerequisites 01, 02, 03, 07

This is the most important file in Section III. Everything in Section V depends on understanding it mechanically, not just conceptually.


1. What is it?#

Attention lets each position in a sequence look at every other position and build its representation from a weighted mixture of them, where the weights are computed from the data itself.

For each token:
   1. Ask a question           (query)
   2. Compare it to what every other token offers  (keys)
   3. Take a weighted average of what those tokens contain (values)

In one line:

Attention(Q, K, V) = softmax(Q Kᵀ / √d_head) V

2. Why does it exist?#

Before attention, sequence models were recurrent: process token 1, then 2, then 3, carrying a fixed-size hidden state. Two fatal problems:

  1. The bottleneck. All information about a 1,000-token history had to fit in one vector.
  2. No parallelism. Token 500 could not be processed until token 499 was done.

Attention solves both: every token can directly access every other token (no bottleneck), and all positions can be computed simultaneously during training and prefill (parallel).

The cost is quadratic: S tokens each looking at S tokens is S² comparisons. That quadratic is the defining engineering constraint of modern LLM inference, and most of Sections V, VII, and XIII are responses to it.


3. Simple analogy#

A research library with a very efficient reference desk.

You arrive with a question (query): “what does this pronoun refer to?”

Every document in the library has an index card describing what it’s about (key) and actual contents (value).

You compare your question to every index card, score the match, convert the scores to weights (softmax), and then read a blend of the documents in proportion to their weight. Highly relevant documents contribute most; irrelevant ones contribute ~0.

Three refinements that map exactly onto the real mechanism:

  • You do this for every question simultaneously (all queries at once = a matrix multiply).
  • You consult several specialists in parallel, each looking for different things — one for grammar, one for topic, one for entity references. Those are heads.
  • In a causal model, you may only read documents filed before yours. That’s the causal mask.

4. Tiny example#

Full attention, worked by hand. Sequence of 3 tokens, d_head = 2.

Q = [1  0]      K = [1  0]      V = [10  0]
    [0  1]          [0  1]          [ 0 10]
    [1  1]          [1  1]          [ 5  5]

Step 1: scores = Q Kᵀ

  Q[0]·K[0] = 1·1+0·0 = 1     Q[0]·K[1] = 0     Q[0]·K[2] = 1
  Q[1]·K[0] = 0               Q[1]·K[1] = 1     Q[1]·K[2] = 1
  Q[2]·K[0] = 1               Q[2]·K[1] = 1     Q[2]·K[2] = 2

  scores = [1  0  1]
           [0  1  1]
           [1  1  2]

Step 2: scale by 1/√2 = 0.7071

  scores = [0.707  0      0.707]
           [0      0.707  0.707]
           [0.707  0.707  1.414]

Step 3: causal mask (token i may only see tokens ≤ i)

  scores = [0.707  -inf   -inf ]
           [0      0.707  -inf ]
           [0.707  0.707  1.414]

Step 4: softmax each row

  row 0: exp([0.707]) normalized → [1.000, 0, 0]
  row 1: exp([0, 0.707]) = [1, 2.028]; sum 3.028 → [0.330, 0.670, 0]
  row 2: exp([0.707,0.707,1.414]) = [2.028, 2.028, 4.113]; sum 8.169
         → [0.248, 0.248, 0.504]

Step 5: multiply by V

  out[0] = 1.000·[10,0]                                    = [10.00, 0.00]
  out[1] = 0.330·[10,0] + 0.670·[0,10]                     = [ 3.30, 6.70]
  out[2] = 0.248·[10,0] + 0.248·[0,10] + 0.504·[5,5]       = [ 5.00, 5.00]

Read what happened. Token 0 could only see itself, so it got V[0] unchanged. Token 1 mixed tokens 0 and 1, weighted toward 1 (its query matched key 1 better). Token 2, whose query pointed “diagonally,” matched key 2 best and blended all three.

Do this by hand once. It converts attention from a formula you’ve memorized into a mechanism you understand.


5. Technical explanation#

The full multi-head computation#

# x: (B, S, d)
q = x @ Wq.T          # (B, S, h·d_head)
k = x @ Wk.T          # (B, S, h_kv·d_head)
v = x @ Wv.T          # (B, S, h_kv·d_head)

# split into heads
q = q.view(B, S, h,    d_head).transpose(1, 2)    # (B, h,    S, d_head)
k = k.view(B, S, h_kv, d_head).transpose(1, 2)    # (B, h_kv, S, d_head)
v = v.view(B, S, h_kv, d_head).transpose(1, 2)

# RoPE applied to q and k here (position information)
q, k = apply_rope(q, k, positions)

# if GQA: repeat k,v to match h query heads
k = k.repeat_interleave(h // h_kv, dim=1)
v = v.repeat_interleave(h // h_kv, dim=1)

scores = q @ k.transpose(-2, -1) / math.sqrt(d_head)   # (B, h, S, S)   ← the big one
scores = scores + causal_mask
attn   = softmax(scores, dim=-1)
out    = attn @ v                                       # (B, h, S, d_head)

out = out.transpose(1, 2).reshape(B, S, d)
out = out @ Wo.T                                        # (B, S, d)

Nine steps. Memorize this. Every optimization in Sections VII and XIII modifies one of them.

Why multiple heads#

With h heads, each of dimension d_head = d/h, you run h independent attention computations on different learned projections and concatenate. Cost is identical to one head of dimension d (the projections are just split), but the model can attend to h different kinds of relationship simultaneously.

Typical: d=4096, h=32, d_head=128.

The causal mask#

A decoder-only LM must not see the future:

        k0  k1  k2  k3
  q0 [  ✓   ✗   ✗   ✗ ]
  q1 [  ✓   ✓   ✗   ✗ ]
  q2 [  ✓   ✓   ✓   ✗ ]
  q3 [  ✓   ✓   ✓   ✓ ]

Consequences:

  • Half the S² score matrix is thrown away — good kernels skip computing it entirely.
  • The KV for position i never changes once computed. This is the property that makes the KV cache possible (Section V.05). If attention were bidirectional, you’d have to recompute everything every step.

MHA, GQA, MQA — the KV cache lever#

MHA (Multi-Head Attention):   h query heads, h key/value heads
GQA (Grouped-Query):          h query heads, h_kv key/value heads (h_kv < h)
MQA (Multi-Query):            h query heads, 1 key/value head
          queries          keys/values
MHA:   [q0 q1 q2 q3]    [k0 k1 k2 k3]      KV cache: 4 units
GQA:   [q0 q1 q2 q3]    [k0    k1   ]      KV cache: 2 units   (groups of 2)
        \__/  \__/
MQA:   [q0 q1 q2 q3]    [k0         ]      KV cache: 1 unit
        \________/

For Llama-3-70B: h=64, h_kv=8 → 8x smaller KV cache than MHA. That directly multiplies your maximum concurrency by 8. This is one of the highest-impact architecture decisions for serving, made at training time.

Complexity#

                        FLOPs                 Memory (naive)
Projections (q,k,v,o):  O(S · d²)             O(S · d)
Score computation:      O(S² · d)             O(S²) per head   ← the problem
Softmax:                O(S²)                 O(S²)
attn @ V:               O(S² · d)             O(S · d)

At S = 8192, h = 32, B = 8, the score tensor alone in FP16:

8 · 32 · 8192 · 8192 · 2 bytes = 34.4 GB      for ONE layer

Impossible. Hence FlashAttention, which computes it in tiles and never stores it (Section VII.10) using the online softmax you built in file 07.


6. Under the hood#

During decode, the shapes change dramatically:

Prefill (S=2048):
   q: (B, h, 2048, 128)
   k: (B, h_kv, 2048, 128)
   scores: (B, h, 2048, 2048)      ← big, compute-bound

Decode (one new token, context 2048):
   q: (B, h, 1, 128)               ← ONE query row
   k: (B, h_kv, 2049, 128)         ← from the cache
   scores: (B, h, 1, 2049)         ← a vector, not a matrix

Attention during decode is a matrix-vector operation against the whole KV cache. Per step it reads the entire KV cache and does very little arithmetic:

FLOPs  ≈ 2 · B · h · S · d_head · 2       (scores + attn@V)
Bytes  ≈ 2 · B · h_kv · S · d_head · 2    (read K and V)

intensity ≈ h / h_kv          ← for MHA that's 1; for GQA with 8:1, it's 8

Decode attention has arithmetic intensity equal to the GQA ratio. That is a beautiful and underappreciated fact: GQA doesn’t just shrink the cache, it raises attention’s arithmetic intensity by the grouping factor, making it less memory-bound. Two wins from one change.


7. Performance implications#

PhaseAttention characterBottleneck
Prefill, short (S<1k)GEMM-ish, compute-boundtensor cores
Prefill, long (S>8k)quadratic, compute-bound, memory-hungryFLOPs + activation memory
Decode, small batchGEMV against KV cacheKV read bandwidth
Decode, long contextKV read dominates everythingKV bandwidth

At 128k context, decode attention can read more bytes than the model weights:

Llama-3-8B, 128k context, GQA(8), FP16:
  KV bytes = 2 · 32 · 8 · 128 · 2 · 131072 = 17.2 GB per sequence
  weights  = 16 GB

The cache is bigger than the model. Every long-context technique in Section XIII exists to attack this number.


8. Production implications#

  • Always use a fused attention kernel (FlashAttention, FlashInfer, xFormers memory-efficient attention, or your engine’s built-in). Never materialize scores.
  • GQA/MQA are first-class serving considerations. When choosing between models of similar quality, the one with fewer KV heads is dramatically cheaper to serve.
  • Long context is not free. Advertise it, price it, and cap it. A 128k-context request can consume 20x the resources of a 8k one.
  • Different kernels for prefill and decode. They have opposite shapes; engines dispatch to different implementations.
  • Watch for attention kernel fallbacks. If your sequence length, head dim, or dtype falls outside the fast kernel’s support matrix, you may silently get a 5x slower path.

9. Common mistakes#

Materializing the score matrix. Fine at S=512, fatal at S=8192.

Forgetting the causal mask — the model sees the future and produces impossibly good perplexity in testing and garbage in generation.

Applying RoPE after caching, or to values. RoPE goes on q and k, before caching k.

Confusing h and h_kv. GQA needs the repeat step; forgetting it gives a shape error, or worse, silently broadcasts.

Assuming attention is the FLOP bottleneck. Below ~30k context, the FFN dominates FLOPs. Attention dominates memory at long context.

Softmax over the wrong axis. Over queries instead of keys.

Ignoring that decode attention is memory-bound. Optimizing its FLOPs achieves nothing.


10. Hands-on exercise#

A. By hand. Redo the section 4 example with V = [[1,2],[3,4],[5,6]] and no causal mask. Then with the mask. Compare.

B. Implement it. Write multi-head attention in Go: projections, head split, RoPE, causal mask, softmax, output projection. Check it against the block function in file 09.

C. GQA. Extend B to support h_kv < h. Verify that with h_kv = h it matches MHA. Compute the KV memory saving for Llama-3-70B’s configuration.

D. Measure the quadratic. Time attention for S ∈ {128, 256, 512, 1024, 2048, 4096} at fixed batch. Plot time vs S on log-log axes. Confirm the slope is 2 for the score computation. At what S does memory become the limit for a naive implementation?

E. Decode vs prefill. Implement both shapes explicitly (S_q = S vs S_q = 1) and measure achieved TFLOP/s and GB/s for each. Confirm one is compute-bound and one is memory-bound.

F. Intensity of GQA. Empirically verify that decode attention’s arithmetic intensity equals h/h_kv by measuring bytes read and FLOPs for h_kv ∈ {1, 8, 32} with h = 32.


11. Interview questions#

  1. Explain attention in three sentences with no equations.
  2. Write the attention formula and explain every term including the scale factor.
  3. Why is the causal mask what makes the KV cache possible?
  4. What are MHA, GQA, and MQA? Quantify the KV cache difference.
  5. Why does GQA improve decode arithmetic intensity, not just memory?
  6. Why can’t you materialize the attention score matrix at long context? Do the arithmetic.
  7. How do prefill and decode attention differ in shape and in bottleneck?
  8. At what context length does the KV cache exceed the model weights for an 8B GQA model?

12. Further reading#

  • [FUNDAMENTAL] Vaswani et al., “Attention Is All You Need” (2017)
  • [FUNDAMENTAL] Jay Alammar, “The Illustrated Transformer”
  • [ESTABLISHED] Shazeer, “Fast Transformer Decoding: One Write-Head is All You Need” (MQA, 2019)
  • [ESTABLISHED] Ainslie et al., “GQA” (2023)
  • [FUNDAMENTAL] Karpathy, “Let’s build GPT” — build it yourself
  • Next: 09 — The transformer

↑↓ navigate↵ openesc close