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
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.
Objectives
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.
Section
Part 1 of 4
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 \]
| resource | standard self-attention | bottleneck |
|---|---|---|
| time | O(n² · d) | matrix multiply QK^T |
| memory | O(n²) | storing the n×n attention matrix |
| example (BERT-base) | n=512, n_heads=12 | attn matrix = 12 MB total |
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.
| resource | standard self-attention | bottleneck |
|---|---|---|
| time | O(n² · d) | matrix multiply QK^T |
| memory | O(n²) | storing the n×n attention matrix |
| example (BERT-base) | n=512, n_heads=12 | attn matrix = 12 MB total |
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.
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.
Section
Part 2 of 4
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).
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.
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).
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.
| step | K_cache shape | score shape | out 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) |
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.
| step | K_cache shape | score shape | out 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) |
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} \]
| model | L | d_model | n | KV cache (float32) |
|---|---|---|---|---|
| BERT-base | 12 | 768 | 512 | 36 MB |
| BERT-base | 12 | 768 | 2048 | 144 MB |
| LLaMA-style (32L, d=4096) | 32 | 4096 | 2048 | 2048 MB |
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:
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.
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.
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).
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.
Section
Part 3 of 4
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=64 | standard attention | flash attention |
|---|---|---|
| peak HBM allocation | 4.00 MB (n×n float32) | O(n) ≈ 256 KB |
| per block in SRAM | n/a | 64×64 = 16 KB |
| num blocks | n/a | 16×16 = 256 |
| output correctness | exact | exact (same result) |
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.
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]} \]
| n | O(n²) ops | O(n·w) ops (w=512) | reduction |
|---|---|---|---|
| 512 | 262,144 | 262,144 | 1x (n=w) |
| 1024 | 1,048,576 | 524,288 | 2x |
| 2048 | 4,194,304 | 1,048,576 | 4x |
| 4096 | 16,777,216 | 2,097,152 | 8x |
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?
Concept
Sparse attention patterns (Longformer, BigBird) combine three attention types to achieve O(n) complexity while maintaining global context.
| attention type | which tokens | purpose |
|---|---|---|
| local (sliding window) | ±w neighbors | capture adjacent syntax |
| global tokens | a few designated [CLS]/task tokens | broadcast document-level signal |
| random | r random tokens | long-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.
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.
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.
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 i | attend to cols | count |
|---|---|---|
| 0 | 0,1,2 | 3 |
| 1 | 0,1,2,3 | 4 |
| 2 | 0,1,2,3,4 | 5 |
| 3 | 1,2,3,4,5 | 5 |
| 4 | 2,3,4,5 | 4 |
| 5 | 3,4,5 | 3 |
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?
Section
Part 4 of 4
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} \]
| variant | H_Q | H_KV | KV cache vs MHA | quality |
|---|---|---|---|---|
| MHA | 32 | 32 | 1x (baseline) | highest |
| GQA-8 | 32 | 8 | 4x smaller | ≈MHA |
| GQA-4 | 32 | 4 | 8x smaller | slight loss |
| MQA | 32 | 1 | 32x smaller | noticeable drop |
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.
| variant | H_Q | H_KV | KV cache vs MHA | quality |
|---|---|---|---|---|
| MHA | 32 | 32 | 1x (baseline) | highest |
| GQA-8 | 32 | 8 | 4x smaller | ≈MHA |
| GQA-4 | 32 | 4 | 8x smaller | slight loss |
| MQA | 32 | 1 | 32x smaller | noticeable drop |
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.
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.
| variant | n_kv | KV cache (MB) | reduction |
|---|---|---|---|
| MHA | 32 | 2048 | 1x |
| GQA-8 | 8 | 512 | 4x |
| GQA-4 | 4 | 256 | 8x |
| MQA | 1 | 64 | 32x |
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?
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).
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).
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.
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.
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.
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.
Pattern
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:
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.
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.
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)?
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.
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.
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?
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.
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.
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.
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)?
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.
Section
Project
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.
| # | milestone | key tool |
|---|---|---|
| 1 | Implement cache-append and attend new Q over full cache | torch.cat, softmax |
| 2 | Add causal mask so future tokens are blocked | torch.tril, masked_fill |
| 3 | Implement sliding window mask and compare active entries | manual 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.
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.
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}')| step | K_cache shape | scores shape |
|---|---|---|
| 0 | (129, 64) | (1, 129) |
| 1 | (130, 64) | (1, 130) |
| 2 | (131, 64) | (1, 131) |
| 3 | (132, 64) | (1, 132) |
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?
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 i | non-zero cols (can attend to) | zero cols (masked) |
|---|---|---|
| 0 | 0 | 1, 2, 3, 4, 5 |
| 1 | 0, 1 | 2, 3, 4, 5 |
| 3 | 0, 1, 2, 3 | 4, 5 |
| 5 | 0, 1, 2, 3, 4, 5 | (none) |
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 i | non-zero cols (can attend to) | zero cols (masked) |
|---|---|---|
| 0 | 0 | 1, 2, 3, 4, 5 |
| 1 | 0, 1 | 2, 3, 4, 5 |
| 3 | 0, 1, 2, 3 | 4, 5 |
| 5 | 0, 1, 2, 3, 4, 5 | (none) |
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%}')| n | w | active entries | fraction | vs O(n²) |
|---|---|---|---|---|
| 6 | 2 | 26 | 72.2% | boundary effects |
| 1024 | 256 | 523,264 | 50.0% | 2x reduction |
| 4096 | 512 | 4,190,208 | 25.0% | 4x reduction |
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.
| n | w | active entries | fraction | vs O(n²) |
|---|---|---|---|---|
| 6 | 2 | 26 | 72.2% | boundary effects |
| 1024 | 256 | 523,264 | 50.0% | 2x reduction |
| 4096 | 512 | 4,190,208 | 25.0% | 4x reduction |
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.
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.
Recap
| mechanism | key formula / insight |
|---|---|
| Standard attn | O(n²·d) time, O(n²) memory — bottleneck at long n |
| KV cache | O(t·d) per step; memory = 2·L·n·d·4 bytes |
| Flash Attention | same output; O(n) SRAM; block-tile + online softmax |
| Sliding window | O(n·w); each token attends ±w; ratio n/w speedup at large n |
| GQA | H_KV < H_Q; KV cache = H_Q/H_KV × smaller; MQA extreme (H_KV=1) |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.