Skip to main content

dsv41_engram_proof/
dsv41_engram_proof.rs

1//! Bounded exact proof for the released V4.1 Engram hash and component path.
2//!
3//! The full CMF is opened only for its embedded tokenizer/header. Engram
4//! component weights are the 664-row compact fixtures; no full table is
5//! decoded or sent to a device.
6//!
7//! Usage:
8//!   dsv41_engram_proof <full.cmf> <engram-reference.json> <component-root> <projected-root> <map.bin>
9
10use cortiq_core::CmfModel;
11use cortiq_engine::dsv41::{
12    Dsv41Cfg, Dsv41Engram, EngramHash, RawFp8Rows, dsv41_apply_engram_for_test,
13    token_map_from_model,
14};
15use cortiq_engine::qtensor::QTensor;
16use serde_json::Value;
17use std::collections::HashMap;
18use std::error::Error;
19use std::fs;
20use std::path::{Path, PathBuf};
21use std::process::Command;
22use std::sync::Arc;
23
24const DIM: usize = 5120;
25const HC_MULT: usize = 4;
26const HASH_COLS: usize = 24;
27
28fn read_bf16(path: &Path) -> Result<Vec<f32>, Box<dyn Error>> {
29    let bytes = fs::read(path)?;
30    if bytes.len() % 2 != 0 {
31        return Err(format!("{} has odd BF16 byte length", path.display()).into());
32    }
33    Ok(bytes
34        .chunks_exact(2)
35        .map(|b| f32::from_bits(u32::from(u16::from_le_bytes([b[0], b[1]])) << 16))
36        .collect())
37}
38
39fn read_i64(path: &Path) -> Result<Vec<usize>, Box<dyn Error>> {
40    let bytes = fs::read(path)?;
41    if bytes.len() % 8 != 0 {
42        return Err(format!("{} has non-i64 length", path.display()).into());
43    }
44    Ok(bytes
45        .chunks_exact(8)
46        .map(|b| i64::from_le_bytes(b.try_into().unwrap()) as usize)
47        .collect())
48}
49
50#[inline]
51fn bf16_roundtrip(value: f32) -> f32 {
52    let bits = value.to_bits();
53    let round = 0x7fff + ((bits >> 16) & 1);
54    f32::from_bits(bits.wrapping_add(round) & 0xffff_0000)
55}
56
57fn compare_bf16(label: &str, got: &[f32], expected: &[f32]) -> Result<(), Box<dyn Error>> {
58    if got.len() != expected.len() {
59        return Err(format!(
60            "{label} length mismatch: runtime={} oracle={}",
61            got.len(),
62            expected.len()
63        )
64        .into());
65    }
66    let mut max_abs = 0.0f32;
67    let mut sum_abs = 0.0f64;
68    let mut exact = 0usize;
69    let mut max_ulp = 0u32;
70    let mut over_one_ulp = 0usize;
71    for (&actual, &want) in got.iter().zip(expected) {
72        let rounded = bf16_roundtrip(actual);
73        let diff = (rounded - want).abs();
74        max_abs = max_abs.max(diff);
75        sum_abs += diff as f64;
76        exact += usize::from(rounded.to_bits() == want.to_bits());
77        let actual_bf16 = (rounded.to_bits() >> 16) as i32;
78        let expected_bf16 = (want.to_bits() >> 16) as i32;
79        let ulp = actual_bf16.abs_diff(expected_bf16);
80        max_ulp = max_ulp.max(ulp);
81        over_one_ulp += usize::from(ulp > 1);
82    }
83    let mean_abs = sum_abs / got.len().max(1) as f64;
84    println!(
85        "{label} values={} exact_bf16={}/{} max_abs={max_abs:.8e} mean_abs={:.8e} max_bf16_ulp={max_ulp} over_1ulp={over_one_ulp}",
86        got.len(),
87        exact,
88        got.len(),
89        mean_abs
90    );
91    // The oracle is a CUDA FP8 GEMM with a tile-dependent FP32 reduction;
92    // this CPU path uses the released F32 container and a scalar reduction.
93    // Compare the externally recorded BF16 oracle with a tight absolute bound
94    // while retaining exact gathered/hash checks and reporting BF16 drift.
95    if max_abs > 2.0e-2 || mean_abs > 1.0e-6 {
96        return Err(format!(
97            "{label} exceeds numeric oracle bound: max_abs={max_abs:.8e} mean_abs={mean_abs:.8e} max_bf16_ulp={max_ulp} over_1ulp={over_one_ulp}"
98        ).into());
99    }
100    Ok(())
101}
102
103fn cfg() -> Dsv41Cfg {
104    Dsv41Cfg {
105        dim: DIM,
106        n_heads: 20,
107        head_dim: 128,
108        rope_head_dim: 64,
109        q_lora_rank: 1536,
110        o_lora_rank: 512,
111        o_groups: 4,
112        hc_mult: HC_MULT,
113        hc_sinkhorn_iters: 20,
114        hc_eps: 1e-6,
115        norm_eps: 1e-20,
116        n_routed_experts: 256,
117        top_k: 8,
118        moe_inter: 1536,
119        gate_temp: 1.0,
120        norm_topk_prob: true,
121        route_scale: 2.5,
122        swiglu_limit: 7.0,
123        window: 128,
124        rope_theta: 10_000.0,
125        compress_rope_theta: 160_000.0,
126        rope_factor: 1.0,
127        original_seq_len: 65_536,
128        beta_fast: 32.0,
129        beta_slow: 1.0,
130        index_heads: 32,
131        index_head_dim: 64,
132        index_topk: 64,
133        candidate_source: 3,
134        candidate_topk_blocks: 2,
135        candidate_block_size: 8,
136        kv_sources: vec![1, 3],
137        index_sources: vec![1, 3],
138        compress_ratios: vec![0, 2, 2, 1, 1],
139        engram_layers: vec![1, 14],
140        engram_vocab: 16_000_000,
141        engram_embeddings: vec![16_000_000, 16_000_000],
142        engram_max_ngram: 4,
143        engram_heads: 8,
144        engram_head_dim: 256,
145        engram_compressed_vocab: 99_092,
146        engram_pad_id: 2,
147        vocab: 129_280,
148    }
149}
150
151fn sha256sum(path: &Path) -> Result<String, Box<dyn Error>> {
152    let output = Command::new("sha256sum").arg(path).output()?;
153    if !output.status.success() {
154        return Err(format!("sha256sum failed for {}", path.display()).into());
155    }
156    let text = String::from_utf8(output.stdout)?;
157    Ok(text
158        .split_whitespace()
159        .next()
160        .ok_or("sha256sum produced no digest")?
161        .to_string())
162}
163
164fn verify_hash_reference(
165    model_path: &Path,
166    reference_path: &Path,
167    map_path: &Path,
168) -> Result<(Value, Vec<Vec<Vec<Vec<usize>>>>), Box<dyn Error>> {
169    let reference: Value = serde_json::from_slice(&fs::read(reference_path)?)?;
170    let model = CmfModel::open(model_path)?;
171    let vocab = model.header.arch.vocab_size;
172    let token_map = token_map_from_model(&model, vocab);
173    let mut bytes = Vec::with_capacity(token_map.len() * 4);
174    for value in &token_map {
175        bytes.extend_from_slice(&value.to_le_bytes());
176    }
177    fs::write(map_path, &bytes)?;
178    let got_sha = sha256sum(map_path)?;
179    let expected_sha = reference["token_map_le_u32_sha256"]
180        .as_str()
181        .ok_or("reference token map SHA is missing")?;
182    println!(
183        "hash token_map vocab={} compressed_vocab={} sha256={got_sha}",
184        token_map.len(),
185        token_map.iter().copied().max().unwrap_or(0) + 1
186    );
187    if got_sha != expected_sha {
188        return Err(
189            format!("token map SHA mismatch: runtime={got_sha} oracle={expected_sha}").into(),
190        );
191    }
192
193    let compressed_vocab = reference["compressed_vocab_size"]
194        .as_u64()
195        .ok_or("reference compressed vocab is missing")? as usize;
196    let hash = EngramHash::new(
197        vec![1, 14],
198        4,
199        8,
200        16_000_000,
201        compressed_vocab,
202        2,
203        token_map,
204    )?;
205    let expected_pad = reference["compressed_pad_id"]
206        .as_i64()
207        .ok_or("reference compressed pad is missing")?;
208    if hash.pad_id != expected_pad {
209        return Err(format!(
210            "pad id mismatch: runtime={} oracle={expected_pad}",
211            hash.pad_id
212        )
213        .into());
214    }
215    let expected_multipliers: Vec<[u64; 4]> =
216        serde_json::from_value(reference["multipliers"].clone())?;
217    let expected_primes: Vec<Vec<Vec<u64>>> = serde_json::from_value(reference["primes"].clone())?;
218    let expected_offsets: Vec<Vec<u64>> = serde_json::from_value(reference["offsets"].clone())?;
219    if hash.multipliers != expected_multipliers {
220        return Err(format!(
221            "multiplier mismatch: runtime={:?} oracle={expected_multipliers:?}",
222            hash.multipliers
223        )
224        .into());
225    }
226    if hash.primes != expected_primes {
227        return Err("prime layout mismatch".into());
228    }
229    let actual_offsets: Vec<Vec<u64>> = hash
230        .offsets
231        .iter()
232        .zip(&hash.primes)
233        .map(|(starts, per_ngram)| {
234            per_ngram
235                .iter()
236                .zip(starts)
237                .flat_map(|(primes, &start)| {
238                    let mut offset = start;
239                    primes.iter().map(move |&prime| {
240                        let current = offset;
241                        offset += prime;
242                        current
243                    })
244                })
245                .collect()
246        })
247        .collect();
248    if actual_offsets != expected_offsets {
249        return Err("offset layout mismatch".into());
250    }
251    println!(
252        "hash layout layers={} primes={} offsets={} multipliers=exact",
253        hash.layer_ids.len(),
254        hash.primes
255            .iter()
256            .map(|x| x.iter().map(Vec::len).sum::<usize>())
257            .sum::<usize>(),
258        hash.offsets.iter().map(Vec::len).sum::<usize>()
259    );
260
261    let cases = reference["cases"]
262        .as_array()
263        .ok_or("reference cases missing")?;
264    let mut all_hashes = Vec::with_capacity(cases.len());
265    for (case_no, case) in cases.iter().enumerate() {
266        let ids: Vec<u32> = serde_json::from_value(case["input_ids"].clone())?;
267        let mask: Option<Vec<bool>> = if case["token_mask"].is_null() {
268            None
269        } else {
270            Some(serde_json::from_value(case["token_mask"].clone())?)
271        };
272        let expected_hashes: Vec<Vec<Vec<usize>>> = serde_json::from_value(case["hashes"].clone())?;
273        let mut state = hash.clone();
274        state.reset();
275        let mut got_hashes = Vec::with_capacity(ids.len());
276        for (pos, &id) in ids.iter().enumerate() {
277            got_hashes.push(state.push(id, mask.as_ref().map(|m| m[pos]).unwrap_or(true)));
278        }
279        if got_hashes != expected_hashes {
280            let first = got_hashes
281                .iter()
282                .zip(&expected_hashes)
283                .enumerate()
284                .find(|(_, (a, b))| a != b)
285                .map(|(i, (a, b))| (i, a, b));
286            return Err(format!("hash case {case_no} mismatch: {first:?}").into());
287        }
288        let cuts: Vec<usize> = serde_json::from_value(case["chunk_ends"].clone())?;
289        let mut chunk_state = hash.clone();
290        chunk_state.reset();
291        let mut chunk_hashes = Vec::with_capacity(ids.len());
292        let mut start = 0;
293        for end in cuts {
294            for pos in start..end {
295                chunk_hashes.push(
296                    chunk_state.push(ids[pos], mask.as_ref().map(|m| m[pos]).unwrap_or(true)),
297                );
298            }
299            start = end;
300        }
301        if start != ids.len() || chunk_hashes != expected_hashes {
302            return Err(format!("hash case {case_no} chunked sequence mismatch").into());
303        }
304        println!(
305            "hash case={} name={} seq={} exact_indices=true chunked=true",
306            case_no,
307            case["name"].as_str().unwrap_or("?"),
308            ids.len()
309        );
310        all_hashes.push(expected_hashes);
311    }
312    Ok((reference, all_hashes))
313}
314
315fn f32_tensor(model: &CmfModel, name: &str) -> Result<Vec<f32>, Box<dyn Error>> {
316    let entry = model
317        .tensor(name)
318        .ok_or_else(|| format!("missing {name}"))?;
319    let bytes = model.entry_bytes(entry);
320    if bytes.len() % 4 != 0 {
321        return Err(format!("{name} is not F32 bytes").into());
322    }
323    Ok(bytes
324        .chunks_exact(4)
325        .map(|b| f32::from_le_bytes(b.try_into().unwrap()))
326        .collect())
327}
328
329fn verify_component_layer(
330    root: &Path,
331    projected_root: &Path,
332    layer: usize,
333    reference_hashes: &[Vec<Vec<Vec<usize>>>],
334) -> Result<(), Box<dyn Error>> {
335    let dir = root.join(format!("raw-layer-{layer}"));
336    let layer_meta: Value =
337        serde_json::from_slice(&fs::read(root.join(format!("layer-{layer}.json")))?)?;
338    let original_rows: Vec<usize> = serde_json::from_value(layer_meta["original_row_ids"].clone())?;
339    let row_to_compact: HashMap<usize, usize> = original_rows
340        .iter()
341        .copied()
342        .enumerate()
343        .map(|(compact, original)| (original, compact))
344        .collect();
345    let prefix = format!("model.layers.{layer}.engram");
346    let cmf = dir.join("engram-component.cmf");
347    let model = Arc::new(CmfModel::open(&cmf)?);
348    let embed = RawFp8Rows::from_model(
349        &model,
350        &format!("{prefix}.embed.weight"),
351        &format!("{prefix}.embed.scale"),
352    )?;
353    let wkv = QTensor::from_model(&model, &format!("{prefix}.wkv.weight"))?;
354    let q_weight = f32_tensor(&model, &format!("{prefix}.q_weight"))?;
355    let k_weight = f32_tensor(&model, &format!("{prefix}.k_weight"))?;
356    let engram = Dsv41Engram {
357        embed,
358        wkv,
359        q_weight,
360        k_weight,
361    };
362    let cases = layer_meta["cases"]
363        .as_array()
364        .ok_or("layer cases missing")?;
365    for (case_no, case) in cases.iter().enumerate() {
366        let original_hashes: Vec<Vec<usize>> =
367            serde_json::from_value(case["original_hash_ids"].clone())?;
368        if original_hashes.len() != reference_hashes[case_no].len() {
369            return Err(format!("layer {layer} case {case_no} sequence length mismatch").into());
370        }
371        let mut expected_compact = Vec::with_capacity(original_hashes.len() * HASH_COLS);
372        for (pos, rows) in original_hashes.iter().enumerate() {
373            if rows.len() != HASH_COLS || reference_hashes[case_no][pos][0].len() != HASH_COLS {
374                return Err(format!("layer {layer} case {case_no} hash width mismatch").into());
375            }
376            let layer_index = if layer == 1 { 0 } else { 1 };
377            if rows != &reference_hashes[case_no][pos][layer_index] {
378                return Err(format!(
379                    "layer {layer} case {case_no} original indices differ from full hash oracle at position {pos}"
380                )
381                .into());
382            }
383            for &row in rows {
384                expected_compact.push(
385                    *row_to_compact
386                        .get(&row)
387                        .ok_or_else(|| format!("layer {layer} missing compact row for {row}"))?,
388                );
389            }
390        }
391        let got_compact = read_i64(&dir.join(format!("case{case_no}_indices.bin")))?;
392        if got_compact != expected_compact {
393            let first = got_compact
394                .iter()
395                .zip(&expected_compact)
396                .enumerate()
397                .find(|(_, (a, b))| a != b);
398            return Err(format!(
399                "layer {layer} case {case_no} compact indices mismatch: {first:?}"
400            )
401            .into());
402        }
403
404        let seq = original_hashes.len();
405        let mut gathered = vec![0.0f32; seq * HASH_COLS * 256];
406        for token in 0..seq {
407            for col in 0..HASH_COLS {
408                engram.embed.row_into(
409                    got_compact[token * HASH_COLS + col],
410                    &mut gathered
411                        [(token * HASH_COLS + col) * 256..(token * HASH_COLS + col + 1) * 256],
412                );
413            }
414        }
415        gathered.iter_mut().for_each(|v| *v = bf16_roundtrip(*v));
416        compare_bf16(
417            &format!("layer={layer} case={case_no} gathered"),
418            &gathered,
419            &read_bf16(&dir.join(format!("case{case_no}_gathered.bin")))?,
420        )?;
421
422        let mut projected = Vec::with_capacity(seq * (DIM * (HC_MULT + 1)));
423        for token in 0..seq {
424            let input = &gathered[token * HASH_COLS * 256..(token + 1) * HASH_COLS * 256];
425            let mut row = vec![0.0f32; DIM * (HC_MULT + 1)];
426            engram.wkv.matvec(input, &mut row, None);
427            row.iter_mut().for_each(|v| *v = bf16_roundtrip(*v));
428            projected.extend_from_slice(&row);
429        }
430        compare_bf16(
431            &format!("layer={layer} case={case_no} projected"),
432            &projected,
433            &read_bf16(&projected_root.join(format!("layer-{layer}/case{case_no}_projected.bin")))?,
434        )?;
435
436        let input = read_bf16(&dir.join(format!("case{case_no}_input.bin")))?;
437        let expected_output = read_bf16(&dir.join(format!("case{case_no}_output.bin")))?;
438        let mask: Option<Vec<bool>> = if case["token_mask"].is_null() {
439            None
440        } else {
441            Some(serde_json::from_value(case["token_mask"].clone())?)
442        };
443        let mut output = Vec::with_capacity(input.len());
444        for token in 0..seq {
445            let mut h = input[token * HC_MULT * DIM..(token + 1) * HC_MULT * DIM].to_vec();
446            dsv41_apply_engram_for_test(
447                &engram,
448                &mut h,
449                &got_compact[token * HASH_COLS..(token + 1) * HASH_COLS],
450                &cfg(),
451                mask.as_ref().map(|m| m[token]).unwrap_or(true),
452            );
453            output.extend_from_slice(&h);
454        }
455        compare_bf16(
456            &format!("layer={layer} case={case_no} output"),
457            &output,
458            &expected_output,
459        )?;
460        println!(
461            "component layer={} case={} name={} compact_indices=true gathered=true projected=true output=true",
462            layer,
463            case_no,
464            case["name"].as_str().unwrap_or("?")
465        );
466    }
467    Ok(())
468}
469
470fn main() -> Result<(), Box<dyn Error>> {
471    let mut args = std::env::args_os().skip(1);
472    let model = PathBuf::from(args.next().ok_or("missing full CMF path")?);
473    let reference = PathBuf::from(args.next().ok_or("missing Engram reference JSON")?);
474    let component_root = PathBuf::from(args.next().ok_or("missing component root")?);
475    let projected_root = PathBuf::from(args.next().ok_or("missing projected root")?);
476    let map_path = PathBuf::from(args.next().ok_or("missing token map output path")?);
477    if args.next().is_some() {
478        return Err("usage: dsv41_engram_proof <full.cmf> <engram-reference.json> <component-root> <projected-root> <map.bin>".into());
479    }
480    let (_reference, reference_hashes) = verify_hash_reference(&model, &reference, &map_path)?;
481    for layer in [1usize, 14] {
482        verify_component_layer(&component_root, &projected_root, layer, &reference_hashes)?;
483    }
484    println!("ENGRAM_PROOF_PASS layers=2 cases=6 hash_cases=3 full_table_decode=false gpu=false");
485    Ok(())
486}