Introduction
In our last lesson, you implemented a basic Key-Value (KV) cache and measured its linear memory growth. While the KV cache is essential for solving the quadratic compute cost of autoregressive decoding, its unbounded memory consumption presents a new bottleneck, making it impossible to handle indefinitely long sequences.
This lesson directly tackles that problem. Your learning outcome is to implement and compare different KV cache eviction strategies (e.g., sliding window, token dropping) and evaluate their impact on generation quality.
We will move from a simple, "append-only" cache to intelligent, fixed-size caches that must decide which information to keep and which to discard. You will implement several core eviction policies, analyze their trade-offs, and—crucially for a systems engineer—measure their performance not just in terms of memory saved, but also in terms of the quality of the generated output. This will build your mental model for the fundamental trade-off between efficiency and accuracy in LLM serving.
The Opportunity for Compression
Before we start dropping tokens, it's important to understand why we can get away with it. If every token in the context were equally important, any eviction would be catastrophic. Fortunately, this isn't the case. Attention in transformers is highly non-uniform.
To build your intuition on this, let's explore the underlying principles that make cache compression feasible.
KV Cache Compression: Eviction, Quantization & H2O Algorithm
This article, 'KV Cache Compression', provides an excellent overview of why compression is possible. We will focus on the section that details the non-uniform nature of attention.
Please read the section titled 'The Compression Opportunity'. As you read, focus on the four phenomena that create this opportunity: local focus, anchor tokens, attention sinks, and sparse activation. This will provide the conceptual foundation for the strategies we're about to implement.
This non-uniformity is the key. Our goal with eviction strategies is to intelligently discard the tokens that receive near-zero attention while preserving the ones that matter.
Strategy 1: Sliding Window Eviction
The most straightforward eviction strategy is the sliding window. The policy is simple: keep only the k most recent tokens in the cache. This is based on the "local focus" principle you just read about—the assumption that the most recent context is the most relevant.
Let's look at the implementation and its significant drawback.
KV Cache Compression: Eviction, Quantization & H2O Algorithm
The same article provides a clear implementation and analysis of the sliding window approach.
Now, read the section 'Window-Based Eviction'. Pay close attention to the code example, which is a simple tensor slice, and the discussion of its fundamental limitation, particularly the example of a question-answering task.
As the article highlights, while a sliding window guarantees a fixed-size cache, it's a blunt instrument. It will indiscriminately discard important long-range context, such as instructions or questions provided at the beginning of a prompt. This can lead to a catastrophic loss of coherence.
Strategy 2: Protecting Attention Sinks (StreamingLLM)
The simple sliding window has a critical flaw that isn't immediately obvious. Many popular transformer models, when not explicitly designed for it, learn to use the very first few tokens as "attention sinks"—a kind of computational anchor for the attention mechanism. Evicting these initial tokens can destabilize the model and severely degrade generation quality, even if the tokens themselves are semantically trivial (like a start-of-sequence token).
The StreamingLLM approach proposes an elegant solution: a hybrid strategy that combines a sliding window with the preservation of these crucial initial tokens.
To understand the problem and the solution, let's watch a video that explains the discovery of attention sinks.
Efficient Streaming Language Models with Attention Sinks (Paper Explained)
Yannic Kilcher's explanation of the 'Efficient Streaming Language Models with Attention Sinks' paper is one of the best resources for understanding this phenomenon. It clearly visualizes why naive windowing fails and how preserving the first few tokens remedies the issue.
Watch from the beginning to 20:44. Focus on these key points: The problem with naive sliding window attention and its effect on the KV cache (0:00 - 10:34). The hypothesis that the first token acts as an 'attention sink' (10:34 - 15:03). The experimental results showing perplexity spikes when this sink is removed (15:03 - 17:33). The StreamingLLM solution: keeping a few initial tokens (sinks) plus a sliding window of recent tokens (17:33 - 20:44).
The video makes the value of attention sinks clear. By keeping just a few initial tokens, we can maintain model stability while still aggressively compressing the rest of the cache.


