How to Implement KV Caching in Llama 2 for Efficient Inference
Llama 2 implements KV caching inside the Attention class of llama/model.py by storing key and value tensors in pre-allocated cache_k and cache_v buffers, enabling linear-time token generation instead of quadratic recomputation.
The meta-llama/llama repository uses an optimized transformer architecture that relies on KV caching to avoid redundant computation during autoregressive generation. When you implement KV caching in Llama 2, you leverage persistent GPU buffers that survive across forward passes, allowing the model to append new keys and values while reusing cached projections from all previous steps.
How KV Caching Works in Llama 2
The mechanism is encapsulated entirely within the Attention.forward method. During initialization, the model allocates zero-filled tensors for keys and values. During the forward pass, it writes new projections to specific positions and reads the full prefix for attention computation.
Cache Allocation and Initialization
In llama/model.py, the Attention class constructor creates cache_k and cache_v with shape (max_batch_size, max_seq_len, n_local_kv_heads, head_dim). These buffers are registered as persistent tensors on the model's device.
Source: llama/model.py L36-L45
Writing New Keys and Values
During the forward pass, the model receives a start_pos argument indicating where the new token sequence begins. The freshly computed xk and xv tensors are inserted into the cache at positions start_pos:start_pos+seqlen.
self.cache_k[:bsz, start_pos : start_pos + seqlen] = xk
self.cache_v[:bsz, start_pos : start_pos + seqlen] = xv
Source: llama/model.py L82-L88
Retrieving Cached Tensors for Attention
After writing, the model reads the entire prefix—including previously cached entries and the newly inserted ones—by slicing from the beginning up to start_pos + seqlen.
keys = self.cache_k[:bsz, : start_pos + seqlen]
values = self.cache_v[:bsz, : start_pos + seqlen]
Source: llama/model.py L88-L90
Handling Multi-Query Attention with repeat_kv
Llama 2 uses grouped-query attention where the number of KV heads is less than the number of query heads. The repeat_kv function expands the cached keys and values to match the query head count before computing attention scores.
keys = repeat_kv(keys, self.n_rep)
values = repeat_kv(values, self.n_rep)
Source: llama/model.py L91-L94
Performance Benefits of KV Caching
Implementing KV caching in Llama 2 reduces inference complexity from quadratic $O(n^2)$ to linear $O(n)$ with respect to sequence length. Without caching, each new token requires recomputing attention over the entire prefix. With caching, the model only computes projections for the new token and reuses cached tensors.
- Memory Efficiency: The cache stores only key and value projections (two tensors per layer), avoiding storage of full attention matrices.
- Latency Stability: Generation latency remains constant even for long contexts because attention computation scales linearly.
- Throughput: Batched inference benefits significantly because the cache is indexed by batch dimension, allowing parallel processing of independent sequences.
Working with the KV Cache in Practice
While the Llama.generate method in generation.py handles cache updates automatically, you can interact with the cache directly for debugging, multi-turn conversations, or optimization.
Inspecting Cache Contents
Access the cached tensors through the model layers to verify dimensions or debug cache states.
import torch
from llama.generation import Llama
# Initialize model
llama = Llama.build(
ckpt_dir="checkpoints",
tokenizer_path="tokenizer.model",
max_seq_len=2048,
max_batch_size=1,
)
# Encode prompt
prompt = "Explain quantum entanglement."
tokens = [llama.tokenizer.encode(prompt)]
# Populate cache with forward pass
_ = llama.generate(prompt_tokens=tokens, max_gen_len=0)
# Inspect first layer cache
first_attn = llama.model.layers[0].attn
print(f"Key cache shape: {first_attn.cache_k.shape}")
print(f"Value cache shape: {first_attn.cache_v.shape}")
print(f"Cache device: {first_attn.cache_k.device}")
Resetting the Cache for New Conversations
When processing multiple independent conversations in the same process, zero out the cache to prevent context leakage between sessions.
def reset_kv_cache(llama_model):
"""Zero out KV caches in all transformer layers."""
for layer in llama_model.layers:
layer.attn.cache_k.zero_()
layer.attn.cache_v.zero_()
# Usage between conversations
reset_kv_cache(llama.model)
new_output = llama.generate(
prompt_tokens=[[llama.tokenizer.encode("New topic")]],
max_gen_len=50
)
Compiling the Model with torch.compile
The KV cache implementation is compatible with PyTorch 2.0 compilation because the cache tensors are persistent module buffers.
import torch
from llama.generation import Llama
# Build model
llama = Llama.build(
ckpt_dir="checkpoints",
tokenizer_path="tokenizer.model",
max_seq_len=2048,
max_batch_size=1,
)
# Compile the transformer (preserves cache buffers)
llama.model = torch.compile(llama.model, mode="max-autotune")
# Inference proceeds normally with cached keys/values
output = llama.generate(
prompt_tokens=[[llama.tokenizer.encode("Tell me a joke.")]],
max_gen_len=20,
)
print(llama.tokenizer.decode(output[0]))
Key Files and Functions
| File | Symbol | Description |
|---|---|---|
[llama/model.py](https://github.com/meta-llama/llama/blob/main/llama/model.py) |
Attention.__init__ |
Allocates cache_k and cache_v buffers during model construction. |
[llama/model.py](https://github.com/meta-llama/llama/blob/main/llama/model.py) |
Attention.forward |
Writes new keys/values to the cache and reads the full prefix for attention computation. |
[llama/model.py](https://github.com/meta-llama/llama/blob/main/llama/model.py) |
repeat_kv |
Expands cached KV heads to match query head count for grouped-query attention. |
[llama/generation.py](https://github.com/meta-llama/llama/blob/main/llama/generation.py) |
Llama.generate |
High-level API that iteratively calls the model; cache updates happen automatically inside Attention.forward. |
Summary
- KV caching in Llama 2 is implemented inside the
Attentionclass inllama/model.pyusing persistent bufferscache_kandcache_v. - The cache is allocated once at model initialization with shape
(max_batch_size, max_seq_len, n_local_kv_heads, head_dim). - During inference,
Attention.forwardwrites new projections at positionsstart_pos:start_pos+seqlenand reads the entire prefix up tostart_pos+seqlen. - The
repeat_kvfunction handles grouped-query attention by expanding cached KV heads to match query head counts. - Users interact with the cache transparently through
Llama.generate, but can manually inspect or reset caches for multi-turn conversations.
Frequently Asked Questions
Does KV caching work automatically when using Llama.generate?
Yes. The Llama.generate method in generation.py automatically manages the KV cache by passing the current token position (start_pos) to the model's forward pass. The Attention module updates its internal cache_k and cache_v tensors without requiring explicit user intervention.
How much GPU memory does the KV cache consume?
The cache consumes 2 × max_batch_size × max_seq_len × n_local_kv_heads × head_dim × sizeof(dtype) bytes. For Llama-2-7B with max_seq_len=2048, n_local_kv_heads=32, head_dim=128, and float16 precision, this equals approximately 33.5 MB per layer, or roughly 400 MB for the full 32-layer model.
Can I disable KV caching to reduce memory usage?
Disabling the cache is not supported in the reference implementation because the Attention class always allocates cache_k and cache_v during initialization. To minimize memory, reduce max_seq_len or max_batch_size when building the model with Llama.build, which proportionally shrinks the cache buffers.
Why does the cache use grouped-query attention with repeat_kv?
Llama 2 uses grouped-query attention (GQA) to reduce memory bandwidth during inference by sharing key and value heads across multiple query heads. The repeat_kv function expands the cached KV tensors from n_local_kv_heads to n_heads by repeating along the head dimension, ensuring compatible tensor shapes for attention score computation without duplicating underlying memory.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →