Lesson 90: KV Cache & Attention Efficiency

USAAIO Lesson 90, from Week 31 of Phase 3. It covers the KV cache for autoregressive inference, which costs O(n·d) memory and O(1) per step; Flash Attention's block-wise computation, which needs only O(n) SRAM; sliding-window attention, at O(n·w) rather than O(n²); the sparse attention patterns of BigBird and Longformer; and Group Query Attention along with Multi-Query Attention. All the complexity claims and memory numbers were verified with torch 2.7.1+cpu and numpy 2.2.6. The lesson runs to 30 slides.

Subject: Machine Learning · 59 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. KV Cache & Attention Efficiency

Title

USAAIO · Lesson 90 · Week 31 (Phase 3)

Why naive attention is O(n²), how KV caching makes autoregressive inference O(1) per step, and four mechanisms — Flash Attention, sliding window, sparse patterns, and GQA — that make Transformers viable at scale.

2. By the end of this lesson you can

Objectives

  1. State the time and memory complexity of standard self-attention and identify where it breaks at long sequences
  2. Explain KV cache: what is stored, why recomputation is avoided, and the exact memory formula
  3. Describe Flash Attention's block-wise tiling strategy and why it achieves O(n) SRAM
  4. Compare sliding window and sparse attention (BigBird/Longformer) for long-document workloads
  5. Analyze Group Query Attention (GQA/MQA): how shared K/V heads reduce KV cache and at what accuracy cost

3. What survived from GPT & Decoder-Only Transformers?

Warm-up

Discussion prompt

Before we open Lesson 90: KV Cache & Attention Efficiency: without looking back, what was the main idea of GPT & Decoder-Only Transformers, and what could you do by the end of it that you could not do before?

Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.

Answer:

GPT as a decoder-only causal transformer — no cross-attention, causal self-attention mask, learned positional embeddings, FFN at 4× expansion, GPT-2 scaling family (12–48 layers, 768–1600 dim), Chinchilla scaling laws, emergent abilities (few-shot, chain-of-thought), and autoregressive inference. Build TinyGPT from scratch in PyTorch and trace a CLM training step.

4. The O(n²) bottleneck

Section

Part 1 of 4

5. Why attention scales badly

Concept

Scaled dot-product attention (Lesson 74) computes a full n×n score matrix — every query token attends to every key token. Both time and memory blow up with sequence length.

\[ \text{Attention}(Q,K,V) = \text{softmax}\!\left(\frac{QK^\top}{\sqrt{d_k}}\right)V \]

resourcestandard self-attentionbottleneck
timeO(n² · d)matrix multiply QK^T
memoryO(n²)storing the n×n attention matrix
example (BERT-base)n=512, n_heads=12attn matrix = 12 MB total

6. Fill in: standard self-attention for Why attention scales badly

Comparison

Comparison matrix

From Why attention scales badly: refill the standard self-attention column from what you know. The rest of the table is as it appeared.

resourcestandard self-attentionbottleneck
timeO(n² · d)matrix multiply QK^T
memoryO(n²)storing the n×n attention matrix
example (BERT-base)n=512, n_heads=12attn matrix = 12 MB total

7. Autoregressive inference without a cache

Concept

During generation, the model produces one token at a time. Without a cache, at step t the model recomputes K and V for all t prior tokens from scratch.

\[ \text{cost at step } t = O(t \cdot d) \quad \Rightarrow \quad \text{total for } n \text{ tokens} = O(n^2 \cdot d) \]

For n=1024 tokens and d=768 (BERT-base), that is over 805 million redundant multiply-adds per generation step. The same K,V vectors are recomputed at every new token.

8. Break it if you can: Autoregressive inference without a cache

Counterexample

Discussion prompt

During generation, the model produces one token at a time. Without a cache, at step t the model recomputes K and V for all t prior tokens from scratch.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

For n=1024 tokens and d=768 (BERT-base), that is over 805 million redundant multiply-adds per generation step. The same K,V vectors are recomputed at every new token.

9. KV cache: O(1) per step

Section

Part 2 of 4

10. What the KV cache stores

Concept

The KV cache stores the K and V projections of all previously seen tokens. At step t, the model only computes Q, K, V for the one new token and appends to the cache.

\[ \text{cache}_t = \{K_1, K_2, \ldots, K_{t-1}\} \cup \{V_1, V_2, \ldots, V_{t-1}\} \]

At each step: compute Q_t, K_t, V_t for the new token only, append K_t / V_t to cache, then attend Q_t over the full cache. One forward pass costs O(t·d) instead of O(t²·d).

