Serving · Memory
- model
- GPT-2 small
- per token
- 36 KiB 2 × 12 × 768 × 2 B
- at 1,024 tokens
- 36 MiB
- head shown
- layer 2 · head 11 64 dims, 8 drawn
All steps
One token at a time
GPT-2 writes one token at a time, and each new token attends to every earlier one. Those tokens’ keys and values never change once computed, so a serving system keeps them in a KV cache instead of recomputing the whole prefix at every step.
K, V rows computed: n(n + 1) / 2 → nPrefill
The prompt is processed in one pass: all its tokens go through each layer together, as matrix products, and every layer writes their keys and values to the cache. These are real GPT-2 values for one head (the first 8 of its 64 numbers).
K, V = X · WK, X · WV for the whole promptA decode step
Each decode step runs one token through the model. It computes its own q, k and v, appends k and v to the cache, and attends over all the cached keys: one row of scores instead of a matrix. Real GPT-2 numbers from layer 2, head 11; hover the cells.
q · Kᵀ → softmax → · VHow big the cache gets
The cache holds two vectors (K and V) per token, per layer, per key/value head. It grows with every token of every sequence being served, and at long contexts it outgrows the model’s weights; grouped-query attention, latent attention and paging all attack it.
2 × layers × kv heads × dhead × 2 bytesCompute-bound and memory-bound
Prefill does many operations for each byte it reads, so it is limited by compute. A decode step reads every weight, and the cache, to produce a single token, so it is limited by memory bandwidth; serving many sequences in one batch is how a GPU’s compute gets used.
FLOPs per byte ≈ tokens per weight read
Code
# prefill: the whole prompt in one pass; every layer keeps its K and V
k, v = ln_1(x) @ W_k + b_k, ln_1(x) @ W_v + b_v # (prompt_len, 64) per head
cache[layer] = (k, v)
# decode: one token per step
q, k_new, v_new = (ln_1(x_t) @ W) .split(768) # just the new token
K = torch.cat([cache_k, k_new]); V = torch.cat([cache_v, v_new])
w = F.softmax(q @ K.T / 8, dim=-1) # one row of scores
out = w @ V