PidokuInfra

Project 05 — KV Cache

Intermediate 6h Difficulty 3/5 Topic 05 of 15

Prerequisites Project 04, Section V (03, 05, 06)

★ 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 neutral

2. 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 token

Two 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#

  1. 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.
  2. Prefill returns K, V. Modify attention to return its K and V tensors.
  3. Decode step. Input shape (B, 1). The causal mask disappears — one query attends to everything cached. Remember the position embedding index is pos, not 0.
  4. Equivalence test. For 20 prompts × 100 tokens, cached and uncached greedy outputs must be identical. Logits within 1e-4.
  5. Layout B. Pre-allocate to max_len. Compare per-step time against layout A at long lengths.
  6. Measure memory. cache.nbytes vs your prediction. They must agree exactly.

6. Starter skeleton#

Go
// 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#

MeasurementExpectation to write down first
Per-token latency vs position, cached vs uncachedUncached climbs; cached nearly flat
Total time for 500 tokens, bothSpeedup grows with length
Cache bytes per token73,728 exactly (FP32)
Decode step at position 50 vs 1000Slightly slower — the attention read grows
Layout A vs B at length 1000A pays O(n) copy per step → O(n²) total
Prefill tokens/s vs decode tokens/sPrefill 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 pos hits max_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#

  1. What exactly is stored in the KV cache, and why not Q?
  2. Derive bytes per token of KV cache for a given config.
  3. Why is decode memory-bound even though the GPU is “100% utilized”?
  4. Why does a GQA model have a smaller cache than an MHA model of the same width?
  5. Why does caching make concurrency a memory problem?

12. Next#

Project 06 — LLM inference server

↑↓ navigate↵ openesc close