11. By analogy: What the KV cache stores

Analogy

Discussion prompt

Explain What the KV cache stores by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

The KV cache stores the K and V projections of all previously seen tokens. At step t, the model only computes Q, K, V for the one new token and appends to the cache.

12. Guess the shape of the answer: KV cache: implement and verify shapes

Estimation

Predict first

Simulate autoregressive generation with and without a cache. Each step appends a new K,V row. Verify the cache grows linearly and the new-token attention cost is O(t).

Commit before you compute: what does KV cache: implement and verify shapes come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Cache grows by one row per step; attention score vector grows by one column per step

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. K_full and V_full each have shape (T+step+1, d_k).

13. KV cache: implement and verify shapes

Worked example

Simulate autoregressive generation with and without a cache. Each step appends a new K,V row. Verify the cache grows linearly and the new-token attention cost is O(t).

import torch
from torch.nn.functional import softmax
torch.manual_seed(42)

d_k = 64
T_context = 128  # tokens already generated
gen_steps = 4    # generate 4 more tokens

# Simulate context already in the cache
K_cache = torch.randn(T_context, d_k)
V_cache = torch.randn(T_context, d_k)

for step in range(gen_steps):
    # New token: compute only Q, K, V for one token
    Q_new = torch.randn(1, d_k)
    K_new = torch.randn(1, d_k)
    V_new = torch.randn(1, d_k)
    # Append to cache
    K_full = torch.cat([K_cache, K_new], dim=0)  # (T+step+1, d_k)
    V_full = torch.cat([V_cache, V_new], dim=0)
    # Attend: Q_new over all keys
    scores = (Q_new @ K_full.T) / (d_k ** 0.5)   # (1, T+step+1)
    attn_w = softmax(scores, dim=-1)
    out = attn_w @ V_full                          # (1, d_k)
    K_cache, V_cache = K_full, V_full
    print(f'step {step}: K_cache {K_cache.shape}, score_row {scores.shape}, out {out.shape}')

Cache grows by one row per step; attention score vector grows by one column per step

Why: K_full and V_full each have shape (T+step+1, d_k). The score computation is O((T+step)·d_k) per step — linear in context length, not quadratic.

stepK_cache shapescore shapeout shape
0(129, 64)(1, 129)(1, 64)
1(130, 64)(1, 130)(1, 64)
2(131, 64)(1, 131)(1, 64)
3(132, 64)(1, 132)(1, 64)

14. What each one costs: KV cache: implement and verify shapes

Trade off

Comparison matrix

From KV cache: implement and verify shapes: every row here is a choice with a cost. Fill the out shape column, then say which row you would actually pick and what you give up for it.

stepK_cache shapescore shapeout shape
0(129, 64)(1, 129)(1, 64)
1(130, 64)(1, 130)(1, 64)
2(131, 64)(1, 131)(1, 64)
3(132, 64)(1, 132)(1, 64)

15. KV cache memory formula

Concept

The KV cache occupies memory proportional to sequence length — not n², but still significant at long context. The exact formula for float32:

\[ \text{KV cache} = 2 \times L \times n \times d_{\text{model}} \times 4\text{ bytes} \]

modelLd_modelnKV cache (float32)
BERT-base1276851236 MB
BERT-base127682048144 MB
LLaMA-style (32L, d=4096)32409620482048 MB

16. By analogy: KV cache memory formula

Analogy

Discussion prompt

Explain KV cache memory formula by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

The KV cache occupies memory proportional to sequence length — not n², but still significant at long context. The exact formula for float32:

17. Something is wrong here: KV cache eliminates the O(n²) term entirely

Anomaly

Predict first

A student writes this, and it looks reasonable:

With a KV cache, autoregressive inference is O(1) per step — like a recurrent model.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: at step t, Q_t must attend over all t cached keys/values.

KV cache reduces each step from O(t² · d) to O(t · d): recomputation eliminated, attention over cache still linear in t.

Why: at step t, Q_t must attend over all t cached keys/values. The attention dot product is O(t · d_k). Each step is O(t · d) — linear in t, not O(1). The cache saves O(t · d) recomputation of K/V but not the attention score computation itself.

18. Trap: KV cache eliminates the O(n²) term entirely

Trap

The trap

With a KV cache, autoregressive inference is O(1) per step — like a recurrent model.

Claim the cache makes each generation step constant-time

Why: Wrong: at step t, Q_t must attend over all t cached keys/values. The attention dot product is O(t · d_k). Each step is O(t · d) — linear in t, not O(1). The cache saves O(t · d) recomputation of K/V but not the attention score computation itself.

