Skip to main content

dsv41_engram_component/
dsv41_engram_component.rs

1//! Compare the isolated V4.1 Engram path with the released component oracle.
2//!
3//! The remote fixture directory is produced from the pinned official
4//! safetensors by `export_engram_component_raw.py`. It contains only the
5//! Engram tensors and the official per-case input/output arrays; the helper
6//! writes a small CMF once, then exercises the same mmap-backed row lookup
7//! used by the full loader.
8
9use cortiq_core::{CmfModel, TensorDtype, TensorSpec};
10use cortiq_engine::dsv41::{Dsv41Cfg, Dsv41Engram, RawFp8Rows, dsv41_apply_engram_for_test};
11use cortiq_engine::qtensor::QTensor;
12use serde_json::Value;
13use std::error::Error;
14use std::fs;
15use std::path::{Path, PathBuf};
16use std::sync::Arc;
17
18fn read_bf16(path: &Path) -> Result<Vec<f32>, Box<dyn Error>> {
19    let bytes = fs::read(path)?;
20    if bytes.len() % 2 != 0 {
21        return Err(format!("{} has odd BF16 byte length", path.display()).into());
22    }
23    Ok(bytes
24        .chunks_exact(2)
25        .map(|b| f32::from_bits(u32::from(u16::from_le_bytes([b[0], b[1]])) << 16))
26        .collect())
27}
28
29fn write_fixture(dir: &Path, cmf: &Path) -> Result<(), Box<dyn Error>> {
30    if cmf.exists() {
31        return Ok(());
32    }
33    let base = CmfModel::open(dir.parent().unwrap().join("../tiny-reference-f16.cmf"))?;
34    let layer = dir
35        .file_name()
36        .and_then(|s| s.to_str())
37        .and_then(|s| s.strip_prefix("raw-layer-"))
38        .ok_or("fixture directory must be named raw-layer-N")?;
39    let prefix = format!("model.layers.{layer}.engram");
40    let mut tensors = Vec::new();
41    let push = |tensors: &mut Vec<TensorSpec>,
42                name: &str,
43                dtype: TensorDtype,
44                shape: &[usize],
45                file: &str| {
46        tensors.push(TensorSpec {
47            name: name.to_string(),
48            dtype,
49            shape: shape.to_vec(),
50            data: fs::read(dir.join(file)).expect("raw Engram fixture is readable"),
51        });
52    };
53    push(
54        &mut tensors,
55        &format!("{prefix}.embed.weight"),
56        TensorDtype::U8,
57        &[664, 256],
58        "embed_weight.bin",
59    );
60    push(
61        &mut tensors,
62        &format!("{prefix}.embed.scale"),
63        TensorDtype::U8,
64        &[664, 8],
65        "embed_scale.bin",
66    );
67    push(
68        &mut tensors,
69        &format!("{prefix}.wkv.weight"),
70        TensorDtype::F32,
71        &[25600, 6144],
72        "wkv_weight.bin",
73    );
74    push(
75        &mut tensors,
76        &format!("{prefix}.q_weight"),
77        TensorDtype::F32,
78        &[4, 5120],
79        "q_weight.bin",
80    );
81    push(
82        &mut tensors,
83        &format!("{prefix}.k_weight"),
84        TensorDtype::F32,
85        &[4, 5120],
86        "k_weight.bin",
87    );
88    CmfModel::write(cmf, &base.header, &tensors, None, None)?;
89    Ok(())
90}
91
92fn f32_tensor(model: &CmfModel, name: &str) -> Result<Vec<f32>, Box<dyn Error>> {
93    let entry = model
94        .tensor(name)
95        .ok_or_else(|| format!("missing {name}"))?;
96    let bytes = model.entry_bytes(entry);
97    if bytes.len() % 4 != 0 {
98        return Err(format!("{name} is not F32 bytes").into());
99    }
100    Ok(bytes
101        .chunks_exact(4)
102        .map(|b| f32::from_le_bytes([b[0], b[1], b[2], b[3]]))
103        .collect())
104}
105
106fn cfg() -> Dsv41Cfg {
107    Dsv41Cfg {
108        dim: 5120,
109        n_heads: 20,
110        head_dim: 128,
111        rope_head_dim: 64,
112        q_lora_rank: 1536,
113        o_lora_rank: 512,
114        o_groups: 4,
115        hc_mult: 4,
116        hc_sinkhorn_iters: 20,
117        hc_eps: 1e-6,
118        norm_eps: 1e-20,
119        n_routed_experts: 256,
120        top_k: 8,
121        moe_inter: 1536,
122        gate_temp: 1.0,
123        norm_topk_prob: true,
124        route_scale: 2.5,
125        swiglu_limit: 7.0,
126        window: 128,
127        rope_theta: 10_000.0,
128        compress_rope_theta: 160_000.0,
129        rope_factor: 1.0,
130        original_seq_len: 65_536,
131        beta_fast: 32.0,
132        beta_slow: 1.0,
133        index_heads: 32,
134        index_head_dim: 64,
135        index_topk: 64,
136        candidate_source: 3,
137        candidate_topk_blocks: 2,
138        candidate_block_size: 8,
139        kv_sources: vec![1, 3],
140        index_sources: vec![1, 3],
141        compress_ratios: vec![0, 2, 2, 1, 1],
142        engram_layers: vec![1, 14],
143        engram_vocab: 16_000_000,
144        engram_embeddings: vec![16_000_000, 16_000_000],
145        engram_max_ngram: 4,
146        engram_heads: 8,
147        engram_head_dim: 256,
148        engram_compressed_vocab: 99_092,
149        engram_pad_id: 2,
150        vocab: 129_280,
151    }
152}
153
154fn main() -> Result<(), Box<dyn Error>> {
155    let mut args = std::env::args_os().skip(1);
156    let dir = PathBuf::from(args.next().ok_or("missing raw fixture directory")?);
157    let layer = dir
158        .file_name()
159        .and_then(|s| s.to_str())
160        .and_then(|s| s.strip_prefix("raw-layer-"))
161        .ok_or("fixture directory must be named raw-layer-N")?;
162    let prefix = format!("model.layers.{layer}.engram");
163    let cmf = dir.join("engram-component.cmf");
164    write_fixture(&dir, &cmf)?;
165    let model = Arc::new(CmfModel::open(&cmf)?);
166    let embed = RawFp8Rows::from_model(
167        &model,
168        &format!("{prefix}.embed.weight"),
169        &format!("{prefix}.embed.scale"),
170    )?;
171    let wkv = QTensor::from_model(&model, &format!("{prefix}.wkv.weight"))?;
172    let q_weight = f32_tensor(&model, &format!("{prefix}.q_weight"))?;
173    let k_weight = f32_tensor(&model, &format!("{prefix}.k_weight"))?;
174    let engram = Dsv41Engram {
175        embed,
176        wkv,
177        q_weight,
178        k_weight,
179    };
180    let cfg = cfg();
181    let meta: Value = serde_json::from_slice(&fs::read(dir.join("manifest.json"))?)?;
182    let mut max_abs = 0.0f32;
183    let mut sum_abs = 0.0f64;
184    let mut count = 0usize;
185    let mut cases = 0usize;
186    for case in 0..3 {
187        let prefix = format!("case{case}");
188        let input_shape = meta[&format!("{prefix}.input")]["shape"]
189            .as_array()
190            .unwrap();
191        let seq = input_shape[1].as_u64().unwrap() as usize;
192        let input = read_bf16(&dir.join(format!("{prefix}_input.bin")))?;
193        let expected = read_bf16(&dir.join(format!("{prefix}_output.bin")))?;
194        let index_bytes = fs::read(dir.join(format!("{prefix}_indices.bin")))?;
195        let mut indices = Vec::with_capacity(index_bytes.len() / 8);
196        for b in index_bytes.chunks_exact(8) {
197            indices.push(u64::from_le_bytes(b.try_into().unwrap()) as usize);
198        }
199        let token_mask: Vec<bool> = if case == 1 {
200            vec![
201                true, true, true, true, true, false, false, false, true, true, true, true, true,
202                true, true, true, true, true, true, true,
203            ]
204        } else {
205            vec![true; seq]
206        };
207        for token in 0..seq {
208            let mut h = input[token * 4 * 5120..(token + 1) * 4 * 5120].to_vec();
209            let hashes = &indices[token * 24..(token + 1) * 24];
210            dsv41_apply_engram_for_test(&engram, &mut h, hashes, &cfg, token_mask[token]);
211            for (&got, &want) in h
212                .iter()
213                .zip(&expected[token * 4 * 5120..(token + 1) * 4 * 5120])
214            {
215                let d = (got - want).abs();
216                max_abs = max_abs.max(d);
217                sum_abs += d as f64;
218                count += 1;
219            }
220        }
221        cases += 1;
222        println!("case={case} seq={seq} cumulative_max_abs={max_abs:.8e}",);
223    }
224    println!(
225        "summary cases={cases} values={count} max_abs={max_abs:.8e} mean_abs={:.8e}",
226        sum_abs / count.max(1) as f64
227    );
228    Ok(())
229}