25. Paged attention and prefix caching
In this chapter
- Why reserving a maximum-length cache per request wastes most of the GPU's memory.
- Paging, borrowed from operating systems: fixed-size blocks, a shared pool and per-sequence block tables.
- Sharing blocks between sequences with reference counts and copy-on-write: parallel sampling and beam search for free.
- Prefix caching: reusing the KV cache of a shared system prompt across requests.
- A decode attention kernel that reads keys and values straight through the block table.
You will build
BlockAllocator and PagedKVCache in engine/paged.py, and a paged decode attention kernel in engine/kernels/triton_paged.py.
Time: 5-7 hours. GPU: optional (the Triton kernel runs in the interpreter).
Reserved memory is wasted memory
Chapter 24’s engine gives each request a slot of fixed capacity, say 1,024 tokens. A request that uses 50 tokens still reserves 1,024. With requests of varied lengths, most of the reserved cache is empty:
python run.py paged
{"requests": 32, "tokens_stored": 9924, "static_slots_reserved": 32768, "paged_slots_reserved": 10144,
"static_utilization": 0.303, "paged_utilization": 0.978}
Thirty-two live requests of 20-500 tokens fill 30% of their reserved slots. The vLLM paper (Kwon et al., 2023) measured existing systems using only 20-38% of their KV memory for actual tokens. Since the cache limits how many requests fit (Chapter 24), that waste directly cuts throughput.
Allocating exactly the right amount isn’t possible either: you don’t know how long an answer will be until it’s done. And growing a contiguous buffer means copying it, or fragmenting memory into holes that no request fits.
Paging
Operating systems solved the same problem for process memory decades ago. PagedAttention applies the solution to the KV cache:
- Split every sequence’s cache into fixed-size blocks of $b$ tokens (16 is typical).
- Keep all blocks of all sequences in one pool: per layer, tensors of shape
[num_blocks, kv_heads, b, head_dim]. - Give each sequence a block table mapping its logical blocks to physical blocks in the pool, which need not be contiguous or in order.
Logical position $p$ lives in physical block table[p // b] at offset p % b. For $b = 4$ and table [5, 1, 7], position 9 is in logical block 2, physical block 7, offset 1.
print(table[9 // 4], 9 % 4) # (7, 1)
inline std::pair<size_t, size_t> physical_slot(const std::vector<size_t>& table, size_t pos, size_t block) {
return {table[pos / block], pos % block};
}
#![allow(unused)]
fn main() {
pub fn physical_slot(table: &[usize], position: usize, block_size: usize) -> (usize, usize) {
(table[position / block_size], position % block_size)
}
}
A sequence allocates a new block only when it fills its last one, so it wastes at most $b - 1$ slots, under 4% in the demo above. Blocks come from a free list and go back when the request finishes.
The allocator
class BlockAllocator:
"""Free list plus reference counts. (Your engine: Chapter 25)"""
def __init__(self, num_blocks):
self.free = deque(range(num_blocks))
self.refs = [0] * num_blocks
def allocate(self):
"""(Your engine: Chapter 25)"""
if not self.free:
raise MemoryError("KV block pool exhausted")
block = self.free.popleft()
self.refs[block] = 1
return block
def share(self, block):
self.refs[block] += 1
def release(self, block):
"""(Your engine: Chapter 25)"""
self.refs[block] -= 1
if self.refs[block] == 0:
self.free.append(block)
elif self.refs[block] < 0:
raise RuntimeError(f"Block {block} released more times than it was referenced")
@property
def num_free(self):
return len(self.free)
Every block carries a reference count: the number of block tables (and caches, below) that point to it. release decrements it, and only a count of zero returns the block to the free list.
The paged cache
class PagedKVCache:
"""Implements the book's cache protocol over a block pool. (Your engine: Chapter 25)
Before a forward pass the caller names the sequences in the batch (set_batch) and
reserves room for the new tokens (reserve); allocation failures therefore happen
before any GPU work is launched.
"""
def __init__(self, layers, num_blocks, block_size, kv_heads, head_dim, device="cpu", dtype=torch.float32):
shape = (num_blocks, kv_heads, block_size, head_dim)
self.k_pool = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.v_pool = [torch.zeros(shape, device=device, dtype=dtype) for _ in range(layers)]
self.block_size, self.device = block_size, device
self.allocator = BlockAllocator(num_blocks)
self.tables, self.lengths = {}, {}
self.batch = []
def blocks_needed(self, tokens):
return -(-tokens // self.block_size)
def create(self, seq):
if seq in self.tables:
raise ValueError(f"Sequence {seq} already exists")
self.tables[seq], self.lengths[seq] = [], 0
def reserve(self, seq, new_tokens):
"""Make blocks exist for positions [length, length + new_tokens), copying any shared
block that is about to be written (copy-on-write). (Your engine: Chapter 25)"""
table, start = self.tables[seq], self.lengths[seq]
end = start + new_tokens
first_written = start // self.block_size
for index in range(first_written, min(len(table), self.blocks_needed(end))):
block = table[index]
if self.allocator.refs[block] > 1: # shared: copy before writing
fresh = self.allocator.allocate()
for pool in self.k_pool + self.v_pool:
pool[fresh].copy_(pool[block])
self.allocator.release(block)
table[index] = fresh
while len(table) < self.blocks_needed(end):
table.append(self.allocator.allocate())
def set_batch(self, seqs):
self.batch = list(seqs)
width = max(len(self.tables[s]) for s in self.batch)
rows = [self.tables[s] + [0] * (width - len(self.tables[s])) for s in self.batch]
self.table_tensor = torch.tensor(rows, device=self.device, dtype=torch.long)
def update(self, layer, k, v, positions, rows=None):
"""Scatter new K/V into their blocks, then gather each sequence's history (reference path). (Your engine: Chapter 25)"""
if positions.ndim == 1:
positions = positions.expand(k.shape[0], -1)
batch_index = torch.arange(k.shape[0], device=k.device)[:, None]
blocks = self.table_tensor[batch_index, positions // self.block_size] # [B, T]
offsets = positions % self.block_size
self.k_pool[layer][blocks, :, offsets] = k.transpose(1, 2)
self.v_pool[layer][blocks, :, offsets] = v.transpose(1, 2)
if layer == len(self.k_pool) - 1: # the last layer commits the lengths
for b, seq in enumerate(self.batch):
self.lengths[seq] = max(self.lengths[seq], int(positions[b].max()) + 1)
span = int(positions.max()) + 1
logical = torch.arange(span, device=k.device)
gather_blocks = self.table_tensor[:, logical // self.block_size] # [B, span]
keys = self.k_pool[layer][gather_blocks, :, logical % self.block_size] # [B, span, H, D]
values = self.v_pool[layer][gather_blocks, :, logical % self.block_size]
return keys.transpose(1, 2), values.transpose(1, 2), logical
def fork(self, parent, child):
"""child starts as an alias of parent's blocks (beam search, n>1 sampling, shared prompts). (Your engine: Chapter 25)"""
self.create(child)
self.tables[child] = list(self.tables[parent])
self.lengths[child] = self.lengths[parent]
for block in self.tables[child]:
self.allocator.share(block)
def free(self, seq):
for block in self.tables.pop(seq):
self.allocator.release(block)
self.lengths.pop(seq)
def truncate(self, seq, length):
"""Roll back to `length` tokens, returning now-unused blocks."""
keep = self.blocks_needed(length)
for block in self.tables[seq][keep:]:
self.allocator.release(block)
del self.tables[seq][keep:]
self.lengths[seq] = length
The cache implements the same update protocol as every cache since Chapter 16, so your Qwen3 runs on it unchanged. The flow per step:
reserve(seq, n)before the forward pass makes sure blocks exist for the $n$ new positions. Allocation failures happen here, before any GPU work is launched, where the scheduler can handle them (by waiting, or preempting another request).set_batch(seqs)builds the block-table tensor for this batch, padded to the longest table.updatescatters the new keys and values to(table[p // b], p % b)and returns each sequence’s history. This reference version gathers the history into a contiguous tensor for the existing attention code, which costs a copy; the kernel at the end of the chapter removes it.
Sharing blocks: forks and copy-on-write
Many requests share content. Parallel sampling (n = 4 answers to one prompt) and beam search need several sequences with the same prompt. With block tables, a fork copies the table, not the data, and increments each block’s reference count.
The sequences can then diverge. When one of them is about to write into a block whose count is above 1, reserve first copies that block to a fresh one and points this sequence’s table at the copy: copy-on-write. Full blocks of the prompt stay shared forever; only the partially filled last block gets copied.
{"parent_blocks": 3, "new_blocks_after_4_forks": 0, "new_blocks_after_each_sample_writes_a_token": 4}
A 40-token prompt takes 3 blocks of 16 (the last one half full). Four forks cost no new blocks. When each sample writes its first token, each copies the shared partial block: 4 new blocks in total, instead of 12 for four independent copies of the prompt.
The milestone test checks the subtle part: after a fork, the parent and the child each continue with a different token, and both must produce exactly the logits of running their full sequence from scratch. If copy-on-write is missing, one sequence’s write corrupts the other’s history.
Prefix caching
Most production traffic shares long prefixes: the same system prompt, the same few-shot examples, the same document with different questions, the growing history of a multi-turn chat. Their KV cache is identical, so it can be computed once.
class PrefixCache:
"""Reuse full blocks whose token contents (and everything before them) match exactly.
Keys are the whole token prefix up to the end of a block, so two blocks with the same
text but different histories never collide. The cache holds its own reference on each
block; least-recently-used entries are evicted when the pool runs low.
"""
def __init__(self, cache):
self.cache = cache
self.entries = OrderedDict() # tuple(prefix tokens) -> block
def attach(self, seq, prompt):
"""Create seq with the longest cached prefix of prompt. Returns tokens reused."""
self.cache.create(seq)
size, reused = self.cache.block_size, 0
# Keep at least one prompt token to prefill: its logits predict the first output.
while (reused + size) <= len(prompt) - 1:
key = tuple(prompt[:reused + size])
block = self.entries.get(key)
if block is None:
break
self.entries.move_to_end(key)
self.cache.allocator.share(block)
self.cache.tables[seq].append(block)
reused += size
self.cache.lengths[seq] = reused
return reused
def remember(self, seq, tokens):
"""After prefill, publish seq's full blocks for later requests."""
size = self.cache.block_size
for index in range(len(tokens) // size):
key = tuple(tokens[:(index + 1) * size])
if key not in self.entries:
block = self.cache.tables[seq][index]
self.cache.allocator.share(block)
self.entries[key] = block
def evict(self, blocks_wanted):
"""Drop LRU entries that nobody else is using until enough blocks are free."""
for key in list(self.entries):
if self.cache.allocator.num_free >= blocks_wanted:
break
block = self.entries[key]
if self.cache.allocator.refs[block] == 1:
del self.entries[key]
self.cache.allocator.release(block)
The key for a cached block is the entire token prefix up to the end of that block, not just the block’s own tokens. A block’s keys and values depend on every earlier token, so two blocks with identical text but different histories must never match. (Production engines hash the chain of block contents instead of storing the tuples, and the key must also include anything else that changes activations: the model revision, LoRA adapter, and multimodal inputs.)
attach walks the prompt one full block at a time while the prefix is cached, shares those blocks, and leaves at least one token to prefill, because the prompt’s last position must run to produce the first output’s logits. remember publishes a sequence’s full blocks after prefill. The cache holds its own reference to each block, so cached blocks survive their requests, and evict drops least-recently-used entries nobody else is using when the pool runs low.
Eight requests sharing a 96-token system prompt, each with 8 tokens of its own:
{"prompt_tokens": 832, "tokens_prefilled": 160, "blocks_in_use": 14, "blocks_without_sharing": 56}
The first request prefills all 104 tokens; the other seven prefill only their own 8. That’s 81% less prefill compute, which lowers time to first token directly, and a quarter of the memory. SGLang’s RadixAttention generalizes this to a radix tree over all cached prefixes and schedules requests to maximize hits.
Attention that reads through the table
The gather in update builds a contiguous copy of every sequence’s history at every layer, every step. A paged attention kernel reads K and V directly from the pool instead. Here’s a decode kernel: one program per (sequence, query head), walking the sequence’s block table and applying Chapter 15’s online softmax one block at a time:
@triton.jit
def paged_decode_kernel(q_ptr, k_pool, v_pool, tables_ptr, lengths_ptr, out_ptr,
max_blocks, heads, group, scale,
s_pool_block, s_pool_head, s_pool_slot,
D: tl.constexpr, BLOCK: tl.constexpr):
"""(Your engine: Chapter 25)"""
seq, h = tl.program_id(0), tl.program_id(1)
kvh = h // group
rd = tl.arange(0, D)
slots = tl.arange(0, BLOCK)
q = tl.load(q_ptr + (seq * heads + h) * D + rd).to(tl.float32)
length = tl.load(lengths_ptr + seq)
m = tl.full((), -float("inf"), tl.float32)
l = tl.zeros((), tl.float32)
acc = tl.zeros((D,), tl.float32)
for i in range(0, tl.cdiv(length, BLOCK)):
block = tl.load(tables_ptr + seq * max_blocks + i) # logical block i -> physical block
base = block * s_pool_block + kvh * s_pool_head
valid = i * BLOCK + slots < length
k = tl.load(k_pool + base + slots[:, None] * s_pool_slot + rd[None, :], mask=valid[:, None], other=0.0)
v = tl.load(v_pool + base + slots[:, None] * s_pool_slot + rd[None, :], mask=valid[:, None], other=0.0)
s = tl.sum(k.to(tl.float32) * q[None, :], axis=1) * scale
s = tl.where(valid, s, -float("inf"))
m_new = tl.maximum(m, tl.max(s, axis=0))
alpha = tl.exp(m - m_new)
p = tl.exp(s - m_new)
l = l * alpha + tl.sum(p, axis=0)
acc = acc * alpha + tl.sum(p[:, None] * v.to(tl.float32), axis=0)
m = m_new
tl.store(out_ptr + (seq * heads + h) * D + rd, (acc / l).to(out_ptr.dtype.element_ty))
def paged_decode_attention(q, k_pool, v_pool, block_tables, lengths, block_size):
"""q [B, Hq, D]; pools [N, Hkv, block_size, D]; block_tables int [B, max_blocks]; lengths int [B] (>= 1)."""
check_device(q, k_pool, v_pool, block_tables, lengths)
batch, heads, d = q.shape
out = torch.empty_like(q)
k_pool, v_pool = k_pool.contiguous(), v_pool.contiguous()
paged_decode_kernel[(batch, heads)](q.contiguous(), k_pool, v_pool, block_tables.contiguous(),
lengths.contiguous(), out, block_tables.shape[1], heads,
heads // k_pool.shape[1], 1.0 / math.sqrt(d),
k_pool.stride(0), k_pool.stride(1), k_pool.stride(2),
D=d, BLOCK=block_size)
return out
Compared with the flash kernel, the only new thing is the indirection: tl.load(tables_ptr + seq * max_blocks + i) turns logical block i into a physical block number, and the K/V addresses come from that. The lengths tensor masks the tail of the last block. Grouped-query attention works the same way as before: query head h reads KV head h // group.
The reference paged_decode_attention in paged.py gathers and computes the same result; the test checks the kernel against it with scrambled block tables.
What production paged kernels add: splitting long sequences across programs with a final merge (flash-decoding, Chapter 15), processing all query heads of a KV group together so K/V is loaded once, and FP8 KV caches. FlashInfer and vLLM’s attention backends implement these for prefill, decode and mixed batches.
Choosing the block size
| smaller blocks (8) | larger blocks (32-128) |
|---|---|
| less waste in each sequence’s last block | longer contiguous reads, better kernel efficiency |
| finer-grained prefix sharing | smaller block tables, fewer allocations |
16 is vLLM’s default. Some FlashAttention-3-based paged kernels prefer larger blocks.
Build it
Engine milestone 25: a paged KV cache. Implement BlockAllocator.allocate and release, and PagedKVCache.reserve, update and fork in engine/paged.py, and paged_decode_kernel in engine/kernels/triton_paged.py (set_batch, free, truncate, the prefix cache and the launcher are provided).
pytest tests/test_ch25_paged.py
python run.py paged --impl engine
The tests check reference counts and pool exhaustion, that a Qwen3 running on the paged cache (prefill, then token-by-token) matches the full forward, that a forked parent and child diverge correctly through copy-on-write and return every block when freed, that a second request reuses three cached prefix blocks and still computes the right logits, and your Triton kernel against the reference with non-contiguous block tables.
Stretch exercises
- ★ Plot
paged_utilizationfromrun.py pagedagainst block sizes 1-128 for the same request lengths. Where:experiments/ch25.py(create it), adaptingrun.py’scmd_pagedwith different block sizes. - ★★ Replace
StaticKVCachein Chapter 24’s engine withPagedKVCacheand aPrefixCache, callingreservefor every planned chunk. WhenreserveraisesMemoryError, tryevict, then preempt the newest request. Verify outputs still match solo runs. Where:ContinuousBatchingEngine.__init__/stepinengine/scheduler.py, usingengine.paged. - ★★ Implement $n$-way parallel sampling in the batching engine with
fork: one prefill, then $n$ independently seeded samples. Where: request/slot handling inengine/scheduler.pyandPagedKVCache.forkinengine/paged.py. - ★★★ Write a paged prefill kernel: a block of queries from one sequence attends causally to that sequence’s paged history plus the new tokens. Where: add a prefill kernel and launcher beside
paged_decode_kernelinengine/kernels/triton_paged.py.
Check your understanding
- Why can’t an engine simply allocate exactly the cache each request will need?
- Where does the logical position $p$ of a sequence live in the pool?
- Why is the key of a cached prefix block the whole prefix, not just the block’s tokens?
- When does copy-on-write copy a block, and which block is it usually?
- Why does
attachalways leave at least one prompt token to prefill?
Going deeper
- Kwon et al., Efficient Memory Management for Large Language Model Serving with PagedAttention (SOSP 2023), the vLLM paper; Zheng et al., SGLang: Efficient Execution of Structured Language Model Programs (2024) for RadixAttention.
- GPU Mode L40 (FlashInfer): paged and ragged attention kernels in production; L35 (SGLang).
- PMPP §20.7 (Alleviating the memory requirements of the attention mechanism).
- vLLM’s
vllm/v1/core/block_pool.pyandkv_cache_manager.py, and its paged attention kernels, which follow the structure of this chapter.