KV Cache Explained: Efficient Attention for LLM Generation

Michael BrenndoerferJanuary 6, 202656 min read

Part of Language AI Handbook

Explains how KV cache eliminates redundant attention computations in transformers. Topics include memory requirements, cache structure.

Choose your expertise level to adjust how many terms are explained. Beginners see more tooltips, experts see fewer to maintain reading flow. Hover over underlined terms for instant definitions.

Article links

Make inline references clickable

KV Cache

When you ask a language model to complete the sentence "The capital of France is," it generates tokens one at a time: first "Paris," then perhaps a comma, then additional context. At each step, the model computes attention over all previous tokens to decide what comes next. Without optimization, this means the model recomputes the same attention calculations for "The," "capital," "of," "France," "is" every single time it generates a new token. For a 100-token response, that's nearly 5,000 redundant computations just for the prompt tokens. This cost does not just slow things down a little; it makes practical deployment nearly impossible.

The key-value cache, commonly called the KV cache, eliminates this redundancy by storing intermediate attention computations and reusing them across generation steps. This optimization reduces the computational complexity of autoregressive generation from quadratic to linear, making inference with large language models feasible at production scale. Think of the KV cache as the model's notepad: instead of re-reading the entire novel every time it needs to write the next sentence, the model keeps its reading notes and consults them directly. The notes grow with each sentence written, but the act of re-reading the full novel from scratch never happens.

Understanding the KV cache is necessary for reasoning about modern LLM serving. Paged attention, continuous batching, speculative decoding, and grouped-query attention all manage or extend the KV cache more effectively. Before you can reason about why these techniques matter, you need to understand what the cache is, why it works, what it costs, and where it breaks down.

This chapter begins by dissecting exactly where redundancy arises during autoregressive generation, building up from the attention formula itself. We then formalize how caching eliminates that redundancy, examine the per-layer and per-head structure of the cache in real architectures, and work through the memory arithmetic that determines whether a deployment is feasible. A fully worked numerical example shows each computation step-by-step before we move into a complete PyTorch implementation. We close with an honest examination of the limitations that have motivated years of follow-on research.

Historical Context

The KV cache is so basic that it is difficult to pinpoint a single paper that "introduced" it. Caching key-value pairs in autoregressive Transformers was a natural consequence of how the original Transformer architecture (Vaswani et al., 2017) was adapted for generation. The idea that intermediate representations could be reused across decoding steps appeared in early inference implementations of the original Transformer, and became standard practice as soon as researchers began deploying decoder-only models for text generation. GPT-2's publicly released code already included generation loops structured to exploit this property. The term "KV cache" became widespread as researchers began to study its memory costs formally, particularly with the rise of long-context models like GPT-3 and the subsequent explosion of 7B-to-70B open-weight models where cache memory began competing seriously with model weight memory for limited GPU resources.

The Redundancy Problem in Autoregressive Generation

As we explored in Part XIII, decoder-only transformers generate text one token at a time. At each step, the model takes all tokens generated so far (including the prompt) and predicts the next token. The computational core of this process is the self-attention mechanism, which allows each token to gather information from all tokens that came before it. To understand why caching gives such large benefits, we must first examine exactly how attention operates during generation and identify where the redundant work occurs. The redundancy is not an accident or a design flaw; it is an inherent consequence of computing context-aware representations at every position.

How Attention Creates Redundancy

Recall from Part XIII that self-attention computes queries, keys, and values from the input. The basic insight behind these three components is that attention works like an information retrieval system: queries stand for what a token is looking for, keys stand for what each token offers, and values contain the actual information to be retrieved. Think of it like a library search: the query is your search term, the keys are the catalog entries for every book on the shelf, and the values are the books themselves. When you search, you compare your query against every catalog entry and retrieve the books whose entries best match. We express these three components through linear projections:

Q=XWQ,K=XWK,V=XWVQ = XW_Q, \quad K = XW_K, \quad V = XW_V

where:

  • XX: the input embedding matrix, containing the vector representations of all tokens in the sequence
  • WQ,WK,WVW_Q, W_K, W_V: the learned projection matrices that transform inputs into query, key, and value subspaces, each capturing a different aspect of the token's meaning
  • Q,K,VQ, K, V: the resulting matrices, where QQ contains vectors representing what each token seeks, KK contains vectors representing what each token can be matched against, and VV contains the actual content to be aggregated

Each of these projection matrices is learned during training to extract the most useful representations for the attention mechanism. Once training is complete, these matrices remain fixed during inference. This means that given the same input token, the projections will always produce the same key and value vectors. This determinism is the property that makes caching possible: if the input and the weight matrix are unchanged, the output cannot change either.

The attention mechanism then computes a weighted sum of values based on the compatibility between queries and keys. This compatibility is measured through dot products, which capture how well a query's "question" matches each key's "answer." The full attention formula is:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V

where:

  • Q,K,VQ, K, V: the query, key, and value matrices computed from the input
  • KTK^T: the transpose of the key matrix, arranged so that the dot product between queries and keys can be computed efficiently through matrix multiplication
  • dkd_k: the dimension of the key vectors, used for scaling to prevent the dot products from growing too large
  • softmax\text{softmax}: the function that converts raw compatibility scores into probability weights summing to 1

The scaling factor dk\sqrt{d_k} is necessary. Without this normalization, the dot products between high-dimensional vectors can become very large in magnitude. When these large values pass through the softmax function, the result becomes extremely peaked, meaning almost all attention weight concentrates on a single token. This causes gradients to vanish for all other positions, making training unstable. The square root scaling keeps the variance of the dot products roughly constant regardless of the dimension, and we explored this phenomenon in depth in Part XIII.

Tracing the Redundancy Step by Step

During generation, let's trace what happens at each step to understand precisely where the wasted work accumulates. Suppose we're generating a response to a 10-token prompt. The model must produce one token at a time, with each new token depending on everything that came before it.

Step 1 (Prefill): The model processes all 10 prompt tokens simultaneously. It computes the QQ, KK, VV matrices, each of shape (10,dk)(10, d_k) or (10,dv)(10, d_v) depending on the specific projection. The attention mechanism allows each prompt token to attend to all previous prompt tokens, establishing the initial context representation.

Step 2: After generating the first new token, we now have 11 tokens total. Without any optimization, the model would need to process the entire 11-token sequence from scratch. This means computing QQ, KK, VV for all 11 tokens, even though 10 of them are unchanged from the previous step.

Step 3: With 12 tokens total, we compute QQ, KK, VV for all 12 tokens, once again recomputing the projections for the original prompt tokens that haven't changed at all.

Notice the wasteful pattern: at step tt, we recompute the keys and values for the first 10 prompt tokens, even though nothing about them has changed since step 1. The same tokens, passing through the same fixed projection matrices WKW_K and WVW_V, produce the same KK and VV vectors every single time. This redundancy stems from a basic property of the attention mechanism: the key and value for a given token depend only on that token's embedding and the projection weights, not on what comes after it in the sequence. The causal mask ensures that earlier tokens never attend to later ones, so later tokens can never influence the keys and values of earlier tokens.

This redundancy has concrete costs that grow sharply with sequence length. For a model with LL layers, HH attention heads, and head dimension dhd_h, generating TT tokens requires computing KK and VV projections repeatedly at every step. The total number of projection operations across all generation steps sums to:

∑t=1Tt=T(T+1)2≈T22\sum_{t=1}^{T} t = \frac{T(T+1)}{2} \approx \frac{T^2}{2}

where:

  • TT: the final sequence length, representing the total number of tokens processed including both the prompt and generated tokens
  • tt: the current step number, which also is the sequence length at that particular step
  • ∑t=1Tt\sum_{t=1}^{T} t: the sum of the arithmetic progression 1+2+3+…+T1 + 2 + 3 + \ldots + T, capturing the total number of token-projection operations across all steps

