Introduction
In our last lesson, you implemented several KV cache eviction strategies, optimizing how we manage memory within a single, long-running generation. We saw a direct trade-off between memory footprint and model quality, with techniques like StreamingLLM offering a strong balance.
Now, we shift our focus from optimizing a single request to optimizing a stream of multiple, concurrent requests. This is a more realistic production scenario. Consider a chatbot application: every new conversation starts with the same system prompt. Or a few-shot learning task where every request includes the same set of examples. Re-computing the KV cache for these common prefixes for every single request is a significant source of wasted computation.
This lesson addresses that inefficiency head-on. Your learning outcome is to implement prefix caching to reuse KV cache entries for requests that share a common prompt. You will learn the concept, explore the data structures that enable it, and implement a functional prefix cache to measure the performance gains yourself. This is a fundamental technique for building high-throughput LLM serving systems.
The Opportunity: Reusing Prefixes
In many LLM applications, requests are not entirely unique. They often share a common starting sequence, or prefix.
- Chatbots:
System: You are a helpful assistant. User: Hello. Assistant: Hi! User: [New Question]The entire chat history is a prefix for the next turn. - Few-Shot Prompting: A set of examples provided for in-context learning is a prefix for every query.
- RAG (Retrieval-Augmented Generation): The retrieved documents, added to the prompt as context, form a prefix.
Let's watch a brief clip that introduces how modern serving engines like vLLM leverage this redundancy.
Optimize LLM inference with vLLM
This short segment from a video on vLLM provides a concise explanation of the concept of prefix caching.
Watch from 03:01 to 03:30. The video explains how identifying and reusing common prefixes across different queries streamlines processing.
The key idea is simple: if we have already computed the KV cache for the token sequence "What is the capital of", we should not re-compute it for a new request that also starts with that phrase. By reusing the pre-computed KV cache, we can skip the expensive prefill stage for that shared portion and immediately start generating the new tokens.
From Simple Reuse to a General-Purpose System: RadixAttention
While the concept is simple, managing potentially thousands of requests with varying shared prefixes requires a more sophisticated data structure than a simple dictionary. An ideal system would automatically detect and share any common prefix among all active requests.
This is precisely what SGLang accomplishes with a technique it calls RadixAttention. It manages the entire KV cache of all requests in a server using a radix tree (also known as a prefix tree or trie). Given your computer science background, you'll recognize this as a highly efficient data structure for prefix-based lookups.
To understand how this works in detail, let's turn to a video from one of SGLang's creators.
Efficient LLM Inference with SGLang, Lianmin Zheng, xAI
Lianmin Zheng, one of the creators of SGLang, provides a deep dive into the motivation and mechanics of RadixAttention. This will give you a systems-level view of prefix caching.
Please watch from 05:40 to 13:09. Focus on these key points: Use Cases (05:40 - 10:06): Pay attention to the examples of multi-turn chat, few-shot learning, and tree search. Notice how they all create prefix-sharing opportunities. The Core Idea (10:06 - 10:51): Understand why existing systems discard the KV cache and how SGLang's approach of retaining it in a radix tree is the key innovation. Building the Radix Tree (10:51 - 13:09): Follow the step-by-step visual example of how the tree is built. Observe how new requests can match existing paths (cache hit), cause a split to create a new branch, and how eviction is handled.
The radix tree provides an elegant solution. Each path from the root represents a unique sequence of tokens, and the KV cache for that sequence is stored along the path.

