16. Prefill, decode and the KV cache
In this chapter
- Why the generation loop from Chapter 8 does quadratic work, and the observation that removes it.
- The KV cache: what to store, how to lay it out, and how positions and masks continue across calls.
- Prefill versus decode as two different workloads, and chunked prefill.
- KV-cache memory arithmetic, and why it shapes modern architectures.
You will build
engine/kv_cache.py: a preallocated per-layer cache that both your GPT and (next chapter) your Qwen3 use, with exact full / chunked / token-by-token equivalence.
Time: 4-6 hours. GPU: not needed.
Generation without a cache repeats itself
Chapter 8’s loop with cached=False re-runs the entire sequence for every new token. With a prompt of $P$ tokens, generating $N$ tokens processes $P$, then $P+1$, …, up to $P+N-1$ positions. Each of those passes recomputes the same projections, MLPs and attention for every earlier position, getting exactly the same numbers as last time. The total work grows like $N \cdot P + N^2/2$, so it’s quadratic in the output length.
Measured on the reference engine (a tiny 4-layer Qwen3 on a laptop CPU, python run.py cache):
| new tokens | uncached | cached |
|---|---|---|
| 16 | 0.10 s | 0.05 s |
| 64 | 0.50 s | 0.19 s |
| 128 | 1.90 s | 0.33 s |
The uncached time grows quadratically and the cached time roughly linearly. Real models with thousand-token prompts make the gap enormous.
The observation: the past doesn’t change
In a causal model, position $t$’s hidden states depend only on tokens $\le t$. Appending a token can’t change anything computed for earlier positions. So their keys and values, the only things later tokens read from them, are fixed once computed.
A new token needs, at every layer: its own query, and the keys and values of all positions so far. So store each layer’s K and V as they’re computed, and on the next step process only the new token: compute its Q, K and V, append K and V to the cache, and attend over the whole cache. That’s the KV cache. Two things it deliberately does not store:
- Old queries. A query is used only once, by its own position.
- Old hidden states or MLP results. They’re never read again; only K and V are.
Prefill and decode
Generation now has two phases, as Chapter 1 promised:
prefill: input IDs at positions 0,1,2 write K/V at 0,1,2 last logits -> token 3
decode 1: input token 3 at position 3 write K/V at 3 logits -> token 4
decode 2: input token 4 at position 4 write K/V at 4 logits -> token 5
Two off-by-one facts trip people up:
- The first new token comes from prefill. The prompt’s last position already predicts it. No separate decode step is needed before the first sample.
- The last generated token is never fed back unless you want another prediction. After $N$ new tokens, decode has run $N-1$ times.
For a prompt of 3 tokens and 3 generated tokens, the uncached loop projects $3 + 4 + 5 = 12$ token positions; the cached one projects $3 + 1 + 1 = 5$.
Layout and validity
For each layer, store keys and values as [B, H_kv, capacity, D_h], preallocated to the maximum length you’ll need, plus a count of valid positions:
class KVCache:
"""Contiguous preallocated cache; every row in the batch has the same length. (Your engine: Chapter 16)"""
def __init__(self, layers, batch, kv_heads, capacity, head_dim, device="cpu", dtype=torch.float32):
self.capacity = capacity
shape = (batch, kv_heads, capacity, head_dim)
self.keys = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.values = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.lengths = [0] * layers
@property
def length(self):
"""Tokens already processed into every layer."""
if len(set(self.lengths)) != 1:
raise RuntimeError("A forward pass stopped part-way; reset or truncate the cache")
return self.lengths[0]
def update(self, layer, k, v, positions=None, rows=None):
"""Write k/v after the existing entries and return views of the whole valid history. (Your engine: Chapter 16)"""
if rows is not None:
raise ValueError("KVCache rows always advance together; use StaticKVCache for slots")
start = self.lengths[layer]
end = start + k.shape[2]
if end > self.capacity:
raise ValueError(f"KV capacity {self.capacity} exceeded")
if k.shape[:2] != self.keys[layer].shape[:2] or k.shape[3] != self.keys[layer].shape[3]:
raise ValueError("Cache batch/head/width shapes differ")
self.keys[layer][:, :, start:end].copy_(k)
self.values[layer][:, :, start:end].copy_(v)
self.lengths[layer] = end
return self.keys[layer][:, :, :end], self.values[layer][:, :, :end], None
def truncate(self, length):
"""Forget everything after `length` tokens (speculative-decoding rollback). (Your engine: Chapter 16)"""
if not 0 <= length <= self.length:
raise ValueError("Invalid truncation length")
self.lengths = [length] * len(self.lengths)
def reset(self):
self.truncate(0)
@property
def bytes(self):
return sum(t.numel() * t.element_size() for t in self.keys + self.values)
struct KvCache {
std::vector<std::vector<float>> keys, values; // per layer: [positions, kv_heads*head_dim]
size_t width, len = 0;
KvCache(size_t layers, size_t width_, size_t capacity) : keys(layers), values(layers), width(width_) {
for (size_t l = 0; l < layers; ++l) { keys[l].reserve(capacity * width); values[l].reserve(capacity * width); }
}
void append(size_t layer, const std::vector<float>& k, const std::vector<float>& v) {
keys[layer].insert(keys[layer].end(), k.begin(), k.end());
values[layer].insert(values[layer].end(), v.begin(), v.end());
}
void truncate(size_t n) {
for (size_t l = 0; l < keys.size(); ++l) { keys[l].resize(n * width); values[l].resize(n * width); }
len = n;
}
};
#![allow(unused)]
fn main() {
/// Rows are positions; each row holds kv_heads * head_dim values. Capacity is reserved up
/// front so appending never reallocates (and never copies the history).
pub struct KvCache {
pub keys: Vec<Vec<f32>>, // one buffer per layer
pub values: Vec<Vec<f32>>,
pub width: usize, // kv_heads * head_dim
pub len: usize, // positions already processed
}
impl KvCache {
pub fn new(layers: usize, width: usize, capacity: usize) -> Self {
KvCache {
keys: (0..layers).map(|_| Vec::with_capacity(capacity * width)).collect(),
values: (0..layers).map(|_| Vec::with_capacity(capacity * width)).collect(),
width,
len: 0,
}
}
/// Append one position's K and V for one layer; return the whole history for that layer.
pub fn append(&mut self, layer: usize, k: &[f32], v: &[f32]) -> (&[f32], &[f32]) {
self.keys[layer].extend_from_slice(k);
self.values[layer].extend_from_slice(v);
(&self.keys[layer], &self.values[layer])
}
/// Roll back to `len` positions (speculative decoding).
pub fn truncate(&mut self, len: usize) {
for layer in 0..self.keys.len() {
self.keys[layer].truncate(len * self.width);
self.values[layer].truncate(len * self.width);
}
self.len = len;
}
pub fn bytes(&self) -> usize {
2 * self.keys.len() * self.len * self.width * 4
}
}
}
Design choices worth noticing:
- Preallocate; don’t concatenate.
torch.caton every step reallocates and copies the whole history, which is $O(T)$ work per token and $O(T^2)$ overall, the very cost the cache was meant to remove. Writing into a preallocated slice is $O(1)$ per token. (Chapter 25’s paged cache gets the same benefit without reserving the maximum up front.) - Capacity versus valid length. Slots beyond
lengthhold zeros or stale data.updatereturns a view of just the valid prefix, so attention never sees them. - Per-layer lengths. Every layer appends during a forward pass. If an exception interrupts a pass halfway, some layers are one token ahead of others. The
lengthproperty detects this and refuses to continue, rather than silently misaligning positions. - Truncate for rollback.
truncate(n)forgets everything after $n$ tokens without copying. Speculative decoding (Chapter 26) depends on this.
The cache follows a one-method protocol: update(layer, k, v, positions, rows) writes new entries and returns everything attention may read. Your models never touch cache internals. Later caches (static for CUDA graphs in Chapter 19, slot-based for batching in Chapter 24, paged in Chapter 25) implement the same method and plug into the same model code unchanged.
Positions continue, and masks become rectangles
When a decode step feeds token 57, it must be treated as position 57, not position 0. GPT adds position_embedding[57], and Qwen rotates its query and key by angle 57 (Chapter 17). Your models compute positions as cache.length + arange(T).
The mask changes too. Without a cache, queries and keys are the same $T$ positions, and the mask is a square lower triangle. With a cache, $T$ new queries attend to $S$ keys ($S > T$): a rectangle. For a cached prefix of 3 tokens and a chunk of 2 new tokens:
key 0 1 2 3 4
query pos 3 1 1 1 1 0
query pos 4 1 1 1 1 1
A square-triangle shortcut applied to this 2-row chunk would let query 3 see only key 0 and query 4 only keys 0-1, hiding the most recent history. This is exactly why your causal_attention takes positions (Chapter 5): the rule key_pos <= query_pos produces the right rectangle automatically. (Watch out for PyTorch’s scaled_dot_product_attention(is_causal=True): with $T \ne S$ it uses a top-left-aligned triangle, which is the wrong rectangle for cached decoding.)
Prove equivalence
The cache is an optimization, so it must change nothing except speed. The strongest test feeds one fixed sequence three ways:
- the whole sequence at once, with no cache;
- a prefix, then a multi-token chunk, then single tokens, all through one cache;
- compare all logits position by position.
They must agree to floating-point tolerance (about 1e-5 in FP32). Use fixed token IDs, not generated ones: once two runs pick different tokens, they’re processing different inputs and the comparison is meaningless.
How big is the cache?
Each layer stores K and V for every position, so:
$$ \text{KV bytes} = 2 \times \text{layers} \times \text{batch} \times \text{tokens} \times H_{kv} \times D_h \times \text{bytes per value}. $$
For Qwen3-0.6B (28 layers, 8 KV heads of dimension 128, BF16):
- 112 KiB per token;
- 224 MiB for a 2,048-token conversation;
- 3.5 GiB for a 32k-token context, about three times the model’s own weights.
def kv_cache_bytes(layers, batch, tokens, kv_heads, head_dim, bytes_per_value=2):
"""2 (K and V) x layers x batch x tokens x kv_heads x head_dim x bytes."""
return 2 * layers * batch * tokens * kv_heads * head_dim * bytes_per_value
Why the cache shapes model design
Two consequences follow from this formula, and they explain several architecture choices you’ll meet:
- Capacity. In serving, the cache limits how many conversations fit on a GPU at once (Chapters 24-25). Note that the formula uses KV heads, not query heads. Grouped-query attention with 8 KV heads instead of 32 shrinks the cache 4x: for an 8B-class model (32 layers, head dimension 128) at 8k tokens, that’s 1 GiB per sequence instead of 4 GiB. Multi-query attention (one KV head) and DeepSeek’s multi-head latent attention push further.
- Bandwidth. Each decode step reads the whole cache of each sequence once, in addition to the weights. Attention during decode has an arithmetic intensity of about 1 FLOP per byte per query head (PMPP §20.6), as memory-bound as the weight reads. At long contexts, the KV reads dominate the step time. That’s what motivates KV-cache quantization (Chapter 20), sparse attention that reads only part of the cache (Chapter 29), and linear attention with a fixed-size state instead of a cache (Chapter 28).
Build it
Engine milestone 16: a KV cache. Implement KVCache.update and KVCache.truncate in engine/kv_cache.py (the constructor, length, reset and bytes are provided). Your GPT from Chapter 6 already calls cache.update(...).
pytest tests/test_ch16_kv_cache.py
python run.py cache --impl engine
The tests run full, chunked and token-by-token forwards through GPT and Qwen3 models and compare all logits, check rollback with truncate, check that capacity is enforced, and check the memory formula against cache.bytes.
Stretch exercises
- ★ Replace the preallocated write with
torch.catand measure generation time for 64, 256 and 1,024 new tokens. Where does the quadratic copy cost become visible? Where: a separate concatenating-cache variant ofKVCache.updateinengine/kv_cache.py; compare it inexperiments/ch16.py(create it). - ★★ Add a
layer_bytes()report and print how much of a 32k-token Qwen3-0.6B decode step’s traffic is KV cache versus weights. Where: addKVCache.layer_bytesinengine/kv_cache.py. - ★★ Implement a sliding-window cache that keeps only the last $w$ positions in a ring buffer. What must change in the positions passed to attention? Where: add a ring-buffer cache class in
engine/kv_cache.py; wire its absolute positions intoQwen3Attention.forwardinengine/qwen3.py. - ★★★ Store the cache in FP8 (E4M3, with one scale per head) and dequantize on read. Measure the logit error against a BF16 cache for 1,000 decode steps. Where: add an FP8 cache variant in
engine/kv_cache.py; integrate its dequantized reads inengine/qwen3.py.
Check your understanding
- Why cache keys and values, but not queries?
- Which forward pass produces the logits for the first generated token?
- Why can a lower-triangular mask be wrong for a 2-token chunk after a 3-token prefix?
- Why does grouped-query attention reduce cache size, but not the number of attention score computations?
- Why does decode at long context become dominated by cache reads?
Going deeper
- PMPP §20.4 (pp. 488-492): KV caching; §20.6 (pp. 504-508): the arithmetic intensity and memory requirement of the KV cache; §20.7: MQA and GQA.
- Pope et al., Efficiently Scaling Transformer Inference (2022): the analysis of cache memory and bandwidth that much of modern serving builds on.
- Shazeer, Fast Transformer Decoding: One Write-Head is All You Need (MQA, 2019); DeepSeek-V2 (2024) for multi-head latent attention.