Chapter 17 β KV-Cache: Mechanics & Memory
βThe KV-cache is the single most important optimization in LLM inference β without it, generating each token would require recomputing attention over the entire sequence from scratch.β
Why KV-Cache Exists
During autoregressive generation, each new token needs to attend to all previous tokens. Without caching, weβd recompute Q, K, V for every previous token at every step β O(TΒ²) total computation to generate T tokens.
The KV-cache stores the Key and Value projections from all previous tokens, so each decode step only computes Q, K, V for the new token and reuses everything else.
- Step 1: compute K, V for tokens [1]
- Step 2: recompute K, V for tokens [1, 2]
- Step 3: recompute K, V for tokens [1, 2, 3]
- Step T: recompute all T tokens
- Total compute: O(TΒ²) per layer
- Step 1: compute & cache Kβ, Vβ
- Step 2: compute Kβ, Vβ, append to cache
- Step 3: compute Kβ, Vβ, append to cache
- Step T: compute Kβ, Vβ, attend to cached [1..T]
- Total compute: O(T) per layer β
Step-by-Step Visualization
At each decode step, the computation is:
- Query: only the new token β
[B, H, 1, d_k] - Key/Value: read from cache β
[B, H, T_total, d_k] - Attention:
Q Γ K^Tβ[B, H, 1, T_total]β softmax βΓ Vβ[B, H, 1, d_k]
This is a matrix-vector product (not matrix-matrix), making decode steps memory-bandwidth bound rather than compute-bound.
Tensor Shapes in Detail
Memory Calculation
The KV-cache memory for the entire model:
\[\text{KV memory} = 2 \times L \times G \times d_k \times T \times B \times \text{bytes}\]where: L = layers, G = KV heads, d_k = head dimension, T = sequence length, B = batch size, bytes = 2 (FP16/BF16).
Worked Examples
| Model | L | G (KV heads) | d_k | Formula | Per Token | At 8K | At 128K |
|---|---|---|---|---|---|---|---|
| LLaMA-2 7B (MHA) | 32 | 32 | 128 | 2Γ32Γ32Γ128Γ2 | 512 KB | 4 GB | 64 GB |
| LLaMA-3 8B (GQA) | 32 | 8 | 128 | 2Γ32Γ8Γ128Γ2 | 128 KB | 1 GB | 16 GB |
| LLaMA-3 70B (GQA) | 80 | 8 | 128 | 2Γ80Γ8Γ128Γ2 | 320 KB | 2.5 GB | 40 GB |
| Mistral 7B (GQA) | 32 | 8 | 128 | 2Γ32Γ8Γ128Γ2 | 128 KB | 1 GB | 16 GB |
| LLaMA-3.1 405B | 126 | 8 | 128 | 2Γ126Γ8Γ128Γ2 | 504 KB | 3.9 GB | 63 GB |
Key insight: The KV-cache for LLaMA-3.1 405B at 128K context exceeds 63 GB per request in FP16. This is why KV-cache optimization (Chapter 18) is critical.
GQA/MQA Impact on KV-Cache
The attention variant directly determines KV-cache size (see Chapter 5):
512 KB/token
4 GB at 8K
128 KB/token
1 GB at 8K
4Γ savings
64 KB/token
512 MB at 8K
8Γ savings
16 KB/token
128 MB at 8K
32Γ savings
Sliding Window KV-Cache
With Sliding Window Attention (Chapter 5), the KV-cache is a fixed-size rolling buffer:
Inspecting KV-Cache in Transformers
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
from transformers import AutoModelForCausalLM, AutoTokenizer
import torch
model = AutoModelForCausalLM.from_pretrained(
"meta-llama/Llama-3.1-8B", torch_dtype=torch.bfloat16, device_map="auto"
)
tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-3.1-8B")
input_ids = tokenizer("Hello, world!", return_tensors="pt").input_ids.to(model.device)
# First forward pass (prefill)
outputs = model(input_ids, use_cache=True)
past_kv = outputs.past_key_values
# Inspect shapes
for layer_idx, (k, v) in enumerate(past_kv):
if layer_idx == 0:
print(f"Layer {layer_idx}:")
print(f" K shape: {k.shape}") # [1, 8, T, 128] for GQA with 8 KV heads
print(f" V shape: {v.shape}") # [1, 8, T, 128]
# Total KV-cache memory
total_bytes = sum(k.nbytes + v.nbytes for k, v in past_kv)
print(f"Total KV-cache: {total_bytes / 1024**2:.1f} MB")
# Decode one more token β only feed the last token
next_token = torch.argmax(outputs.logits[:, -1, :], dim=-1, keepdim=True)
outputs2 = model(next_token, past_key_values=past_kv, use_cache=True)
past_kv2 = outputs2.past_key_values
# Cache grew by 1 token
k0_old, k0_new = past_kv[0][0], past_kv2[0][0]
print(f"Cache grew: {k0_old.shape[2]} β {k0_new.shape[2]} tokens")
KV-Cache vs Model Weights Memory
At long context lengths, KV-cache can exceed the model weights:
| Context | Model Weights (8B, BF16) | KV-Cache (8B, GQA-8) | KV-Cache % |
|---|---|---|---|
| 2K | 16 GB | 0.25 GB | 1.5% |
| 8K | 16 GB | 1 GB | 6% |
| 32K | 16 GB | 4 GB | 20% |
| 128K | 16 GB | 16 GB | 50% |
| 512K | 16 GB | 64 GB | 80% |
At 128K context, the KV-cache equals the model size. At 512K, itβs 4Γ the model. This is why context extension (Chapter 12) and KV-cache optimization (Chapter 18) go hand in hand.
Batching and KV-Cache
When serving multiple requests, each request has its own KV-cache of potentially different length. Managing memory across a dynamic batch is non-trivial β this is the core problem that PagedAttention and other optimizations solve (see Chapter 18).
Whatβs Next
The KV-cache memory problem is the central challenge of LLM serving. The next chapter covers the full landscape of KV-cache optimization β from PagedAttention to quantized KV-cache to token eviction strategies.
β Previous: Chapter 16 β Inference & Sampling Β· Next: Chapter 18 β KV-Cache β Optimization β
Last updated: April 2026