Serving · Memory

KV Cache

prefill once, then one token per step

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

  1. 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 → n

  2. Prefill

    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 prompt

  3. A 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 → · V

  4. How 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 bytes

  5. Compute-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

Go deeper