The KV Cache
Trading GPU VRAM memory for compute by caching historical Key and Value tensors during autoregressive generation.
The Autoregressive Bottleneck Without KV Cache
In Causal Self-Attention:
At step , we generate token 100. To compute attention for token 100, we need Key vectors and Value vectors .
- Without KV Cache: Re-compute for tokens from scratch. Total operations over tokens = .
- With KV Cache: Fetch pre-computed from GPU VRAM; compute for token 100 only. Total operations over tokens = .
WITHOUT KV CACHE (Redundant Matrix Re-computation)
Step 1: Compute [T1]
Step 2: Re-compute [T1], Compute [T2]
Step 3: Re-compute [T1], Re-compute [T2], Compute [T3] <-- O(N³) Waste!
WITH KV CACHE (Append to VRAM Memory)
Step 1: Compute [T1] -> Store K1, V1 in VRAM
Step 2: Fetch K1, V1 -> Compute [T2] -> Append K2, V2 to VRAM
Step 3: Fetch K1..2, V1..2 -> Compute [T3] -> Append K3, V3 to VRAM <-- O(N²) Efficient!
KV Cache VRAM Formula
For a batch of sequences, sequence length , layers, Key-Value heads, and head dimension :
Concrete Example (LLaMA-3 70B FP16)
- , (GQA), , Precision = 16-bit (2 bytes).
- Per token per sequence: .
- For batch size at sequence length :
The KV Cache size quickly exceeds the base model weights size at high batch sizes and long contexts!
Architectural Mitigations
Multi-Head Attention (MHA) Grouped-Query Attention (GQA) Multi-Query Attention (MQA)
N_q Query Heads, N_kv = N_q KV Heads N_q Query Heads, N_kv = 8 KV Heads N_q Query Heads, N_kv = 1 KV Head
[Q1 Q2 Q3 Q4] [K1 K2 K3 K4] [Q1 Q2 Q3 Q4] ──► [K1] [Q1 Q2 Q3 Q4] ──► [K1]
(Large KV Cache) (8x Memory Reduction! - LLaMA-3) (32x Memory Reduction!)
Say this out loud
"The KV Cache stores Key and Value projection tensors of past tokens in GPU VRAM during autoregressive generation. This eliminates redundant self-attention re-computations, reducing generation complexity per token from O(N²) to O(N). However, KV cache VRAM footprint scales linearly with context length and batch size, requiring Grouped-Query Attention (GQA) and PagedAttention to fit long context serving in GPU memory."
Follow-ups to expect
- What is KV Cache Quantization (INT8 / INT4 KV Cache)? Quantizes cached K and V tensors from FP16 to INT8 or FP8 before storing in VRAM, reducing KV Cache VRAM footprint by 50-75% with minimal accuracy degradation.
- How does Chunked Prefill work with KV Cache? Splits long input prompt prefill processing into smaller chunks interleaved with decode steps, preventing long prompt prefills from starving ongoing decode token latency.
Check yourself
Why does autoregressive token generation require re-computing attention scores over historical tokens if KV Cache is NOT used?