Skip to content

Paged attention for decode in C

ModuleL9.4 · side · C · Pass 6 · 4 to 6 h
You buildc/src/kernels/paged_attn.c: tl_paged_attn_decode_f32, one decode step of attention for a batch of sequences whose K and V live in blocks of your rt.04 pool, reached through block tables, read as f16, with GQA and a sliding window, allocating nothing
Contractcourse/contracts/c/include/tinyllm/attention.h · the pool: kv_pool.h · rules: c/ABI.md
Testscourse/tests/L9.4/: test_paged_attn.c (C, under ASan and UBSan, and TSan for the pool case, against a naive oracle and, bitwise, against your L9.3) and shared file fixtures (what they check: section 4)
Needsrt.02 the loader · rt.03 the pool · rt.04 the block pool · M09.7 tl_f16_to_f32 · M09.6 tl_expf · L9.2 tl_softmax_f32 (the C oracle) · L9.3 FlashAttention (the bitwise oracle) · L8.3 PagedKVCache · L7.7 windowed_attention (or --ref-deps). Reading: L7.5 (GQA)
Used byNone: this optional C exercise is tested as a standalone binary; the Rust engine implements decode independently
MilestoneMS-L9
Optional depthKwon et al., “Efficient Memory Management for Large Language Model Serving with PagedAttention” (vLLM, 2023); Dao et al., “Flash-Decoding” (2023)
  • A sequence’s cache is a list of blocks, not one buffer: position jj lives in slot j mod bj \bmod b of block table[⌊j/b⌋]\mathit{table}[\lfloor j / b \rfloor], so the kernel reads keys in position order through the table, whatever the block ids (hand_example).
  • Decode is FlashAttention with one query at the newest position: the same online softmax over the blocks, with the running output kept in the caller’s out row, so the kernel allocates nothing (matches_naive_over_the_gathered_cache).
  • Blocks are tiles aligned to absolute positions, so a paged decode gives the same bits as your L9.3 kernel on the gathered cache with BcB_c = the block size: prefill with one kernel and decode with the other, and the tokens agree (equals_flash_attention_bitwise).
  • Shared prefix blocks just work: after a fork two tables name the same blocks, each read independently, and a sequence’s output has the same bits alone or in a batch of 16 (shared_blocks_and_batch_invariance).
  • The whole table is validated before the first write: an empty context, a table too short, a block id outside the pool, or a layer out of range is TL_EINVAL (bad_tables_and_shapes).
Terminal window
ol start L9.4 # stubs c/src/kernels/paged_attn.c into your repo
ol tests L9.4 # read the test catalog first
ol check L9.4 # exit code is the verdict
ol check L9.4 --ref-deps # only if rt.04, L8.3, L9.3, or another dependency is not passing yet
ol diff L9.4 # after passing: your code against the reference

L8.3 defines the logical paged-cache behavior in Python: each sequence owns a block table, prefix blocks are shared after a fork, and a block is copied only when someone writes into a shared one. This optional C exercise implements the same layout for a standalone decode kernel. Its tests compare the C result with a naive oracle and the C prefill kernel, so the implementation can be checked without binding C into Python or Rust. The Rust engine owns an independent Rust KV block manager and attention implementation.

SymbolMeaningType / shape
bbblock_tokens: positions per block (16 in the engine)uint32_t
tables\mathit{table}_ssequence ss‘s block ids in position order, row ss of block_tablesuint32_t[max_blocks]
nsn_sctx_lens[s]: positions in the cache, the newest includedint32_t
ps=ns−1p_s = n_s - 1the query’s position (the newest token)
qs,h∈RDq_{s,h} \in \mathbb{R}^Dthe query of sequence ss, head hhfloat[D]
kj,vjk_j, v_jkey and value at position jj, stored as f16uint16_t[D]
wwthe window (w≤0w \le 0: none)int64_t
m,ℓ,am, \ell, arunning maximum, denominator, output (as in L9.3)float, float, float[D]

kv_pool.h lays out one block as, for each layer, a K slab and a V slab, each [n_kv_heads][block_tokens][head_dim] elements of the pool’s dtype (f16 in format 1). So the key of KV head gg at position jj of sequence ss is at

tl_kv_block_ptr(kv,tables[⌊j/b⌋],layer,0)+(g⋅b+(j mod b))⋅D,\texttt{tl\_kv\_block\_ptr}(\mathit{kv}, \mathit{table}_s[\lfloor j/b \rfloor], \mathit{layer}, 0) + \big(g \cdot b + (j \bmod b)\big) \cdot D ,

and the value is the same with is_v = 1. The table maps logical blocks (positions 0..b−10..b-1, b..2b−1b..2b-1, …) to physical blocks anywhere in the pool. That indirection is the whole point of paging: memory is allocated a block at a time, no sequence needs a contiguous region, and two sequences can name the same physical block.