This quadratic scaling has severe implications for practical deployment. If generating 100 tokens requires 100×1012=5,050\frac{100 \times 101}{2} = 5,050 projection operations, then generating 1,000 tokens requires 1000×10012=500,500\frac{1000 \times 1001}{2} = 500,500 operations, a 100-fold increase in work for only a 10-fold increase in output length. This contrasts with the ideal case where each token's keys and values are computed only once, requiring TT projection operations. The KV cache reaches exactly this ideal.

Out[3]:
Visualization
Line chart showing total projection operations vs sequence length. The red quadratic curve (without KV cache) rises steeply to ~125,000 at T=500, while the green linear curve (with KV cache) remains flat near zero.
Computational complexity comparison between generation with and without KV cache. Without caching, the number of projection operations grows quadratically with sequence length. With KV cache, the growth is linear, giving large savings for longer sequences. The shaded red region is the eliminated redundant operations.

The KV Cache Solution

The insight behind KV caching is straightforward once we recognize where the redundancy lies: since each token's key and value vectors depend only on that token's representation and the fixed projection weights, we can compute them once and store them for reuse in all subsequent generation steps. The query vectors, by contrast, must be recomputed because they are used differently at each step. A query asks "what information should I attend to?" and the answer depends on the token's position in the evolving sequence. The key insight is that keys and values are determined entirely by the token that produced them, while queries are only real when there is a full, growing context to attend to.

This asymmetry between queries and key-values is basic to understanding the cache. When token 5 generates a query, it needs to attend to tokens 1 through 4. When token 10 generates a query, it needs to attend to tokens 1 through 9. The keys and values for tokens 1 through 4 are the same in both cases since they haven't changed. But the query's role is different at each step: it always asks about everything that came before, and "everything that came before" grows with each step. Think of queries as questions that depend on who is asking and when, while keys and values are stable answers that remain constant once a token has been processed.

During generation with KV caching, the process changes fundamentally from the naive approach. The two phases, prefill and decode, serve very different computational roles:

Prefill phase (Step 1): Process all prompt tokens at once, computing queries, keys, and values for the entire prompt. Store the computed KK and VV matrices in the cache, as these will be reused throughout generation. Use the queries to compute attention within the prompt, establishing the initial hidden states. This phase is compute-intensive but happens only once.

Decode phase (Steps 2+): For each new token, the process is efficient:

  1. Compute QQ, KK, VV only for the single new token
  2. Append the new KK and VV vectors to the end of the cache
  3. Compute attention using the new QQ against all cached KK values
  4. Weight all cached VV values by the resulting attention scores
  5. Produce the output for this position

The important difference is that we only compute projections for one token per step (the newly generated one), while the attention computation can freely access all previously computed keys and values through the cache. This transforms the projection work from quadratic to linear in the total sequence length. The attention computation itself remains linear in the cache size at each step (since we compute qtKcacheTq_t K_{\text{cache}}^T for a single query), but that is unavoidable because the query must attend over all prior context.

Mathematical Formulation

Let's formalize this caching mechanism to understand exactly how it reaches the computational savings. The formalization reveals that caching is not an approximation but an algebraically equivalent reformulation of the attention computation. The outputs are bit-for-bit identical to non-cached attention; only the order and grouping of computations changes.

At generation step tt, let xtx_t denote the embedding of the newly generated token. This embedding captures all the information the model has about this token at the input to the attention layer. We project this single token's embedding to obtain its query, key, and value vectors:

qt=xtWQ,kt=xtWK,vt=xtWVq_t = x_t W_Q, \quad k_t = x_t W_K, \quad v_t = x_t W_V

where:

  • xtx_t: the input embedding for the current token at step tt, a vector of dimension dmodeld_{\text{model}}
  • WQ,WK,WVW_Q, W_K, W_V: the fixed projection matrices shared across all positions and all generation steps
  • qtq_t: the query vector representing what the current token is seeking to attend to
  • ktk_t: the key vector representing how this token can be matched by future queries
  • vtv_t: the value vector containing the information this token contributes when attended to

Note the important distinction from the non-cached case: each of these is a single vector of dimension dkd_k or dvd_v, not a full matrix covering all positions. We compute projections for exactly one token, not the entire sequence.

The cache is the model's memory of past tokens, maintaining the history of key and value vectors computed in all previous steps. At the beginning of step tt, before processing the new token, the cache contains concatenated matrices:

Kcache=[k1;k2;…;kt−1]∈R(t−1)×dkVcache=[v1;v2;…;vt−1]∈R(t−1)×dv\begin{aligned} K_{\text{cache}} &= [k_1; k_2; \ldots; k_{t-1}] \in \mathbb{R}^{(t-1) \times d_k} \\ V_{\text{cache}} &= [v_1; v_2; \ldots; v_{t-1}] \in \mathbb{R}^{(t-1) \times d_v} \end{aligned}

where:

  • Kcache,VcacheK_{\text{cache}}, V_{\text{cache}}: matrices storing the complete history of keys and values for all previous t−1t-1 tokens
  • ki,vik_i, v_i: the key and value vectors computed when token ii was first processed
  • [;][;]: the concatenation operation along the sequence dimension, stacking vectors as rows
  • dk,dvd_k, d_v: the dimensions of the key and value vectors, often equal in practice

This cache is all the "memory" the attention mechanism has of the sequence so far. Each row corresponds to a token's contribution to the attention computation.

After computing the new token's projections, we update the cache by appending the new key and value vectors:

Kcache←[Kcache;kt]Vcache←[Vcache;vt]\begin{aligned} K_{\text{cache}} &\leftarrow [K_{\text{cache}}; k_t] \\ V_{\text{cache}} &\leftarrow [V_{\text{cache}}; v_t] \end{aligned}

where:

  • Kcache,VcacheK_{\text{cache}}, V_{\text{cache}}: the cached key and value matrices being extended with new entries
  • ←\leftarrow: the assignment operator showing an in-place update to the stored cache state
  • kt,vtk_t, v_t: the newly computed key and value vectors being appended to preserve the sequence history

After this update, the cache contains keys and values for all tt tokens, ready for the attention computation.

With the cache updated, the model computes attention for the new token by having its query vector interact with the entire history:

at=softmax(qtKcacheTdk)∈R1×tot=atVcache∈R1×dv\begin{aligned} a_t &= \text{softmax}\left(\frac{q_t K_{\text{cache}}^T}{\sqrt{d_k}}\right) \in \mathbb{R}^{1 \times t} \\ o_t &= a_t V_{\text{cache}} \in \mathbb{R}^{1 \times d_v} \end{aligned}

where:

  • ata_t: the attention weights for the current step, a vector of tt probabilities representing how much the new token should attend to each previous position
  • qtq_t: the query vector for the current token, seeking needed information from the context
  • KcacheTK_{\text{cache}}^T: the transpose of the cached key matrix, shaped for efficient dot product computation with the query
  • VcacheV_{\text{cache}}: the cached value matrix containing the actual content to be retrieved and aggregated
  • dkd_k: the dimension of key vectors, giving the scaling factor for numerical stability
  • dvd_v: the dimension of value vectors, determining the output size
  • oto_t: the final attention output for step tt, a weighted combination of all cached values

The computation qtKcacheTq_t K_{\text{cache}}^T produces a vector of tt scores, one for each cached position. These scores measure how needed each previous token is to the current query. After scaling and applying softmax, these become proper attention weights that sum to 1. Finally, multiplying by VcacheV_{\text{cache}} aggregates the cached values according to these weights, creating the output vector that captures all the needed information from the sequence history.

This formulation reaches efficiency because the matrix multiplications involve a single query vector against the cached matrices, requiring O(t⋅d)O(t \cdot d) operations rather than O(t2⋅d)O(t^2 \cdot d) that would be needed to recompute attention from scratch at every step.

Out[4]:
Visualization
Bar chart showing cached token count per generation step. The bar labeled Prefill starts at 8 tokens, then 12 decode-step bars grow from 9 to 20 tokens.
Cache growth during autoregressive generation. Starting with an 8-token prompt processed during prefill, the cache grows by one token at each decode step. The prefill phase processes all prompt tokens simultaneously, while each subsequent step adds exactly one new token's keys and values to the cache. The orange dashed line marks the initial prompt length boundary.