The fix

KV cache reduces each step from O(t² · d) to O(t · d): recomputation eliminated, attention over cache still linear in t.

Each step: compute (Q_t, K_t, V_t) for ONE token O(d), then attend Q_t over t cached keys: O(t · d)

Why: The cache removes quadratic recomputation of K and V. It does NOT remove the O(t) dot-product scan — total generation is still O(n² · d) summed over all steps, but each individual step is O(t · d) rather than O(t² · d).

19. Break it on purpose: KV cache eliminates the O(n²) term entirely

Break the constraint

Discussion prompt

The rule this trap just fixed:

KV cache reduces each step from O(t² · d) to O(t · d): recomputation eliminated, attention over cache still linear in t.

Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?

Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.

Answer:

at step t, Q_t must attend over all t cached keys/values. The attention dot product is O(t · d_k). Each step is O(t · d) — linear in t, not O(1). The cache saves O(t · d) recomputation of K/V but not the attention score computation itself.

20. Flash Attention & sparse variants

Section

Part 3 of 4

21. Flash Attention: IO-aware block tiling

Concept

Standard attention materializes the full n×n matrix in GPU HBM (slow, ~GB). Flash Attention (Dao et al., 2022) never writes the full matrix — it tiles Q, K, V into blocks that fit in the fast on-chip SRAM.

n=1024, d_k=64standard attentionflash attention
peak HBM allocation4.00 MB (n×n float32)O(n) ≈ 256 KB
per block in SRAMn/a64×64 = 16 KB
num blocksn/a16×16 = 256
output correctnessexactexact (same result)

22. Which is which, by standard attention

Discrimination

Sort into buckets

Sort these by standard attention, from memory, without looking back at Flash Attention: IO-aware block tiling. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

4.00 MB (n×n float32)
peak HBM allocation
n/a
per block in SRAM; num blocks
exact
output correctness
g1
standard attention is "4.00 MB (n×n float32)" for peak HBM allocation — that is what the table on "Flash Attention: IO-aware block tiling" records, and it is the single property separating this group from the rest.
g2
standard attention is "n/a" for per block in SRAM, num blocks — that is what the table on "Flash Attention: IO-aware block tiling" records, and it is the single property separating this group from the rest.
g3
standard attention is "exact" for output correctness — that is what the table on "Flash Attention: IO-aware block tiling" records, and it is the single property separating this group from the rest.

23. Sliding window attention

Concept

In sliding window attention each token attends only to a local window of w neighboring tokens. This reduces attention from O(n²) to O(n·w) — linear in n for fixed window size.

\[ \text{Attention}_{\text{sliding}}[i] = \text{softmax}\!\left(\frac{Q_i K_{[i-w, i]}^\top}{\sqrt{d_k}}\right)V_{[i-w, i]} \]

nO(n²) opsO(n·w) ops (w=512)reduction
512262,144262,1441x (n=w)
10241,048,576524,2882x
20484,194,3041,048,5764x
409616,777,2162,097,1528x

24. Watch it run: Sliding window attention

Pattern

Step through it

Step through Sliding window attention one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: n is 512
  2. Step 2: n is 1024
  3. Step 3: n is 2048
  4. Step 4: n is 4096

25. Sparse attention: BigBird & Longformer

Concept

Sparse attention patterns (Longformer, BigBird) combine three attention types to achieve O(n) complexity while maintaining global context.

attention typewhich tokenspurpose
local (sliding window)±w neighborscapture adjacent syntax
global tokensa few designated [CLS]/task tokensbroadcast document-level signal
randomr random tokenslong-range coverage (BigBird only)

BigBird/Longformer achieve O(n·(w + g + r)) complexity — linear in n. This enabled processing documents of 4,096–16,384 tokens vs BERT's 512 limit.

26. Where does each piece belong: Lesson 90: KV Cache & Attention Efficiency

Sorting

Sort into buckets

These are the pieces of Lesson 90: KV Cache & Attention Efficiency, out of order. Put each one back under the part of the lesson it belongs to.