By managing the KV cache in this way, the system can achieve:
- Automatic Prefix Sharing: Any new request can be matched against the tree to find the longest possible prefix to reuse.
- Fine-Grained Caching: Since the tree is token-based, sharing can happen at any level of granularity.
- Efficient Memory Management: The tree structure allows for efficient eviction policies, as we saw in the video.
To bridge the gap between this high-level concept and a concrete implementation, let's look at the underlying data structure from the SGLang source code.
SGLang Deep Dive: Inside SGLang
The SugiV Blog provides an excellent deep dive into SGLang's source code. We'll look at the data structure for the tree nodes and the core matching algorithm.
Please review the following two sections: Find the TreeNode implementation under the heading 'Radix Tree Data Structure: The Foundation'. Examine the attributes of the class (children, parent, value for the KV cache indices, etc.). This is the concrete representation of the nodes in the diagram you just saw. Next, look at the code block under 'Prefix Matching Algorithm: The Heart of RadixAttention'. You don't need to parse every line, but understand its purpose: it's a tree traversal algorithm that finds the longest common prefix between a new request and the sequences already in the cache.
Hands-On: Implementing and Measuring Prefix Caching
While a full radix tree is complex, we can implement the core logic of prefix caching in a more direct way to prove its effectiveness. We will simulate a common scenario: a server that uses a fixed system prompt for all requests.
Our strategy will be:
- Pre-compute: Perform a single prefill pass for the shared system prompt and save the resulting KV cache.
- Reuse: For each incoming user query, load the pre-computed cache and start generation from there.
- Measure: Compare the time taken with and without prefix caching.
This approach is sometimes called "Cache-Augmented Generation" (CAG), especially when the prefix is a large knowledge base. Let's look at a practical guide that uses the Hugging Face transformers library to do exactly this.
Cache-Augmented Generation (CAG) from Scratch
This article, 'Cache-Augmented Generation (CAG) from Scratch', provides a clear, hands-on tutorial for pre-loading a KV cache. While it's framed as 'loading knowledge,' the technique is identical to prefix caching.
Skim through the article, focusing on the code blocks. Pay attention to: The preprocess_knowledge function: This is our 'pre-compute' step. It runs the model on a prompt and returns the past_key_values. The prepare_kvcache function: This wraps the previous function and formats the prompt. The generate function and the final query example: Notice how the pre-computed knowledge_cache is passed directly into the generation loop for a new query. Also, note the clean_up function, which is essential for reusing the cache for multiple, separate queries.
Now, let's implement this ourselves.
1. Setup
Copy the following code into a new Python script. It includes our minimal transformer model from the previous lesson and a timer for benchmarking.
import torch
import torch.nn as nn
import torch.nn.functional as F
import time
import copy
# --- Minimal Transformer Model (from previous lessons) ---
# Note: A real implementation requires careful handling of positional embeddings
# when using a KV cache. This is simplified for demonstration.
class CausalSelfAttention(nn.Module):
def __init__(self, d_model, num_heads, head_dim):
super().__init__()
self.d_out = d_model
self.num_heads = num_heads
self.head_dim = head_dim
self.W_query = nn.Linear(d_model, d_model, bias=False)
self.W_key = nn.Linear(d_model, d_model, bias=False)
self.W_value = nn.Linear(d_model, d_model, bias=False)
self.out_proj = nn.Linear(d_model, d_model, bias=False)
def forward(self, x, past_k=None, past_v=None, use_cache=False):
b, num_tokens, d_in = x.shape
q = self.W_query(x)
k_new = self.W_key(x)
v_new = self.W_value(x)
if use_cache and past_k is not None:
k = torch.cat([past_k, k_new], dim=1)
v = torch.cat([past_v, v_new], dim=1)
else:
k = k_new
v = v_new
# Reshape for multi-head attention
q = q.view(b, num_tokens, self.num_heads, self.head_dim).transpose(1, 2)
k_reshaped = k.view(b, k.shape[1], self.num_heads, self.head_dim).transpose(1, 2)
v_reshaped = v.view(b, v.shape[1], self.num_heads, self.head_dim).transpose(1, 2)
attn_scores = q @ k_reshaped.transpose(2, 3)
mask = torch.triu(torch.ones(num_tokens, k.shape[1]), diagonal=k.shape[1] - num_tokens + 1).bool()
attn_scores = attn_scores.masked_fill(mask, -torch.inf)
attn_weights = F.softmax(attn_scores / self.head_dim**0.5, dim=-1)
context_vec = (attn_weights @ v_reshaped).transpose(1, 2).reshape(b, num_tokens, self.d_out)
return self.out_proj(context_vec), k, v
class TransformerBlock(nn.Module):
def __init__(self, d_model, num_heads, head_dim):
super().__init__()
self.attn = CausalSelfAttention(d_model, num_heads, head_dim)
# FFN, LayerNorm etc. omitted for brevity
def forward(self, x, past_k=None, past_v=None, use_cache=False):
attn_output, k, v = self.attn(x, past_k, past_v, use_cache)
return attn_output, k, v
class GPTModel(nn.Module):
def __init__(self, vocab_size, d_model, num_layers, num_heads, head_dim):
super().__init__()
self.tok_emb = nn.Embedding(vocab_size, d_model)
# For simplicity, we omit positional embeddings in this hands-on.
self.trf_blocks = nn.ModuleList([
TransformerBlock(d_model, num_heads, head_dim) for _ in range(num_layers)
])
self.out_head = nn.Linear(d_model, vocab_size, bias=False)
self.num_layers = num_layers
def forward(self, idx, past_kv=None, use_cache=False):
x = self.tok_emb(idx)
new_kv_cache = []
for i, block in enumerate(self.trf_blocks):
past_k, past_v = (past_kv[i] if use_cache and past_kv is not None else (None, None))
x, k, v = block(x, past_k, past_v, use_cache)
new_kv_cache.append((k,v))
logits = self.out_head(x)
return logits, new_kv_cache
2. Generation Functions
Now, add the functions to handle generation, both with and without prefix caching.
def generate_no_cache(model, tokenizer, prompt):
"""Generates text without any caching."""
input_ids = tokenizer(prompt, return_tensors="pt").input_ids
start_time = time.time()
with torch.no_grad():
logits, _ = model(input_ids, use_cache=False)
prefill_time = time.time() - start_time
# Simple greedy decoding for one step for measurement
next_token = torch.argmax(logits[:, -1, :], dim=-1).unsqueeze(0)
# In a real scenario, this would be a loop. We just measure one step.
return prefill_time
def precompute_prefix_cache(model, tokenizer, prefix_prompt):
"""Performs the prefill for the prefix and returns the KV cache."""
print(f"Pre-computing cache for prefix of length {len(tokenizer(prefix_prompt).input_ids)}...")
prefix_ids = tokenizer(prefix_prompt, return_tensors="pt").input_ids
with torch.no_grad():
_, prefix_cache = model(prefix_ids, use_cache=True)
print("Pre-computation complete.")
return prefix_cache
def generate_with_prefix_cache(model, tokenizer, query, prefix_cache):
"""Generates text using the pre-computed prefix cache."""
# IMPORTANT: In a concurrent server, you would deepcopy the cache for each request
# to avoid race conditions.
request_cache = copy.deepcopy(prefix_cache)
query_ids = tokenizer(query, return_tensors="pt").input_ids
start_time = time.time()
with torch.no_grad():
# The model only processes the short query
logits, _ = model(query_ids, past_kv=request_cache, use_cache=True)
prefill_time = time.time() - start_time
# Simple greedy decoding for one step for measurement
next_token = torch.argmax(logits[:, -1, :], dim=-1).unsqueeze(0)
return prefill_time
A note on copy.deepcopy(): In a real high-performance system, repeatedly copying large tensors is inefficient. Production engines like vLLM use more sophisticated reference counting and memory management (which we'll cover soon) to share memory blocks without copying. For this exercise, deepcopy correctly isolates request states and demonstrates the logical principle.
3. Run the Comparison
Finally, let's set up the experiment and measure the difference. We'll use a long system prompt and two short, different queries.
# Mock tokenizer for demonstration
class MockTokenizer:
def __init__(self, vocab_size=1000):
self.vocab_size = vocab_size
def __call__(self, text, return_tensors="pt"):
# Super simple tokenizer: just map chars to ints
ids = [ord(c) for c in text]
if return_tensors == "pt":
return {"input_ids": torch.tensor([ids])}
return {"input_ids": ids}
# --- Configuration ---
VOCAB_SIZE = 1000
D_MODEL = 512
NUM_LAYERS = 8
NUM_HEADS = 8
HEAD_DIM = D_MODEL // NUM_HEADS
# --- Setup ---
model = GPTModel(VOCAB_SIZE, D_MODEL, NUM_LAYERS, NUM_HEADS, HEAD_DIM)
tokenizer = MockTokenizer()
# --- Define Prompts ---
system_prompt = "You are a master AI systems engineer. You are an expert in CUDA, PyTorch, and LLM serving infrastructure. You provide detailed, accurate, and concise answers to questions about performance optimization. " * 5
query1 = "Explain the concept of tensor parallelism."
query2 = "How does continuous batching improve throughput?"
full_prompt1 = system_prompt + query1
full_prompt2 = system_prompt + query2
# --- BENCHMARK 1: No Prefix Caching ---
print("--- Running without Prefix Caching ---")
t1_no_cache = generate_no_cache(model, tokenizer, full_prompt1)
print(f"Time for Request 1 (full prompt): {t1_no_cache * 1000:.2f} ms")
t2_no_cache = generate_no_cache(model, tokenizer, full_prompt2)
print(f"Time for Request 2 (full prompt): {t2_no_cache * 1000:.2f} ms")
print(f"Total Time: {(t1_no_cache + t2_no_cache) * 1000:.2f} ms\n")
# --- BENCHMARK 2: With Prefix Caching ---
print("--- Running with Prefix Caching ---")
# Step 1: Pre-compute the shared prefix (one-time cost)
start_precompute = time.time()
prefix_cache = precompute_prefix_cache(model, tokenizer, system_prompt)
precompute_time = time.time() - start_precompute
print(f"One-time cost to create prefix cache: {precompute_time * 1000:.2f} ms")
# Step 2: Reuse the cache for both requests
t1_with_cache = generate_with_prefix_cache(model, tokenizer, query1, prefix_cache)
print(f"Time for Request 1 (query only): {t1_with_cache * 1000:.2f} ms")
t2_with_cache = generate_with_prefix_cache(model, tokenizer, query2, prefix_cache)
print(f"Time for Request 2 (query only): {t2_with_cache * 1000:.2f} ms")
print(f"Total Time (excluding one-time pre-computation): {(t1_with_cache + t2_with_cache) * 1000:.2f} ms")
When you run this code, you will see a dramatic difference. The time taken for the "query only" steps with prefix caching will be significantly lower than the "full prompt" steps. This is because the model only has to perform the prefill pass on the short query, having already processed the long system prompt. The performance benefit scales directly with the length of the shared prefix.
Conclusion
In this lesson, you moved from single-request to multi-request optimization by implementing prefix caching. This is a cornerstone of all modern, high-throughput LLM inference servers.
Key Takeaways:
- Prefix Caching avoids redundant computation by reusing the KV cache for shared leading portions of prompts.
- It provides significant latency reduction and throughput improvement, especially for applications with long system prompts, few-shot examples, or multi-turn chat histories.
- Production-grade systems like SGLang use efficient data structures like radix trees to automatically find and share all possible prefixes among concurrent requests.
- The core implementation involves a pre-compute step for the shared prefix and a reuse step where new queries are processed starting from the cached state.
Preview of the Next Lesson:
We have now seen how to manage the KV cache logically—by evicting old tokens and reusing shared prefixes. But our implementation still relies on torch.cat, which involves creating new, larger tensors in memory. In a highly concurrent server, this constant allocation and de-allocation of variably sized tensors leads to severe memory fragmentation, much like frequent malloc/free calls in C. This wastes precious VRAM and can slow down the system.
In the next lesson, we will address this fundamental systems problem. You will learn to explain the memory fragmentation problem in naive KV cache allocation and the core concept of paged memory management, which is the key idea behind vLLM's PagedAttention.