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.

each earlier token offers a key (what I am) and a value (what I carry)TheKV0.05catKV0.62satKV0.12onKV0.06theKV0.15new token asks: Qattention weights(sum to 1)answer = 0.05·V₁ + 0.62·V₂ + ...mostly the value of "cat"
The newest token compares its query with the key of every earlier token, turns the scores into weights, and takes a weighted mix of their values. Here it mostly pulls from "cat".

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.

no cache: recompute K and V for all tokensKV cache: compute only the new tokensteps 1 to 6 (rows), tokens (columns)steps 1 to 6 (rows), tokens (columns)orange = computed this stepaqua = read from the cache
Without a cache, every step recomputes keys and values for every token so far. With a cache, every step computes one new row and reads the rest.

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:

0100200300400500128 tokens, no cache: 13.6 ms13.6128 tokens, KV cache: 9.9 ms9.87128 tokens so far512 tokens, no cache: 27.3 ms27.3512 tokens, KV cache: 10.2 ms10.2512 tokens so far2048 tokens, no cache: 94.9 ms94.92048 tokens, KV cache: 10.8 ms10.82,048 tokens so far4096 tokens, no cache: 200.7 ms2014096 tokens, KV cache: 11.0 ms11.04,096 tokens so far8192 tokens, no cache: 449.1 ms4498192 tokens, KV cache: 11.5 ms11.58,192 tokens so farno cacheKV cachetime to produce the next token, milliseconds
Time to produce the next token. Without a cache the cost climbs with the length of the text; with a cache it stays almost flat.
Tokens so farNo cacheKV cacheSlower without
12813.6 ms9.9 ms1.4x
51227.3 ms10.3 ms2.7x
2,04894.9 ms10.9 ms8.7x
4,096200.7 ms11.0 ms18x
8,192449.1 ms11.5 ms39x

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:

Llama-3.1-8B config.json (excerpt){ "hidden_size": 4096, "num_attention_heads": 32, "num_hidden_layers": 32, "num_key_value_heads": 8, "head_dim": 128, "torch_dtype": "bfloat16", "max_position_embeddings": 131072}2K and V× 32layers× 8KV heads× 128head_dim× 2 bytesbfloat16= 131,072 bytes= 128 KiB for every token
The four fields in a model config that decide the size of its KV cache: the number of layers, the number of key/value heads, the head dimension and the data type.
KV bytes per token=2×layers×KV heads×head dim×bytes per number\text{KV bytes per token} = 2 \times \text{layers} \times \text{KV heads} \times \text{head dim} \times \text{bytes per number}

For Llama 3.1 8B that is 2×32×8×128×2=131,0722 \times 32 \times 8 \times 128 \times 2 = 131{,}072 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:

Qwen2.5-0.5B: 12.0 KiBQwen2.5-0.5B12.0 KiB1.5 GiB for a 128K-token contextLlama-2-7B (MHA): 512 KiBLlama-2-7B (MHA)512 KiB64.0 GiB for a 128K-token contextLlama-3.1-8B (GQA): 128 KiBLlama-3.1-8B (GQA)128 KiB16.0 GiB for a 128K-token contextLlama-3.1-70B (GQA): 320 KiBLlama-3.1-70B (GQA)320 KiB40.0 GiB for a 128K-token context
KV cache per token, and for one 128,000-token conversation, from each model's config.
ModelLayersKV headsHead dimPer tokenOne 128K-token conversation
Qwen2.5-0.5B2426412 KiB1.5 GiB
Llama 3.1 8B328128128 KiB16 GiB
Llama 3.1 70B808128320 KiB40 GiB
Llama 2 7B3232128512 KiB64 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 (MHA)8 KV heads to cachegrouped-query (GQA)KVKV2 KV heads to cachemulti-query (MQA)KV1 KV head to cacheorange = query heads (not cached). aqua = key/value heads (cached for every token).
In multi-head attention every query head has its own key and value head. Grouped-query attention lets several query heads share one; multi-query attention shares a single one across all.
  • 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:

01020304002,0484,0966,1448,192batch 1, 128 tokens of context: 10.2 ms per stepbatch 1, 1,024 tokens of context: 11.4 ms per stepbatch 1, 4,096 tokens of context: 10.9 ms per stepbatch 1, 8,192 tokens of context: 11.5 ms per stepbatch 1: 11.5 msbatch 16, 128 tokens of context: 11.3 ms per stepbatch 16, 1,024 tokens of context: 15.4 ms per stepbatch 16, 4,096 tokens of context: 23.4 ms per stepbatch 16, 8,192 tokens of context: 35.9 ms per stepbatch 16: 35.9 mstokens already in the context (each sequence)ms per decode step
Time per decode step as the context grows. One sequence barely notices; sixteen long sequences triple the step time.
Context per sequence1 sequence16 sequences
128 tokens10.2 ms11.3 ms
1,02411.5 ms15.4 ms
4,09610.9 ms23.4 ms
8,19211.5 ms35.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.

contiguous allocation: each request reserves room for its maximum length up frontrequest Arequest Bgap too small for Cprompt tokensgenerated so farreserved but empty (reservation, internal and external fragmentation)
Contiguous allocation. Each request reserves space for its maximum length up front, so most of the reservation sits empty, and the leftover gaps between requests are often too small to use.

That wastes memory in three ways, which the vLLM paper names:

  1. Reserved slots: space held for tokens that will be generated later, empty for now.
  2. Internal fragmentation: space reserved for a maximum length the answer never reaches, empty forever.
  3. 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 =2×layers×KV heads×head dim×bytes= 2 \times \text{layers} \times \text{KV heads} \times \text{head dim} \times \text{bytes}. 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
python
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

  1. A. Vaswani et al. Attention Is All You Need. NeurIPS 2017.
  2. 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.)
  3. J. Ainslie et al. GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints. EMNLP 2023.
  4. N. Shazeer. Fast Transformer Decoding: One Write-Head is All You Need. 2019.
  5. DeepSeek-AI. DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model. 2024.
  6. Model configs: Llama 3.1 8B, Llama 3.1 70B, Qwen2.5-0.5B.