The O(n²) bottleneck
Why attention scales badly; Autoregressive inference without a cache
KV cache: O(1) per step
What the KV cache stores; KV cache: implement and verify shapes; KV cache memory formula
Flash Attention & sparse variants
Flash Attention: IO-aware block tiling; Sliding window attention; Sparse attention: BigBird & Longformer
s1
The O(n²) bottleneck is where Lesson 90: KV Cache & Attention Efficiency puts Why attention scales badly, Autoregressive inference without a cache. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
KV cache: O(1) per step is where Lesson 90: KV Cache & Attention Efficiency puts What the KV cache stores, KV cache: implement and verify shapes, KV cache memory formula. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
Flash Attention & sparse variants is where Lesson 90: KV Cache & Attention Efficiency puts Flash Attention: IO-aware block tiling, Sliding window attention, Sparse attention: BigBird & Longformer. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

27. Guess the shape of the answer: Sliding window mask: implement and trace

Estimation

Predict first

Build a sliding window attention mask for n=6, window=2 (each token attends ±2 neighbors). Verify the sparsity pattern.

Commit before you compute: what does Sliding window mask: implement and trace come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: Active entries: 26/36 = 72.2% for n=6, w=2; for large n, approaches 2w+1 entries per row → O(n·w)

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. At n=6, boundary effects keep sparsity modest.

28. Sliding window mask: implement and trace

Worked example

Build a sliding window attention mask for n=6, window=2 (each token attends ±2 neighbors). Verify the sparsity pattern.

import torch

def sliding_window_mask(n, w):
    '''True = attend, False = masked. Each token attends to [i-w, i+w].'''
    mask = torch.zeros(n, n, dtype=torch.bool)
    for i in range(n):
        lo = max(0, i - w)
        hi = min(n, i + w + 1)
        mask[i, lo:hi] = True
    return mask

n, w = 6, 2
mask = sliding_window_mask(n, w)
print(mask.int())   # 1=attend, 0=masked
print(f'Active entries: {mask.sum().item()} / {n*n} = {mask.float().mean():.1%}')

Active entries: 26/36 = 72.2% for n=6, w=2; for large n, approaches 2w+1 entries per row → O(n·w)

Why: At n=6, boundary effects keep sparsity modest. At large n (e.g. n=4096, w=256), each row has exactly 513 active entries vs 4096 — 87.5% reduction.

row iattend to colscount
00,1,23
10,1,2,34
20,1,2,3,45
31,2,3,4,55
42,3,4,54
53,4,53

29. Watch it run: Sliding window mask: implement and trace

Pattern

Step through it

Step through Sliding window mask: implement and trace one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: row i is 0
  2. Step 2: row i is 1
  3. Step 3: row i is 2
  4. Step 4: row i is 3
  5. Step 5: row i is 4
  6. Step 6: row i is 5

30. Group Query Attention

Section

Part 4 of 4

31. Multi-head vs Multi-query vs Group Query

Concept

Standard Multi-Head Attention (MHA) has n_q=n_k=n_v=H separate head projections. Multi-Query Attention (MQA) shares one K,V head across all Q heads. GQA generalizes with G groups.

\[ \text{GQA: } H_Q \text{ query heads}, \; H_{KV} \text{ KV heads}, \; G = H_Q / H_{KV} \text{ queries per KV group} \]

variantH_QH_KVKV cache vs MHAquality
MHA32321x (baseline)highest
GQA-83284x smaller≈MHA
GQA-43248x smallerslight loss
MQA32132x smallernoticeable drop

32. Fill in: quality for Multi-head vs Multi-query vs Group Query

Comparison

Comparison matrix

From Multi-head vs Multi-query vs Group Query: refill the quality column from what you know. The rest of the table is as it appeared.

variantH_QH_KVKV cache vs MHAquality
MHA32321x (baseline)highest
GQA-83284x smaller≈MHA
GQA-43248x smallerslight loss
MQA32132x smallernoticeable drop

33. Guess the shape of the answer: GQA: verify KV cache reduction (LLaMA-scale)

Estimation

Predict first

Compute actual KV cache sizes for a 32-layer, d=4096, H_Q=32 model at seq_len=2048, comparing MHA vs GQA-8 vs MQA.

Commit before you compute: what does GQA: verify KV cache reduction (LLaMA-scale) come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: GQA-8 cuts KV cache from 2048 MB → 512 MB (4x) with minimal quality drop; LLaMA-2 and LLaMA-3 use GQA-8

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Each K/V head is shared by G=H_Q/H_KV query heads.

34. GQA: verify KV cache reduction (LLaMA-scale)

Worked example

Compute actual KV cache sizes for a 32-layer, d=4096, H_Q=32 model at seq_len=2048, comparing MHA vs GQA-8 vs MQA.

import numpy as np

# LLaMA-style model params
d_model, n_q_heads, n_layers = 4096, 32, 32
d_head = d_model // n_q_heads   # 128
seq_len = 2048

for label, n_kv in [('MHA', 32), ('GQA-8', 8), ('GQA-4', 4), ('MQA', 1)]:
    # KV cache: 2 matrices of (n_kv_heads, seq_len, d_head) per layer
    kv_elements = 2 * n_layers * n_kv * d_head * seq_len
    kv_MB = kv_elements * 4 / (1024**2)  # float32
    reduction = 32 // n_kv
    print(f'{label:6s} n_kv={n_kv:2d}: {kv_MB:6.0f} MB  ({reduction:2d}x reduction)')

GQA-8 cuts KV cache from 2048 MB → 512 MB (4x) with minimal quality drop; LLaMA-2 and LLaMA-3 use GQA-8

Why: Each K/V head is shared by G=H_Q/H_KV query heads. With H_KV=8 instead of 32, the number of K and V projections stored in the cache drops by 4x. The Q projections are unaffected.

variantn_kvKV cache (MB)reduction
MHA3220481x
GQA-885124x
GQA-442568x
MQA16432x

35. Watch it run: GQA: verify KV cache reduction (LLaMA-scale)

Pattern

Step through it

Step through GQA: verify KV cache reduction (LLaMA-scale) one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: variant is MHA
  2. Step 2: variant is GQA-8
  3. Step 3: variant is GQA-4
  4. Step 4: variant is MQA

36. Something is wrong here: Flash Attention changes the attention output

Anomaly

Predict first

A student writes this, and it looks reasonable:

Flash Attention is an approximate method — it trades a small amount of precision for speed and memory savings by working in smaller blocks.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: Flash Attention is mathematically exact.

Flash Attention is exact — it produces the same output as standard attention, just without materializing the n×n matrix.

Why: Flash Attention is mathematically exact. It uses the online softmax (log-sum-exp) trick to accumulate the numerically identical result block-by-block. The only difference from standard attention is IO cost — the output is bit-for-bit the same (up to floating-point associativity).

37. Trap: Flash Attention changes the attention output

Trap

The trap

Flash Attention is an approximate method — it trades a small amount of precision for speed and memory savings by working in smaller blocks.

Accept that FA outputs differ slightly from standard attention due to block-level approximations

Why: Wrong: Flash Attention is mathematically exact. It uses the online softmax (log-sum-exp) trick to accumulate the numerically identical result block-by-block. The only difference from standard attention is IO cost — the output is bit-for-bit the same (up to floating-point associativity).

The fix

Flash Attention is exact — it produces the same output as standard attention, just without materializing the n×n matrix.

FA tiles Q, K, V into blocks; uses running log-sum-exp to maintain exact softmax numerics; merges results at the end

Why: The online softmax identity: softmax([a,b,c]) can be computed as a running update without storing all logits simultaneously. This is a numerical identity, not an approximation — FA's memory saving is purely about IO scheduling.

38. Which of these survive contact with Lesson 90: KV Cache & Attention Efficiency?

Two truths and a lie

Sort into buckets

Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.

Holds up
Scaled dot-product attention (Lesson 74) computes a full n×n score matrix — every query token attends to every key token. Both time and memory blow up with sequence length.; During generation, the model produces one token at a time. Without a cache, at step t the model recomputes K and V for all t prior tokens from scratch.; The KV cache occupies memory proportional to sequence length — not n², but still significant at long context. The exact formula for float32:
Breaks
With a KV cache, autoregressive inference is O(1) per step — like a recurrent model.; Flash Attention is an approximate method — it trades a small amount of precision for speed and memory savings by working in smaller blocks.
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 90: KV Cache & Attention Efficiency puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

39. Rebuild the recipe: Attention efficiency decision map

Ranking

Put in order

These are the steps of Attention efficiency decision map, scrambled. Put them back in order before the next slide shows you.

  1. KV cache (inference): store K,V per layer; cost per step = O(t·d) not O(t²·d); memory = 2·L·n·d·4 bytes
  2. Flash Attention (training + inference): block-tile to fit SRAM; exact output; O(n) HBM memory vs O(n²)
  3. Sliding window (long docs): each token attends ±w; O(n·w) time; use when local context suffices
  4. Sparse (BigBird/Longformer): local + global + random; O(n·(w+g+r)); add global tokens for cross-doc signals
  5. GQA/MQA: share K,V heads across G query heads; KV cache shrinks G×; combine with KV cache for inference

Why: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.

40. Attention efficiency decision map

