Skip to content

KV cache, incremental decode, incremental UTF-8 detokenizer, generate

ModuleL8.2 · build · Python · Pass 6 · 5 to 7 h
You buildpython/tinyllm/infer/kvcache.py: KVCache and LatentCache (update, seq_len, positions, mask, truncate, nbytes) · python/tinyllm/infer/generate.py: cache_dims, IncrementalDecoder (push, flush), Generation, generate
Contractcourse/contracts/py/tinyllm/infer/kvcache.pyi · course/contracts/py/tinyllm/infer/generate.pyi
Testscourse/tests/L8.2/ (what they check: section 4; the test models are in _tinylm.py) · your own tests in python/tests/l8-2-generate/, rung R5, graded by mutation (threshold 0.80, every pitfall mutant required)
NeedsL8.1 (sample, request_rng, sampled_entropy) · L5.2 (causal_mask with q_offset) · L1.2 (the byte-level BPE the detokenizer streams) · L7.9 (LlamaForCausalLM, the model generate drives) · L7.5 (GQAttention, whose cache hook this cache implements) · L7.3 (RopeSpec, to build that attention) · L0.1 (Tensor, its input) · reading: M05.1 KV bytes per token
Used byL8.4, L8.5, L8.6 (speculative decoding rolls the cache back with truncate), and L10.5 ports the detokenizer to Rust
MilestoneMS-L8 (step 1: generate --cache none,contiguous,paged give the same greedy text)
Optional depthPope et al., “Efficiently Scaling Transformer Inference” (2022); Kwon et al., “Efficient Memory Management for Large Language Model Serving with PagedAttention” (2023), section 2; the Unicode Standard, ch. 3.9 (UTF-8 and the replacement of ill-formed sequences)
  • Attention at position tt needs the keys and values of positions 0..t0..t; caching them turns each decode step into one token’s work against the stored keys, and the cached logits equal a full recompute to 10−510^{-5} at every step (test_cached_logits_equal_full_recompute).
  • A decode chunk sits at absolute positions s,s+1,…s, s + 1, \dots with ss the committed length, and its mask is the causal mask aligned to the bottom right, causal_mask(T, s + T, q_offset=s); both come from the cache before any layer appends (test_mask_is_l52_causal_with_offset, test_seq_len_commits_after_every_layer).
  • The cache returns what it stores: a float16 cache hands back the float16-rounded keys, the new chunk included, so a paged cache in float16 (L8.3) can be compared with it exactly (test_float16_cache_rounds_what_it_stores).
  • Byte-level tokens split characters; the incremental detokenizer emits text only when it no longer ends in U+FFFD, and the pieces always concatenate to decode(all ids) (test_incremental_decode_concat_equals_decode).
  • Streaming must hold back any tail that could still become a stop string, so a client never receives part of the text that is later cut (test_stop_strings_cut_and_stream_safely).
Terminal window
ol start L8.2 # stubs kvcache.py and generate.py into your repo
ol tests L8.2 # read the test catalog first
ol check L8.2 # course tests, then your tests graded by mutation
ol mutate L8.2 # the full mutation grade of your tests
ol check L8.2 --ref-deps # only if you skipped a dependency
ol diff L8.2 # after passing: your code against the reference

L7.9 gives you a Llama-family model that reproduces SmolLM2’s logits, and L8.1 a sampler that picks a token from them. Put the two in a loop and generation works, slowly: to produce token 200 the model re-runs over all 200 positions, so 256 new tokens cost about 2562/2256^2/2 token-forwards instead of 256, and decode speed falls as the text grows. The milestone measures it (MS-L8 step 4 wants a speedup of at least 5). The fix is the KV cache, and it has to be exact: cached generation must produce the same tokens as recomputation, or every later optimization (paging, prefix sharing, speculative decoding) is measured against a moving target. Text must also stream out as it is generated, and byte-level BPE splits “é” and ”🙂” across tokens, so naive per-token decoding prints U+FFFD to the user. This module writes the cache, the loop, and the streaming detokenizer, each checked against its slow, obviously correct counterpart.

SymbolMeaningType / shape
LLnumber of layersinteger
BBbatch sizeinteger
HkvH_{kv}, dhd_hkey/value heads and head width (GQA, L7.5)integers
TTpositions in the current chunk (the prompt at prefill, 1 at decode)integer
sscommitted length: positions every layer holdsinteger
K(ℓ),V(ℓ)K^{(\ell)}, V^{(\ell)}layer ℓ\ell‘s cached keys and values, [B,Hkv,s,dh][B, H_{kv}, s, d_h]arrays
Nmax⁡N_{\max}preallocated capacity (max_len)integer
qoffq_{\text{off}}absolute position of a chunk’s first query, =s= sinteger
ids, PPgenerated ids and prompt idslists

Encode the prompt to PP, run the model once over all of it (prefill: T=∣P∣T = \lvert P\rvert, positions 0..∣P∣−10..\lvert P\rvert - 1), take the last position’s logits. Then repeat: sample with L8.1 (sample(logits, p, history=ids, rng, prompt=P), rng = request_rng(seed)), stop on EOS, a stop string, or the budget, and otherwise run the model on that one token (decode: T=1T = 1, position ∣P∣+∣ids∣−1\lvert P\rvert + \lvert\text{ids}\rvert - 1). The model only ever sees the new tokens; the cache supplies the past. Without a seed the request uses seed 0, so runs are reproducible.

In layer ℓ\ell the key and value of position jj depend only on tokens 0..j0..j (causal attention, L5.2), so they never change once computed. Caching them and appending new ones produces exactly the same KK and VV a recompute would build; the query of the new token attends to the same keys with the same mask. The only difference is floating-point order: the recompute evaluates QK⊤QK^\top for all rows at once, the decode step for one row, and BLAS may sum differently, which is why the comparison is to 10−510^{-5} and not bitwise. Feeding the whole context back into a cache that already holds it double-counts every key and overflows the storage.

2.3 Committed length, positions, and the mask

Section titled “2.3 Committed length, positions, and the mask”

A forward pass appends layer by layer. If ss moved as soon as layer 0 appended, layer 1 would place the same chunk one chunk later. So each layer keeps its own fill count and seq_len() returns s=min⁡ℓfillℓs = \min_\ell \text{fill}_\ell: it advances only after the last layer. seq_len(layer) returns one layer’s count, the convention of L7.5’s ConcatKVCache that L7.7’s mask and L7.9’s forward read; between forward passes all counts are equal. A model reads positions(T) =[s,s+T)= [s, s + T) (what RoPE rotates by) and mask(T) once, at the start of the pass. Query ii of the chunk sits at absolute position s+is + i and may see keys j≤s+ij \le s + i:

mask(T)ij=[ j≤s+i ],0≤i<T,  0≤j<s+T,\text{mask}(T)_{ij} = [\, j \le s + i \,], \qquad 0 \le i < T,\; 0 \le j < s + T,

L5.2’s causal_mask(T, s + T, q_offset=s). With qoff=0q_{\text{off}} = 0 instead, a decode query sees only key 0.

Speculative decoding (L8.6) appends draft positions, asks the target to verify them, and keeps only an accepted prefix. truncate(n) sets every layer’s fill count to nn; the next append overwrites from nn, and the result is the cache that never saw the rejected positions. Preallocated storage makes this free: nothing is freed or copied.

The cache preallocates 2LBHkvNmax⁡dh2 L B H_{kv} N_{\max} d_h elements (nbytes), so a 135M-parameter model at 2048 tokens holds 94 MB of float32 per sequence. A float16 cache halves it. Values are converted on the way in, and update returns the stored values as float32: the past and the new chunk both rounded. Returning the raw new chunk next to a rounded past mixes two precisions in one attention row, and the paged cache of L8.3, which stores float16 blocks, would no longer match this one.

MLA (L7.6) caches one compressed latent c∈Rrc \in \mathbb{R}^{r} and one shared RoPE key kR∈RdRk^R \in \mathbb{R}^{d_R} per position instead of per-head KK and VV. LatentCache follows the same rules (per-layer fill, committed length, rollback) on [B,T,r][B, T, r] and [B,T,dR][B, T, d_R] chunks; its nbytes is LBNmax⁡(r+dR)L B N_{\max} (r + d_R) times the item size.

A byte-level BPE token is a byte string; “é” is C3 A9, and a tokenizer may put C3 and A9 in different tokens. Decoding a prefix that ends inside a character produces U+FFFD. The decoder keeps the ids and a window: prefix (context already emitted) and read (end of emitted text). On each push it decodes ids[prefix:read] and ids[prefix:]; if the second is longer and does not end in U+FFFD, the new text is final: emit the difference and slide the window. Otherwise wait. flush emits whatever is left, so a sequence cut off at the end becomes one U+FFFD, and resets. Invariant: the pieces plus the flush equal decode(all ids). A lone continuation byte is never completed; it is emitted as U+FFFD as soon as the next character shows it is invalid.

