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
KVCacheclass inpkg/llms_from_scratch/kv_cache/utils.pyprovides simple list-based storage, while the batched variant supports parallel sequence generation. - The
MultiHeadAttentionlayer concatenates new K/V tensors with cached values usingtorch.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_posandself.ptr_current_posindices.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →