How to Implement KV Cache for Efficient LLM Inference: A Complete Guide to the LLMs-from-Scratch Implementation

Implementing a KV cache stores the key and value tensors from previously processed tokens, reducing autoregressive generation complexity from O(T·d) to O(d) by recomputing only the query projection for new tokens while attending over cached K/V matrices.

During autoregressive text generation, standard transformer architectures redundantly recompute self-attention over the entire sequence history for every new token. The rasbt/LLMs-from-scratch repository demonstrates how to implement a KV cache—a mechanism that persists intermediate key and value tensors across forward passes—to eliminate redundant projections and dramatically accelerate inference while maintaining identical outputs.

Why KV Cache Eliminates Redundant Computation

In vanilla transformer inference, each generation step computes query (Q), key (K), and value (V) projections for the full sequence length. Since K and V tensors for existing tokens remain constant after initial computation, recalculating them wastes floating-point operations. By caching these tensors per layer, the per-step cost drops from O(T·d) (where T is the current sequence length and d is the hidden dimension) to O(d), as only the new token’s query requires fresh computation.

Core Components in the Repository

The implementation spans utility classes, modified attention layers, and generation helpers, all located in the pkg/llms_from_scratch/kv_cache directory.

The KVCache Storage Classes

Two variants handle different use cases. The single-sample implementation in pkg/llms_from_scratch/kv_cache/utils.py manages a list of (K, V) tuples per transformer layer, providing get(), update(), and reset() methods. For production workloads, the batched variant in pkg/llms_from_scratch/kv_cache_batched/utils.py adds a batch dimension, storing a matrix per batch element to enable parallel sequence generation.

Cache-Aware Multi-Head Attention

In pkg/llms_from_scratch/kv_cache/gpt2.py, the MultiHeadAttention class accepts a use_cache boolean parameter. When enabled, the layer concatenates newly computed K/V tensors with cached values using torch.cat, rather than recomputing projections for the entire sequence. This modification appears in the forward method (lines 30–55), where the attention mechanism attends over the concatenated [cached_K, new_K] and [cached_V, new_V] tensors.

Model-Level Cache Orchestration

The TransformerBlock and GPTModel classes in the same file manage cache state across layers. GPTModel exposes reset_kv_cache() (lines 62–86), which clears per-layer caches and resets the internal current_pos counter used for positional embeddings. This ensures correct alignment when starting new independent generations.

Architectural Flow and Implementation Steps

According to the source code in ch04/03_kv-cache/gpt_with_kv_cache.py, implementing KV caching follows a four-phase pattern.

Step 1: Cache Initialization

Instantiate the cache container before the first forward pass, specifying the number of transformer layers.

from pkg.llms_from_scratch.kv_cache.utils import KVCache

cache = KVCache(n_layers=model.cfg["n_layers"])
model.reset_kv_cache()  # Clear any stale state

Step 2: Prompt Priming

Feed the entire prompt through the model once with use_cache=True. The attention layers automatically populate the cache with K/V tensors for every token in the input sequence.

prompt_ids = torch.tensor(tokenizer.encode(prompt)).unsqueeze(0)
_ = model(prompt_ids, cache=cache, use_cache=True)

Step 3: Incremental Generation with Tensor Concatenation

For each subsequent token, the model receives only the newly generated token ID. Inside MultiHeadAttention, the implementation concatenates the cached tensors with the new token’s projections:


# Simplified logic from gpt2.py

if use_cache and cache is not None:
    k_cached, v_cached = cache.get(layer_idx)
    k = torch.cat([k_cached, k_new], dim=-2)
    v = torch.cat([v_cached, v_new], dim=-2)
    cache.update(layer_idx, k, v)

The attention scores are computed between the new token’s query and the full cached keys, then the cache is updated in-place for the next iteration.

Step 4: Cache Reset for New Sequences

When switching to a new independent prompt, call model.reset_kv_cache() to clear per-layer tensors and reset the positional embedding index (self.current_pos), preventing contamination between unrelated sequences.

End-to-End Code Example

The following runnable script demonstrates the full inference loop using the repository’s generation helpers:

import torch
import tiktoken
from pkg.llms_from_scratch.kv_cache.generate import generate_text_simple_cached
from pkg.llms_from_scratch.kv_cache.utils import KVCache
from pkg.llms_from_scratch.kv_cache.gpt2 import GPTModel

# Configuration matching GPT-2 small

cfg = {
    "vocab_size": 50257,
    "context_length": 1024,
    "emb_dim": 768,
    "n_heads": 12,
    "n_layers": 12,
    "drop_rate": 0.1,
    "qkv_bias": False,
}

model = GPTModel(cfg).eval()
tokenizer = tiktoken.get_encoding("gpt2")

# Initialize cache

cache = KVCache(n_layers=cfg["n_layers"])
model.reset_kv_cache()

# Encode prompt

prompt = "Once upon a time"
prompt_ids = torch.tensor(tokenizer.encode(prompt)).unsqueeze(0)

# Generate 50 tokens with caching enabled

output_ids = generate_text_simple_cached(
    model=model,
    idx=prompt_ids,
    max_new_tokens=50,
    context_size=cfg["context_length"],
    use_cache=True,
)

print(tokenizer.decode(output_ids.squeeze().tolist()))

Handling Batched KV Caching for Parallel Generation

For batch inference, substitute the cache class with kv_cache_batched.KVCache and specify the batch_size during construction. The attention mechanism in the batched variant retrieves cache entries via cache[layer_idx, batch_idx], maintaining separate K/V tensors for each sequence in the batch while still concatenating along the token dimension.

Critical Implementation Details

Causal Mask Management

As implemented in gpt_with_kv_cache.py (lines 78–84), the attention layer tracks self.ptr_current_pos to slice the causal mask correctly. The mask slice [ptr:ptr+Q, :K] ensures that the new query attends only to positions up to the current cache length, maintaining autoregressive constraints even as the cache grows.

Positional Embedding Alignment

When use_cache=True, GPTModel.forward generates position indices starting from self.current_pos rather than zero, incrementing this counter after each forward pass (lines 21–24 in gpt2.py). This ensures that cached tokens retain their original positional encodings while new tokens receive incremental positions.

Memory Efficiency Characteristics

The implementation allocates memory only for the new token’s K/V projections at each step, keeping memory consumption roughly linear in the number of layers (O(L·d)) rather than quadratic in sequence length. Previously cached tensors are reused by reference, minimizing unnecessary memory copies.

Summary

  • KV caching eliminates redundant computation during autoregressive generation by storing per-layer key and value tensors across forward passes.
  • The KVCache class in pkg/llms_from_scratch/kv_cache/utils.py provides simple list-based storage, while the batched variant supports parallel sequence generation.
  • The MultiHeadAttention layer concatenates new K/V tensors with cached values using torch.cat, reducing per-step complexity from O(T·d) to O(d).
  • Proper cache management requires initializing the cache, priming it with the prompt, and resetting via model.reset_kv_cache() between independent generations.
  • Positional embeddings and causal masks must align with the growing cache using self.current_pos and self.ptr_current_pos indices.

Frequently Asked Questions

What is the memory complexity trade-off of KV caching?

While KV caching reduces computational complexity from linear in sequence length to constant per step, it increases memory usage linearly with sequence length. The cache stores two tensors (K and V) per layer per token, requiring approximately 2 × L × d × T additional parameters, where L is the number of layers and T is the maximum sequence length.

How does the KV cache handle positional embeddings?

The implementation tracks self.current_pos in GPTModel to offset positional indices. When caching is enabled, the model generates position IDs starting from this counter rather than zero, ensuring that cached tokens retain their original positional encodings while newly generated tokens receive subsequent positions.

Can I use KV caching with batch size greater than one?

Yes, by using KVCache from pkg/llms_from_scratch/kv_cache_batched/utils.py. This variant maintains separate cache entries for each batch element via cache[layer_idx, batch_idx], allowing simultaneous generation of multiple sequences with independent context lengths while sharing the same model weights.

When should I reset the KV cache during inference?

Always call model.reset_kv_cache() when beginning a new independent generation sequence. Failure to reset causes the model to attend to K/V tensors from previous conversations, resulting in context contamination and degraded output quality. The repository provides this method to clear all per-layer caches and reset the positional counter atomically.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →