Keyboard shortcuts

Press ← or → to navigate between chapters

Press S or / to search in the book

Press ? to show this help

Press Esc to hide this help

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.

Explore: block tables

Add tokens to three sequences and watch blocks being allocated from the shared pool. Fork a sequence to see shared blocks and reference counts, then write to the fork to trigger a copy-on-write.

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:

  1. 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).
  2. set_batch(seqs) builds the block-table tensor for this batch, padded to the longest table.
  3. update scatters 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 blocklonger contiguous reads, better kernel efficiency
finer-grained prefix sharingsmaller 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

  1. ★ Plot paged_utilization from run.py paged against block sizes 1-128 for the same request lengths. Where: experiments/ch25.py (create it), adapting run.py’s cmd_paged with different block sizes.
  2. ★★ Replace StaticKVCache in Chapter 24’s engine with PagedKVCache and a PrefixCache, calling reserve for every planned chunk. When reserve raises MemoryError, try evict, then preempt the newest request. Verify outputs still match solo runs. Where: ContinuousBatchingEngine.__init__ / step in engine/scheduler.py, using engine.paged.
  3. ★★ Implement $n$-way parallel sampling in the batching engine with fork: one prefill, then $n$ independently seeded samples. Where: request/slot handling in engine/scheduler.py and PagedKVCache.fork in engine/paged.py.
  4. ★★★ 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_kernel in engine/kernels/triton_paged.py.

Check your understanding

  1. Why can’t an engine simply allocate exactly the cache each request will need?
  2. Where does the logical position $p$ of a sequence live in the pool?
  3. Why is the key of a cached prefix block the whole prefix, not just the block’s tokens?
  4. When does copy-on-write copy a block, and which block is it usually?
  5. Why does attach always 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.py and kv_cache_manager.py, and its paged attention kernels, which follow the structure of this chapter.