Worked Example: Tracing a 4-Token Generation

To make the abstract formulas concrete, let's trace through a minimal numerical example from start to finish. We will use a toy model with dmodel=4d_{\text{model}} = 4, one attention head (H=1H = 1), and dk=dv=4d_k = d_v = 4. The prompt is three tokens: "The cat sat." We will then generate one additional token and track every vector and matrix involved.

Setup

Suppose the three prompt tokens produce input embeddings (after the embedding layer and any positional encoding):

x1=[1.0,0.5,0.2,0.8],x2=[0.3,1.0,0.6,0.1],x3=[0.7,0.4,1.0,0.3]x_1 = [1.0, 0.5, 0.2, 0.8], \quad x_2 = [0.3, 1.0, 0.6, 0.1], \quad x_3 = [0.7, 0.4, 1.0, 0.3]

For the projection matrices, we use simplified illustrative weights. In practice these are dense learned matrices, but here we treat them as identity-like for clarity:

WK=WV=I4(identity matrix)W_K = W_V = I_4 \quad \text{(identity matrix)} WQ=0.5⋅I4W_Q = 0.5 \cdot I_4

This means the key and value for each token are just that token's embedding, and the query is half the embedding. The simplification does not change the structural argument at all.

Prefill Phase: All Three Prompt Tokens

We compute the keys and values for all three tokens simultaneously:

k1=x1WK=[1.0,0.5,0.2,0.8]k2=x2WK=[0.3,1.0,0.6,0.1]k3=x3WK=[0.7,0.4,1.0,0.3]\begin{aligned} k_1 &= x_1 W_K = [1.0, 0.5, 0.2, 0.8] \\ k_2 &= x_2 W_K = [0.3, 1.0, 0.6, 0.1] \\ k_3 &= x_3 W_K = [0.7, 0.4, 1.0, 0.3] \end{aligned} v1=x1WV=[1.0,0.5,0.2,0.8]v2=x2WV=[0.3,1.0,0.6,0.1]v3=x3WV=[0.7,0.4,1.0,0.3]\begin{aligned} v_1 &= x_1 W_V = [1.0, 0.5, 0.2, 0.8] \\ v_2 &= x_2 W_V = [0.3, 1.0, 0.6, 0.1] \\ v_3 &= x_3 W_V = [0.7, 0.4, 1.0, 0.3] \end{aligned}

We store these in the cache immediately:

Kcache=[1.00.50.20.80.31.00.60.10.70.41.00.3],Vcache=[1.00.50.20.80.31.00.60.10.70.41.00.3]K_{\text{cache}} = \begin{bmatrix} 1.0 & 0.5 & 0.2 & 0.8 \\ 0.3 & 1.0 & 0.6 & 0.1 \\ 0.7 & 0.4 & 1.0 & 0.3 \end{bmatrix}, \quad V_{\text{cache}} = \begin{bmatrix} 1.0 & 0.5 & 0.2 & 0.8 \\ 0.3 & 1.0 & 0.6 & 0.1 \\ 0.7 & 0.4 & 1.0 & 0.3 \end{bmatrix}

(In this toy example the keys and values are equal since WK=WV=IW_K = W_V = I; in real models they differ entirely.)

Decode Phase: Generating Token 4

Suppose the model samples token 4 from the distribution produced after the prefill phase, and that token 4 has embedding x4=[0.5,0.9,0.3,0.6]x_4 = [0.5, 0.9, 0.3, 0.6]. To predict token 5, we need to compute attention from position 4 over all previous positions.

Step 1: Compute projections for token 4 only.

q4=x4WQ=0.5⋅[0.5,0.9,0.3,0.6]=[0.25,0.45,0.15,0.30]q_4 = x_4 W_Q = 0.5 \cdot [0.5, 0.9, 0.3, 0.6] = [0.25, 0.45, 0.15, 0.30] k4=x4WK=[0.5,0.9,0.3,0.6],v4=x4WV=[0.5,0.9,0.3,0.6]k_4 = x_4 W_K = [0.5, 0.9, 0.3, 0.6], \quad v_4 = x_4 W_V = [0.5, 0.9, 0.3, 0.6]

The key insight is that we compute exactly three vectors (q4q_4, k4k_4, v4v_4) for this step, not 4×3=124 \times 3 = 12 vectors. The keys and values for tokens 1 through 3 remain exactly as computed during prefill.

Step 2: Append k4k_4 and v4v_4 to the cache.

Kcache←[1.00.50.20.80.31.00.60.10.70.41.00.30.50.90.30.6]K_{\text{cache}} \leftarrow \begin{bmatrix} 1.0 & 0.5 & 0.2 & 0.8 \\ 0.3 & 1.0 & 0.6 & 0.1 \\ 0.7 & 0.4 & 1.0 & 0.3 \\ 0.5 & 0.9 & 0.3 & 0.6 \end{bmatrix}

Step 3: Compute attention scores for token 4's query against all four cached keys.

The raw dot products are:

s1=q4⋅k1=(0.25)(1.0)+(0.45)(0.5)+(0.15)(0.2)+(0.30)(0.8)=0.25+0.225+0.03+0.24=0.745s2=q4⋅k2=(0.25)(0.3)+(0.45)(1.0)+(0.15)(0.6)+(0.30)(0.1)=0.075+0.45+0.09+0.03=0.645s3=q4⋅k3=(0.25)(0.7)+(0.45)(0.4)+(0.15)(1.0)+(0.30)(0.3)=0.175+0.18+0.15+0.09=0.595s4=q4⋅k4=(0.25)(0.5)+(0.45)(0.9)+(0.15)(0.3)+(0.30)(0.6)=0.125+0.405+0.045+0.18=0.755\begin{aligned} s_1 &= q_4 \cdot k_1 = (0.25)(1.0) + (0.45)(0.5) + (0.15)(0.2) + (0.30)(0.8) = 0.25 + 0.225 + 0.03 + 0.24 = 0.745 \\ s_2 &= q_4 \cdot k_2 = (0.25)(0.3) + (0.45)(1.0) + (0.15)(0.6) + (0.30)(0.1) = 0.075 + 0.45 + 0.09 + 0.03 = 0.645 \\ s_3 &= q_4 \cdot k_3 = (0.25)(0.7) + (0.45)(0.4) + (0.15)(1.0) + (0.30)(0.3) = 0.175 + 0.18 + 0.15 + 0.09 = 0.595 \\ s_4 &= q_4 \cdot k_4 = (0.25)(0.5) + (0.45)(0.9) + (0.15)(0.3) + (0.30)(0.6) = 0.125 + 0.405 + 0.045 + 0.18 = 0.755 \end{aligned}

Step 4: Scale by dk=4=2.0\sqrt{d_k} = \sqrt{4} = 2.0.

s^=[0.745/2, 0.645/2, 0.595/2, 0.755/2]=[0.3725, 0.3225, 0.2975, 0.3775]\hat{s} = [0.745 / 2, \ 0.645 / 2, \ 0.595 / 2, \ 0.755 / 2] = [0.3725, \ 0.3225, \ 0.2975, \ 0.3775]

Step 5: Apply softmax to get attention weights.

First compute the exponentials:

e0.3725≈1.451,e0.3225≈1.380,e0.2975≈1.346,e0.3775≈1.459e^{0.3725} \approx 1.451, \quad e^{0.3225} \approx 1.380, \quad e^{0.2975} \approx 1.346, \quad e^{0.3775} \approx 1.459

Sum: 1.451+1.380+1.346+1.459=5.6361.451 + 1.380 + 1.346 + 1.459 = 5.636

Normalized weights:

a=[1.451/5.636, 1.380/5.636, 1.346/5.636, 1.459/5.636]≈[0.257, 0.245, 0.239, 0.259]a = [1.451/5.636, \ 1.380/5.636, \ 1.346/5.636, \ 1.459/5.636] \approx [0.257, \ 0.245, \ 0.239, \ 0.259]

Notice that token 4's own key (position 4) receives slightly more weight than the others, which is typical for recent tokens in causal attention.

Step 6: Compute the attention output as a weighted sum of cached values.