An id in eos_ids ends generation and is in neither the ids nor the text. A stop string ends it too, and the text is cut before its first occurrence, even when it spans tokens. When streaming, the last characters of the text may be the beginning of a stop string (“here” while waiting for “here!”): they are held back until the next token proves otherwise, so a client never receives text that is later removed. The budget (max_tokens) and the model’s context (max_len) end generation with finish_reason = "length".

The cache. One layer, one kv head of width 2, max_len 4. Prefill a chunk of two positions with keys [[1,2],[3,4]][[1, 2], [3, 4]]: storage positions 0 and 1 are filled, the fill count is 2, update returns both rows, and s=2s = 2; the next token’s position is 2. Decode one position with key [5,6][5, 6]: it is written at position 2, update returns all three rows [[1,2],[3,4],[5,6]][[1, 2], [3, 4], [5, 6]], and s=3s = 3. The next query sits at position 3 and its mask row over keys 0..30..3 is [1,1,1,1][1, 1, 1, 1]. Storage is 2×1×1×1×4×2×4=642 \times 1 \times 1 \times 1 \times 4 \times 2 \times 4 = 64 bytes.

The detokenizer. Push the byte token C3: decode([C3]) is U+FFFD, so nothing is emitted. Push A9: decode([C3, A9]) is “é”, longer than the empty window and not ending in U+FFFD, so “é” is emitted and the window moves past both ids. Push “x”: “x” is emitted at once. flush returns "".

These are the first cases in section 4: test_hand_example and test_hand_example_detokenizer.

python/tinyllm/infer/kvcache.py
class KVCache:
def __init__(self, n_layers, n_kv_heads, d_head, max_len, batch=1, dtype=np.float32)
def seq_len(self, layer=None) -> int # committed (min); or one layer's count
def update(self, layer, k_new, v_new) -> tuple[NDArray, NDArray] # [B, Hkv, T, dh] in, all stored out
def positions(self, t_new) -> NDArray # [seq_len, seq_len + t_new)
def mask(self, t_new) -> NDArray # causal_mask(t, s + t, q_offset=s)
def truncate(self, n) -> None
def nbytes(self) -> int
class LatentCache: ... # the same for MLA's (c, k_rope)
# python/tinyllm/infer/generate.py
def cache_dims(model) -> tuple[int, int, int, int] # n_layers, n_kv_heads, d_head, max_len
class IncrementalDecoder:
def __init__(self, tok); def push(self, token_id) -> str; def flush(self) -> str
@dataclass
class Generation: text; ids; logprobs; timings; stats; finish_reason = "length"
def generate(model, tok, prompt, p, cache="contiguous", kv_dtype=np.float32,
on_text=None, eos_ids=()) -> Generation

model.forward(ids, positions, cache) returns logits [B,T,V][B, T, V] and calls cache.update(layer, k, v) once per layer: the CausalLM protocol in the contract, which L7.9’s model follows. cache is "none" (recompute), "contiguous" (this cache), or a cache object (the paged cache of L8.3).

TestKINDChecksWhy it matters downstream
test_hand_exampleunitsection 3’s cache: returned rows, ss, positions, mask, bytesyou and the test agree on the definitions
test_hand_example_detokenizerunitsection 3’s “é” from two byte tokensstreaming text never shows U+FFFD
test_cached_logits_equal_full_recomputedifferential16 decode steps, logits within 10−510^{-5}, same ids, one prefill then one token per callthe cache is exact
test_chunked_prefill_equals_wholedifferentialchunks of 3, 1, 5, 2 positions equal one forwardchunked prefill (L10.2)
test_l75_attention_through_the_cachedifferentialL7.5’s attention with this cache as its hook equals the full forwardthe seam L7.9 uses
test_llama_generates_the_same_with_and_without_cachedifferentialL7.9’s tiny Llama: same greedy ids, logits within 10−510^{-5} every stepthe model of MS-L8
test_update_copies_and_checksboundaryinputs copied, shapes and capacity checked, nothing changed on errora caller’s buffer reuse is safe
test_seq_len_commits_after_every_layerboundaryss moves only after the last layerevery layer places a chunk at the same position
test_mask_is_l52_causal_with_offsetunitmask equals L5.2’s offset causal maskdecode queries see every cached key
test_truncate_rolls_backpropertytruncate then append equals never having appendedL8.6 rollback
test_float16_cache_rounds_what_it_storesunithalf the bytes; stored values returned, new chunk tooL8.3 compares in float16
test_float16_generation_stays_closedifferentialfloat16 cache keeps logits within 2×10−22 \times 10^{-2} and greedy idsthe memory saving is safe
test_latent_cacheunitMLA’s latent cache follows the same rulesL7.6’s hook
test_incremental_decode_concat_equals_decodepropertypieces plus flush equal decode(all); every piece is finalthe streaming invariant
test_detokenizer_waits_for_whole_charactersboundarya 4-byte emoji, a lone continuation byte, a cut-off sequenceill-formed UTF-8 handled like decode
test_generate_greedy_follows_the_scriptunitids, text, logprobs, counts, finish_reasonthe loop end to end
test_eos_stops_and_is_not_emittedunitEOS stops and is not in ids or textOpenAI semantics
test_stop_strings_cut_and_stream_safelyboundarystops across tokens are cut; streamed pieces never contain a stop’s beginningclients see only final text
test_max_len_limits_generationboundarythe context limit stops with “length”; bad prompts raiseno position past the model’s range
test_seeded_sampling_uses_the_request_streamdifferentialreplaying L8.1 by hand gives the same ids, logprobs, entropyseeded requests reproduce
test_cache_modes_and_dimsboundary“paged” and unknown modes raise; a cache object is used; cache_dims reads HF namesL8.3 and L7.9 plug in

