inferencelayer 0.2.9

Kortexya's engine-native inference layer — LLM generation + embedding/encoder family on wgpu (WGSL kernels, any adapter) with a pure-Rust CPU fallback
Documentation
//! `qwen35-cpu` — proves the pure-CPU Qwen3.5 decoder ([`inferencelayer::reference_qwen35`]) matches
//! the transformers gold for NuExtract-3, with NO wgpu/Vulkan involved.
//!
//! ```text
//! qwen35-cpu --model <dir> --gold <gold.json>
//! ```
//!
//! The gold (scratchpad/qwen35_gold.json + _last_logits.f32) is produced by gen_qwen35_gold.py
//! (transformers, float32/eager/CPU). Checks: prefill next-token argmax, top-8 overlap, full
//! last-position logit cosine + max-abs diff, per-decoder-layer last-hidden norm, and end-to-end
//! greedy continuation parity.

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()
    );

    // ---- Prefill: replay prompt tokens, capture the last-token layer trace + logits. ----
    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();

    // Per-layer last-hidden norm vs gold (localizes any divergence to a layer).
    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
    );

    // Full logit-vector comparison.
    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
    );

    // ---- Greedy generation, compare to gold continuation. ----
    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")
    }
}