OpenAI · ML & AI Fundamentals
Explain KV cache in Transformer inference
TrueInterview
October 7, 2026 · 10 min read
Question
During inference with a Transformer-based large language model, what does a key-value (KV) cache refer to? Provide a thorough, systems-oriented account that addresses:
- Which tensors get stored — roughly what shapes they take, and where inside the model they reside.
- The reason KV caching accelerates autoregressive decoding, and which asymptotic behavior it alters.
- How the prefill stage (running the prompt) differs from the decode stage (emitting one token at a time), and how differently each of them performs.
- The principal tradeoffs and pitfalls: memory growth, managing batched / variable-length requests, the multi-head-attention variants (MHA vs. MQA vs. GQA), consistency of positional encoding, and handling long contexts.
- No fewer than two practical optimizations that production serving systems actually use (for instance paged attention, a quantized KV cache, sliding-window / streaming attention, GQA).
hint Where to start Start from what self-attention redoes on every decoding step. Consider the query of a freshly generated token: among the per-token projections belonging to tokens that came before, which ones depend solely on hidden states that are already settled, and therefore stay constant once computed?
hint The key invariant Across steps, only the and of earlier tokens get reused; a serves its own token once and is then thrown away. Consider which quantity drops to per step and which remains once the prefix projections are no longer recomputed.
hint Two regimes Treat the prompt pass and the per-token loop as separate things. The first is a tall matrix–matrix product (many query rows at once); the second is a thin matrix–vector product (a single query row). Ask which of the two is constrained by GPU FLOPs and which by HBM bandwidth — that determines where batching pays off.
hint Memory and pitfalls
Express the cache size as a product of the evident factors (layers, batch, sequence, heads, head dimension, bytes per element, times two for and ) and observe which of them increase while serving. The production remedies then follow: reduce (GQA/MQA), reduce bytes per element (quantization), cap the sequence term (sliding window), or quit reserving max_len up front (paged / block-wise allocation).
Constraints & Assumptions
- Take a standard decoder-only Transformer performing autoregressive generation (causal self-attention), served on a GPU.
- = layer count, = batch size, = prompt length, = current sequence length, = query heads, = KV heads, = per-head dimension.
- The subject is inference, not training — there is no backward pass and the weights are frozen.
- "Production" here means a multi-tenant serving system juggling many concurrent requests of differing lengths, not a single-sequence toy script.
Clarifying Questions to Ask
Someone scoping this problem would ask:
- Is the target a decoder-only model (the usual case), or does the scope also include caching cross-attention for encoder–decoder models?
- Are we tuning for time-to-first-token (TTFT), for inter-token latency / throughput (TPOT, tokens/sec), or for maximum concurrent requests? These pull the design in different directions.
- Which context lengths and batch sizes are in play? That decides whether the cache or the weights dominate HBM.
- Which positional scheme is being used (RoPE, learned absolute, ALiBi)? It changes what has to be stored and how eviction or sliding interacts with positions.
- Is the architecture fixed, or may we assume or choose GQA/MQA (a decision made at architecture time, not an inference-time knob)?
What a Strong Answer Covers
These are the dimensions a strong answer is judged against (not the answers themselves):
- Precision about the cached object: that per-layer and for every past position are stored, that is not cached, and that the cache exists only in attention layers (not in embeddings, MLP or the LM head).
- Correct asymptotics: naming the redundant recomputation the cache removes (per-step projection cost , eliminating the projection work across tokens) while acknowledging honestly that the attention scan over cached keys remains .
- Prefill and decode as distinct regimes: the compute-bound prompt pass → TTFT versus the bandwidth-bound per-token loop → TPOT, and the reason batching amortizes the weight read during decode.
- The memory model: a correct size formula, the observation that it is linear in , and the consequence that throughput is usually limited by KV memory.
- Attention variants and their cache effect: MHA versus MQA versus GQA, and how scales the footprint.
- At least two concrete production optimizations named together with their mechanism, not merely buzzwords (for example paged attention's block table plus prefix sharing; quantization's reduction of bytes per element; sliding window's bounded sequence term).
- Awareness of the sharp edges: ragged / variable-length batching, per-sequence causal masking when packing, positional consistency under eviction, and cache duplication for beam or multi-sample decoding.
Follow-up Questions
- Derive the KV-cache size in bytes for a concrete model (given , , , FP16) at batch and context , and compare it against the weight memory — at what context length does the cache take over?
- How does paged attention make prefix sharing possible across requests that share a system prompt, and what must happen when a copy-on-write divergence occurs?
- Why are keys more sensitive than values to low-bit quantization, and which scaling granularity (per-tensor vs. per-channel vs. per-token) mitigates it?
- With a sliding-window cache, what breaks if the oldest tokens are simply dropped, and how do "attention-sink" schemes keep streaming generation stable at fixed memory?
- In disaggregated serving, why might prefill and decode be placed on separate hardware, and what has to be transferred between them?
Overview: This question gauges how well a candidate understands KV cache mechanics in Transformer inference: caching of attention state, the memory-versus-latency balance, and the engineering techniques that make autoregressive decoding practical.
Solution
What a KV cache is
Generation in a decoder-only Transformer proceeds autoregressively: producing token means attending over every earlier token . At each layer, self-attention forms three projections from the hidden state :
after which , where is the causal mask and is the per-head dimension.
The crucial point: appending a token leaves the and vectors of every preceding token untouched, since they are determined solely by those tokens' hidden states, which are already settled. Only the incoming token's has to attend over them. Rather than recompute and for the entire prefix on each step, you keep them around.
A KV cache stores, for each layer and each head, the and tensors of all positions encountered so far. On every decode step, , and are computed for just the one new token; its and are appended to the cache, and that single query attends against the complete cached and .
Worth noting: is never cached — a query serves its own token once and is then dropped. Reuse across later steps applies only to and .
What gets cached and the shapes
For each of the transformer layers, two tensors are kept:
K_cache: [batch, n_kv_heads, seq_len, head_dim] V_cache: [batch, n_kv_heads, seq_len, head_dim]
seq_lenincrements by one on each decode step, and equals the prompt length once prefill finishes.n_kv_headscounts the KV heads, and under grouped-query attention it can be smaller than the query-head count (covered below).- Caching happens in every attention layer; the embedding layer, the MLPs and the final LM head are left out, since nothing there is reused across tokens.
Exact memory cost
The total number of bytes held by the KV cache:
Here the leading accounts for and , is the layer count, the batch, the sequence length, the KV heads and the head dimension. The cache is linear in sequence length and batch size, which is exactly why it becomes the biggest memory consumer at long context. For one sequence it only approaches the weight footprint at very long context, but with a large batch the combined cache commonly surpasses the weights — and that is the situation that limits how many requests fit into HBM.
Why it speeds up decoding
Absent a cache, every new token forces attention to be re-run over the whole prefix, redoing the and projections for each of the earlier tokens. Across generated tokens this amounts to duplicated projection work.
With a cache in place:
- Per-step projection work becomes — only the new token is projected.
- The new query's attention still reads cached keys and values (those dot products are unavoidable), but the past projections are no longer recomputed.
The cache thus reduces the per-step projection cost from down to . What remains, the attention scan, is inexpensive next to the projections and MLP, and it is bound by memory bandwidth rather than compute. The upshot is a decode loop whose per-token time stays roughly flat instead of worsening as the prefix grows.
Prefill vs. decode
Operationally these are two separate phases, and their performance characteristics differ.
| Prefill | Decode | |
|---|---|---|
| Input per forward pass | Entire prompt, tokens | A single token |
| Cache action | Fill all positions across every layer | Add one position |
| Matrix shape | matrix–matrix ( is tall) | matrix–vector ( has one row) |
| Bottleneck | Limited by compute (GPU FLOPs) | Limited by memory bandwidth |
| Parallelism | Every prompt position handled at once | Sequential by nature |
Prefill pushes the whole prompt through the model in one pass. Since has rows, attention and the FFN become large dense matmuls that keep the GPU's compute units busy — this phase favors throughput and batches efficiently. Its latency sets the time-to-first-token (TTFT).
Decode handles a single token per step, making each step a thin matrix–vector product. Arithmetic intensity is low: the full set of model weights (plus the expanding KV cache) must be streamed out of HBM to support a tiny amount of computation, so memory bandwidth, not FLOPs, is the limit. Decode latency sets the inter-token latency (TPOT, time per output token) — the per-token latency whose reciprocal is the tokens-per-second throughput. That is the reason serving systems group many concurrent decodes into a batch: it spreads one weight read from HBM across numerous requests.
Because of this divide, large deployments frequently schedule the two phases separately and even place them on distinct hardware ("disaggregated" prefill and decode) — see the follow-up question below.
Tradeoffs and pitfalls
1. Memory growth is the dominant concern. As derived above, KV memory grows linearly with . At long context the cache may far exceed the weights, which caps the number of concurrent requests (batch size) that fit in HBM. In production, throughput is typically constrained by KV memory rather than by FLOPs.
2. Ragged, variable-length batching. Sequences within one batch carry different prompt lengths and complete at different moments. Reserving max_seq_len for each slot in a naive way leaves most of the cache unused. Completed sequences create gaps, so some bookkeeping is needed to reclaim and reuse that memory (the main motivation behind paged attention, discussed below).
3. Attention variants alter the cache footprint. Cache size scales directly with :
- MHA (multi-head): — the largest cache.
- MQA (multi-query): one shared KV head, — the smallest cache, at some cost in quality.
- GQA (grouped-query): lies between and , with query heads sharing KV heads in groups — the usual modern compromise. GQA cuts the cache by the grouping factor while barely affecting quality, which explains its current status as the default.
4. Consistency of positional encoding. Under RoPE, the rotation applied to depends on absolute position and happens before caching, so stored keys already embed their positional phase — fine provided the rotation uses the correct index. Trouble arises when positions move, as with a sliding window or cache eviction: dropping entries naively throws off the position arithmetic unless you re-anchor or adopt a relative scheme. With learned absolute positions, the embedding has to be added at the right index.
5. Masking when batches are packed. The causal mask has to restrict every token to earlier positions within its own sequence. If several sequences share one tensor, per-sequence (block-diagonal) masks are required so that tokens cannot leak across sequence boundaries.
6. Beam search and multi-sample decoding. Beam search or parallel sampling scales the cache by the beam width or sample count, unless the common prompt prefix is stored a single time and only the divergent branches are duplicated (prefix sharing).
Practical production optimizations
1. Paged attention (block-wise KV management). The cache is kept in fixed-size blocks ("pages"), with a per-sequence block table translating logical positions into physical pages, much like virtual memory in an operating system. The gains: almost no internal fragmentation, growth one page at a time rather than reserving max_len in advance, and prefix sharing — several requests sharing a system prompt reference the same physical pages, with copy-on-write when they diverge. In practice this permits much larger effective batch sizes within the same HBM. (Popularized by vLLM.)
2. Quantized KV cache. and are held at reduced precision (INT8 or FP8, say) while matmul accumulation stays in BF16/FP16. This roughly halves or quarters both the cache footprint and the bandwidth required to read it during decode, which lifts batch capacity and decode throughput. It calls for per-channel or per-token scaling — keys are more sensitive than values, since a corrupted outlier channel in a key distorts every attention score computed through the dot product.