o4=0.257⋅v1+0.245⋅v2+0.239⋅v3+0.259⋅v4o_4 = 0.257 \cdot v_1 + 0.245 \cdot v_2 + 0.239 \cdot v_3 + 0.259 \cdot v_4 o4[0]=0.257(1.0)+0.245(0.3)+0.239(0.7)+0.259(0.5)=0.257+0.073+0.167+0.130=0.627o4[1]=0.257(0.5)+0.245(1.0)+0.239(0.4)+0.259(0.9)=0.129+0.245+0.096+0.233=0.703o4[2]=0.257(0.2)+0.245(0.6)+0.239(1.0)+0.259(0.3)=0.051+0.147+0.239+0.078=0.515o4[3]=0.257(0.8)+0.245(0.1)+0.239(0.3)+0.259(0.6)=0.206+0.025+0.072+0.155=0.458\begin{aligned} o_4[0] &= 0.257(1.0) + 0.245(0.3) + 0.239(0.7) + 0.259(0.5) = 0.257 + 0.073 + 0.167 + 0.130 = 0.627 \\ o_4[1] &= 0.257(0.5) + 0.245(1.0) + 0.239(0.4) + 0.259(0.9) = 0.129 + 0.245 + 0.096 + 0.233 = 0.703 \\ o_4[2] &= 0.257(0.2) + 0.245(0.6) + 0.239(1.0) + 0.259(0.3) = 0.051 + 0.147 + 0.239 + 0.078 = 0.515 \\ o_4[3] &= 0.257(0.8) + 0.245(0.1) + 0.239(0.3) + 0.259(0.6) = 0.206 + 0.025 + 0.072 + 0.155 = 0.458 \end{aligned}

So o4≈[0.627,0.703,0.515,0.458]o_4 \approx [0.627, 0.703, 0.515, 0.458].

What did we compute in this decode step? Three vector projections (for token 4 only), four dot products, a softmax over four values, and four weighted additions. Without the cache, we would have recomputed all projections for tokens 1, 2, and 3 as well. That is 9 additional projection operations for a sequence of length 4. For a sequence of length 100, the savings would be 297 additional projections avoided. For length 1000, the savings grow to 2997 projections avoided in that single step alone.

Cache Structure

The KV cache must be maintained separately for each layer and each attention head in the transformer architecture. This requirement arises because different layers and heads learn different attention patterns and have different projection weights. Layer 1's keys and values are fundamentally different from layer 32's, as each layer captures different levels of abstraction: early layers tend to track syntactic patterns like word boundaries and part-of-speech relationships, while later layers tend to encode more semantic, task-needed information. Similarly, within a single layer, head 1 might learn to track coreference (which pronoun refers to which entity), while head 4 tracks positional relationships, each requiring its own separate cache.

Think of the cache as a filing cabinet with LL drawers (one per layer), and within each drawer, HH hanging folders (one per head). Every hanging folder contains two documents: the key history and the value history. When the model processes a new token, it opens every drawer, pulls out all the folders, appends one line to each key document and one line to each value document, then puts everything back. The output depends on reading every folder in the needed drawer at every layer.

Per-Layer, Per-Head Organization

For a transformer with LL layers and HH attention heads per layer, the complete cache consists of 2×L×H2 \times L \times H tensors: one KK cache and one VV cache for each attention head at each layer. This combinatorial structure means that even modest models require tracking many separate cache tensors. For LLaMA-2 7B with 32 layers and 32 heads, this is 2×32×32=2,0482 \times 32 \times 32 = 2,048 individual tensors per sequence in the batch.

In practice, deep learning frameworks typically consolidate this organization into two tensors per layer, using an additional dimension to index across heads:

Kcache(l)∈RB×H×T×dhVcache(l)∈RB×H×T×dh\begin{aligned} K_{\text{cache}}^{(l)} &\in \mathbb{R}^{B \times H \times T \times d_h} \\ V_{\text{cache}}^{(l)} &\in \mathbb{R}^{B \times H \times T \times d_h} \end{aligned}

where:

  • Kcache(l),Vcache(l)K_{\text{cache}}^{(l)}, V_{\text{cache}}^{(l)}: the key and value cache tensors for layer ll, containing caches for all heads
  • BB: the batch size, letting multiple sequences to be processed in parallel
  • HH: the number of attention heads, each with its own cached key-value pairs
  • TT: the current sequence length, growing as generation proceeds
  • dhd_h: the head dimension, typically dmodel/Hd_{\text{model}} / H

Some implementations combine keys and values into a single tensor of shape (B,H,T,2,dh)(B, H, T, 2, d_h) or (B,2,H,T,dh)(B, 2, H, T, d_h) for improved memory locality. This layout keeps each position's key and value adjacent in memory, which can improve cache efficiency during the attention computation. However, the logical structure remains the same regardless of the physical memory layout.

How Cached Values Flow Through the Model

Understanding how the cache integrates with the transformer's layer-by-layer computation clarifies why the optimization is both correct and efficient. During the decode phase, when generating token by token, each transformer layer performs the following sequence of operations:

  1. Receives the hidden state for only the new token: ht(l)∈RB×1×dmodelh_t^{(l)} \in \mathbb{R}^{B \times 1 \times d_{\text{model}}}. This is the representation of the new token as computed by the previous layer.
  2. Projects this hidden state to obtain qtq_t, ktk_t, vtv_t for the new token only.
  3. Retrieves the cached Kcache(l)K_{\text{cache}}^{(l)} and Vcache(l)V_{\text{cache}}^{(l)} containing all previously computed keys and values for this layer.
  4. Appends the new ktk_t and vtv_t to the cache, extending the history by one position.
  5. Computes attention between the new query qtq_t and the full cache, letting the new token to attend to all previous positions.
  6. Passes the attention output through the feed-forward network, which processes each position independently.
  7. Returns the hidden state for the new token: ht(l+1)h_t^{(l+1)}, ready for the next layer.

Notice that operations 1, 2, 6, and 7 all operate on just a single position. The feed-forward network, layer normalization, and residual connections process only the new token's hidden state, making them computationally trivial during decoding. The attention computation in step 5 is the only operation that touches the full sequence length, and even there, the work is linear in tt rather than quadratic because we compute attention for just one query position.

This also explains why the KV cache only caches keys and values, not queries: queries are consumed immediately within the layer to compute attention, and a query for position tt can never be needed again because future positions t′,t′>tt', t' > t will generate their own queries using their own embeddings. The query is inherently step-local; the key and value are inherently step-global.

The Role of Causal Masking

An important subtlety is that causal masking interacts with the cache in a way that is often glossed over. During prefill, the model processes the full prompt as a matrix operation and applies a causal mask so that position ii cannot attend to positions j>ij > i. During decode, however, the new token at position tt can attend to all prior positions 1 through t−1t-1 without any masking restriction, because all those positions precede it in the sequence. The decode phase therefore applies no mask at all for the single-token query case. This simplification is valid and saves computation: there is no mask matrix to construct or apply when the query has length 1.

Cache Memory Requirements

Understanding KV cache memory consumption is needed for deployment planning, as the cache often becomes the primary memory bottleneck in production systems. Unlike model weights that have a fixed size once the model is loaded, the cache grows dynamically during generation and can easily exceed the model weights in memory usage for long sequences or large batches. A system that can comfortably load and run a 7B model may be unable to serve multiple concurrent users at long context lengths, not because the model is too large, but because the caches are.

Memory Formula

To derive the memory requirements, we must account for all the tensors that constitute the complete cache. For a single sequence of length TT, we need to store keys and values for every layer and every head. The basic memory formula is:

Memory=2×L×H×dh×T×bytes per element\text{Memory} = 2 \times L \times H \times d_h \times T \times \text{bytes per element}

where:

  • 22: accounts for storing both keys and values, as each requires a separate tensor
  • LL: the number of transformer layers, each maintaining its own independent cache
  • HH: the number of attention heads per layer
  • dhd_h: the dimension of each head, determining the size of individual key and value vectors
  • TT: the sequence length in tokens, the dynamic factor that grows during generation
  • bytes per element\text{bytes per element}: the memory size required for a single floating-point number (e.g., 2 bytes for FP16, 4 bytes for FP32)

Since the total model dimension is typically expressed as dmodel=H×dhd_{\text{model}} = H \times d_h, we can simplify this formula by substituting:

Memory=2×L×dmodel×T×bytes per element\text{Memory} = 2 \times L \times d_{\text{model}} \times T \times \text{bytes per element}

where:

  • dmodeld_{\text{model}}: the total model dimension, equal to H×dhH \times d_h
  • LL: the number of layers in the model
  • TT: the sequence length
  • bytes per element\text{bytes per element}: the numerical precision in bytes

This simplified formula reveals an important insight: the cache memory scales with the product of model depth and width, multiplied by sequence length. Doubling any of these factors doubles the memory requirement. This linear relationship is the cache's great advantage over the naive quadratic scheme, but it still produces very large absolute values for modern models.

Example: LLaMA 2 7B

To make these formulas concrete, let's calculate the cache requirements for a widely deployed model. The LLaMA 2 7B model has the following specifications:

  • L=32L = 32 layers
  • dmodel=4096d_{\text{model}} = 4096
  • Typically stored in FP16 (2 bytes per element)

For a single sequence of length T=2048T = 2048, we can compute the cache size step by step:

Memory=2×32×4096×2048×2=1,073,741,824 bytes=1 GB\begin{aligned} \text{Memory} &= 2 \times 32 \times 4096 \times 2048 \times 2 \\ &= 1,073,741,824 \text{ bytes} \\ &= 1 \text{ GB} \end{aligned}

A single 2048-token sequence requires a full gigabyte just for the KV cache. This is separate from the memory needed for model weights and activations, along with the rest of the inference overhead.

For a batch of 8 sequences at the 4096-token context length, the memory requirement grows proportionally:

Memory=8×2×32×4096×4096×2=17,179,869,184 bytes=16 GB\begin{aligned} \text{Memory} &= 8 \times 2 \times 32 \times 4096 \times 4096 \times 2 \\ &= 17,179,869,184 \text{ bytes} \\ &= 16 \text{ GB} \end{aligned}

For comparison, the model weights themselves require about 14 GB in FP16. This means that for a batch of 8 long-context requests, the KV cache consumes more memory than the entire model. At longer context lengths or larger batch sizes, the cache can substantially exceed the model size, making memory management the necessary challenge for deployment.

Scaling with Modern Models

Larger models and longer contexts exacerbate the memory pressure, creating significant challenges for deployment at scale. The following table shows KV cache sizes for various model configurations, illustrating how the requirements grow across different architectural choices:

KV cache memory requirements for different model sizes at their respective context lengths in FP16 precision.
ModelLayersdmodeld_{\text{model}}ContextCache per Sequence
LLaMA 7B3240964K2 GB
LLaMA 13B4051204K3.1 GB
LLaMA 70B8081924K10.0 GB
GPT-4 (estimated)12012288128K590 GB

The cache memory scales linearly with both model size (via the L×dmodelL \times d_{\text{model}} product) and context length. This linear scaling in context length, combined with the quadratic scaling of the attention computation itself, makes long-context inference particularly challenging. A model like GPT-4 with a 128K context window requires hundreds of gigabytes per sequence, necessitating distributed systems and advanced memory management strategies. Understanding this table should give you a visceral sense of why every modern LLM serving system devotes significant engineering effort to cache optimization.

Out[5]:
Visualization
Log-log line chart of KV cache memory in GB versus context length in thousands of tokens for LLaMA 7B, 13B, and 70B. Dashed gray reference lines mark 24 GB and 80 GB GPU limits.
KV cache memory scaling with context length for LLaMA 7B, 13B, and 70B models, plotted on a log-log scale. Memory grows linearly with context length (appearing as a straight line on this scale), and larger models require proportionally more cache memory. The dashed horizontal lines indicate typical GPU memory capacities. LLaMA 70B reaches 80 GB (the A100/H100 limit) at 32K tokens and exceeds it at longer contexts, making single-GPU serving beyond that point impossible without memory optimizations.

Cache Management

Efficient cache management involves necessary decisions about memory allocation strategies, handling variable sequence lengths across concurrent requests, and managing batched generation effectively. These operational considerations often determine whether a deployment can reach its throughput and latency targets. Getting the compute right is necessary but not sufficient; you also need to get the memory management right.

Think of cache management like hotel room allocation. Model weights are the building itself: fixed overhead that is always there. Each user request is a guest who needs a room (cache memory) for the duration of their stay (generation). Some guests stay for 10 minutes (short generations), others for hours (long generations). If you pre-assign the biggest possible room to everyone, you waste capacity. If you assign rooms dynamically but poorly, you get fragmentation where no contiguous block is large enough for the next guest even though total free memory is sufficient.

Static vs Dynamic Allocation

Static allocation pre-allocates cache tensors for the maximum sequence length at the start of generation. This approach reserves all memory upfront, avoiding the overhead of repeated memory allocation during the generation process. Static allocation is preferred when the maximum sequence length is known in advance, when memory fragmentation must be avoided, and when consistent, predictable latency is required. The downside is waste: a request that only generates 100 tokens uses the same amount of pre-allocated memory as one that generates 4,096 tokens.

Dynamic allocation grows the cache as needed, typically by allocating larger tensors and copying existing data when the current allocation is exhausted. This saves memory for shorter sequences but introduces allocation overhead and potential fragmentation. Modern frameworks often use a hybrid approach, allocating in chunks of 128 or 256 tokens at a time to balance these tradeoffs. The chunk size is a compromise between allocation frequency (smaller chunks mean more frequent allocations) and waste (larger chunks mean more potential waste at the end).

The management challenge compounds when serving multiple requests concurrently. A request that finishes early frees its cache memory, but that freed memory may be fragmented: a 500-token hole here, a 200-token hole there. A new incoming request wanting 1,000 tokens cannot use these fragments even though total free memory is sufficient. This is the classical memory fragmentation problem, and it is why vLLM's Paged Attention, which we will cover in the next chapter, attracted such attention: it applies the operating system's paging approach to eliminate fragmentation entirely.

Batched Generation

When generating for multiple sequences simultaneously, cache management becomes more complex because different sequences may have different lengths. The typical approaches each involve different tradeoffs:

  • Padding: Allocate cache based on the longest sequence in the batch, padding shorter sequences. Simple to implement and preserves batched matrix operations, but wasteful when sequence lengths vary materially. A batch where one sequence has 4,096 tokens and five others have 50 tokens wastes nearly all the allocated memory for those five.
  • Separate caches: Maintain independent cache tensors per sequence. Avoids wasted memory but complicates the attention computation and prevents batched matrix operations, hurting GPU utilization.
  • Paged attention: Allocate cache in fixed-size blocks and track which blocks belong to which sequence using a block table (in effect a page table). This approach lets efficient memory utilization with variable-length sequences and is the basis of the vLLM serving system.

Cache Clearing and Context Window Management

When a sequence reaches the model's maximum context length, the system must decide how to proceed. Three main strategies exist, each with different tradeoffs:

Truncation simply stops accepting new tokens or drops the oldest tokens from the sequence. The model loses access to early context, which can cause coherence problems in long conversations.

Sliding window maintains only the most recent WW tokens, discarding older ones. This is the approach used in Mistral's architecture, as we discussed in Part XXIX. The advantage is a constant memory footprint; the disadvantage is that the model cannot reference information from the distant past.

Attention sinks keeps initial tokens (which empirically receive disproportionate attention regardless of content, the "sink" phenomenon we explored in Part XVIII) plus a window of recent tokens. This preserves the model's numerical stability while using a bounded cache.

Each strategy has tradeoffs between memory usage, coherence over long conversations, and computational complexity. Production systems often need to support multiple strategies and allow users or operators to choose based on their application requirements.

Implementation

Let's build a simple KV cache implementation to see how these concepts work in practice. The implementation below prioritizes clarity over production-level efficiency, but it captures all the needed logic.

In[6]:
Code
import torch


class KVCache:
    """
    Simple KV cache for a single attention layer.
    """

    def __init__(
        self,
        batch_size: int,
        num_heads: int,
        head_dim: int,
        max_seq_len: int,
        dtype=torch.float32,
        device="cpu",
    ):
        self.batch_size = batch_size
        self.num_heads = num_heads
        self.head_dim = head_dim
        self.max_seq_len = max_seq_len

        # Pre-allocate cache tensors
        cache_shape = (batch_size, num_heads, max_seq_len, head_dim)
        self.k_cache = torch.zeros(cache_shape, dtype=dtype, device=device)
        self.v_cache = torch.zeros(cache_shape, dtype=dtype, device=device)

        # Track current sequence length
        self.seq_len = 0

    def update(self, k_new: torch.Tensor, v_new: torch.Tensor) -> tuple:
        """
        Append new key-value pairs to cache and return full cache.

        Args:
            k_new: New keys of shape (batch, num_heads, new_len, head_dim)
            v_new: New values of shape (batch, num_heads, new_len, head_dim)

        Returns:
            Tuple of (all_keys, all_values) including new additions
        """
        new_len = k_new.shape[2]

        # Store new keys and values
        self.k_cache[:, :, self.seq_len : self.seq_len + new_len, :] = k_new
        self.v_cache[:, :, self.seq_len : self.seq_len + new_len, :] = v_new

        self.seq_len += new_len

        # Return only the valid portion of the cache
        return (
            self.k_cache[:, :, : self.seq_len, :],
            self.v_cache[:, :, : self.seq_len, :],
        )

    def get_seq_len(self) -> int:
        return self.seq_len

The KVCache class pre-allocates tensors at initialization and uses an integer counter to track how much of the pre-allocated space has been used. The update method writes new entries into the pre-allocated buffer using slice assignment, which is an in-place operation that avoids memory allocation during the hot loop of generation. Returning only the valid slice [:, :, :self.seq_len, :] ensures the attention computation only sees actual cached entries, not the zero-padding at the end.

In[7]:
Code
from __future__ import annotations

import math
from typing import TYPE_CHECKING

import torch
import torch.nn as nn
import torch.nn.functional as F

if TYPE_CHECKING:
    pass


class CausalSelfAttentionWithCache(nn.Module):
    """
    Causal self-attention module that supports KV caching.
    """

    def __init__(self, d_model: int, num_heads: int):
        super().__init__()
        assert d_model % num_heads == 0

        self.d_model = d_model
        self.num_heads = num_heads
        self.head_dim = d_model // num_heads

        # Combined QKV projection for efficiency
        self.qkv_proj = nn.Linear(d_model, 3 * d_model, bias=False)
        self.out_proj = nn.Linear(d_model, d_model, bias=False)

        self.scale = 1.0 / math.sqrt(self.head_dim)

    def forward(
        self,
        x: torch.Tensor,
        kv_cache: "KVCache | None" = None,
        use_cache: bool = False,
    ) -> "tuple[torch.Tensor, KVCache | None]":
        """
        Forward pass with optional KV caching.

        Args:
            x: Input tensor of shape (batch, seq_len, d_model)
            kv_cache: Optional KVCache object for incremental decoding
            use_cache: Whether to use/update the cache

        Returns:
            Tuple of (output, updated_cache)
        """
        batch_size, seq_len, _ = x.shape

        # Compute Q, K, V projections
        qkv = self.qkv_proj(x)
        qkv = qkv.reshape(batch_size, seq_len, 3, self.num_heads, self.head_dim)
        qkv = qkv.permute(2, 0, 3, 1, 4)  # (3, batch, heads, seq, head_dim)
        q, k, v = qkv[0], qkv[1], qkv[2]

        # Handle KV cache
        if use_cache and kv_cache is not None:
            k, v = kv_cache.update(k, v)
            cache_len = kv_cache.get_seq_len()
        else:
            cache_len = seq_len

        # Compute attention scores
        # q: (batch, heads, seq_len, head_dim)
        # k: (batch, heads, cache_len, head_dim)
        scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale

        # Apply causal mask
        # For cached attention: new queries can attend to all cached keys
        if seq_len == 1 and cache_len > 1:
            # Single token attending to full cache - no masking needed
            pass
        else:
            # Create causal mask for the needed portion
            mask = torch.triu(
                torch.ones(
                    seq_len, cache_len, device=x.device, dtype=torch.bool
                ),
                diagonal=cache_len - seq_len + 1,
            )
            scores = scores.masked_fill(
                mask.unsqueeze(0).unsqueeze(0), float("-inf")
            )

        attn_weights = F.softmax(scores, dim=-1)

        # Apply attention to values
        out = torch.matmul(attn_weights, v)

        # Reshape and project output
        out = out.permute(0, 2, 1, 3).reshape(batch_size, seq_len, self.d_model)
        out = self.out_proj(out)

        return out, kv_cache

The combined QKV projection (qkv_proj) computes all three projections in a single matrix multiply, which is faster than three separate operations on modern hardware due to better memory bandwidth utilization. The reshaped qkv tensor is then split into separate query, key, and value tensors before the cache update.

Now let's show how the cache accelerates generation:

In[8]:
Code
# Set up a small model setup
d_model = 64
num_heads = 4
batch_size = 1
max_seq_len = 32

# Create attention module
attention = CausalSelfAttentionWithCache(d_model, num_heads)

# Simulate a prompt of 8 tokens
prompt_len = 8
prompt = torch.randn(batch_size, prompt_len, d_model)

