# How to Implement KV Caching in Llama 2 for Efficient Inference

> Discover how Llama 2 implements KV caching in its Attention class for faster inference. Learn to optimize token generation and avoid quadratic recomputation for efficient LLM performance.

- Repository: [Meta Llama/llama](https://github.com/meta-llama/llama)
- Tags: performance
- Published: 2026-03-05

---

**Llama 2 implements KV caching inside the `Attention` class of [`llama/model.py`](https://github.com/meta-llama/llama/blob/main/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`](https://github.com/meta-llama/llama/blob/main/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](https://github.com/meta-llama/llama/blob/main/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`.

```python
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](https://github.com/meta-llama/llama/blob/main/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`.

```python
keys = self.cache_k[:bsz, : start_pos + seqlen]
values = self.cache_v[:bsz, : start_pos + seqlen]

```

[Source: llama/model.py L88-L90](https://github.com/meta-llama/llama/blob/main/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.

```python
keys = repeat_kv(keys, self.n_rep)
values = repeat_kv(values, self.n_rep)

```

[Source: llama/model.py L91-L94](https://github.com/meta-llama/llama/blob/main/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`](https://github.com/meta-llama/llama/blob/main/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.

```python
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.

```python
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.

```python
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)](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)](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)](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)](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 `Attention` class in [`llama/model.py`](https://github.com/meta-llama/llama/blob/main/llama/model.py) using persistent buffers `cache_k` and `cache_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.forward` writes new projections at positions `start_pos:start_pos+seqlen` and reads the entire prefix up to `start_pos+seqlen`.
- The `repeat_kv` function 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`](https://github.com/meta-llama/llama/blob/main/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.