Candle model runner (Llama + bigram, int4), Rust KV cache, sampler and PCG32
Overview
Section titled “Overview”| Module | L10.1 · build · Rust · Pass 7 · 14 to 20 h |
| You build | rust/crates/tl-engine/src/kv.rs: the bounded Rust KV cache · rust/crates/tl-engine/src/quant.rs (f16, bf16, int4) · model.rs (config.json, the memory-mapped safetensors reader, the weights) · forward.rs (the Llama and bigram forwards built with Candle tensors) · runner.rs (ModelRunner, EngineConfig, the KV pool) · sample.rs (PCG32 and the sampler) · the module lines of the tl-engine/src/lib.rs crate root |
| Contract | files: formats/safetensors.md, formats/config.schema.json, formats/kv-block.md · determinism: spec/sampling.md, spec/pcg32.md |
| Tests | course/tests/rust/l10_1.rs, 24 tests (what they check: section 4); parity suites ol parity sampler rng (your Rust against your Python, through one golden file) |
| Needs | L10.0 the checkpoint contract (chapter) · reading: lang.04 Rust (primer), L8.1 the Python sampler you port, M06.3 PCG32, L7.9 the Llama model you port, L8.5 the int4 scheme, M09.4 f16 rounding, L9.7 the Python twin of this runner, L1.5 tl-tok · or --ref-deps |
| Used by | L10.2 (the scheduler runs this runner), L10.3 (chunked prefill feeds it), L10.4 (the block manager hands it blocks), L10.5 (the engine loop and the server) · later: L10.6, L10.8, L10.9 |
| Milestone | MS-L10 (the engine’s greedy stream equals your Python’s; ol parity sampler) |
| Optional depth | Rust reference, Drop (free); safetensors format (free); Hugging Face modeling_llama.py (free); Kwon et al. 2023, PagedAttention (free) |
Key Takeaways
Section titled “Key Takeaways”- The Rust engine owns the model graph and uses Candle for tensor operations. The learner writes the Llama layers and owns a bounded Rust KV cache.
KvPoolowns fixed-size f16 K and V slabs, block references, the prefix hash index, and LRU eviction (kv_pool_bounds_and_cache_lifecycle).- One step runs a batch of sequences, each a run of new tokens at absolute positions; their K and V go into the pool as f16 and are read back for attention, so the logits match Hugging Face’s float32 forward within the f16 bound (
tiny_llama_logits_match_hf). - The model forward is batch-invariant bit for bit and chunk-invariant within the test tolerance, so batching sequences together or splitting a prompt into steps preserves logits (
batched_forward_equals_single,incremental_decode_equals_full_prefill). - The sampler follows spec/sampling.md step by step in f64, so the same logits and seed give your Python’s id and logprob exactly (
sampler_matches_l81_golden,hand_example_sampling).
How to work this chapter
Section titled “How to work this chapter”ol start L10.1 # stubs tl-engine kv, quant, model, forward, runner, sampleol tests L10.1 # read the test catalog first: rung R0, you write no graded tests hereol check L10.1 # exit code is the verdictol check L10.1 --ref-deps # only if a referenced module is not passing yetol parity sampler rng # your Rust sampler and PCG32 against your Python, through the golden filesol start never rewrites your files. Add one pub mod line per new file to tl-engine/src/lib.rs (yours since L8.4). Add Candle, anyhow, memmap2, and serde_json from contracts/allowed-deps.toml to tl-engine/Cargo.toml.
1. Why now
Section titled “1. Why now”Your tracer engine (L10.0) serves one model, a byte bigram, with one request per thread and a sampler that is not the one your Python uses. Pass 7 turns it into a real inference engine, and everything after it in this part (scheduling, chunked prefill, prefix caching, the OpenAI server) needs one thing first: a component that loads a Llama checkpoint, runs a batch of sequences through it with Candle tensor operations, keeps each sequence’s past in the Rust KV pool, and turns logits into tokens exactly as your Python does. That component is the model runner. Rust owns the graph and the production math path. Candle supplies tensor operations, while this module implements the layers and cache.
2. Principles
Section titled “2. Principles”2.1 Three layers in one engine
Section titled “2.1 Three layers in one engine”tl-engine runner.rs ModelRunner::forward(&ForwardBatch) -> Logits forward.rs llama_forward: embedding, per layer [norm, q k v, rope, write KV, attention, o, norm, swiglu], norm, LM head model.rs config.json, model.safetensors (mmap), Weights quant.rs f16 / bf16 / int4 numbers sample.rs Pcg32, SamplingParams, sample()candle-core + nn tensor operations: matmul, embedding, RMSNorm, softmaxkv.rs Rust-owned f16 blocks, references, hash index, LRU2.2 Candle tensors and Rust-owned state
Section titled “2.2 Candle tensors and Rust-owned state”Candle provides tensor storage, device placement, matrix multiplication, softmax, RMSNorm, and embedding lookup. The engine uses candle-core and candle-nn 0.10.2. The learner still writes the model structure: the layer order, RoPE positions, grouped-query attention, residual connections, SwiGLU, and final norm. CI uses the CPU device; macOS can select Metal.
The KV cache is a separate Rust data structure. KvPool allocates fixed-size K and V slabs as f16 values, tracks references, and keeps released full blocks in an LRU cache when prefix caching is enabled. Its methods return Result for invalid dimensions, exhausted capacity, invalid block ids, and format errors. No C library or foreign-function boundary is part of the model runner.
2.3 The model directory
Section titled “2.3 The model directory”config.json says what to build (tl_arch: bigram or llama; tl_tokenizer; the shape fields of Hugging Face’s Llama config). Anything this engine cannot run correctly, a scaled RoPE (rope_type other than default), a sliding window, attention biases, is refused at load with a message, never served wrong.
model.safetensors is read through a memory map (memmap2): the operating system maps the file into the address space and pages bytes in on first touch, so a 270 MB checkpoint is not copied before the first request. The header rules are those of L10.0: a u64 little-endian header length , bytes of JSON, then the data buffer, with every data_offsets range counted from the start of the data buffer and the ranges tiling it exactly.
Weights arrive as F32, F16, or BF16 and are widened to f32 at load:
with subnormals for . The reverse direction (f32 to f16, for the KV cache) rounds to nearest, ties to even: keep the top 10 mantissa bits, and add one when the dropped bits are above half, or exactly half with the kept value odd.
2.4 The Llama forward, one step
Section titled “2.4 The Llama forward, one step”| Symbol | Meaning | Type / shape |
|---|---|---|
| sequences in the step | integer | |
| , | new tokens of sequence and the absolute position of its first one | integers |
| tokens in the step | integer | |
hidden size (hidden_size) | integer | |
| , , | query heads, key/value heads, head size | integers |
| hidden states of every new token, in batch order | f32 | |
attention projections, stored [out, in] | f32 or int4 | |
| absolute position of token : for the -th new token of | i32 | |
rms_norm_eps | f32 | |
token positions per KV block (block_size) | integer |
Each step computes, for every layer:
then rotates and by RoPE at (half-split layout), writes each token’s and into its sequence’s blocks (position lives in block of the block table, slot ), and for each sequence gathers its keys and values for positions to and uses Candle matrix multiplication and softmax with q_offset and a causal mask:
with before the MLP. After the last layer, only each sequence’s last token goes through the final norm and the LM head (the embedding matrix itself when tie_word_embeddings is true), giving one row of logits per sequence.
Two properties carry the rest of this part. Batch invariance: a token’s logits must not depend on which other tokens share its step. Candle picks its reduction path from the shape: a one-row product takes a matrix-vector kernel and a taller one a blocked kernel that adds in another order, so on x86 the same row comes out a few ulps apart depending on the batch. The runner therefore sends every row through the same product (matmul_rows), and query attends over exactly its visible keys instead of a masked full row. Then a row is computed by the same operations in the same order whether it runs alone or in a batch, and batched_forward_equals_single compares bit for bit. Chunk invariance: positions are absolute and the cache holds the same f16 keys either way, so prefilling a prompt in one step or in pieces agrees within the tested tolerance. Without row-wise products this fails on Linux, not just in the last bits: a one-ulp change in a key can round to a different f16, which moves the logits by about .
The f16 bound. K and V pass through f16, which rounds with relative error at most . Through two layers this perturbs the logits of the tiny test model (magnitudes up to about 20) by about , measured; the tests allow of the largest logit, . Your Python reference (L7.9) keeps f32 keys, so this bound, not 1e-4, is what an f16 cache can promise.
2.5 Int4 weights
Section titled “2.5 Int4 weights”A quantized linear weight with groups of columns stores, per group, one f16 scale and four-bit integers (L8.5, formats/safetensors.md):
with rint rounding half to even and the f16 value actually stored, so . Two values share a byte: column in the low nibble, in the high nibble, each in 4-bit two’s complement. QLinear::forward unpacks the quantized values, widens the effective weight to f32, and uses Candle matrix multiplication. The runner loads int4 two ways: from an int4 file (<name>.qweight, <name>.scales, metadata quant = "int4-g<g>-sym"), or by quantizing f32 weights at load (EngineConfig.quant = Some(Quant::Int4 { group })).
2.6 The sampler and PCG32
Section titled “2.6 The sampler and PCG32”sample.rs ports your L8.1 sampler and M06.3 generator, and parity means bit-identical: the same token id and the same f64 logprob from the same logits and seed. spec/sampling.md fixes the order: widen to f64; repetition penalty over the distinct ids of prompt and output ( when positive, otherwise); presence and frequency over the output only (); the logprob is at this point; greedy () takes the argmax with ties to the lowest id and no draw; otherwise divide by , keep the top by , keep the nucleus up to and including the token that reaches , keep ids with , softmax over the kept ids in ascending id order, draw one , and walk the cumulative sum in ascending id order to the first id with . Every sum is a left-to-right loop, so Rust and Python add the same numbers in the same order.
Each request gets its own generator, stream(seed, sample): , then pcg32_srandom_r(child_seed, 4). One uniform takes two outputs : .
3. Worked example by hand
Section titled “3. Worked example by hand”One sampled token (spec/sampling.md, hand_example_sampling). Logits , no history, , top-k 3, top-p 0.8, seed 0. Logprobs first: , , so ids 1 and 3 have logprob . Top-k: order by (logit desc, id asc) is ; keep . Top-p over the kept: , , ; walking : , then , so keep and renormalize to . The generator: , and the first uniform of stream(0, sample) is . Walk ids in ascending order: after id 1, , not above ; after id 3, : the token is id 3, logprob .
Int4 (int4_hand_example). Weights in one group of 4: , , and (0.2 is not a binary fraction; f16 keeps 10 mantissa bits). Dividing by the stored scale: , , , . With the unrounded the last would be : the stored scale changes a value. Packing : is and is , so byte 0 is ; and give .
f16 rounding (f16_rounds_to_nearest_even). lies exactly halfway between (0x3C00) and (0x3C01); ties go to the even mantissa, so the result is 0x3C00. lies halfway between 0x3C01 and 0x3C02: the even one is 0x3C02.
Where a key lands. Blocks of , a sequence with block table : position 37 is block of the table, pool block 9, slot . Its K for KV head starts at element of layer ‘s K slab of block 9.
4. The interface
Section titled “4. The interface”pub struct KvCfg { pub n_blocks: u32, pub block_tokens: u32, pub n_layers: u32, pub n_kv_heads: u32, pub head_dim: u32, pub dtype: i32, pub format: u32 }pub struct KvPool { /* Rust-owned K/V slabs, references, hash index, LRU */ }impl KvPool { pub fn new(cfg: KvCfg) -> Result<KvPool, KvError>; pub fn alloc(&mut self, n: usize) -> Result<Vec<u32>, KvError>; pub fn retain(&mut self, id: u32) -> Result<(), KvError>; pub fn release(&mut self, id: u32) -> Result<(), KvError>; pub fn set_fill(&mut self, id: u32, n: u32) -> Result<(), KvError>; pub fn slab(&self, id: u32, layer: u32, is_v: bool) -> Result<&[u16], KvError>; pub fn slab_mut(&mut self, id: u32, layer: u32, is_v: bool) -> Result<&mut [u16], KvError>; pub fn export(&self, ids: &[u32]) -> Result<Vec<u8>, KvError>; pub fn import(&mut self, buf: &[u8]) -> Result<Vec<u32>, KvError>;}// rust/crates/tl-engine/src/{quant,model,forward,runner,sample}.rspub fn f16_to_f32(h: u16) -> f32; pub fn f32_to_f16(x: f32) -> u16; pub fn bf16_to_f32(h: u16) -> f32;pub fn pack_int4(q: &[i8], rows: usize, cols: usize) -> Result<Vec<u8>, String>; pub fn unpack_int4(packed: &[u8]) -> Vec<i8>;pub struct QLinear { pub out: usize, pub inp: usize, pub group: usize, pub qweight: Vec<u8>, pub scales: Vec<u16> }impl QLinear { pub fn quantize(w: &[f32], out: usize, inp: usize, group: usize) -> Result<QLinear, String>; pub fn dequantize(&self) -> Vec<f32>; pub fn forward(&self, x: &[f32], m: usize, y: &mut [f32]) -> Result<(), TlError>; }
pub struct ModelConfig { pub arch: Arch, pub tokenizer: TokenizerKind, pub vocab_size: usize, pub hidden_size: usize, /* ... */ pub eos_token_ids: Vec<u32> }impl ModelConfig { pub fn from_json(text: &str) -> Result<ModelConfig, String>; pub fn load(dir: &Path) -> Result<ModelConfig, String>; }pub fn parse_layout(bytes: &[u8]) -> Result<Layout, String>;pub struct SafeTensors { pub layout: Layout /* + the map */ }impl SafeTensors { pub fn open(path: &Path) -> Result<SafeTensors, String>; pub fn f32(&self, name: &str, shape: &[usize]) -> Result<Vec<f32>, String>; pub fn q4(&self, name: &str, out: usize, inp: usize, group: usize) -> Result<QLinear, String>; }pub enum Linear { F32 { w: Vec<f32>, out: usize, inp: usize }, Q4(QLinear) } // forward: y = x @ W^Tpub enum Weights { Bigram(Vec<f32>), Llama(LlamaWeights) }
pub struct ForwardSeq<'a> { pub tokens: &'a [u32], pub start: usize, pub blocks: &'a [u32] }pub struct ForwardBatch<'a> { pub seqs: Vec<ForwardSeq<'a>> }pub struct Logits { pub vocab: usize, pub data: Vec<f32> } // rows(), row(i): one per sequence
pub struct EngineConfig { pub max_batch_tokens: usize, pub max_seqs: usize, pub prefill_chunk: usize, pub prefix_cache: PrefixCache, pub kv: KvConfig, pub policy: SchedPolicy, pub threads: usize, pub quant: Option<Quant> }pub type SharedPool = Arc<Mutex<KvPool>>;impl ModelRunner { pub fn load(dir: &Path, cfg: &EngineConfig) -> anyhow::Result<Self>; pub fn forward(&mut self, batch: &ForwardBatch) -> anyhow::Result<Logits>; pub fn run_once(&mut self, tokens: &[u32]) -> anyhow::Result<(Vec<f32>, Vec<f32>)>; // hidden [T, d], last logits pub fn pool(&self) -> SharedPool; pub fn block_tokens(&self) -> usize; pub fn config(&self) -> &ModelConfig;}
pub struct Pcg32 { /* state, inc */ }impl Pcg32 { pub fn new(seed: u64, seq: u64) -> Pcg32; pub fn next_u32(&mut self) -> u32; pub fn uniform_f64(&mut self) -> f64; }pub fn child_seed(seed: u64, purpose: u64) -> u64; pub fn stream(seed: u64, purpose: u64) -> Pcg32; // PURPOSE_SAMPLE = 4pub struct SamplingParams { pub temperature: f64, pub top_k: usize, pub top_p: f64, pub min_p: f64, pub repetition_penalty: f64, pub presence_penalty: f64, pub frequency_penalty: f64 }pub fn sample(logits: &[f32], p: &SamplingParams, prompt: &[u32], output: &[u32], rng: &mut Pcg32) -> (u32, f64);forward writes K and V for the new tokens of each sequence into the blocks its table names; the caller allocates the blocks (L10.2 and L10.4 do from here on). run_once takes temporary blocks and gives them back, for embeddings and tests.
What the tests check
Section titled “What the tests check”| Test | KIND | Checks | Why it matters downstream |
|---|---|---|---|
hand_example_sampling | unit | section 3: top-k 3 and top-p 0.8 keep ids 1 and 3, , token 3 with logprob | the spec’s worked example, end to end |
pcg32_reference_vectors | golden | pcg32(42) demo line, uniform_f64, child_seed, stream outputs from pcg32.vectors.json | one generator in every language (D10) |
sampler_matches_l81_golden | differential, golden | 12 cases x 12 tokens: ids and f64 logprobs bit for bit against L8.1’s golden | ol parity sampler; the engine’s seeded streams equal your Python’s |
greedy_takes_no_draw_and_sampling_takes_one | property | greedy leaves the generator untouched; a sampled token takes one uniform_f64 | draw accounting for the disaggregated hand-off (L10.6) |
top_p_keeps_the_crossing_token | boundary | mass exactly at keeps the crossing id | nucleus edge, the most common sampler bug |
penalties_follow_hf_and_openai_semantics | unit | repetition divides positives and multiplies negatives over prompt and output; presence and frequency over output only | the OpenAI request fields mean what clients expect |
seeded_sampling_matches_the_distribution | statistical | 4000 draws fit 0.4, 0.3, 0.2, 0.1 (chi-square below 16.27); a seed repeats | the sampler draws from the right distribution |
kv_pool_bounds_and_cache_lifecycle | boundary | capacity, references, registration, release, and lookup preserve pool invariants | block manager and prefix cache rely on these transitions |
kv_pool_evicts_the_oldest_cached_block | boundary | exhausted allocation evicts the oldest cached prefix | cache policy preserves recent prefixes |
kv_pool_rejects_zero_capacity | boundary | zero blocks is rejected during construction | invalid engine configuration fails early |
kv_pool_exports_and_imports_rust_owned_blocks | differential | f16 cache data survives the transfer envelope round trip | disaggregated decode resumes the same state |
kv_pool_bounds_and_cache_lifecycle | boundary | capacity is all-or-nothing, refs balance, full blocks register, release caches, lookup reacquires | scheduler and prefix cache share one pool |
kv_pool_exports_and_imports_rust_owned_blocks | differential | a written f16 value survives an export/import round trip | disaggregated decode resumes with the same cache |
f16_rounds_to_nearest_even | boundary | ties, overflow, subnormals, NaN; every f16 pattern round-trips | the KV cache stores f16 |
int4_hand_example | unit | section 3: nibble order and the stored-scale rule | int4 files written by your Python load here |
q4_linear_matches_its_dequantized_weights | differential | Candle output equals x times the dequantized weights within float tolerance | the packed layout means what the format says |
safetensors_reader_rules | boundary | offsets from the data buffer, F16 and BF16 widening, gaps and trailing bytes refused | every checkpoint goes through this reader |
config_refuses_what_the_engine_cannot_run | boundary | scaled RoPE, sliding window, unknown arch, missing tokenizer refused at load | wrong models fail loudly at start |
bigram_checkpoint_still_serves | conformance | the tracer checkpoint loads and greedy continues bcd | the Pass 1 smoke stays green (D32) |
tiny_llama_logits_match_hf | golden | last-token logits within of HF’s float32 forward | the whole graph: GQA, tied embeddings, BF16 file |
tiny_llama_greedy_matches_hf | golden | 32 greedy tokens per prompt equal HF’s under the near-tie rule | what MS-L10 compares with your Python |
incremental_decode_equals_full_prefill | differential | prefill then one token per step equals one prefill, within floating-point tolerance | chunked prefill and preemption by recompute |
batched_forward_equals_single | differential | two sequences in one step agree with their alone logits within float tolerance | continuous batching changes nothing |
run_once_returns_its_blocks | unit | five one-off forwards leave every block free | embeddings do not leak KV |
int4_runner_matches_its_dequantized_model | differential | quantize-at-load equals an int4 file bit for bit and its dequantized f32 twin closely | both int4 paths agree |
forward_refuses_bad_batches | boundary | a short block table or a position past the context is Err | the scheduler’s mistakes surface, not corrupt KV |
5. Pitfalls
Section titled “5. Pitfalls”| Pitfall | Symptom | Caught by |
|---|---|---|
Swapping the shifts of the two draws in uniform_f64 | uniforms that look fine and match nothing | pcg32_reference_vectors (mutant s01) |
| Seeding the sample stream on sequence 54 instead of the purpose id | every seeded stream differs from Python’s | hand_example_sampling (mutant s02) |
| Taking the logprob after temperature | logprobs change with ; OpenAI clients get wrong values | sampler_matches_l81_golden (mutant s03) |
> instead of >= in the nucleus walk | the token that reaches is dropped | top_p_keeps_the_crossing_token (mutant s04) |
| Walking the CDF in probability order | correct distribution, different ids for the same seed | sampler_matches_l81_golden (mutant s05) |
| Presence and frequency penalties over the prompt | prompt words become unlikely in the answer | penalties_follow_hf_and_openai_semantics (mutant s06) |
| Dividing negative logits by the repetition penalty | repeated unlikely tokens become MORE likely | penalties_follow_hf_and_openai_semantics (mutant s07) |
| Truncating instead of rounding to f16 | a drift of half an ulp per cached value | f16_rounds_to_nearest_even (mutant s08) |
| The odd column in the low nibble | int4 files from your Python decode to noise | int4_hand_example (mutant s09) |
| Rounding int4 with the f32 scale, storing the f16 one | some weights off by one level | int4_hand_example (mutant s10) |
data_offsets from the start of the file | garbage weights | safetensors_reader_rules (mutant s11) |
| Reading BF16 as F16 | SmolLM2 weights come out wildly wrong | safetensors_reader_rules (mutant s12) |
| RoPE positions relative to the step | decode tokens all at position 0; output degrades after the prompt | incremental_decode_equals_full_prefill (mutant s13) |
q_offset 0 for a decode step | the new token sees only key 0 | incremental_decode_equals_full_prefill (mutant s14) |
| K written to the V slab | logits far from HF | tiny_llama_logits_match_hf (mutant s15) |
| Skipping the final norm | logits scaled wrong | tiny_llama_logits_match_hf (mutant s16) |
| Not giving temporary blocks back | the pool drains a little per embedding call | run_once_returns_its_blocks (mutant s17) |
| Accepting a zero-block pool | the first allocation fails far from the invalid configuration | kv_pool_rejects_zero_capacity (mutant s19) |
| Evicting the newest prefix instead of the oldest | a useful recent prefix disappears early | kv_pool_evicts_the_oldest_cached_block (mutant s18) |
| Reversing RoPE frequency order | Llama logits differ from the reference | tiny_llama_logits_match_hf (mutant s20) |
| Not checking token ids against the vocabulary | an invalid id reaches the embedding lookup | forward_refuses_bad_batches (mutant s21) |
| silu applied to the up projection | close-looking, wrong logits | tiny_llama_logits_match_hf (mutant s22) |
6. Where it’s used next
Section titled “6. Where it’s used next”| Forward | L10.6 | Registered call site uses this module. |
| Forward | L10.8 | Registered call site uses this module. |
| Forward | L10.9 | Registered call site uses this module. |
| Direction | Module | How it uses this |
|---|---|---|
| Back | L10.0 | checkpoint loading and the byte bigram tracer |
| Back | L7.9, L8.5, M09.4 | the model graph, int4 layout, and f16 conversions reimplemented here |
| Back | L8.1, M06.3 | the sampler and generator this one reproduces bit for bit |
| Back | L7.9 | the Llama graph this one reproduces |
| Forward | L10.2 | the scheduler forms each step’s ForwardBatch |
| Forward | L10.3 | chunked prefill relies on chunk invariance |
| Forward | L10.4 | the block manager owns the pool’s blocks; this runner writes into them |
| Forward | L10.5 | the engine loop samples with sample and a per-request stream(seed, sample) |
If you skip this module, the engine has no model to run past the tracer bigram.
Going further
Section titled “Going further”| Your piece | Production equivalent | What it adds | Where to look |
|---|---|---|---|
| a gather of K and V per step | vLLM PagedAttention | attention reads the blocks in place, no copy | vllm/attention/ |
| f16 KV | fp8 KV with per-head scales | half the memory, a calibrated error | craft.13, formats/kv-block.md v2 |
| int4 groups of 32 | GPTQ, AWQ | error-aware rounding, activation-aware scales | L8.5 going further |
| CPU tensor operations | candle Metal backend | device execution with the same layer graph | Candle documentation |