use anyhow::{Context, Result};
use inferencelayer::reference_qwen35::{Qwen35Ref, argmax};
use std::path::PathBuf;
fn read_f32_le(path: &str) -> Result<Vec<f32>> {
let bytes = std::fs::read(path).with_context(|| format!("reading {path}"))?;
Ok(bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect())
}
fn topk(v: &[f32], k: usize) -> Vec<usize> {
let mut idx: Vec<usize> = (0..v.len()).collect();
idx.sort_unstable_by(|&a, &b| v[b].partial_cmp(&v[a]).unwrap());
idx.truncate(k);
idx
}
fn main() -> Result<()> {
let scratch = "/private/tmp/claude-501/-Users-dlo-code-crawl-bee/\
6a3b906d-8642-4873-90fc-2aab8224d53c/scratchpad";
let mut model = PathBuf::from(format!("{scratch}/models/nuextract-3"));
let mut gold = PathBuf::from(format!("{scratch}/qwen35_gold.json"));
let mut args = std::env::args().skip(1);
while let Some(a) = args.next() {
match a.as_str() {
"--model" => model = args.next().context("--model needs a path")?.into(),
"--gold" => gold = args.next().context("--gold needs a path")?.into(),
other => anyhow::bail!("unknown arg {other}"),
}
}
let g: serde_json::Value = serde_json::from_slice(&std::fs::read(&gold)?)?;
let prompt_ids: Vec<u32> = g["prompt_ids"]
.as_array()
.context("prompt_ids")?
.iter()
.map(|x| x.as_u64().unwrap() as u32)
.collect();
let gold_argmax = g["prefill_argmax"].as_u64().unwrap() as u32;
let gold_cont: Vec<u32> = g["continuation"]
.as_array()
.context("continuation")?
.iter()
.map(|x| x.as_u64().unwrap() as u32)
.collect();
let gold_logits = read_f32_le(g["last_logits_bin"].as_str().context("last_logits_bin")?)?;
let gold_layers = g["layer_hidden"].as_array().context("layer_hidden")?;
eprintln!("loading CPU reference from {} ...", model.display());
let t0 = std::time::Instant::now();
let mut m = Qwen35Ref::load(&model)?;
eprintln!(
"loaded {} layers ({} full-attn) in {:.1}s",
m.cfg.n_layers,
m.cfg.layer_is_full.iter().filter(|&&x| x).count(),
t0.elapsed().as_secs_f32()
);
let t1 = std::time::Instant::now();
let mut last_logits = Vec::new();
for (pos, &tok) in prompt_ids.iter().enumerate() {
m.capture_trace = pos == prompt_ids.len() - 1;
last_logits = m.forward(tok, pos);
}
let prefill_dt = t1.elapsed().as_secs_f32();
println!("\n== per-layer last-token hidden (CPU vs gold) ==");
let mut worst_layer = (0usize, 0f32);
for (li, gl) in gold_layers.iter().enumerate() {
let gnorm = gl["norm"].as_f64().unwrap() as f32;
let h = &m.trace[li];
let cnorm = h.iter().map(|v| v * v).sum::<f32>().sqrt();
let rel = (cnorm - gnorm).abs() / gnorm.max(1e-6);
if rel > worst_layer.1 {
worst_layer = (li, rel);
}
if li < 2 || li >= gold_layers.len() - 2 || rel > 0.02 {
println!(
" layer {li:2}: cpu_norm={cnorm:9.4} gold_norm={gnorm:9.4} rel={:.2e}",
rel
);
}
}
println!(
" worst layer: {} (rel norm diff {:.2e})",
worst_layer.0, worst_layer.1
);
let cpu_argmax = argmax(&last_logits);
let dot: f64 = last_logits
.iter()
.zip(&gold_logits)
.map(|(&a, &b)| a as f64 * b as f64)
.sum();
let na: f64 = last_logits.iter().map(|&a| (a as f64).powi(2)).sum::<f64>().sqrt();
let nb: f64 = gold_logits.iter().map(|&a| (a as f64).powi(2)).sum::<f64>().sqrt();
let cos = dot / (na * nb);
let maxdiff = last_logits
.iter()
.zip(&gold_logits)
.map(|(&a, &b)| (a - b).abs())
.fold(0f32, f32::max);
let cpu_top = topk(&last_logits, 8);
let gold_top = topk(&gold_logits, 8);
let overlap = cpu_top.iter().filter(|t| gold_top.contains(t)).count();
println!("\n== prefill next-token ==");
println!(" cpu argmax = {cpu_argmax}");
println!(" gold argmax = {gold_argmax}");
println!(" logit cosine = {cos:.6} max|Δlogit| = {maxdiff:.4}");
println!(" top-8 overlap = {overlap}/8");
println!(" cpu top8 = {cpu_top:?}");
println!(" gold top8 = {gold_top:?}");
println!(
" prefill: {} tokens in {:.1}s ({:.2}s/tok)",
prompt_ids.len(),
prefill_dt,
prefill_dt / prompt_ids.len() as f32
);
println!("\n== greedy continuation ==");
let mut produced = Vec::new();
let mut tok = cpu_argmax;
let mut pos = prompt_ids.len();
m.capture_trace = false;
let t2 = std::time::Instant::now();
for _ in 0..gold_cont.len() {
produced.push(tok);
let logits = m.forward(tok, pos);
tok = argmax(&logits);
pos += 1;
}
let gen_dt = t2.elapsed().as_secs_f32();
let match_len = produced
.iter()
.zip(&gold_cont)
.take_while(|(a, b)| a == b)
.count();
println!(" cpu = {produced:?}");
println!(" gold = {gold_cont:?}");
println!(
" matching prefix = {match_len}/{} ({:.2}s/tok)",
gold_cont.len(),
gen_dt / gold_cont.len() as f32
);
let argmax_ok = cpu_argmax == gold_argmax;
let cont_ok = match_len == gold_cont.len();
println!("\n== VERDICT ==");
println!(" argmax match : {}", if argmax_ok { "PASS" } else { "FAIL" });
println!(" cosine>0.999 : {}", if cos > 0.999 { "PASS" } else { "FAIL" });
println!(" continuation : {}", if cont_ok { "PASS" } else { "FAIL" });
if argmax_ok && cont_ok && cos > 0.999 {
println!("\nALL CHECKS PASS — pure-CPU Qwen3.5 decode matches transformers gold.");
Ok(())
} else {
anyhow::bail!("parity check failed")
}
}