Write python/tests/l8-2-generate/ against the contract only. The oracles are the slow versions: the same model without a cache, decode(all ids) for the detokenizer, and L8.1’s sample replayed by hand for seeded generation. Write a small numpy attention model with rotary positions so that a wrong position or mask changes the logits, and a scripted model whose greedy output you choose, for stops and EOS.

PitfallSymptomCaught by
1. decode tokens at the wrong position (0, or one too far)text degrades after the prompt; RoPE angles wrongtest_cached_logits_equal_full_recompute (mutants s01, s13)
2. committing the length after the first layerlayer 1 places the chunk one chunk latetest_seq_len_commits_after_every_layer (mutants s02, s03)
3. a decode mask without the offsetthe new token sees only the first keytest_mask_is_l52_causal_with_offset (mutant s04)
4. update returning only the new chunkattention over one key: the cache is ignoredtest_hand_example (mutant s05)
5. feeding the whole context into the cache each stepkeys counted twice, then a capacity errortest_cached_logits_equal_full_recompute (mutant s06)
6. truncating one layer onlyrejected drafts survive in deeper layerstest_truncate_rolls_back (mutant s09)
7. emitting text that ends in U+FFFD, decoding tokens one by one, or a flush that keeps the window“é” or U+FFFD on screentest_detokenizer_waits_for_whole_characters (mutants s10, s11, s14)
8. a float16 cache returning the raw new chunkfloat16 paged and contiguous caches disagreetest_float16_cache_rounds_what_it_stores (mutant s12)
9. EOS in the outputan end-of-text marker printed to the usertest_eos_stops_and_is_not_emitted (mutant s16)
10. a stop string left in, streamed early, or held back one character shortthe client sees text that is then removedtest_stop_strings_cut_and_stream_safely (mutants s17, s18, s19)
11. the prompt missing from the repetition penalty, or counted as generatedseeded output differs from the engine’stest_seeded_sampling_uses_the_request_stream (mutants s21, s22)
DirectionModuleHow it uses this
BackL8.1sample, request_rng, and sampled_entropy every step
BackL5.2causal_mask(T, s + T, q_offset=s) is the decode mask
BackL1.2the byte-level BPE whose tokens split characters
BackL7.9LlamaForCausalLM.forward(ids, positions, cache) is the model generate drives; cache_dims reads its config
BackL7.5GQAttention calls cache.update(layer, k, v): this cache is its hook
BackL7.3RopeSpec configures that attention’s rotary positions
BackL0.1Tensor wraps the attention’s input
ForwardL8.3optional: the pure-Python paged cache, compared with this one in float16
ForwardL8.6speculative decoding verifies drafts and rolls back with truncate
ForwardL8.5quantized models generate through the same loop
ForwardL10.5the Rust server ports the incremental detokenizer for SSE

If you skip this module, ol check L8.3 stops with L8.3 needs L8.2; --ref-deps substitutes the reference.

Your pieceProduction equivalentWhat it addsWhere to look
KVCacheHugging Face StaticCachepreallocated per-layer buffers for torch.compile; DynamicCache grows insteadtransformers/cache_utils.py
IncrementalDecoderHugging Face tokenizers DecodeStream, vLLM detokenize_incrementallythe same prefix/read window, in Rust, with special-token handlingtokenizers/src/tokenizer/mod.rs, vllm/transformers_utils/detokenizer_utils.py
generatevLLM LLMEnginecontinuous batching: many requests share one forward pass, each with its own cachevllm/v1/engine/
stop stringsvLLM StopCheckerthe same hold-back of partial stop strings in streamed outputvllm/v1/engine/detokenizer.py