# PREFILL PHASE: Process entire prompt at once
cache = KVCache(batch_size, num_heads, d_model // num_heads, max_seq_len)
output, cache = attention(prompt, cache, use_cache=True)
Out[9]:
Console
After prefill phase:
  Cache sequence length: 8
  K cache shape (used portion): torch.Size([1, 4, 8, 16])

The output confirms that the cache has been initialized with the 8 prompt tokens. The key cache shape shows we have stored the projection results for these tokens, ready to be attended to by future generation steps.

In[10]:
Code
# DECODE PHASE: Generate tokens one at a time
num_new_tokens = 5
generation_outputs = []

for step in range(num_new_tokens):
    # Simulate embedding of newly generated token
    new_token = torch.randn(batch_size, 1, d_model)

    # Process only the new token, reusing cached K,V
    output, cache = attention(new_token, cache, use_cache=True)
    generation_outputs.append(output)
Out[11]:
Console

After generating 5 tokens:
  Cache sequence length: 13
  Total tokens processed: 13
  K cache used: torch.Size([1, 4, 13, 16])

The cache has grown by 5 tokens, now containing the history for 13 tokens total.

Let's compare the computational savings. The important detail is that we only processed the 5 new tokens through the model's projection layers, yet the attention mechanism had access to the full 13-token history via the cache.

In[12]:
Code
def count_qkv_projections(
    prompt_len: int, gen_len: int, use_cache: bool
) -> dict:
    """
    Count Q, K, V projection operations with and without caching.
    """
    if use_cache:
        # Prefill: compute Q, K, V for all prompt tokens
        prefill_projections = prompt_len * 3  # Q, K, V for each token

        # Decode: compute Q, K, V only for new tokens
        decode_projections = gen_len * 3

        total = prefill_projections + decode_projections
    else:
        # Without cache: recompute all projections at each step
        total = 0
        for t in range(1, gen_len + 1):
            seq_len = prompt_len + t
            total += seq_len * 3  # Q, K, V for entire sequence

    return {"total_projections": total, "with_cache": use_cache}


# Compare for realistic generation scenario
prompt_len = 100
gen_len = 200

with_cache = count_qkv_projections(prompt_len, gen_len, use_cache=True)
without_cache = count_qkv_projections(prompt_len, gen_len, use_cache=False)

speedup = without_cache["total_projections"] / with_cache["total_projections"]
Out[13]:
Console
Generating 200 tokens from 100-token prompt:
  With KV cache: 900 projection operations
  Without cache: 120,300 projection operations
  Reduction factor: 133.7x fewer operations

The savings are substantial. For generation, the projection computation savings grow quadratically with sequence length.

Out[14]:
Visualization
Single-row heatmap with 13 columns (positions 0-12) showing attention weights for a new token during decode. Blue color scale indicates weight magnitude.
Attention weights for a new token (position 12) attending to all cached positions (0-12) during decode. The new token computes attention scores against the entire cache, letting it gather needed information from the full context. Darker blue cells receive more attention weight. The orange dashed line separates prompt positions (0-7) from generated positions (8-12).

Verifying Cache Correctness

A necessary property of KV caching is that it must produce identical outputs to non-cached attention. This guarantee makes caching a valid optimization rather than an approximation. If caching changed the output even slightly, every generated token would be slightly different, and those small errors would compound over the course of a generation, creating wildly different sequences. Let's verify this mathematically and empirically.

The algebraic argument is straightforward. In standard full-sequence attention, the output for position tt in a sequence of length TT is:

ot=softmax(qtK[1:t]Tdk)V[1:t]o_t = \text{softmax}\left(\frac{q_t K[1:t]^T}{\sqrt{d_k}}\right) V[1:t]

where K[1:t]K[1:t] denotes the submatrix of KK containing only the first tt rows. In cached attention, we have:

ot=softmax(qtKcacheTdk)Vcacheo_t = \text{softmax}\left(\frac{q_t K_{\text{cache}}^T}{\sqrt{d_k}}\right) V_{\text{cache}}

where KcacheK_{\text{cache}} at step tt contains exactly the first tt rows of KK. Since ki=xiWKk_i = x_i W_K is computed identically whether we process the full sequence at once or token-by-token, the two expressions are algebraically identical. The numerical values in KcacheK_{\text{cache}} are bitwise identical to the corresponding rows of KK, so the outputs are numerically equivalent up to floating-point precision.

In[15]:
Code
def verify_cache_equivalence():
    """
    Verify that cached and non-cached attention produce identical outputs.
    """
    torch.manual_seed(42)

    d_model = 64
    num_heads = 4
    batch_size = 2

    attention = CausalSelfAttentionWithCache(d_model, num_heads)

    # Create a sequence
    prompt = torch.randn(batch_size, 5, d_model)
    next_tokens = torch.randn(batch_size, 3, d_model)
    full_sequence = torch.cat([prompt, next_tokens], dim=1)

    # Method 1: Process entire sequence without cache
    output_no_cache, _ = attention(full_sequence, use_cache=False)

    # Method 2: Process with cache (prefill + decode)
    cache = KVCache(batch_size, num_heads, d_model // num_heads, max_seq_len=32)

    # Prefill with prompt
    output_prefill, cache = attention(prompt, cache, use_cache=True)

    # Decode remaining tokens one at a time
    cached_outputs = [output_prefill]
    for i in range(next_tokens.shape[1]):
        token = next_tokens[:, i : i + 1, :]
        output_step, cache = attention(token, cache, use_cache=True)
        cached_outputs.append(output_step)

    output_with_cache = torch.cat(cached_outputs, dim=1)

    return output_no_cache, output_with_cache


output_no_cache, output_with_cache = verify_cache_equivalence()

# Compare outputs
max_diff = (output_no_cache - output_with_cache).abs().max().item()
are_equal = torch.allclose(output_no_cache, output_with_cache, atol=1e-5)
Out[16]:
Console
Maximum difference between cached and non-cached: 1.19e-07
Outputs are equivalent: True

The numerical equivalence confirms that caching is purely an optimization: it changes how we compute the result, not what we compute. The tiny residual difference (typically around 10−710^{-7}) shows floating-point rounding from the different order of operations, not a systematic error.

Profiling Memory Usage

Let's measure actual memory consumption for different configurations:

In[17]:
Code
def calculate_cache_memory(
    num_layers: int,
    d_model: int,
    max_seq_len: int,
    batch_size: int,
    dtype_bytes: int = 2,  # FP16
) -> dict:
    """
    Calculate KV cache memory requirements.
    """
    # Memory per layer: 2 (K and V) x batch x seq_len x d_model x bytes
    memory_per_layer = 2 * batch_size * max_seq_len * d_model * dtype_bytes
    total_memory = num_layers * memory_per_layer

    return {
        "memory_per_layer_mb": memory_per_layer / (1024**2),
        "total_memory_mb": total_memory / (1024**2),
        "total_memory_gb": total_memory / (1024**3),
    }


# Model configurations
models = {
    "LLaMA-7B": {"layers": 32, "d_model": 4096},
    "LLaMA-13B": {"layers": 40, "d_model": 5120},
    "LLaMA-70B": {"layers": 80, "d_model": 8192},
}

# Calculate for different context lengths
context_lengths = [2048, 4096, 8192, 16384]
batch_size_1 = 1
memory_table_1 = []

for name, config in models.items():
    row = []
    for ctx in context_lengths:
        mem = calculate_cache_memory(
            config["layers"], config["d_model"], ctx, batch_size_1
        )
        row.append(mem["total_memory_gb"])
    memory_table_1.append((name, row))

# Calculate for different batch sizes (Context 4096)
target_ctx = 4096
batch_sizes = [1, 8, 32]
memory_table_2 = []

for name, config in models.items():
    row = []
    for bs in batch_sizes:
        mem = calculate_cache_memory(
            config["layers"], config["d_model"], target_ctx, bs
        )
        row.append(mem["total_memory_gb"])
    memory_table_2.append((name, row))
Out[18]:
Console
KV Cache Memory (GB) - Batch Size 1, FP16

Model        |   2048 |   4096 |   8192 |  16384
-------------------------------------------------------
LLaMA-7B     |   1.0G |   2.0G |   4.0G |   8.0G |
LLaMA-13B    |   1.6G |   3.1G |   6.2G |  12.5G |
LLaMA-70B    |   5.0G |  10.0G |  20.0G |  40.0G |

As context length increases, memory usage grows linearly. For the LLaMA-70B model, a single sequence at 8192 context requires 20 GB of cache memory, which is significant even for high-end hardware.

Out[19]:
Console


KV Cache Memory (GB) - Context 4096, FP16

Model        | Batch 1 | Batch 8 | Batch 32
--------------------------------------------------
LLaMA-7B     |    2.0G |   16.0G |   64.0G |
LLaMA-13B    |    3.1G |   25.0G |  100.0G |
LLaMA-70B    |   10.0G |   80.0G |  320.0G |

These numbers reveal why KV cache memory management is necessary for production systems. A LLaMA-70B model serving 32 concurrent requests at 4K context requires 320 GB just for the KV cache, exceeding what most single GPUs can give.

Out[20]:
Visualization
Line chart of KV cache memory in GB versus batch size (1-32) at 4096 context length for LLaMA 7B, 13B, and 70B. All three lines increase linearly.
KV cache memory scaling with batch size at fixed 4096 context length in FP16 for LLaMA 7B, 13B, and 70B models. All three curves increase linearly with batch size. The dashed horizontal reference line at 80 GB marks the capacity limit of an A100/H100 GPU. LLaMA 70B reaches 80 GB at batch size 8 and exceeds it beyond that point, meaning a single 80 GB GPU cannot serve more than eight concurrent 4K-context requests with this model.
Out[21]:
Visualization
Grouped bar chart comparing KV cache memory in GB at batch sizes 1, 8, and 16 for LLaMA 7B, 13B, and 70B at 4096 context length.
Comparison of KV cache memory versus model weights for batch sizes 1, 8, and 16 at 4096 context length. Bars show cache memory; dashed horizontal lines mark the model weight sizes for reference. For LLaMA 70B at batch size 16, the 160 GB cache exceeds the 140 GB model-weight footprint, and at batch size 32 (not shown) it grows to 320 GB. This comparison illustrates that cache memory, not model weight memory, often becomes the bottleneck in production inference.

Key Parameters

Understanding each parameter's role helps you reason about deployment decisions when configuring inference servers:

  • d_model: The dimensionality of the model's hidden states. This appears directly in the memory formula and in the projection computation cost.
  • num_heads: The number of attention heads. Along with head_dim, it determines how the model dimension is partitioned. More heads mean more granular attention patterns but the same total cache size (since H×dh=dmodelH \times d_h = d_{\text{model}}).
  • head_dim: The dimension of each attention head, equal to dmodel/num_headsd_{\text{model}} / \text{num\_heads}. Smaller head dimensions mean more heads for the same model width.
  • max_seq_len: The maximum sequence length the cache can store. This determines the upper bound on context length and the maximum cache memory for a given model.
  • batch_size: The number of sequences processed simultaneously. Larger batches improve GPU utilization and throughput but multiply cache memory linearly.

Limitations and Practical Considerations

KV caching introduces several challenges that drive ongoing research in inference optimization. These limitations directly motivate major inference optimization techniques that have emerged since 2020, from grouped-query attention to paged attention to speculative decoding.

The most basic limitation is that memory consumption scales linearly with sequence length. While this is better than the quadratic scaling we would face without caching, it still creates a hard constraint on context length and batch size. A server with 80 GB of GPU memory might fit the model weights comfortably at 14 GB, but may be unable to serve multiple long-context requests concurrently because the caches fill the remaining 66 GB faster than expected. This tension between throughput and context length is basic to LLM deployment and not solvable by any single optimization. The KV cache memory is the single most important factor determining how many requests a given GPU can serve simultaneously. Every percentage reduction in cache memory directly translates to higher concurrency and lower cost per query.

Memory fragmentation becomes increasingly problematic as the number of concurrent requests grows. When different sequences in a batch have different lengths, naive approaches either waste memory through padding or sacrifice batching efficiency. Consider a batch where one sequence is actively using 3,000 tokens of cache and another finished early and freed 2,000 tokens. A new incoming request wanting 2,500 tokens cannot use the freed space because it is not contiguous. The server either rejects the request, queues it until more contiguous space is available, or wastes memory by padding smaller requests to the maximum. Production systems must carefully manage cache allocation to maximize GPU utilization. The Paged Attention chapter (coming next) addresses this with a memory management approach inspired by operating system virtual memory, dividing the cache into fixed-size pages that can be assigned non-contiguously.

The cache must persist across the entire generation process. Unlike model weights that are read-only during inference, the KV cache is constantly growing and being read from. This means the cache cannot be easily offloaded to CPU during generation without incurring significant latency for memory transfers (CPU-GPU bandwidth is typically 10-30 GB/s, while GPU memory bandwidth is 1-3 TB/s, a gap of two orders of magnitude). Systems that need to handle many concurrent requests must carefully orchestrate which caches are active on GPU at any moment. Some systems implement "cache swapping" that moves inactive caches to CPU memory between generation steps, accepting the bandwidth cost in exchange for higher concurrency. This tradeoff is only viable when GPU processing time per step is long enough to hide the transfer latency.

A fourth, less obvious limitation is that the decode phase becomes memory-bandwidth-bound rather than compute-bound. During decode, each step processes a single token's worth of activations through the model's projection matrices. Modern GPUs have large FLOPs but moderate memory bandwidth. The weight matrices for a large model contain billions of parameters, and at each decode step the model must read in effect all of them from DRAM to compute a single output token. For LLaMA-2 7B, this means reading ~14 GB of weights and up to several GB of cache per generated token. The GPU's arithmetic units are starved waiting for memory, resulting in low hardware utilization. This is why batching is so important: serving 32 concurrent requests at once amortizes the weight-reading cost over 32 tokens, sharply improving utilization. The cache then becomes the bottleneck that limits how large a batch can be.

Grouped-Query Attention (GQA), which we discussed in Part XXIX, directly addresses the cache size problem by having multiple query heads share a single key-value head. For example, if 8 query heads share 1 KV head, the cache is 8 times smaller while the model's representational capacity for queries is unchanged. LLaMA-2 70B uses GQA with a 2:1 reduction (8 KV heads for 64 query heads), cutting the cache size by 8 times. This architectural choice was motivated specifically by the memory constraints we have described here: making the cache smaller is the most direct way to let larger batches or longer contexts.

Summary

The KV cache is a basic optimization that makes autoregressive generation practical. By storing and reusing the key and value projections from attention, we avoid recomputing the same values at every generation step. This reduces the computational overhead from quadratic to linear in the number of generated tokens, and the reduction is not an approximation but an algebraically exact reformulation of the same computation.

The key concepts covered in this chapter are:

  • Redundancy problem: Without caching, each generation step recomputes KK and VV for all previous tokens, yielding O(T2)O(T^2) total projection operations for a TT-token generation
  • Cache structure: Separate KK and VV tensors per layer per head, growing with sequence length, storing projections for all processed tokens
  • Prefill vs decode: The initial prompt processing (prefill) populates the cache in a single parallel pass, while subsequent generation steps (decode) each process only one token at a time
  • Memory requirements: Cache size scales as 2×L×dmodel×T×bytes2 \times L \times d_{\text{model}} \times T \times \text{bytes}, easily reaching multiple gigabytes for long sequences with modern models
  • Cache management: Decisions about static vs dynamic allocation, batching strategies, and context length limits materially impact deployment efficiency
  • Limitations: Memory pressure, fragmentation, bandwidth bottlenecks, and the memory-bandwidth-bound decode phase motivate a detailed ecosystem of follow-on optimizations

Understanding KV cache mechanics is needed for working with modern LLMs, as cache memory often becomes the primary constraint on throughput and context length. The next chapter examines KV cache memory arithmetic in greater detail, including how grouped-query attention changes the calculation. Following that, Paged Attention introduces an operating-system-inspired approach to cache memory management that lets the vLLM serving system's large throughput improvements. After that, we explore techniques for compressing the cache through quantization and eviction to extend context lengths further, all building directly on the foundation established here.

Quiz

Ready to test your understanding? Take this quick quiz to reinforce what you've learned about KV cache optimization in transformer inference.

KV Cache

Question 1 of 70 of 7 completed
Without KV caching, how does the number of key-value projection operations scale with generating T tokens?

Comments

No comments yet. Be the first to share your thoughts!

Reference

Citation details

Cite or share this article.

BIBTEXAcademic
@misc{brenndoerfer2026kvcache, author = {Michael Brenndoerfer}, title = {KV Cache Explained: Efficient Attention for LLM Generation}, year = {2026}, url = {https://mbrenndoerfer.com/writing/kv-cache-transformer-attention-optimization}, organization = {mbrenndoerfer.com}, note = {Accessed: 2026-09-27} }
APAAcademic
Michael Brenndoerfer (2026). KV Cache Explained: Efficient Attention for LLM Generation. Retrieved from https://mbrenndoerfer.com/writing/kv-cache-transformer-attention-optimization
MLAAcademic
Michael Brenndoerfer. "KV Cache Explained: Efficient Attention for LLM Generation." 2026. Web. September 27, 2026. <https://mbrenndoerfer.com/writing/kv-cache-transformer-attention-optimization>.
CHICAGOAcademic
Michael Brenndoerfer. "KV Cache Explained: Efficient Attention for LLM Generation." Accessed September 27, 2026. https://mbrenndoerfer.com/writing/kv-cache-transformer-attention-optimization.
HARVARDAcademic
Michael Brenndoerfer (2026) 'KV Cache Explained: Efficient Attention for LLM Generation'. Available at: https://mbrenndoerfer.com/writing/kv-cache-transformer-attention-optimization (Accessed: September 27, 2026).
SimpleBasic
Michael Brenndoerfer (2026). KV Cache Explained: Efficient Attention for LLM Generation. https://mbrenndoerfer.com/writing/kv-cache-transformer-attention-optimization

About the author

Continue with the full handbook

This chapter is part of Language AI Handbook. Use the handbook page to browse the complete table of contents and continue reading in sequence.

Explore Language AI Handbook
Newsletter

Stay up to date

Get articles, book updates, and news delivered to your inbox.

No spam, unsubscribe anytime.

or

Join the community

Sign in to remove popups, track your reading progress, and join the discussion.