Pattern

  1. KV cache (inference): store K,V per layer; cost per step = O(t·d) not O(t²·d); memory = 2·L·n·d·4 bytes
  2. Flash Attention (training + inference): block-tile to fit SRAM; exact output; O(n) HBM memory vs O(n²)
  3. Sliding window (long docs): each token attends ±w; O(n·w) time; use when local context suffices
  4. Sparse (BigBird/Longformer): local + global + random; O(n·(w+g+r)); add global tokens for cross-doc signals
  5. GQA/MQA: share K,V heads across G query heads; KV cache shrinks G×; combine with KV cache for inference

41. Where does it stop working: Attention efficiency decision map

Edge cases

Discussion prompt

Attention efficiency decision map works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".

Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.

Answer:

  1. KV cache (inference): store K,V per layer; cost per step = O(t·d) not O(t²·d); memory = 2·L·n·d·4 bytes
  2. Flash Attention (training + inference): block-tile to fit SRAM; exact output; O(n) HBM memory vs O(n²)
  3. Sliding window (long docs): each token attends ±w; O(n·w) time; use when local context suffices
  4. Sparse (BigBird/Longformer): local + global + random; O(n·(w+g+r)); add global tokens for cross-doc signals
  5. GQA/MQA: share K,V heads across G query heads; KV cache shrinks G×; combine with KV cache for inference

42. Rule out three: Check yourself — KV cache cost per step

Elimination

Eliminate the wrong options

With a KV cache, what is the time complexity for the attention computation at generation step t (attending over t context tokens)?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. O(t · d_k)
  • B. O(1)
  • C. O(t² · d_k)
  • D. O(n · d_k) where n is the total vocabulary size

Survives elimination: A

Why: At step t, Q_t (shape 1 × d_k) dots with K_cache (shape t × d_k): one dot product per cached key → O(t · d_k). The cache eliminates recomputing K,V but the Q-to-K attention scan is still linear in t.

43. Check yourself — KV cache cost per step

Check

Work out the cost before clicking.

Check your understanding

With a KV cache, what is the time complexity for the attention computation at generation step t (attending over t context tokens)?

  • A. O(t · d_k) (correct)
  • B. O(1)
  • C. O(t² · d_k)
  • D. O(n · d_k) where n is the total vocabulary size

Answer: A

Why: At step t, Q_t (shape 1 × d_k) dots with K_cache (shape t × d_k): one dot product per cached key → O(t · d_k). The cache eliminates recomputing K,V but the Q-to-K attention scan is still linear in t.

Why B tempts people
O(1) would require the model to ignore context entirely. The attention score must be computed against all t cached tokens — that is O(t·d_k), not constant.
Why C tempts people
O(t²·d_k) is the cost WITHOUT a cache — the model would recompute K and V for all t tokens each step. The cache converts that recomputation into O(t·d_k) by storing K,V from prior steps.
Why D tempts people
Vocabulary size is irrelevant to attention computation. Attention operates over sequence positions (t context tokens), not over the vocabulary.

44. Answer it before you see the options: Check yourself — Flash Attention memory

Prediction

Predict first

Flash Attention achieves O(n) peak GPU SRAM usage (vs O(n²) for standard attention) without approximation. What is the core technique that enables this?

Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.

Correct: Block-tiling Q, K, V with an online log-sum-exp softmax accumulation

Why: FA tiles Q, K, V into blocks that fit in fast SRAM. The online softmax trick (running log-sum-exp) lets it accumulate the exact output for each output block without ever materializing the full n×n score matrix.

45. Check yourself — Flash Attention memory

Check

Recall Flash Attention's key property.

Check your understanding

Flash Attention achieves O(n) peak GPU SRAM usage (vs O(n²) for standard attention) without approximation. What is the core technique that enables this?

  • A. Block-tiling Q, K, V with an online log-sum-exp softmax accumulation (correct)
  • B. Quantizing attention weights to int8 to reduce memory
  • C. Sharing K and V heads across multiple Q heads (like GQA)
  • D. Only computing attention within a fixed window of w tokens

Answer: A

Why: FA tiles Q, K, V into blocks that fit in fast SRAM. The online softmax trick (running log-sum-exp) lets it accumulate the exact output for each output block without ever materializing the full n×n score matrix.

Why B tempts people
Quantization reduces per-element byte cost but still requires the full n×n matrix to be allocated — memory is O(n²) with int8 just as with float32. FA's O(n) comes from never materializing that matrix at all.
Why C tempts people
GQA/MQA reduce KV cache size during inference but do not change the O(n²) attention matrix at training time. They are orthogonal to Flash Attention.
Why D tempts people
Windowed (local) attention changes the sparsity pattern and produces an approximate/different output. Flash Attention has no window — it computes full attention exactly.