The implementation combines these two ideas: preserving a prefix and a suffix of the cache.
Strategy 3: Attention-Based Eviction (Heavy-Hitter Oracle)
Our strategies so far have been position-based. A more sophisticated approach is to be content-aware. Instead of just keeping recent tokens, what if we keep the tokens that the model has paid the most attention to in the past? This is the idea behind "Heavy-Hitter" eviction policies.
The Heavy-Hitter Oracle (H2O) algorithm formalizes this. It maintains a cache of:
- The
k_recentmost recent tokens. - The
k_heavytokens that have the highest cumulative attention scores from past generation steps.
This hybrid approach aims to get the best of both worlds: maintaining local coherence (recency) and preserving semantically important, long-range dependencies (heavy-hitters).
KV Cache Compression: Eviction, Quantization & H2O Algorithm
Let's return to our main article, which details the H2O algorithm.
Read the section 'The H2O Algorithm'. Study the description of the two categories of tokens (recent vs. heavy hitters) and examine the H2OCache class implementation. Note the added complexity: this approach requires tracking cumulative attention and original token positions.
The H2O approach is powerful but introduces significant systems-level considerations:
- Computational Overhead: Tracking and updating attention scores at each step adds latency.
- Incompatibility with Fused Kernels: As noted in the "Cold Compress" paper, this strategy requires access to the full attention matrix. This makes it incompatible with optimized attention implementations like FlashAttention, which are designed specifically to avoid materializing this matrix. This is a critical trade-off for a systems engineer: a smarter eviction policy might prevent you from using a faster attention kernel.
Hands-On: Implementing and Evaluating Eviction Strategies
Now, you will implement these strategies and, more importantly, evaluate their impact on generation quality. We will use perplexity as our metric. Perplexity measures how well a probability model predicts a sample; in our case, a lower perplexity means the model with the compressed cache is doing a better job of predicting the next token, indicating higher quality.
1. Setup
First, let's set up a few components. We'll use a slightly modified version of the GPTModel from the previous lesson. We'll also add a function to calculate perplexity. Copy this code into a new Python script.
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
import matplotlib.pyplot as plt
# --- Minimal Transformer Model (from previous lesson) ---
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):
b, num_tokens, d_in = x.shape
queries = self.W_query(x)
keys_new = self.W_key(x)
values_new = self.W_value(x)
if past_k is not None:
keys = torch.cat([past_k, keys_new], dim=1)
values = torch.cat([past_v, values_new], dim=1)
else:
keys = keys_new
values = values_new
# Reshape for multi-head attention
keys = keys.view(b, keys.shape[1], self.num_heads, self.head_dim).transpose(1, 2)
values = values.view(b, values.shape[1], self.num_heads, self.head_dim).transpose(1, 2)
queries = queries.view(b, num_tokens, self.num_heads, self.head_dim).transpose(1, 2)
attn_scores = queries @ keys.transpose(2, 3)
mask = torch.triu(torch.ones(num_tokens, keys.shape[2]), diagonal=keys.shape[2] - 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 @ values).transpose(1, 2).reshape(b, num_tokens, self.d_out)
return self.out_proj(context_vec), keys_new, values_new, attn_weights
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):
attn_output, k_new, v_new, attn_weights = self.attn(x, past_k, past_v)
return attn_output, k_new, v_new, attn_weights
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)
self.pos_emb = nn.Embedding(4096, d_model)
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)
def forward(self, idx, past_kv=None, current_pos=0):
seq_len = idx.shape[1]
pos_ids = torch.arange(current_pos, current_pos + seq_len).unsqueeze(0)
x = self.tok_emb(idx) + self.pos_emb(pos_ids)
new_kvs = []
all_attn_weights = []
for i, block in enumerate(self.trf_blocks):
past_k, past_v = (past_kv[i] if past_kv is not None else (None, None))
x, k_new, v_new, attn_weights = block(x, past_k, past_v)
new_kvs.append((k_new, v_new))
all_attn_weights.append(attn_weights)
logits = self.out_head(x)
return logits, new_kvs, all_attn_weights
def calculate_perplexity(model, text_ids, cache_manager):
"""Calculates perplexity of a model on a text sequence using a given cache strategy."""
model.eval()
total_neg_log_likelihood = 0.0
total_tokens = 0
cache_manager.reset()
# Process text in chunks to simulate streaming generation
chunk_size = 1
for i in range(0, text_ids.shape[1] - 1, chunk_size):
input_chunk = text_ids[:, i : i + chunk_size]
target_chunk = text_ids[:, i + 1 : i + chunk_size + 1]
with torch.no_grad():
past_kv = cache_manager.get_cache()
logits, new_kvs, attn_weights = model(input_chunk, past_kv, current_pos=i)
cache_manager.update(new_kvs, attn_weights)
# Calculate loss for this step
log_probs = F.log_softmax(logits, dim=-1)
neg_log_likelihood = F.nll_loss(log_probs.view(-1, log_probs.size(-1)), target_chunk.view(-1), reduction='sum')
total_neg_log_likelihood += neg_log_likelihood.item()
total_tokens += target_chunk.numel()
perplexity = torch.exp(torch.tensor(total_neg_log_likelihood / total_tokens))
return perplexity.item()
2. Cache Manager Implementations
Now, define the classes that will manage our different cache eviction strategies. Each class will have a reset, update, and get_cache method.
class BaseCacheManager:
def __init__(self, model, max_cache_size):
self.model = model
self.max_cache_size = max_cache_size
self.kv_cache = None
def reset(self):
self.kv_cache = None
def get_cache(self):
return self.kv_cache
def update(self, new_kvs, attn_weights):
raise NotImplementedError
class FullCache(BaseCacheManager):
def update(self, new_kvs, attn_weights):
if self.kv_cache is None:
# Prefill
self.kv_cache = [(k, v) for k, v in new_kvs]
else:
# Decode
new_cache = []
for i, (k_new, v_new) in enumerate(new_kvs):
k_old, v_old = self.kv_cache[i]
new_cache.append((torch.cat([k_old, k_new], dim=1), torch.cat([v_old, v_new], dim=1)))
self.kv_cache = new_cache
class SlidingWindowCache(BaseCacheManager):
def update(self, new_kvs, attn_weights):
# This is a simplified implementation for demonstration.
# A real implementation would be more efficient.
if self.kv_cache is None:
self.kv_cache = [(k, v) for k, v in new_kvs]
else:
new_cache = []
for i, (k_new, v_new) in enumerate(new_kvs):
k_old, v_old = self.kv_cache[i]
k_cat = torch.cat([k_old, k_new], dim=1)
v_cat = torch.cat([v_old, v_new], dim=1)
# Evict oldest tokens if over capacity
current_size = k_cat.shape[1]
if current_size > self.max_cache_size:
k_cat = k_cat[:, -self.max_cache_size:, :, :]
v_cat = v_cat[:, -self.max_cache_size:, :, :]
new_cache.append((k_cat, v_cat))
self.kv_cache = new_cache
class StreamingLLMCache(BaseCacheManager):
def __init__(self, model, sink_size, window_size):
super().__init__(model, sink_size + window_size)
self.sink_size = sink_size
self.window_size = window_size
def update(self, new_kvs, attn_weights):
if self.kv_cache is None:
self.kv_cache = [(k, v) for k, v in new_kvs]
else:
new_cache = []
for i, (k_new, v_new) in enumerate(new_kvs):
k_old, v_old = self.kv_cache[i]
k_cat = torch.cat([k_old, k_new], dim=1)
v_cat = torch.cat([v_old, v_new], dim=1)
current_size = k_cat.shape[1]
if current_size > self.max_cache_size:
# Keep sink tokens + recent window
sink_k, sink_v = k_cat[:, :self.sink_size], v_cat[:, :self.sink_size]
window_k, window_v = k_cat[:, -self.window_size:], v_cat[:, -self.window_size:]
k_cat = torch.cat([sink_k, window_k], dim=1)
v_cat = torch.cat([sink_v, window_v], dim=1)
new_cache.append((k_cat, v_cat))
self.kv_cache = new_cache
# H2O is more complex and left as a conceptual exercise.
# Implementing it efficiently requires careful state management across layers.
3. Run the Comparison
Finally, let's run the evaluation. We'll create a dummy text sequence and calculate the perplexity for each cache strategy.
# --- Model Configuration ---
VOCAB_SIZE = 1000
D_MODEL = 128
NUM_LAYERS = 4
NUM_HEADS = 4
HEAD_DIM = D_MODEL // NUM_HEADS
# --- Evaluation Parameters ---
SEQUENCE_LENGTH = 256
CACHE_SIZE = 64 # Let's say our memory budget is 64 tokens
# --- Setup ---
model = GPTModel(VOCAB_SIZE, D_MODEL, NUM_LAYERS, NUM_HEADS, HEAD_DIM)
# Let's create some dummy text data
dummy_text = torch.randint(0, VOCAB_SIZE, (1, SEQUENCE_LENGTH))
# --- Run Evaluations ---
cache_strategies = {
"Full (No Eviction)": FullCache(model, SEQUENCE_LENGTH),
"Sliding Window (64)": SlidingWindowCache(model, max_cache_size=CACHE_SIZE),
"StreamingLLM (4+60)": StreamingLLMCache(model, sink_size=4, window_size=CACHE_SIZE - 4)
}
results = {}
for name, manager in cache_strategies.items():
print(f"Evaluating: {name}...")
ppl = calculate_perplexity(model, dummy_text, manager)
results[name] = ppl
print(f" -> Perplexity: {ppl:.4f}")
# --- Plot Results ---
names = list(results.keys())
perplexities = list(results.values())
plt.figure(figsize=(10, 6))
bars = plt.bar(names, perplexities)
plt.ylabel('Perplexity (Lower is Better)')
plt.title('Impact of KV Cache Eviction on Model Quality')
for bar in bars:
yval = bar.get_height()
plt.text(bar.get_x() + bar.get_width()/2.0, yval, f'{yval:.2f}', va='bottom') # va: vertical alignment
plt.show()
When you run this experiment, you will see a bar chart comparing the perplexity of each method. You should observe:
- Full Cache has the lowest perplexity (it's our baseline with no information loss).
- Sliding Window likely has the highest perplexity, as it naively discards initial tokens which might be attention sinks.
- StreamingLLM performs significantly better than the simple sliding window and is much closer to the full cache baseline, demonstrating the power of preserving attention sinks.
This exercise makes the trade-off tangible: by accepting a small hit in perplexity, StreamingLLM allows the model to operate with a fixed cache that is 4x smaller (256 -> 64) than the full sequence length.
Conclusion
In this lesson, you have explored the critical challenge of managing KV cache memory for long-context generation. You've gone beyond a simple, growing cache to implement and evaluate intelligent eviction policies.
Key Takeaways:
- KV cache eviction is necessary to handle long sequences within a finite memory budget.
- Sliding Window is the simplest strategy but fails when long-range context or attention sinks are important.
- StreamingLLM offers a robust and efficient solution by preserving a few initial "sink" tokens alongside a sliding window of recent tokens, dramatically improving quality over a naive window.
- Heavy-Hitter strategies are content-aware and can preserve important semantic context, but they come with computational overhead and may be incompatible with optimized attention kernels like FlashAttention.
- There is a direct, measurable trade-off between memory compression and generation quality (perplexity). The "best" strategy is application-dependent.
Preview of the Next Lesson:
So far, we've focused on optimizing the cache within a single generation sequence. But in a production server, you have many requests arriving concurrently. What if multiple users start their chats with the same system prompt? Must we compute and store the KV cache for that prompt for every single user? The answer is no.
In the next lesson, you will implement prefix caching to reuse KV cache entries for requests that share a common prompt. This is a powerful technique for improving throughput and reducing latency in multi-user environments.