1use 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 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}