46. Rule out three: Check yourself — GQA reduction

Elimination

Eliminate the wrong options

A model has 32 query heads and uses GQA with 8 KV heads (GQA-8). By what factor does the KV cache shrink compared to standard MHA (32 KV heads)?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. 4x
  • B. 32x
  • C. 8x
  • D. 2x

Survives elimination: A

Why: GQA-8 uses H_KV=8 KV heads vs MHA's H_KV=32. KV cache is proportional to H_KV, so the reduction is 32/8 = 4x. Verified: 2048 MB (MHA) → 512 MB (GQA-8) for the LLaMA-scale example.

47. Check yourself — GQA reduction

Check

Apply the memory formula.

Check your understanding

A model has 32 query heads and uses GQA with 8 KV heads (GQA-8). By what factor does the KV cache shrink compared to standard MHA (32 KV heads)?

  • A. 4x (correct)
  • B. 32x
  • C. 8x
  • D. 2x

Answer: A

Why: GQA-8 uses H_KV=8 KV heads vs MHA's H_KV=32. KV cache is proportional to H_KV, so the reduction is 32/8 = 4x. Verified: 2048 MB (MHA) → 512 MB (GQA-8) for the LLaMA-scale example.

Why B tempts people
32x is the MQA reduction (H_KV=1, one shared KV head). GQA-8 retains 8 KV heads, not 1, so the reduction is only 32/8=4x.
Why C tempts people
8x would require H_KV=4 (GQA-4), not H_KV=8. The reduction factor is H_Q/H_KV = 32/8 = 4, not equal to H_KV.
Why D tempts people
2x would require H_KV=16. With H_KV=8, the factor is 32/8=4x. Confusing the number of KV heads with the reduction ratio.

48. Your turn: implement KV cache

Section

Project

49. Project: autoregressive attention with KV cache

Concept

Build a single-head causal attention module that maintains a KV cache across generation steps. Three milestones: manual cache append → causal masking → profile vs no-cache.

#milestonekey tool
1Implement cache-append and attend new Q over full cachetorch.cat, softmax
2Add causal mask so future tokens are blockedtorch.tril, masked_fill
3Implement sliding window mask and compare active entriesmanual mask loop

Build rules: use fixed torch.manual_seed(42), print shapes at every step, and verify that without the mask the output changes when you permute past tokens.

50. Break it if you can: Project: autoregressive attention with KV cache

Counterexample

Discussion prompt

Build a single-head causal attention module that maintains a KV cache across generation steps. Three milestones: manual cache append → causal masking → profile vs no-cache.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

Build rules: use fixed torch.manual_seed(42), print shapes at every step, and verify that without the mask the output changes when you permute past tokens.

51. Milestone 1 — cache append and attend

Worked example

Your turn: generate 4 tokens autoregressively. At each step, compute Q/K/V for the new token, append to cache, and attend. Predict the shape of the score tensor at step 3.

Hint: K_full = torch.cat([K_cache, K_new], dim=0) gives shape (T+step+1, d_k). Score = Q_new @ K_full.T / sqrt(d_k) is shape (1, T+step+1).

import torch
from torch.nn.functional import softmax
torch.manual_seed(42)
d_k, T = 64, 128
K_cache = torch.randn(T, d_k)
V_cache = torch.randn(T, d_k)
for step in range(4):
    Q_new = torch.randn(1, d_k)
    K_new = torch.randn(1, d_k)
    V_new = torch.randn(1, d_k)
    K_full = torch.cat([K_cache, K_new], dim=0)
    V_full = torch.cat([V_cache, V_new], dim=0)
    scores = Q_new @ K_full.T / (d_k ** 0.5)
    out = softmax(scores, dim=-1) @ V_full
    K_cache, V_cache = K_full, V_full
    print(f'step {step}: K_cache {K_cache.shape}, scores {scores.shape}')
stepK_cache shapescores shape
0(129, 64)(1, 129)
1(130, 64)(1, 130)
2(131, 64)(1, 131)
3(132, 64)(1, 132)

52. Watch it run: Milestone 1 — cache append and attend

Pattern

Step through it

Step through Milestone 1 — cache append and attend one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: step is 0
  2. Step 2: step is 1
  3. Step 3: step is 2
  4. Step 4: step is 3

53. Milestone 2 — causal mask

Worked example

Your turn: add a causal (lower-triangular) mask to the attention so token i cannot attend to j>i. Predict: what does torch.tril(torch.ones(6,6)) look like?

