★ Add the single most important optimization in LLM inference to your own engine, and prove the outputs did not change.
1. What you build#
A KV cache for your Project 04 transformer: prefill computes K and V for the whole prompt once; each decode step computes Q, K, V for one new token, appends K and V to the cache, and attends over the cached history. Output tokens must be bit-for-bit the same as the uncached version.
Diagram — Before and after#
flowchart TB
subgraph NO["Project 04 - no cache"]
direction LR
N1["Step t"] --> N2["Re-run all t tokens<br/>through every layer"] --> N3["Keep only the last logit"]
end
subgraph YES["Project 05 - with cache"]
direction LR
Y1["Step t"] --> Y2["Run ONE token<br/>through every layer"] --> Y3["Append its K, V"]
Y3 --> KV[("KV cache")]
KV --> Y2
end
class N2 warn
class Y2 compute
class Y3,KV memory
class N1,N3,Y1 neutral2. Why it matters#
Decode without a cache recomputes work that cannot have changed. Caching trades compute for memory — and that trade defines the rest of the field: once you cache, memory, not compute, limits how many users fit on a GPU. Projects 08 and 09, and most of Sections X and XIII, are about managing the thing you build here.
3. Read first#
4. Spec#
cache = KVCache(n_layer, n_head, head_dim, max_len, batch)
logits, cache = prefill(prompt_ids, cache) # T tokens in one pass
logits, cache = decode_step(last_token, cache) # exactly 1 tokenTwo storage layouts, both implemented:
A. Concatenate K = np.concatenate([K, k_new], axis=2) each step (simple, reallocates)
B. Pre-allocated K[:, :, pos] = k_new ; pos += 1 (what engines do)5. Milestones#
- Predict the size. GPT-2 small, FP32:
2 (K,V) × 12 layers × 768 (heads × head_dim) × 4 B = 73,728 B ≈ 72 KB per token. At 1024 tokens ≈ 75 MB per sequence. Write this down before coding. - Prefill returns K, V. Modify attention to return its K and V tensors.
- Decode step. Input shape
(B, 1). The causal mask disappears — one query attends to everything cached. Remember the position embedding index ispos, not0. - Equivalence test. For 20 prompts × 100 tokens, cached and uncached greedy outputs must
be identical. Logits within
1e-4. - Layout B. Pre-allocate to
max_len. Compare per-step time against layout A at long lengths. - Measure memory.
cache.nbytesvs your prediction. They must agree exactly.
6. Starter skeleton#
// KVCache holds keys and values for one sequence: [layer][head] -> a growing (t, D) buffer.
type KVCache struct {
K, V [][][]float32 // K[layer][head] is a flat slice of pos*D values
Pos int // tokens cached so far
H, D, Cap int
}
func NewKVCache(layers, heads, maxLen, d int) *KVCache {
c := &KVCache{H: heads, D: d, Cap: maxLen}
for l := 0; l < layers; l++ {
k, v := make([][]float32, heads), make([][]float32, heads)
for h := range k {
k[h], v[h] = make([]float32, 0, maxLen*d), make([]float32, 0, maxLen*d) // allocate ONCE
}
c.K, c.V = append(c.K, k), append(c.V, v)
}
return c
}
// Append stores the keys and values of t new tokens for one layer and head.
func (c *KVCache) Append(layer, head int, k, v []float32) {
c.K[layer][head] = append(c.K[layer][head], k...)
c.V[layer][head] = append(c.V[layer][head], v...)
}
// attentionCached handles both phases: x holds t new tokens — t = T (prefill) or 1 (decode).
func attentionCached(x [][]float32, p *Layer, cache *KVCache, layer int) [][]float32 {
t := len(x)
q, k, v := splitQKV(x, p) // each [head][t*D]
out := make([][]float32, t)
for h := 0; h < cache.H; h++ {
cache.Append(layer, h, k[h], v[h])
K, V := cache.K[layer][h], cache.V[layer][h] // history + new
for i := 0; i < t; i++ {
visible := cache.Pos + i + 1 // causal: token i sees the history and tokens 0..i
out[i] = append(out[i], attendOne(q[h][i*cache.D:(i+1)*cache.D], K, V, visible, cache.D)...)
}
}
return project(out, p) // c_proj
}
// Advance cache.Pos += t ONCE per forward pass, after all layers.7. What to measure#
| Measurement | Expectation to write down first |
|---|---|
| Per-token latency vs position, cached vs uncached | Uncached climbs; cached nearly flat |
| Total time for 500 tokens, both | Speedup grows with length |
| Cache bytes per token | 73,728 exactly (FP32) |
| Decode step at position 50 vs 1000 | Slightly slower — the attention read grows |
| Layout A vs B at length 1000 | A pays O(n) copy per step → O(n²) total |
| Prefill tokens/s vs decode tokens/s | Prefill far higher: parallel over positions |
| Same math for Llama-3-8B (32 L, 8 KV heads × 128, FP16) | 131,072 B = 128 KB/token (Checkpoint C) |
8. Done when#
- Cached and uncached generation produce identical tokens on all test prompts.
- Measured cache bytes equal the formula.
- You have the cached-vs-uncached plot.
- You can explain why cached decode is memory-bandwidth-bound: per token it reads all the weights and the whole cache, for a tiny amount of arithmetic.
- You can do the KV size calculation for an arbitrary model config on a whiteboard.
9. Common pitfalls#
Advancing pos inside the layer loop. Every layer writes at the same position; advance
once per forward.
Wrong position embedding on decode. The new token is at pos, not 0.
Applying a causal mask during decode with the wrong shape — a 1×T attention needs no mask.
Declaring victory on speed before the equivalence test passes. A cache that is fast and subtly wrong is the worst outcome.
Hidden O(n²) from concatenate. It looks constant per step; it isn’t.
10. Stretch goals#
- Store the cache in FP16 and measure the logit drift (preview of VII.12).
- Batch multiple sequences of different lengths: you now need per-sequence lengths and a mask. This pain is exactly what motivates Project 09.
- Implement a sliding window when
poshitsmax_len— and discover why absolute position embeddings make that awkward. - Prefix reuse: snapshot the cache after a shared system prompt, restore it for a second request, and measure the TTFT saved (V.11).
11. Interview questions this project answers#
- What exactly is stored in the KV cache, and why not Q?
- Derive bytes per token of KV cache for a given config.
- Why is decode memory-bound even though the GPU is “100% utilized”?
- Why does a GQA model have a smaller cache than an MHA model of the same width?
- Why does caching make concurrency a memory problem?