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;
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,
);
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,
);
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,
);
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");
let mut got_conf = Vec::with_capacity(b - 1);
let w1_all = {
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 });
}