Hint: mask = torch.tril(torch.ones(n, n, dtype=torch.bool)). Apply: scores = scores.masked_fill(~mask, -1e9) before softmax.

import torch
torch.manual_seed(0)
n = 6
d_k = 8
Q = torch.randn(n, d_k)
K = torch.randn(n, d_k)
V = torch.randn(n, d_k)

scores = Q @ K.T / (d_k ** 0.5)   # (6, 6)
causal_mask = torch.tril(torch.ones(n, n, dtype=torch.bool))
scores = scores.masked_fill(~causal_mask, float('-inf'))
attn = torch.softmax(scores, dim=-1)
print('attn[0]:', attn[0].round(4))  # only token 0 visible
print('attn[5]:', attn[5].round(4))  # all 6 tokens visible
row inon-zero cols (can attend to)zero cols (masked)
001, 2, 3, 4, 5
10, 12, 3, 4, 5
30, 1, 2, 34, 5
50, 1, 2, 3, 4, 5(none)

54. What each one costs: Milestone 2 — causal mask

Trade off

Comparison matrix

From Milestone 2 — causal mask: every row here is a choice with a cost. Fill the non-zero cols (can attend to) column, then say which row you would actually pick and what you give up for it.

row inon-zero cols (can attend to)zero cols (masked)
001, 2, 3, 4, 5
10, 12, 3, 4, 5
30, 1, 2, 34, 5
50, 1, 2, 3, 4, 5(none)

55. Milestone 3 — sliding window mask sparsity

Worked example

Your turn: implement a sliding window mask for n=6, w=2. Count active entries and compare to full attention. Predict the active count for large n.

Hint: for i in range(n): mask[i, max(0,i-w):min(n,i+w+1)] = True. For large n, each row has exactly 2w+1 active entries.

import torch

def sliding_mask(n, w):
    m = torch.zeros(n, n, dtype=torch.bool)
    for i in range(n):
        m[i, max(0, i-w):min(n, i+w+1)] = True
    return m

for n, w in [(6, 2), (1024, 256)]:
    mask = sliding_mask(n, w)
    active = mask.sum().item()
    total = n * n
    print(f'n={n}, w={w}: {active}/{total} active = {active/total:.1%}')
nwactive entriesfractionvs O(n²)
622672.2%boundary effects
1024256523,26450.0%2x reduction
40965124,190,20825.0%4x reduction

56. Fill in: w for Milestone 3 — sliding window mask sparsity

Comparison

Comparison matrix

From Milestone 3 — sliding window mask sparsity: refill the w column from what you know. The rest of the table is as it appeared.

nwactive entriesfractionvs O(n²)
622672.2%boundary effects
1024256523,26450.0%2x reduction
40965124,190,20825.0%4x reduction

57. Show it off

Concept

Out loud, slides closed: (1) explain why the KV cache does NOT make each generation step O(1); (2) state what Flash Attention stores in SRAM and why the output is exact; (3) compute the GQA KV cache size for a 32-layer model with d=4096, n_heads=32, GQA-4, seq=2048.

Stretch (homework): implement a toy single-layer GQA module in PyTorch — H_Q=8 query heads, H_KV=2 KV heads. Expand K/V to match Q via repeat_interleave. Profile the KV cache footprint vs full MHA at seq=1024. Next: Lesson 91 — Vision Transformer (ViT): patch embeddings, positional encodings, and the full ViT forward pass.

58. Connect it up: Lesson 90: KV Cache & Attention Efficiency

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — The O(n²) bottleneck · KV cache: O(1) per step · Flash Attention & sparse variants · Group Query Attention · Your turn: implement KV cache. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

59. What you can do now

Recap

mechanismkey formula / insight
Standard attnO(n²·d) time, O(n²) memory — bottleneck at long n
KV cacheO(t·d) per step; memory = 2·L·n·d·4 bytes
Flash Attentionsame output; O(n) SRAM; block-tile + online softmax
Sliding windowO(n·w); each token attends ±w; ratio n/w speedup at large n
GQAH_KV < H_Q; KV cache = H_Q/H_KV × smaller; MQA extreme (H_KV=1)

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 90 (Week 31 — Attention Variants & Vision Transformer) — Barron · USAAIO Round 2 Preparation, 2026
  2. Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (2022)
  3. Ainslie et al., GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints (2023)
  4. KV cache memory, attention complexity, GQA reduction factors verified with torch 2.7.1+cpu, numpy 2.2.6, June 2026 — Real execution, verified

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108