Transformer Architecture Q&A

Course Content

Transformer Architecture Q&A

6 sections · 60 lessons

What is the KV cache, and how does it speed up autoregressive decoding?


One 80 GB H100 serving Llama-3-8BWeights: 16 GB,shared by allKV cache:about 60 GB free128 KiB pertoken per chat2k-tokenchat: 0.25 GiBAbout 240chats fit at oncetopbottomReserving the full 8k per chat instead would fit only 60.
The weights are paid once and the cache is paid per user, which is why the cache sets how many people one GPU can serve.

What you need to know

Why the cached values never change

In a decoder, token j can only attend to tokens at positions up to j. So its hidden state in every layer — and therefore its key and value — depends only on the prefix up to j. When token 501 is generated, tokens 1 to 500 do not change at all. Recomputing their K and V would give exactly the same numbers.

Only the newest query is needed. Past queries have already produced their outputs, and nothing later looks at them. That is why the cache holds K and V, never Q.

What each step costs, with and without the cache

Let N be the number of parameters and t the current length.

Without cacheWith cache
Tokens pushed through the weights per stepall t1
Weight FLOPs per stepabout 2·N·tabout 2·N
Attention work per stept × t scores1 × t scores
Total for T generated tokensgrows like T² (and T³ for attention)grows like T (and T² for attention)

The cache does not make a step constant-time: the new query still reads all t cached keys and values. It removes the repeated work, not the attention itself.

The memory formula

Text
bytes = 2 (K and V) × n_layers × n_kv_heads × head_dim × seq_len × batch × bytes_per_value
Python
def kv_cache_gib(layers, kv_heads, head_dim, seq_len, batch=1, bytes_per=2):    b = 2 * layers * kv_heads * head_dim * seq_len * batch * bytes_per    return b / 2**30# Llama-3-8B: 32 layers, 8 KV heads, head_dim 128, fp16print(kv_cache_gib(32, 8, 128, 8192))              # 1.0  GiB — one 8k chatprint(kv_cache_gib(32, 8, 128, 8192, batch=40))    # 40.0 GiB — 40 usersprint(kv_cache_gib(32, 8, 128, 8192, bytes_per=1)) # 0.5  GiB — FP8 cache

Per token that is 2 × 32 × 8 × 128 × 2 = 131,072 bytes, or 128 KiB. The 16 GB of weights are shared by every user; the cache is paid again for each one. At 40 users with 8k context, the cache is already bigger than the model.

How serving systems manage it

  • PagedAttention (vLLM, Kwon et al., 2023) stores the cache in small fixed-size blocks, like virtual-memory pages, so a request only holds memory for tokens it has actually produced, not for its maximum length.
  • Prefix caching reuses the K/V of a shared system prompt across requests.
  • Fewer KV heads (GQA, MQA) or a compressed latent (DeepSeek's MLA) shrink the formula's n_kv_heads × head_dim term.
  • Quantised caches (FP8, sometimes 4-bit) halve or quarter bytes_per_value.

The cache only exists for decoder generation. An encoder such as BERT reads the whole input in one pass and produces no tokens, so it has nothing to cache.

A real-life example

A company serves a support chatbot built on Llama-3-8B to about a million users a day, with a peak of 20,000 open conversations. One 80 GB H100 holds the 16 GB of weights and leaves roughly 60 GB for cache.

If the server reserved the full 8k context for every conversation, it could hold 60 conversations per GPU — about 330 GPUs at peak. With paged allocation, each conversation only uses what it has written; the average is 2,000 tokens, or 0.25 GiB, so one GPU holds about 240 conversations and the peak needs about 84 GPUs. Switching the cache to FP8 roughly doubles that again. None of these changes touch the model's weights — they are all about the cache.

Follow-up questions to expect

  • "Why not cache the queries too?" — A past token's query is only used to compute that token's own output, which is already done. Only the newest token's query is ever needed again.
  • "Does the KV cache change the model's output?" — No. It is an exact reuse of values that would be recomputed identically, apart from tiny floating-point differences from a different kernel order.
  • "What is prefix caching?" — Keeping the K/V of a common prefix, such as a 1,500-token system prompt, and reusing it for every request that starts with it, which saves both prefill compute and memory.
  • "How does GQA reduce the cache?" — It lowers n_kv_heads; Llama-3-8B has 32 query heads but only 8 KV heads, a 4× smaller cache than full multi-head attention.