For each sequence ss and head hh (KV head g=⌊h/(H/Hkv)⌋g = \lfloor h / (H/H_{kv}) \rfloor), the query is the newest token, so it is causal by construction and sees positions

V={ j:max⁡(0,ps−w+1)≤j≤ps }(all of 0..ps when w≤0).V = \{\, j : \max(0, p_s - w + 1) \le j \le p_s \,\} \quad (\text{all of } 0..p_s \text{ when } w \le 0).

The kernel walks the logical blocks that intersect VV in increasing order and applies L9.3’s update to each block’s visible positions: scores σ q⋅kj\sigma\, q \cdot k_j (the f16 key decoded by your tl_f16_to_f32), the block’s maximum, m′=max⁡(m,⋅)m' = \max(m, \cdot), α=em−m′\alpha = e^{m - m'}, ℓ←αℓ+∑esj−m′\ell \leftarrow \alpha \ell + \sum e^{s_j - m'}, a←αa+∑esj−m′vja \leftarrow \alpha a + \sum e^{s_j - m'} v_j. At the end o=a⋅(1/ℓ)o = a \cdot (1/\ell). The query sees at least itself, so ℓ>0\ell > 0.

The contract gives this kernel no arena: a decode step is small and frequent, and every byte of state fits elsewhere. The running output aa lives in the caller’s out row (it is overwritten anyway); mm and ℓ\ell are two locals; the scores of a block go into a fixed stack buffer of 256 floats (a block larger than that is scored in aligned pieces of 256). Nothing depends on nsn_s.

L9.3 with one query row at q_offset =ps= p_s, causal, the same window, and Bc=bB_c = b visits key tiles [tb,(t+1)b)[tb, (t+1)b) in increasing tt and performs, for each, exactly the operations above on the same float32 values (the f16 values decoded). The two kernels therefore produce the same bits; equals_flash_attention_bitwise checks it. Anything that changes the grouping of the sums (a piece size that depends on the batch, a different shift, dividing by ℓ\ell instead of multiplying by 1/ℓ1/\ell) keeps the answer close but breaks the bits, and with them the engine’s guarantee that prefill and decode agree.

Each (s,h)(s, h) row reads only its own query, its own table, and the pool, and writes only its own output: a sequence’s result is the same alone or in a test batch of 16, and the same whichever rt.03 worker computes it.

The engine builds the tables, and a bug there would make the kernel read another sequence’s memory or past the pool. So before the first write the kernel checks the shapes against the pool (TL_ESHAPE when HkvH_{kv} or DD differ, or H mod Hkv≠0H \bmod H_{kv} \ne 0), the layer, every context length (≥1\ge 1), that each table holds ⌈ns/b⌉≤\lceil n_s / b \rceil \le max_blocks entries, and that every block id it will read is below n_blocks (TL_EINVAL). Pools in format 2 (fp8, craft.13) are TL_EUNSUPPORTED until that migration.

L9.3’s example in a paged cache: D=2D = 2, b=2b = 2, one sequence of n=3n = 3 positions in a pool of 4 blocks, table [3,1][3, 1]:

Position jjblock ⌊j/2⌋\lfloor j/2 \rfloor → idslotkjk_jvjv_j
00 → 30[1,0][1, 0][1,2][1, 2]
10 → 31[0,1][0, 1][3,4][3, 4]
21 → 10[1,1][1, 1][5,6][5, 6]

All values are exact in f16. The query q=[1,0]q = [1, 0] is position p=2p = 2 and sees positions 0, 1, 2.

Block id 3 (positions 0, 1): scores 1 and 0, m′=1m' = 1, α=e−∞=0\alpha = e^{-\infty} = 0, ℓ=1+e−1=1.3678794\ell = 1 + e^{-1} = 1.3678794, a=[1,2]+0.3678794 [3,4]=[2.1036383,3.4715177]a = [1, 2] + 0.3678794\,[3, 4] = [2.1036383, 3.4715177].

Block id 1 (position 2; slot 1 is empty and not visible): score 1, m′=1m' = 1, α=1\alpha = 1, ℓ=2.3678794\ell = 2.3678794, a=[7.1036383,9.4715177]a = [7.1036383, 9.4715177].

o=a/ℓ=[3,4]o = a / \ell = [3, 4]: the first test, hand_example, computes this exact result through the C interface.

tinyllm/attention.h
tl_status tl_paged_attn_decode_f32(const float *q, const tl_kv_pool *kv, uint32_t layer,
const uint32_t *block_tables, int32_t max_blocks,
const int32_t *ctx_lens, float *out,
int64_t B, int64_t H, int64_t Hkv, int64_t D,
float scale, int64_t window, tl_pool *tp);
/* q, out [B, H, D]. Sequence b's blocks: block_tables[b * max_blocks + i].
TL_EINVAL: NULL pointers, a block id or layer out of range, ctx_lens[b] < 1,
a sequence needing more than max_blocks blocks. TL_ESHAPE: Hkv or D differ
from the pool, or H % Hkv != 0. */

