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