memra-engine 0.93.0

From-scratch CUDA LLM inference engine for NVIDIA RTX 50-series (sm_120a) and Hopper (sm_90a) - custom kernels, no frameworks
Documentation
//! DSpark Q38 drafter parity gate (lane/dspark-q38-recover): the memra loader +
//! forward + markov chain + confidence head on the arm-a HF export must reproduce
//! the SpecForge reference (tools/dspark_q38_oracle.py flat dumps) on fixed-seed
//! synthetic inputs.
//!
//! Stages and bars (calibration inherited from dflash_parity, 2026-07-13: the f32
//! GEMM rides TF32-class compute — rel-to-max 2e-3; integer decisions EXACT):
//!   ctx_features  rel < 2e-3
//!   final         rel < 2e-3
//!   markov tokens EXACT (6/6) — chained greedy decisions
//!   markov logits rel < 2e-3 (bf16 w2 under MEMRA_DFLASH_PREC=bf16)
//!   confidence    rel < 2e-3 (host dot, reference-final inputs isolate the head)
//!
//! Run with MEMRA_DFLASH_PREC=bf16 (parity precision class).
use memra_engine::Engine;
use memra_engine::dflash::DflashDraft;

fn read_f32(p: &str) -> Vec<f32> {
    let b = std::fs::read(p).unwrap_or_else(|e| panic!("{p}: {e}"));
    b.chunks_exact(4)
        .map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
        .collect()
}
fn read_u32(p: &str) -> Vec<u32> {
    let b = std::fs::read(p).unwrap_or_else(|e| panic!("{p}: {e}"));
    b.chunks_exact(4)
        .map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]))
        .collect()
}

fn rel_gate(name: &str, got: &[f32], want: &[f32], bar: f32) -> bool {
    assert_eq!(got.len(), want.len(), "{name}: length mismatch");
    let (mut md, mut mi) = (0f32, 0usize);
    for (i, (a, b)) in got.iter().zip(want).enumerate() {
        let d = (a - b).abs();
        if d > md {
            md = d;
            mi = i;
        }
    }
    let mx = want.iter().fold(0f32, |a, v| a.max(v.abs()));
    let rel = md / mx.max(1e-20);
    let pass = rel < bar;
    println!(
        "{name}: maxdiff {md:.3e} (idx {mi}: got {} want {}), rel-to-max {rel:.3e} -> {}",
        got[mi],
        want[mi],
        if pass { "PASS" } else { "FAIL" }
    );
    pass
}

fn main() -> Result<(), Box<dyn std::error::Error>> {
    let ckpt = std::env::args()
        .nth(1)
        .expect("usage: dspark_q38_parity <export_dir> <dump_dir>");
    let cache = std::env::args().nth(2).expect("dump dir");
    if std::env::var("MEMRA_DFLASH_PREC").as_deref() != Ok("bf16") {
        panic!("parity gate requires MEMRA_DFLASH_PREC=bf16");
    }
    let e = Engine::new(0)?;
    let m = DflashDraft::load(&e, std::path::Path::new(&ckpt))?;
    let c = &m.cfg;
    println!(
        "loaded dspark draft: {} layers, hidden {}, block {}, taps {:?}, markov {}, confidence {}",
        c.n_layer,
        c.hidden,
        c.block_size,
        c.target_layer_ids,
        m.markov.is_some(),
        m.confidence.is_some()
    );
    let mk = m
        .markov
        .as_ref()
        .expect("arm-a export must carry the markov head");
    let ch = m
        .confidence
        .as_ref()
        .expect("arm-a export must carry the confidence head");
    let b = c.block_size;
    let h = c.hidden;
    let n_taps = c.target_layer_ids.len();
    let v = mk.vocab;
    let mut ok = true;

    // ---- stage 1: ctx features ----
    let taps = read_f32(&format!("{cache}/dspark-taps.f32"));
    let ctx = taps.len() / (n_taps * h);
    let taps_d = e.htod(&taps)?;
    let ctxf = m.ctx_features(&e, &taps_d, ctx)?;
    ok &= rel_gate(
        "ctx_features",
        &e.dtoh(&ctxf)?,
        &read_f32(&format!("{cache}/dspark-ctx_features.f32")),
        2e-3,
    );

    // ---- stages 2-3: block forward, final hidden ----
    let noise = read_f32(&format!("{cache}/dspark-noise.f32"));
    assert_eq!(noise.len(), b * h);
    let noise_d = e.htod(&noise)?;
    let pos: Vec<i32> = (0..(ctx + b) as i32).collect();
    let fin = m.forward(&e, &taps_d, &noise_d, &pos, ctx)?;
    ok &= rel_gate(
        "final",
        &e.dtoh(&fin)?,
        &read_f32(&format!("{cache}/dspark-final.f32")),
        2e-3,
    );

    // ---- stage 4: markov chained greedy over the synthetic base logits ----
    // Mirrors generate_spec_dflash's device chain: chain_d[0] = anchor; per step k the
    // bias row gathers from chain_d[k], adds onto logits row k, argmax writes k+1.
    let base = read_f32(&format!("{cache}/dspark-base_logits.f32"));
    assert_eq!(base.len(), (b - 1) * v);
    let anchor = read_u32(&format!("{cache}/dspark-anchor.u32"))[0];
    let mut dl = e.htod(&base)?;
    let mut chain_d = e.stream().alloc_zeros::<u32>(b)?;
    e.set_u32_one(&mut chain_d, anchor)?;
    for k in 0..(b - 1) {
        let mut f = e.uninit(mk.rank)?;
        e.gather_row_bf16(&mk.w1_bf16, &chain_d, k, &mut f, mk.rank)?;
        let bias = e.matmul(&mk.w2, &f, 1)?;
        e.add_row_inplace(&mut dl, &bias, v, k * v)?;
        e.argmax_token_device_col(&dl, k, v, &mut chain_d, k + 1)?;
    }
    let chain = e.dtoh_u32(&chain_d)?;
    let want_tok = read_u32(&format!("{cache}/dspark-markov_tokens.u32"));
    let tok_pass = chain[1..] == want_tok[..];
    println!(
        "markov tokens: got {:?} want {:?} -> {}",
        &chain[1..],
        &want_tok[..],
        if tok_pass { "PASS (EXACT)" } else { "FAIL" }
    );
    ok &= tok_pass;
    ok &= rel_gate(
        "markov logits",
        &e.dtoh(&dl)?,
        &read_f32(&format!("{cache}/dspark-markov_logits.f32")),
        2e-3,
    );

    // ---- stage 5: confidence head (host dot; reference-final rows isolate the head
    // from forward drift; prev ids = [anchor, chain[..-1]] — the gate contract the
    // oracle uses; the reference SERVING loop never consumes this head) ----
    let ref_final = read_f32(&format!("{cache}/dspark-final.f32"));
    let want_conf = read_f32(&format!("{cache}/dspark-confidence.f32"));
    assert!(ch.with_markov, "arm-a confidence head is with_markov");
    // w1 rows on host for the tiny gate dot
    let mut got_conf = Vec::with_capacity(b - 1);
    let w1_all = {
        // gather via device (same primitive the chain uses), one row per prev id
        let mut rows = Vec::new();
        let mut prev_ids: Vec<u32> = vec![anchor];
        prev_ids.extend_from_slice(&want_tok[..b - 2]);
        let mut id_d = e.stream().alloc_zeros::<u32>(1)?;
        for &id in &prev_ids {
            e.set_u32_one(&mut id_d, id)?;
            let mut f = e.uninit(mk.rank)?;
            e.gather_row_bf16(&mk.w1_bf16, &id_d, 0, &mut f, mk.rank)?;
            rows.push(e.dtoh(&f)?);
        }
        rows
    };
    for i in 0..(b - 1) {
        let hrow = &ref_final[(i + 1) * h..(i + 2) * h];
        let emb = &w1_all[i];
        let mut acc = ch.b;
        for (j, x) in hrow.iter().enumerate() {
            acc += ch.w[j] * x;
        }
        for (j, x) in emb.iter().enumerate() {
            acc += ch.w[h + j] * x;
        }
        got_conf.push(acc);
    }
    ok &= rel_gate("confidence", &got_conf, &want_conf, 2e-3);

    println!(
        "== dspark_q38_parity: {} ==",
        if ok { "ALL PASS" } else { "FAIL" }
    );
    std::process::exit(if ok { 0 } else { 1 });
}