From Python, PagedKVCache.pool (L8.3) is the tl_kv_pool * and block_table(seq) the row of ids.

TestKINDChecksWhy it matters downstream
hand_exampleunit, smokesection 3 through table [3,1][3, 1]the indirection and the slot arithmetic
matches_naive_over_the_gathered_cachedifferentialGQA 6:2, D=32D = 32, two layers, lengths 1, 37, 70 in shuffled blocks, window 0 and 20every flag against plain attention
equals_flash_attention_bitwisedifferential, propertyyour L9.3 on the gathered values with Bc=16B_c = 16, bitwiseprefill and decode agree
shared_blocks_and_batch_invarianceproperty16 sequences sharing two prefix blocks; each alone vs in the test batch, bitwiseforked prefix blocks and sequence isolation
pool_result_equals_serial_bitwiseproperty4 threads vs serial (also under TSan)threads never change bits
bad_tables_and_shapesboundaryempty context, short table, id 9 in a 4-block pool, layer out of range (TL_EINVAL); head or width mismatch (TL_ESHAPE); out untoucheda table bug fails loudly
hand_exampleunit, smokesection 3 written through your PagedKVCachethe cache and the kernel share one layout
PitfallSymptomCaught by
not rescaling the running outputblocks before a larger score are overweightedmatches_naive_over_the_gathered_cache (mutant s01)
not rescaling the denominatoroutputs do not average the valuesmatches_naive_over_the_gathered_cache (mutant s02)
using the logical block index as the block idreads whatever block happens to have that idhand_example (mutant s03)
forgetting the KV head’s offset inside the slabevery head reads KV head 0matches_naive_over_the_gathered_cache (mutant s04)
mapping heads round-robin (h mod Hkvh \bmod H_{kv})GQA models attend with the wrong keysmatches_naive_over_the_gathered_cache (mutant s05)
a window of w+1w + 1 positionsdrifts from the model’s trainingmatches_naive_over_the_gathered_cache (mutant s06)
placing the query at nn instead of n−1n - 1reads an unwritten slothand_example (mutant s07)
reading V from the K slaboutputs are averages of keyshand_example (mutant s08)
shifting by a maximum that includes positions outside the windowthe answer is close, but not FlashAttention’s bitsequals_flash_attention_bitwise (mutant s09)
dividing by ℓ\ell instead of multiplying by 1/ℓ1/\ellclose, but not FlashAttention’s bits: prefill and decode disagree on near-tiesequals_flash_attention_bitwise (mutant s10)
a smaller piece size for a lone sequenceits bits change when it joins a batchshared_blocks_and_batch_invariance (mutant s11)
a pooled partition that drops a rowwrong only with threadspool_result_equals_serial_bitwise (mutant s12)
accepting an empty context1/ℓ=∞1/\ell = \infty: NaNbad_tables_and_shapes (mutant s13)
trusting max_blocksreads past the tablebad_tables_and_shapes (mutant s14)
trusting the block idstl_kv_block_ptr returns NULL, then a crashbad_tables_and_shapes (mutant s15)
DirectionModuleHow it uses this
Backrt.02the error slot and allocator support
Backrt.03tl_parallel_for over (sequence, head)
Backrt.04the block pool: tl_kv_pool_cfg and tl_kv_block_ptr
BackM09.7tl_f16_to_f32 decodes every key and value in this standalone C module
BackM09.6tl_expf
BackL9.2the online update; tl_softmax_f32 is the naive oracle’s softmax
BackL9.3the bitwise oracle: decode must equal FlashAttention with one query
BackL8.3PagedKVCache writes the blocks the Python tests read
BackL7.7windowed_attention, the specification
ForwardNonethe production Rust engine uses its own Rust implementation; this optional C exercise has no production caller
Your pieceProduction equivalentWhat it addsWhere to look
one query per (sequence, head)vLLM PagedAttention v1 and v2one thread block per (sequence, head); v2 splits long contexts across blocks and merges partial (m, l, a)vLLM csrc/attention/
a serial walk over a long tableFlash-Decodingsplit the context across workers for one query and merge with the same rescale: parallel at batch size 1, at the price of a fixed split orderDao et al. (2023)
f16 KVfp8 KV cacheshalf the bytes per token; scales per head or per block (your craft.13 migration)vLLM kv_cache_dtype="fp8"
block tables from the engineSGLang RadixAttentionthe tables come from a radix tree over token ids, so shared prefixes are found automaticallySGLang radix_cache.py