The KV Cache: An LLM's Short-Term Memory, and Its Bill
Why every LLM keeps notes on each token it has read, how big those notes get (16 GiB for one 128K-token Llama 3.1 8B conversation), how models shrink them, and why they, not compute, usually decide how many people one GPU can serve.
In Part 1 we saw that an LLM writes one token per step, and that each step is slow because the GPU has to read all of the model's weights to produce a single token.
But there is a second thing every step needs, and we skipped over it. To pick the next word, the model has to know everything that came before: your whole prompt and every word it has written so far. On step 500, that is 500 tokens of context.
Does the model reread all 500 tokens, from scratch, on every step? It could. It would be painfully slow. Instead it keeps notes. Those notes are the KV cache, and by the end of this piece you will see why they, far more than raw compute, decide how many people a GPU can serve.
Attention, in one picture
Each layer of a transformer has a step called attention. It is how a token gathers information from the tokens before it. It works a lot like a library.
- Every token publishes a key: a short description of what it is. ("I am a noun, an animal, the subject of this sentence.")
- Every token also carries a value: the actual information it can hand over.
- The token being processed asks a query: a description of what it is looking for.
The query is compared with every key. Good matches get high scores, the scores are turned into weights that add up to 1, and the token takes a weighted mix of the values.
The keys and values are just vectors, lists of numbers computed from each token. Queries, keys and values are why this is called the KV cache: we are about to cache the K's and V's.
The waste hiding in the loop
Here is the key observation. In a model that writes left to right, each token only looks backwards. So once a token's key and value have been computed, they never change. The key for "cat" on step 3 is exactly the key for "cat" on step 300.
Without a cache, the model would recompute the keys and values of the whole sequence on every step. With a cache, it computes them only for the one new token, stores them, and reads the old ones back.
The difference is not subtle. Here is the time to produce the next token, with and without a cache, measured on Qwen2.5-0.5B:
| Tokens so far | No cache | KV cache | Slower without |
|---|---|---|---|
| 128 | 13.6 ms | 9.9 ms | 1.4x |
| 512 | 27.3 ms | 10.3 ms | 2.7x |
| 2,048 | 94.9 ms | 10.9 ms | 8.7x |
| 4,096 | 200.7 ms | 11.0 ms | 18x |
| 8,192 | 449.1 ms | 11.5 ms | 39x |
Without a cache, producing the next token means running a full prefill over the whole text again (compare the no-cache column with the prefill table in Part 1: they match). So writing a long answer gets slower and slower as it goes, and the gap keeps widening: 39 times at 8,192 tokens. With a cache, each step costs about the same no matter how long the text is. Every serious inference engine uses one.
The catch is the word stores. The cache trades compute for memory, and that memory adds up fast.
What exactly gets stored
For every token, the cache holds one key and one value in every layer, and in every layer there can be several heads (parallel attention units, each with its own keys and values). Each key or value is a vector of head_dim numbers.
You can read everything you need from a model's config.json. Here is the one for Llama 3.1 8B:
For Llama 3.1 8B that is bytes: 128 KiB for every single token in every conversation the GPU is serving.
How big it gets
Run the same formula for a few real models:
| Model | Layers | KV heads | Head dim | Per token | One 128K-token conversation |
|---|---|---|---|---|---|
| Qwen2.5-0.5B | 24 | 2 | 64 | 12 KiB | 1.5 GiB |
| Llama 3.1 8B | 32 | 8 | 128 | 128 KiB | 16 GiB |
| Llama 3.1 70B | 80 | 8 | 128 | 320 KiB | 40 GiB |
| Llama 2 7B | 32 | 32 | 128 | 512 KiB | 64 GiB |
Look at the Llama 3.1 8B row again. A single conversation at its full 128K context needs 16 GiB of cache. The model's weights are also about 16 GB. One long conversation needs as much memory as the entire model.
This is not a new problem. The original vLLM paper (Kwon et al., 2023) worked it through for a 13-billion-parameter model: 800 KB per token, so up to 1.6 GB for a single 2,048-token request. On a 40 GB A100, the weights took about 65% of memory and the KV cache close to 30%. As they put it, even if all memory went to the cache, "only a few tens of requests could be accommodated."
Shrinking the notes
Look at the table once more. Llama 2 7B and Llama 3.1 8B are almost the same size, yet the older model needs four times more cache per token. The difference is one number in the config: 8 KV heads instead of 32.
- Multi-head attention (MHA) gives every query head its own keys and values. Llama 2 7B: 32 query heads, 32 KV heads.
- Grouped-query attention (GQA) lets a group of query heads share one KV head. Llama 3.1 8B: 32 query heads, 8 KV heads, so a quarter of the cache. Ainslie et al. (2023) showed this keeps quality close to full multi-head attention.
- Multi-query attention (MQA) shares a single KV head across all query heads (Shazeer, 2019). Smallest cache, but it can cost quality.
- Multi-head latent attention (MLA), used by DeepSeek-V2, stores a compressed version of the keys and values instead. The DeepSeek-V2 paper reports that it "reduces the KV cache by 93.3%" compared with their earlier dense model.
The other lever is precision. The cache is normally stored at the model's own precision, 2 bytes per number. Storing it in 8-bit FP8 halves it; vLLM exposes this as --kv-cache-dtype fp8. It is a small accuracy risk worth measuring on your own workload.
The cache is read on every step
There is a second cost, and it connects straight back to Part 1. On every decode step, attention has to read the entire cache of every sequence in the batch, not only the weights. So as conversations get longer, each step has more bytes to move.
Here is decode step time as the context grows, for 1 sequence and for 16 sequences at once:
| Context per sequence | 1 sequence | 16 sequences |
|---|---|---|
| 128 tokens | 10.2 ms | 11.3 ms |
| 1,024 | 11.5 ms | 15.4 ms |
| 4,096 | 10.9 ms | 23.4 ms |
| 8,192 | 11.5 ms | 35.9 ms |
For one sequence the line is flat. This model's cache is small (12 KiB per token), so even 8,192 tokens is only 96 MiB, tiny next to the 0.99 GB of weights. For 16 sequences at 8,192 tokens each, the cache is 1.5 GiB, now bigger than the weights, and every step has to read all of it. The step time more than triples.
So decode is memory-bound twice over: it reads the weights once per step, and the KV cache of every sequence once per step. Long contexts and big batches make the second read dominate.
The real problem: nobody knows how long the answer will be
So far this sounds like a sizing problem: buy enough memory. The hard part is that the cache grows while the request runs, and you do not know in advance how far.
When a request arrives, you know its prompt length. You do not know whether the answer will be 20 tokens or 2,000. The simple approach, used by serving systems before vLLM, was to reserve one contiguous chunk of memory per request, big enough for the longest possible answer.
That wastes memory in three ways, which the vLLM paper names:
- Reserved slots: space held for tokens that will be generated later, empty for now.
- Internal fragmentation: space reserved for a maximum length the answer never reaches, empty forever.
- External fragmentation: gaps between reservations that are too small to fit the next request.
How bad is it in practice? The vLLM authors profiled existing systems and found that only 20.4% to 38.2% of the KV cache memory actually stored token data. The rest was reserved, fragmented or unusable. And since KV memory is what caps the batch, and the batch is what drives throughput (Part 1), this waste translated directly into fewer users per GPU.
That is the problem vLLM was built to solve. Its answer came from an idea older than most of the people using it: the way operating systems have managed memory since the 1960s. That is Part 3.
Summary
- Attention lets each token look back at every earlier token, through their keys and values.
- In a left-to-right model those keys and values never change, so they are computed once and cached. Without the cache, the next token at 8,192 tokens took 449 ms instead of 11.5 ms.
- Cache size per token . Llama 3.1 8B: 128 KiB per token, 16 GiB for one 128K-token conversation, as much as its weights.
- GQA, MQA, MLA and FP8 caches shrink it; GQA alone is why Llama 3.1 8B needs a quarter of Llama 2 7B's cache.
- The cache is read on every decode step, so long contexts and large batches slow decode down.
- Answers grow unpredictably, so reserving memory up front wastes most of it: older systems used only 20.4 to 38.2% of their KV memory for real tokens.
The cache vs no-cache measurement
import statistics, time
import torch
from transformers import AutoModelForCausalLM, DynamicCache
DEV, sync = "mps", torch.mps.synchronize # "cuda", torch.cuda.synchronize on NVIDIA
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2.5-0.5B", dtype=torch.bfloat16).to(DEV).eval()
def median_ms(fn, reps=5):
fn(); sync()
ts = []
for _ in range(reps):
sync(); t = time.perf_counter(); fn(); sync(); ts.append(time.perf_counter() - t)
return statistics.median(ts) * 1000
with torch.inference_mode():
for n in (128, 512, 2048, 4096, 8192):
x = torch.randint(0, model.config.vocab_size, (1, n), device=DEV)
# no cache: producing the next token means running the whole sequence again
no_cache = median_ms(lambda: model(input_ids=x, use_cache=False, logits_to_keep=1))
# cache: run the prompt once, then each new token is a single-token step
cache = DynamicCache()
out = model(input_ids=x, past_key_values=cache, logits_to_keep=1)
tok = out.logits[:, -1:].argmax(-1)
def step():
global tok
o = model(input_ids=tok, past_key_values=cache, logits_to_keep=1)
tok = o.logits[:, -1:].argmax(-1)
print(f"{n:5d} tokens: no cache {no_cache:6.1f} ms, cache {median_ms(step):5.2f} ms")
# KV bytes per token from a config
cfg = model.config
head_dim = getattr(cfg, "head_dim", None) or cfg.hidden_size // cfg.num_attention_heads
print(2 * cfg.num_hidden_layers * cfg.num_key_value_heads * head_dim * 2, "bytes per token")References
- A. Vaswani et al. Attention Is All You Need. NeurIPS 2017.
- W. Kwon et al. Efficient Memory Management for Large Language Model Serving with PagedAttention. SOSP 2023. (800 KB per token for OPT-13B; 20.4% to 38.2% of KV memory used in existing systems.)
- J. Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023.
- N. Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019.
- DeepSeek-AI. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model. 2024.
- Model configs: Llama 3.1 8B, Llama 3.1 70B, Qwen2.5-0.5B.