use memra_validate::{maxdiff, pr};
use memra_engine::Engine;
fn kc_model(section: &str, fname: &str, legacy: &[&str], gguf_arg: &Option<String>) -> Option<String> {
if std::env::var("MEMRA_KC_FAST").as_deref() == Ok("1") {
println!("KC-SKIP [{section}] {fname}: MEMRA_KC_FAST=1 (weight-oracle section skipped; \
synthetic arms only — full battery still required before merge/tag)");
return None;
}
if let Ok(filter) = std::env::var("MEMRA_KC_ONLY") {
if !filter.split(',').any(|f| !f.is_empty() && section.contains(f)) {
println!("KC-SKIP [{section}] {fname}: filtered by MEMRA_KC_ONLY={filter} \
(fast-gate change-scoped run — full battery still required before merge/tag)");
return None;
}
}
let mut cands: Vec<String> = Vec::new();
if let Ok(d) = std::env::var("MEMRA_KC_MODELS_DIR") {
cands.push(format!("{}/{fname}", d.trim_end_matches('/')));
}
if let Some(a) = gguf_arg {
if std::path::Path::new(a).file_name().map(|f| f == fname).unwrap_or(false) {
cands.push(a.clone());
}
}
if let Ok(h) = std::env::var("HOME") {
cands.push(format!("{h}/models/{fname}"));
}
cands.push(format!("/opt/dlami/nvme/models/{fname}"));
cands.extend(legacy.iter().map(|s| s.to_string()));
if let Some(p) = cands.iter().find(|p| std::path::Path::new(p).exists()) {
return Some(p.clone());
}
println!(
"KC-SKIP [{section}] {fname}: absent on this box ({} candidates tried) — \
set MEMRA_KC_MODELS_DIR=<dir containing it> to run this section",
cands.len()
);
None
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let e = Engine::new(0)?;
println!("GPU: {}", e.ctx().name()?);
let mut fails = 0;
let gguf_arg: Option<String> = std::env::args().nth(1).filter(|p| {
let is_dir = std::path::Path::new(p).is_dir();
if is_dir {
println!("(arg is an HF safetensors dir — GGUF weight-oracle sections will be skipped; \
pass a GGUF path to run them)");
}
!is_dir
});
{
let (ncols, nrows) = (320usize, 4usize);
let eps = 1e-6f32;
let x: Vec<f32> = (0..ncols * nrows).map(pr).collect();
let w: Vec<f32> = (0..ncols).map(|i| 0.5 + pr(i + 9) * 0.1).collect();
let mut cpu = vec![0f32; ncols * nrows];
for r in 0..nrows {
let xr = &x[r * ncols..r * ncols + ncols];
let ms: f32 = xr.iter().map(|v| v * v).sum::<f32>() / ncols as f32;
let s = 1.0 / (ms + eps).sqrt();
for i in 0..ncols { cpu[r * ncols + i] = xr[i] * s * w[i]; }
}
let xd = e.htod(&x)?; let wd = e.htod(&w)?; let mut dd = e.zeros(ncols * nrows)?;
e.rms_norm(&xd, &wd, &mut dd, ncols, nrows, eps)?;
let gpu = e.dtoh(&dd)?;
let d = maxdiff(&cpu, &gpu);
println!("rms_norm maxdiff={d:.2e} {}", if d < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
{
let (n_embd, t) = (2048usize, 3usize);
let x: Vec<f32> = (0..t * n_embd).map(|i| pr(i + 13) - 0.5).collect();
let w: Vec<f32> = (0..n_embd).map(|i| pr(i + 41) - 0.5).collect();
let mut cpu = vec![0f32; t];
for r in 0..t {
let s: f32 = (0..n_embd).map(|i| x[r * n_embd + i] * w[i]).sum();
cpu[r] = 1.0 / (1.0 + (-s).exp());
}
let xd = e.htod(&x)?; let wd = e.htod(&w)?;
let gd = e.sigmoid_dot_rows(&xd, &wd, n_embd, t)?;
let gpu = e.dtoh(&gd)?;
let d = maxdiff(&cpu, &gpu);
println!("sigmoid_dot maxdiff={d:.2e} {}", if d < 1e-5 { "OK" } else { fails += 1; "FAIL" });
}
{
let (hd, nh, nkv, t) = (512usize, 4usize, 1usize, 32usize);
let eps = 1e-6f32;
let (rq, rk) = (nh * t, nkv * t);
let q: Vec<f32> = (0..rq * hd).map(|i| pr(i + 29)).collect();
let k: Vec<f32> = (0..rk * hd).map(|i| pr(i + 31)).collect();
let v: Vec<f32> = (0..rk * hd).map(|i| pr(i + 37)).collect();
let wq: Vec<f32> = (0..hd).map(|i| 0.5 + pr(i + 41) * 0.1).collect();
let wk: Vec<f32> = (0..hd).map(|i| 0.5 + pr(i + 43) * 0.1).collect();
let wv: Vec<f32> = vec![1.0; hd];
let cpu_norm = |x: &[f32], w: &[f32], rows: usize| -> Vec<f32> {
let mut o = vec![0f32; rows * hd];
for r in 0..rows {
let xr = &x[r * hd..(r + 1) * hd];
let ms: f32 = xr.iter().map(|v| v * v).sum::<f32>() / hd as f32;
let s = 1.0 / (ms + eps).sqrt();
for i in 0..hd { o[r * hd + i] = xr[i] * s * w[i]; }
}
o
};
let (cq, ck, cv) = (cpu_norm(&q, &wq, rq), cpu_norm(&k, &wk, rk), cpu_norm(&v, &wv, rk));
let qd = e.htod(&q)?; let kd = e.htod(&k)?; let vd = e.htod(&v)?;
let wqd = e.htod(&wq)?; let wkd = e.htod(&wk)?; let wvd = e.htod(&wv)?;
let mut dq = e.zeros(rq * hd)?; let mut dk = e.zeros(rk * hd)?; let mut dv = e.zeros(rk * hd)?;
e.rms_norm_qkv(&qd, &kd, &vd, &wqd, &wkd, &wvd, &mut dq, &mut dk, &mut dv, hd, rq, rk, eps)?;
let d = maxdiff(&cq, &e.dtoh(&dq)?)
.max(maxdiff(&ck, &e.dtoh(&dk)?))
.max(maxdiff(&cv, &e.dtoh(&dv)?));
println!("rms_norm_qkv_w4 (prefill rows) maxdiff={d:.2e} {}",
if d < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
{
let (ncols, nrows) = (4096usize, 1usize);
let eps = 1e-6f32;
let a: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 61)).collect();
let b: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 67)).collect();
let w: Vec<f32> = (0..ncols).map(|i| 0.5 + pr(i + 71) * 0.1).collect();
let ad = e.htod(&a)?; let bd = e.htod(&b)?; let wd = e.htod(&w)?;
let mut res_ref = e.zeros(ncols * nrows)?;
e.add(&ad, &bd, &mut res_ref, ncols * nrows)?;
let mut z_ref = e.zeros(ncols * nrows)?;
e.rms_norm(&res_ref, &wd, &mut z_ref, ncols, nrows, eps)?;
let mut res_f = e.zeros(ncols * nrows)?;
let mut z_f = e.zeros(ncols * nrows)?;
e.add_rms_norm(&ad, &bd, &wd, &mut res_f, &mut z_f, ncols, nrows, eps)?;
let rr = e.dtoh(&res_ref)?; let rf = e.dtoh(&res_f)?;
let zr = e.dtoh(&z_ref)?; let zf = e.dtoh(&z_f)?;
let rbad = rr.iter().zip(&rf).filter(|(x, y)| x != y).count();
let zbad = zr.iter().zip(&zf).filter(|(x, y)| x != y).count();
println!("add_rms_norm fused: res_mismatch={rbad} norm_mismatch={zbad} {}",
if rbad == 0 && zbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
for nrows in [1usize, 5] {
let ncols = 4096usize;
let eps = 1e-6f32;
let x: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 31)).collect();
let w: Vec<f32> = (0..ncols).map(|i| 0.5 + pr(i + 41) * 0.1).collect();
let xd = e.htod(&x)?; let wd = e.htod(&w)?;
let mut z_ref = e.zeros(ncols * nrows)?;
e.rms_norm_decode(&xd, &wd, &mut z_ref, ncols, nrows, eps)?;
let (q_ref, d_ref) = e.quantize_q8_1(&z_ref, nrows, ncols)?;
let (q_f, d_f) = e.rms_norm_q8_1(&xd, &wd, ncols, nrows, eps)?;
let qr: Vec<i8> = e.stream().clone_dtoh(&q_ref)?; e.stream().synchronize()?;
let qf: Vec<i8> = e.stream().clone_dtoh(&q_f)?; e.stream().synchronize()?;
let dr = e.dtoh(&d_ref)?; let df = e.dtoh(&d_f)?;
let qbad = qr.iter().zip(&qf).filter(|(x, y)| x != y).count();
let dbad = dr.iter().zip(&df).filter(|(x, y)| x != y).count();
println!("rms_norm_q8_1 fused (nrows={nrows}): q_mismatch={qbad} d_mismatch={dbad} {}",
if qbad == 0 && dbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let (ncols, nrows) = (4096usize, 1usize);
let eps = 1e-6f32;
let a: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 61)).collect();
let b: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 67)).collect();
let w: Vec<f32> = (0..ncols).map(|i| 0.5 + pr(i + 71) * 0.1).collect();
let ad = e.htod(&a)?; let bd = e.htod(&b)?; let wd = e.htod(&w)?;
let mut res_ref = e.zeros(ncols * nrows)?;
let mut z_ref = e.zeros(ncols * nrows)?;
e.add_rms_norm(&ad, &bd, &wd, &mut res_ref, &mut z_ref, ncols, nrows, eps)?;
let (q_ref, d_ref) = e.quantize_q8_1(&z_ref, nrows, ncols)?;
let mut res_f = e.zeros(ncols * nrows)?;
let (q_f, d_f) = e.add_rms_norm_q8_1(&ad, &bd, &wd, &mut res_f, ncols, nrows, eps)?;
let rr = e.dtoh(&res_ref)?; let rf = e.dtoh(&res_f)?;
let qr: Vec<i8> = e.stream().clone_dtoh(&q_ref)?; e.stream().synchronize()?;
let qf: Vec<i8> = e.stream().clone_dtoh(&q_f)?; e.stream().synchronize()?;
let dr = e.dtoh(&d_ref)?; let df = e.dtoh(&d_f)?;
let rbad = rr.iter().zip(&rf).filter(|(x, y)| x != y).count();
let qbad = qr.iter().zip(&qf).filter(|(x, y)| x != y).count();
let dbad = dr.iter().zip(&df).filter(|(x, y)| x != y).count();
println!("add_rms_norm_q8_1 fused: res_mismatch={rbad} q_mismatch={qbad} d_mismatch={dbad} {}",
if rbad == 0 && qbad == 0 && dbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
for nrows in [2usize, 4, 5, 8] {
let ncols = 4096usize;
let eps = 1e-6f32;
let a: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 61)).collect();
let b: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 67)).collect();
let w: Vec<f32> = (0..ncols).map(|i| 0.5 + pr(i + 71) * 0.1).collect();
let ad = e.htod(&a)?; let bd = e.htod(&b)?; let wd = e.htod(&w)?;
let mut res_ref = e.zeros(ncols * nrows)?;
e.add(&ad, &bd, &mut res_ref, ncols * nrows)?;
let mut z_ref = e.zeros(ncols * nrows)?;
e.rms_norm_decode(&res_ref, &wd, &mut z_ref, ncols, nrows, eps)?;
let (q_ref, d_ref) = e.quantize_q8_1(&z_ref, nrows, ncols)?;
let mut res_f = e.zeros(ncols * nrows)?;
let (q_f, d_f) = e.add_rms_norm_q8_1(&ad, &bd, &wd, &mut res_f, ncols, nrows, eps)?;
let rr = e.dtoh(&res_ref)?; let rf = e.dtoh(&res_f)?;
let qr: Vec<i8> = e.stream().clone_dtoh(&q_ref)?; e.stream().synchronize()?;
let qf: Vec<i8> = e.stream().clone_dtoh(&q_f)?; e.stream().synchronize()?;
let dr = e.dtoh(&d_ref)?; let df = e.dtoh(&d_f)?;
let rbad = rr.iter().zip(&rf).filter(|(x, y)| x != y).count();
let qbad = qr.iter().zip(&qf).filter(|(x, y)| x != y).count();
let dbad = dr.iter().zip(&df).filter(|(x, y)| x != y).count();
println!("add_rms_norm_q8_1 batched (T={nrows}): res_mismatch={rbad} q_mismatch={qbad} d_mismatch={dbad} {}",
if rbad == 0 && qbad == 0 && dbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let (ncols, nrows) = (128usize, 6usize);
let eps = 1e-6f32;
let x: Vec<f32> = (0..ncols * nrows).map(|i| pr(i + 3)).collect();
let mut cpu = vec![0f32; ncols * nrows];
for r in 0..nrows {
let xr = &x[r * ncols..r * ncols + ncols];
let ss: f32 = xr.iter().map(|v| v * v).sum();
let s = 1.0 / (ss + eps).sqrt();
for i in 0..ncols { cpu[r * ncols + i] = xr[i] * s; }
}
let xd = e.htod(&x)?; let mut dd = e.zeros(ncols * nrows)?;
e.l2_norm(&xd, &mut dd, ncols, nrows, eps)?;
let gpu = e.dtoh(&dd)?;
let d = maxdiff(&cpu, &gpu);
println!("l2_norm maxdiff={d:.2e} {}", if d < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
{
let (head_dim, n_dims, n_heads, n_tokens) = (128usize, 128usize, 1usize, 3usize);
let freq_base = 1e6f32; let freq_scale = 1.0f32;
let theta_scale = freq_base.powf(-2.0 / n_dims as f32);
let x: Vec<f32> = (0..head_dim * n_heads * n_tokens).map(|i| pr(i + 5)).collect();
let pos: Vec<i32> = (0..n_tokens as i32).collect();
let half = n_dims / 2;
let mut cpu = x.clone();
for tok in 0..n_tokens {
for h in 0..n_heads {
let base = (tok * n_heads + h) * head_dim;
for j in 0..half {
let theta = pos[tok] as f32 * theta_scale.powf(j as f32) * freq_scale;
let (c, s) = (theta.cos(), theta.sin());
let x0 = x[base + j]; let x1 = x[base + j + half];
cpu[base + j] = x0 * c - x1 * s;
cpu[base + j + half] = x0 * s + x1 * c;
}
}
}
let mut xd = e.htod(&x)?; let posd = e.htod_i32(&pos)?;
e.rope_neox(&mut xd, &posd, head_dim, n_dims, n_heads, n_tokens, freq_base, freq_scale)?;
let gpu = e.dtoh(&xd)?;
let d = maxdiff(&cpu, &gpu);
println!("rope_neox maxdiff={d:.2e} {}", if d < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
{
let n = 1024usize;
let g: Vec<f32> = (0..n).map(|i| pr(i)).collect();
let u: Vec<f32> = (0..n).map(|i| pr(i + 1)).collect();
let cpu: Vec<f32> = (0..n).map(|i| (g[i] / (1.0 + (-g[i]).exp())) * u[i]).collect();
let gd = e.htod(&g)?; let ud = e.htod(&u)?; let mut dd = e.zeros(n)?;
e.silu_mul(&gd, &ud, &mut dd, n)?;
let gpu = e.dtoh(&dd)?;
let d = maxdiff(&cpu, &gpu);
println!("silu_mul maxdiff={d:.2e} {}", if d < 1e-5 { "OK" } else { fails += 1; "FAIL" });
}
{
let n = 4099usize; let mask_words = (n - 67).div_ceil(32); let mask: Vec<u32> = (0..mask_words).map(|w| {
let mut bits = 0u32;
for b in 0..32 { if (w * 32 + b) % 7 == 3 { bits |= 1 << b; } }
bits as u32
}).collect();
let allowed = |i: usize| -> bool { i < mask_words * 32 && i % 7 == 3 };
let rows = 2usize;
let x: Vec<f32> = (0..rows * n).map(|i| pr(i) * 8.0).collect();
let mut cpu = x.clone();
for i in 0..n {
if !allowed(i) { cpu[n + i] = f32::MIN; }
}
let mut xd = e.htod(&x)?;
let md = e.htod_u32_v(&mask)?;
e.mask_logits_col(&mut xd, &md, 1, n, mask_words)?;
let gpu = e.dtoh(&xd)?;
let bad = cpu.iter().zip(&gpu).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
let am_cpu = cpu[n..].iter().enumerate()
.max_by(|(i, a), (j, b)| a.partial_cmp(b).unwrap().then(j.cmp(i))).unwrap().0;
let am_gpu = gpu[n..].iter().enumerate()
.max_by(|(i, a), (j, b)| a.partial_cmp(b).unwrap().then(j.cmp(i))).unwrap().0;
println!("mask_logits_col: mismatch={bad} argmax {}=={} {}",
am_cpu, am_gpu,
if bad == 0 && am_cpu == am_gpu { "OK (byte-identical)" }
else { fails += 1; "FAIL" });
}
{
let n = 2048usize; let (gs, us) = (1.31f32, 0.77f32); let g: Vec<f32> = (0..n).map(|i| pr(i + 3)).collect();
let u: Vec<f32> = (0..n).map(|i| pr(i + 5)).collect();
let gd = e.htod(&g)?;
let ud = e.htod(&u)?;
let mut act = e.zeros(n)?;
e.silu_mul_scaled(&gd, &ud, gs, us, &mut act, n)?;
let (aq_ref, ad_ref) = e.quantize_q8_1(&act, 1, n)?;
let (aq_f, ad_f) = e.silu_mul_scaled_q8_1(&gd, &ud, gs, us, n)?;
let q_ref: Vec<i8> = e.stream().clone_dtoh(&aq_ref)?; e.stream().synchronize()?;
let q_f: Vec<i8> = e.stream().clone_dtoh(&aq_f)?; e.stream().synchronize()?;
let d_ref = e.dtoh(&ad_ref)?;
let d_f = e.dtoh(&ad_f)?;
let qbad = q_ref.iter().zip(&q_f).filter(|(a, b)| a != b).count();
let dbad = d_ref.iter().zip(&d_f).filter(|(a, b)| a != b).count();
println!("silu_mul_q8_1 fold: int8_mismatch={qbad} scale_mismatch={dbad} {}",
if qbad == 0 && dbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let (t, n_ff) = (5usize, 2048usize);
let n = t * n_ff;
let g: Vec<f32> = (0..n).map(|i| pr(i + 7)).collect();
let u: Vec<f32> = (0..n).map(|i| pr(i + 11)).collect();
let gd = e.htod(&g)?;
let ud = e.htod(&u)?;
let mut act = e.zeros(n)?;
e.silu_mul(&gd, &ud, &mut act, n)?;
let (aq_ref, ad_ref) = e.quantize_q8_1(&act, t, n_ff)?;
let (aq_f, ad_f) = e.silu_mul_scaled_q8_1(&gd, &ud, 1.0, 1.0, n)?;
let q_ref: Vec<i8> = e.stream().clone_dtoh(&aq_ref)?; e.stream().synchronize()?;
let q_f: Vec<i8> = e.stream().clone_dtoh(&aq_f)?; e.stream().synchronize()?;
let d_ref = e.dtoh(&ad_ref)?;
let d_f = e.dtoh(&ad_f)?;
let qbad = q_ref.iter().zip(&q_f).filter(|(a, b)| a != b).count();
let dbad = d_ref.iter().zip(&d_f).filter(|(a, b)| a != b).count();
println!("silu_mul_q8_1 batched (T={t}): int8_mismatch={qbad} scale_mismatch={dbad} {}",
if qbad == 0 && dbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
for t in [1usize, 5] {
let (d_state, num_v) = (128usize, 16usize);
let nrows = num_v * t;
let eps = 1e-6f32;
let o: Vec<f32> = (0..nrows * d_state).map(|i| pr(i + 83)).collect();
let z: Vec<f32> = (0..nrows * d_state).map(|i| pr(i + 89) - 0.5).collect();
let w: Vec<f32> = (0..d_state).map(|i| 0.5 + pr(i + 97) * 0.1).collect();
let od = e.htod(&o)?; let zd = e.htod(&z)?; let wd = e.htod(&w)?;
let mut gn_ref = e.zeros(nrows * d_state)?;
e.gated_rmsnorm(&od, &wd, &zd, &mut gn_ref, d_state, nrows, eps)?;
let (q_ref, d_ref) = e.quantize_q8_1(&gn_ref, nrows, d_state)?;
let (q_f, d_f) = e.gated_rmsnorm_q8_1(&od, &wd, &zd, d_state, nrows, eps)?;
let qr: Vec<i8> = e.stream().clone_dtoh(&q_ref)?; e.stream().synchronize()?;
let qf: Vec<i8> = e.stream().clone_dtoh(&q_f)?; e.stream().synchronize()?;
let dr = e.dtoh(&d_ref)?; let df = e.dtoh(&d_f)?;
let qbad = qr.iter().zip(&qf).filter(|(x, y)| x != y).count();
let dbad = dr.iter().zip(&df).filter(|(x, y)| x != y).count();
println!("gated_rmsnorm_q8_1 (T={t}): q_mismatch={qbad} d_mismatch={dbad} {}",
if qbad == 0 && dbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
for (name, in_f, n_pairs, act_kind) in [("silu", 768usize, 33usize, 0i32),
("silu", 512, 7, 0),
("gelu", 704, 29, 1)] {
let n = n_pairs * in_f;
let g: Vec<f32> = (0..n).map(|i| pr(i + 17) * 4.0).collect();
let u: Vec<f32> = (0..n).map(|i| pr(i + 29) * 4.0).collect();
let gd = e.htod(&g)?;
let ud = e.htod(&u)?;
let act = if act_kind == 0 { e.moe_pairs_silu_mul(&gd, &ud, n)? }
else { e.moe_pairs_gelu_mul(&gd, &ud, n)? };
let scr_ref = e.mmq_iq_quantize_act(&act, in_f, n_pairs)?;
let scr_f = e.mmq_iq_fused_act_quant(&gd, &ud, in_f, n_pairs, act_kind)?;
let b_ref: Vec<u8> = e.stream().clone_dtoh(&scr_ref)?;
let b_f: Vec<u8> = e.stream().clone_dtoh(&scr_f)?;
e.stream().synchronize()?;
let nbad = b_ref.iter().zip(&b_f).filter(|(a, b)| a != b).count();
println!("iq fused act+quant [{name} in_f={in_f} n_pairs={n_pairs}]: \
byte_mismatch={nbad}/{} {}",
b_ref.len(), if nbad == 0 && b_ref.len() == b_f.len() { "OK" }
else { fails += 1; "FAIL" });
}
{
let (hd, nh, nhkv, t, tkv) = (64usize, 2usize, 1usize, 4usize, 4usize);
let scale = 1.0 / (hd as f32).sqrt();
let q: Vec<f32> = (0..hd * nh * t).map(|i| pr(i) * 0.2).collect();
let k: Vec<f32> = (0..hd * nhkv * tkv).map(|i| pr(i + 7) * 0.2).collect();
let v: Vec<f32> = (0..hd * nhkv * tkv).map(|i| pr(i + 11) * 0.2).collect();
let mut cpu = vec![0f32; hd * nh * t];
for head in 0..nh {
let kvh = head / (nh / nhkv);
for qt in 0..t {
let q_pos = (tkv - t) + qt;
let qv = &q[(qt * nh + head) * hd..][..hd];
let mut sc = vec![0f32; tkv];
for tk in 0..tkv {
let kv = &k[(tk * nhkv + kvh) * hd..][..hd];
let mut acc = 0.0; for d in 0..hd { acc += qv[d] * kv[d]; }
acc *= scale;
if tk > q_pos { acc = -1e30; }
sc[tk] = acc;
}
let mx = sc.iter().cloned().fold(-1e30f32, f32::max);
let mut sum = 0.0; for s in sc.iter_mut() { *s = (*s - mx).exp(); sum += *s; }
for s in sc.iter_mut() { *s /= sum; }
let ov = &mut cpu[(qt * nh + head) * hd..][..hd];
for d in 0..hd {
let mut acc = 0.0;
for tk in 0..tkv { acc += sc[tk] * v[(tk * nhkv + kvh) * hd + d]; }
ov[d] = acc;
}
}
}
let qd = e.htod(&q)?; let kd = e.htod(&k)?; let vd = e.htod(&v)?; let mut od = e.zeros(hd * nh * t)?;
e.sdpa_naive(&qd, &kd, &vd, &mut od, hd, nh, nhkv, t, tkv, scale, true)?;
let gpu = e.dtoh(&od)?;
let d = maxdiff(&cpu, &gpu);
println!("sdpa_naive maxdiff={d:.2e} {}", if d < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
{
let (conv_dim, t, d_conv) = (8usize, 5usize, 4usize);
let tp = t + d_conv - 1;
let x: Vec<f32> = (0..conv_dim * tp).map(|i| pr(i + 13)).collect();
let w: Vec<f32> = (0..d_conv * conv_dim).map(|i| pr(i + 21) * 0.3).collect();
let mut cpu = vec![0f32; conv_dim * t];
for c in 0..conv_dim {
for tt in 0..t {
let mut acc = 0.0;
for j in 0..d_conv { acc += x[c * tp + tt + j] * w[c * d_conv + j]; }
cpu[c * t + tt] = acc / (1.0 + (-acc).exp());
}
}
let xd = e.htod(&x)?; let wd = e.htod(&w)?; let mut yd = e.zeros(conv_dim * t)?;
e.ssm_conv1d(&xd, &wd, &mut yd, conv_dim, t, d_conv, true)?;
let gpu = e.dtoh(&yd)?;
let d = maxdiff(&cpu, &gpu);
println!("ssm_conv1d maxdiff={d:.2e} {}", if d < 1e-5 { "OK" } else { fails += 1; "FAIL" });
}
{
let (conv_dim, d_conv) = (96usize, 4usize);
let pad = d_conv - 1;
let qkv: Vec<f32> = (0..conv_dim).map(|i| pr(i + 31)).collect();
let st0: Vec<f32> = (0..conv_dim * pad).map(|i| pr(i + 41) * 0.7).collect();
let w: Vec<f32> = (0..d_conv * conv_dim).map(|i| pr(i + 51) * 0.3).collect();
let qd = e.htod(&qkv)?;
let wd = e.htod(&w)?;
let mut st_ref = e.htod(&st0)?;
let mut conv_in = e.zeros(conv_dim * (pad + 1))?;
e.conv_assemble_and_roll(&qd, &mut st_ref, &mut conv_in, conv_dim, pad)?;
let mut out_ref = e.zeros(conv_dim)?;
e.ssm_conv1d(&conv_in, &wd, &mut out_ref, conv_dim, 1, d_conv, true)?;
let mut st_f = e.htod(&st0)?;
let mut out_f = e.zeros(conv_dim)?;
e.ssm_conv1d_fused_decode(&qd, &mut st_f, &wd, &mut out_f, conv_dim, d_conv)?;
let or = e.dtoh(&out_ref)?; let of = e.dtoh(&out_f)?;
let sr = e.dtoh(&st_ref)?; let sf = e.dtoh(&st_f)?;
let obad = or.iter().zip(&of).filter(|(a, b)| a != b).count();
let sbad = sr.iter().zip(&sf).filter(|(a, b)| a != b).count();
println!("ssm_conv1d fused: out_mismatch={obad} state_mismatch={sbad} {}",
if obad == 0 && sbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let s_v = 128usize; let h = 1usize; let t = 3usize;
let scale = 1.0 / (s_v as f32).sqrt();
let q: Vec<f32> = (0..s_v * h * t).map(|i| pr(i) * 0.1).collect();
let k: Vec<f32> = (0..s_v * h * t).map(|i| pr(i + 5) * 0.1).collect();
let v: Vec<f32> = (0..s_v * h * t).map(|i| pr(i + 9) * 0.1).collect();
let g: Vec<f32> = (0..h * t).map(|i| -0.05 - pr(i).abs() * 0.1).collect(); let beta: Vec<f32> = (0..h * t).map(|i| 0.5 + pr(i + 3) * 0.2).collect();
let st0 = vec![0f32; s_v * s_v * h];
let mut s = vec![0f32; s_v * s_v]; let mut cpu_o = vec![0f32; s_v * h * t];
for tt in 0..t {
let qt = &q[(tt * h) * s_v..][..s_v];
let kt = &k[(tt * h) * s_v..][..s_v];
let vt = &v[(tt * h) * s_v..][..s_v];
let gv = (g[tt]).exp();
let bv = beta[tt];
let mut new_s = s.clone();
for col in 0..s_v {
let mut kv = 0.0f32;
for i in 0..s_v { kv += s[col * s_v + i] * kt[i]; }
let delta = (vt[col] - gv * kv) * bv;
let mut attn = 0.0f32;
for i in 0..s_v {
let ns = gv * s[col * s_v + i] + kt[i] * delta;
new_s[col * s_v + i] = ns;
attn += ns * qt[i];
}
cpu_o[(tt * h) * s_v + col] = attn * scale;
}
s = new_s;
}
let qd = e.htod(&q)?; let kd = e.htod(&k)?; let vd = e.htod(&v)?;
let gd = e.htod(&g)?; let bd = e.htod(&beta)?; let sid = e.htod(&st0)?;
let mut sod = e.zeros(s_v * s_v * h)?; let mut od = e.zeros(s_v * h * t)?;
e.gdn_scan_s128(&qd, &kd, &vd, &gd, &bd, &sid, &mut sod, &mut od, h, t, scale)?;
let gpu_o = e.dtoh(&od)?;
let d = maxdiff(&cpu_o, &gpu_o);
println!("gdn_scan maxdiff={d:.2e} {}", if d < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
{
let s_v = 128usize; let h = 4usize;
let relerr = |a: &[f64], b: &[f32]| -> f32 {
a.iter().zip(b)
.map(|(x, y)| ((*x - *y as f64).abs() / x.abs().max(*y as f64).max(1e-3)) as f32)
.fold(0.0f32, f32::max)
};
for &(t, c) in &[(200usize, 32usize), (200, 64), (200, 128), (17, 64), (512, 64)] {
let mut q = vec![0f32; s_v * h * t];
let mut k = vec![0f32; s_v * h * t];
for row in 0..h * t {
let (mut nq, mut nk) = (0f32, 0f32);
for i in 0..s_v {
let a = pr(row * s_v + i + 11); let b = pr(row * s_v + i + 17);
q[row * s_v + i] = a; k[row * s_v + i] = b;
nq += a * a; nk += b * b;
}
for i in 0..s_v {
q[row * s_v + i] /= nq.sqrt(); k[row * s_v + i] /= nk.sqrt();
}
}
let v: Vec<f32> = (0..s_v * h * t).map(|i| pr(i + 23)).collect();
let g: Vec<f32> = (0..h * t).map(|i| -0.02 - pr(i + 29).abs() * 0.5).collect();
let beta: Vec<f32> = (0..h * t).map(|i| 0.3 + pr(i + 31).abs() * 0.6).collect();
let st0: Vec<f32> = (0..s_v * s_v * h).map(|i| pr(i + 37) * 0.5).collect(); let scale = 1.0 / (s_v as f32).sqrt();
let mut o64 = vec![0f64; s_v * h * t];
let mut s64 = vec![0f64; s_v * s_v * h];
for hh in 0..h {
let s = &mut s64[hh * s_v * s_v..(hh + 1) * s_v * s_v]; for (i, sv) in s.iter_mut().enumerate() { *sv = st0[hh * s_v * s_v + i] as f64; }
for tt in 0..t {
let base = (tt * h + hh) * s_v;
let gv = (g[tt * h + hh] as f64).exp();
let bv = beta[tt * h + hh] as f64;
for col in 0..s_v {
let mut kv = 0f64;
for i in 0..s_v { kv += s[col * s_v + i] * k[base + i] as f64; }
let delta = (v[base + col] as f64 - gv * kv) * bv;
let mut attn = 0f64;
for i in 0..s_v {
let ns = gv * s[col * s_v + i] + k[base + i] as f64 * delta;
s[col * s_v + i] = ns;
attn += ns * q[base + i] as f64;
}
o64[base + col] = attn * scale as f64;
}
}
}
let qd = e.htod(&q)?; let kd = e.htod(&k)?; let vd = e.htod(&v)?;
let gd = e.htod(&g)?; let bd = e.htod(&beta)?; let sid = e.htod(&st0)?;
let mut so_s = e.zeros(s_v * s_v * h)?; let mut o_s = e.zeros(s_v * h * t)?;
e.gdn_scan_s128(&qd, &kd, &vd, &gd, &bd, &sid, &mut so_s, &mut o_s, h, t, scale)?;
let mut so_c = e.zeros(s_v * s_v * h)?; let mut o_c = e.zeros(s_v * h * t)?;
unsafe { std::env::set_var("MEMRA_GDN_MMA", "0"); }
unsafe { std::env::set_var("MEMRA_GDN_WGMMA", "0"); }
e.gdn_scan_chunked(&qd, &kd, &vd, &gd, &bd, None, None, &sid, &mut so_c, &mut o_c, h, t, scale, c, h)?;
unsafe { std::env::remove_var("MEMRA_GDN_MMA"); }
let (ro_s, rs_s) = (relerr(&o64, &e.dtoh(&o_s)?), relerr(&s64, &e.dtoh(&so_s)?));
let (ro_c, rs_c) = (relerr(&o64, &e.dtoh(&o_c)?), relerr(&s64, &e.dtoh(&so_c)?));
let ok = ro_c < 1e-4 && rs_c < 2.5e-4
&& ro_c <= (ro_s * 32.0).max(1e-6) && rs_c <= (rs_s * 32.0).max(1e-6);
println!("gdn_chunked T={t:3} C={c:3} vs f64-truth: out seq={ro_s:.2e}/chunk={ro_c:.2e} \
state seq={rs_s:.2e}/chunk={rs_c:.2e} {}",
if ok { "OK" } else { fails += 1; "FAIL" });
if c == 32 && cfg!(memra_hopper_mma) {
unsafe { std::env::set_var("MEMRA_GDN_MMA", "1"); }
let mut so_m = e.zeros(s_v * s_v * h)?; let mut o_m = e.zeros(s_v * h * t)?;
e.gdn_scan_chunked(&qd, &kd, &vd, &gd, &bd, None, None, &sid, &mut so_m, &mut o_m, h, t, scale, c, h)?;
let (ro_m, rs_m) = (relerr(&o64, &e.dtoh(&o_m)?), relerr(&s64, &e.dtoh(&so_m)?));
let okm = ro_m < 8e-2 && rs_m < 8e-1;
println!("gdn_chunked T={t:3} C={c:3} MMA config pin: out={ro_m:.2e} state={rs_m:.2e} {}",
if okm { "OK" } else { fails += 1; "FAIL" });
unsafe { std::env::set_var("MEMRA_GDN_WGMMA", "1"); }
let mut so_w = e.zeros(s_v * s_v * h)?; let mut o_w = e.zeros(s_v * h * t)?;
e.gdn_scan_chunked(&qd, &kd, &vd, &gd, &bd, None, None, &sid, &mut so_w, &mut o_w, h, t, scale, c, h)?;
unsafe { std::env::remove_var("MEMRA_GDN_MMA"); }
let (ro_w, rs_w) = (relerr(&o64, &e.dtoh(&o_w)?), relerr(&s64, &e.dtoh(&so_w)?));
let okw = ro_w < 4e-1 && rs_w < 8e-1;
println!("gdn_chunked T={t:3} C={c:3} WGMMA-fused config pin: out={ro_w:.2e} state={rs_w:.2e} {}",
if okw { "OK" } else { fails += 1; "FAIL" });
}
unsafe { std::env::remove_var("MEMRA_GDN_WGMMA"); }
}
}
{
use memra_gguf::{GgmlType, dequant};
use memra_runtime::cpu_linear;
let (in_f, out_f, m, row_bytes) = (256usize, 7usize, 3usize, 84usize);
let mut raw = vec![0u8; out_f * row_bytes];
for row in 0..out_f {
let base = row * row_bytes;
for group in 0..16 {
let scale = 1 + ((row * 3 + group * 5) % 15) as u8;
let min = 1 + ((row * 7 + group * 2) % 15) as u8;
raw[base + group] = scale | (min << 4);
}
for byte in 0..64 {
raw[base + 16 + byte] = ((row * 41 + byte * 17 + 13) & 0xff) as u8;
}
raw[base + 80..base + 82].copy_from_slice(&0x2c00u16.to_le_bytes()); raw[base + 82..base + 84].copy_from_slice(&0x2800u16.to_le_bytes()); }
let weights = dequant::dequantize(GgmlType::Q2_K, &raw, in_f * out_f);
let x: Vec<f32> = (0..m * in_f).map(|i| pr(i + 79) * 0.1).collect();
let cpu = cpu_linear(&x, &weights, m, in_f, out_f);
let wd = e.htod_bytes(&raw)?;
let xd = e.htod(&x)?;
let gpu = e.dtoh(&e.qmatvec(
&wd, &xd, m, in_f, out_f, memra_engine::QT_Q2_K, row_bytes,
)?)?;
let scale = cpu.iter().map(|value| value.abs()).fold(0.0, f32::max).max(1e-3);
let rel = maxdiff(&cpu, &gpu) / scale;
println!("qmatvec Q2_K synthetic Stage-A: rel={rel:.2e} {}",
if rel < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType, dequant};
use memra_runtime::cpu_linear;
let g = GgufFile::open(&path)?;
let cases = [
("blk.0.ffn_gate.weight", memra_engine::QT_Q8_0), ("blk.0.attn_qkv.weight", memra_engine::QT_Q8_0), ("blk.3.attn_q.weight", memra_engine::QT_Q8_0), ("blk.0.attn_v.weight", memra_engine::QT_Q6_K), ("output.weight", memra_engine::QT_Q6_K), ("token_embd.weight", memra_engine::QT_Q8_0),
];
for (tname, _) in cases {
if let Some(t) = g.find(tname) {
let qt = match t.ggml_type {
GgmlType::Q8_0 => memra_engine::QT_Q8_0,
GgmlType::Q4_K => memra_engine::QT_Q4_K,
GgmlType::Q6_K => memra_engine::QT_Q6_K,
other => { println!("qmatvec skip {tname}: {other:?} not in stage-A"); continue; }
};
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t);
let row_bytes = raw.len() / out_f;
let w_f32 = dequant::dequantize(t.ggml_type, raw, in_f * out_f);
let m = 2usize;
let x: Vec<f32> = (0..m * in_f).map(|i| pr(i + 31) * 0.1).collect();
let cpu = cpu_linear(&x, &w_f32, m, in_f, out_f);
let wd = e.htod_bytes(raw)?; let xd = e.htod(&x)?;
let yd = e.qmatvec(&wd, &xd, m, in_f, out_f, qt, row_bytes)?;
let gpu = e.dtoh(&yd)?;
let d = maxdiff(&cpu, &gpu);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1.0);
let rel = d / scale;
println!("qmatvec {tname} [{:?}] rel={rel:.2e} {}", t.ggml_type,
if rel < 1e-4 { "OK" } else { fails += 1; "FAIL" });
}
}
} else {
println!("(pass a GGUF path to also validate qmatvec vs CPU oracle)");
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType};
let g = GgufFile::open(&path)?;
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::Q8_0) {
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let m = 2usize;
let x: Vec<f32> = (0..m * in_f).map(|i| pr(i + 41) * 0.1).collect();
let wd = e.htod_bytes(raw)?; let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec(&wd, &xd, m, in_f, out_f, memra_engine::QT_Q8_0, row_bytes)?)?;
let yb = e.dtoh(&e.qmatvec_q8_0_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("qmatvec_q8_0_fast vs Stage-A: rel={rel:.2e} {}", if rel < 3e-2 { "OK" } else { fails += 1; "FAIL" });
println!(" (ya[0..3]={:?} yb[0..3]={:?})", &ya[..3], &yb[..3]);
}
for (tname, qt) in [("blk.0.attn_q.weight", memra_engine::QT_Q4_K),
("blk.0.attn_v.weight", memra_engine::QT_Q6_K),
("output.weight", memra_engine::QT_Q6_K)] {
if let Some(t) = g.find(tname) {
let gt = match t.ggml_type { GgmlType::Q4_K => memra_engine::QT_Q4_K, GgmlType::Q6_K => memra_engine::QT_Q6_K, _ => continue };
if gt != qt { continue; }
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let m = 2usize;
let x: Vec<f32> = (0..m * in_f).map(|i| pr(i + 51) * 0.1).collect();
let wd = e.htod_bytes(raw)?; let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec(&wd, &xd, m, in_f, out_f, gt, row_bytes)?)?;
let yb = if gt == memra_engine::QT_Q4_K { e.dtoh(&e.qmatvec_q4_K_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)? }
else { e.dtoh(&e.qmatvec_q6_K_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)? };
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("{tname} [{:?}] fast vs Stage-A: rel={rel:.2e} {}", t.ggml_type, if rel < 3e-2 { "OK" } else { fails += 1; "FAIL" });
}
}
}
{
use memra_gguf::{GgufFile, GgmlType, dequant};
use memra_runtime::cpu_linear;
let gguf_9b = kc_model("dtype5", "Qwen3.5-9B-NVFP4-MTP-GGUF.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf"],
&gguf_arg);
let gguf_35b = kc_model("dtype5", "Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf"],
&gguf_arg);
let cases: [(&Option<String>, &str, GgmlType, i32, &str); 5] = [
(&gguf_9b, "blk.0.ffn_gate.weight", GgmlType::NVFP4, memra_engine::QT_NVFP4, "nvfp4"),
(&gguf_9b, "blk.0.attn_gate.weight", GgmlType::Q5_K, memra_engine::QT_Q5_K, "q5k"),
(&gguf_35b, "blk.0.ffn_gate_exps.weight", GgmlType::IQ3_S, memra_engine::QT_IQ3_S, ""),
(&gguf_35b, "blk.0.ffn_down_exps.weight", GgmlType::IQ4_XS, memra_engine::QT_IQ4_XS, "iq4xs"),
(&gguf_35b, "blk.40.ffn_gate_exps.weight",GgmlType::Q3_K, memra_engine::QT_Q3_K, "q3k"),
];
for (path, tname, gty, qt, sel) in cases {
let Some(path) = path.as_deref() else { continue }; let g = GgufFile::open(path)?;
let t = match g.find(tname).filter(|t| t.ggml_type == gty) {
Some(t) => t,
None => match g.tensors.iter()
.filter(|t| t.ggml_type == gty && t.ne.len() >= 2 && t.ne[1] > 1
&& t.name.ends_with(".weight"))
.min_by_key(|t| t.n_bytes) {
Some(t) => {
println!("dtype5 {gty:?}: pinned {tname} absent/re-typed in this artifact \
revision — substituting {}", t.name);
t
}
None => {
println!("KC-SKIP [dtype5] {path}: no {gty:?} .weight tensor at all \
(pinned {tname} absent) — this artifact revision lacks the dtype");
continue;
}
}
};
let in_f = t.ne[0] as usize;
let out_f = t.ne[1] as usize;
let raw_all = g.tensor_data(t);
let n_experts = if t.ne.len() >= 3 { t.ne[2] as usize } else { 1 };
let total_rows = out_f * n_experts;
let row_bytes = raw_all.len() / total_rows;
let raw = &raw_all[..out_f * row_bytes]; let w_f32 = dequant::dequantize(gty, raw, in_f * out_f);
let m = 2usize;
let x: Vec<f32> = (0..m * in_f).map(|i| pr(i + 61) * 0.1).collect();
let cpu = cpu_linear(&x, &w_f32, m, in_f, out_f);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1.0);
let wd = e.htod_bytes(raw)?; let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec(&wd, &xd, m, in_f, out_f, qt, row_bytes)?)?;
let rela = maxdiff(&cpu, &ya) / scale;
println!("dtype5 [{gty:?}] {tname} (in={in_f} out={out_f}) Stage-A: rel={rela:.2e} {}",
if rela < 1e-4 { "OK" } else { fails += 1; "FAIL" });
if sel.is_empty() {
println!("dtype5 [{gty:?}] {tname} Stage-B dp4a: (no fast path — Stage-A only)");
} else {
let yb = match sel {
"nvfp4" => e.dtoh(&e.qmatvec_nvfp4_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)?,
"q5k" => e.dtoh(&e.qmatvec_q5_K_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)?,
"iq4xs" => e.dtoh(&e.qmatvec_iq4_XS_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)?,
"q3k" => e.dtoh(&e.qmatvec_q3_K_fast(&wd, &xd, m, in_f, out_f, row_bytes)?)?,
_ => unreachable!(),
};
let relb = maxdiff(&cpu, &yb) / scale;
println!("dtype5 [{gty:?}] {tname} Stage-B dp4a: rel={relb:.2e} {}",
if relb < 3e-2 { "OK" } else { fails += 1; "FAIL" });
}
}
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType};
let g = GgufFile::open(&path)?;
let gemm_cases: [(&str, i32, &str); 6] = [
("blk.0.ffn_gate.weight", memra_engine::QT_Q8_0, "q8_0"), ("blk.0.attn_qkv.weight", memra_engine::QT_Q8_0, "q8_0"),
("blk.3.attn_q.weight", memra_engine::QT_Q4_K, "q4_K"), ("blk.0.ssm_out.weight", memra_engine::QT_Q5_K, "q5_K"), ("blk.0.attn_v.weight", memra_engine::QT_Q6_K, "q6_K"),
("output.weight", memra_engine::QT_Q6_K, "q6_K"), ];
for (tname, want_qt, sel) in gemm_cases {
let t = match g.find(tname) { Some(t) => t, None => continue };
let gt = match t.ggml_type {
GgmlType::Q8_0 => memra_engine::QT_Q8_0, GgmlType::Q4_K => memra_engine::QT_Q4_K,
GgmlType::Q6_K => memra_engine::QT_Q6_K, GgmlType::NVFP4 => memra_engine::QT_NVFP4,
GgmlType::Q5_K => memra_engine::QT_Q5_K,
_ => continue,
};
if gt != want_qt { continue; }
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
if t.ne.len() > 2 { continue; } let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
let wgmma_mirror = if cfg!(memra_hopper_mma) && gt == memra_engine::QT_Q8_0
&& out_f % 64 == 0 && in_f % 32 == 0 {
Some(e.build_q8_rp4_raw(&wd, in_f, out_f)?)
} else { None };
let f16_mirror = if gt == memra_engine::QT_Q8_0 && in_f % 32 == 0 {
Some(e.build_q8_f16_raw(&wd, in_f, out_f)?)
} else if gt == memra_engine::QT_Q4_K && in_f % 256 == 0 {
Some(e.build_q4k_f16_raw(&wd, in_f, out_f)?)
} else if gt == memra_engine::QT_Q5_K && in_f % 256 == 0 {
Some(e.build_q5k_f16_raw(&wd, in_f, out_f)?)
} else if gt == memra_engine::QT_Q6_K && in_f % 256 == 0 {
Some(e.build_q6k_f16_raw(&wd, in_f, out_f)?)
} else { None };
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 71) * 0.1).collect();
let xd = e.htod(&x)?;
let ydp = match sel {
"q8_0" => e.qmatvec_q8_0_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?,
"q4_K" => e.qmatvec_q4_K_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?,
"q5_K" => e.qmatvec_q5_K_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?,
"q6_K" => e.qmatvec_q6_K_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?,
_ => unreachable!(),
};
let ya = e.dtoh(&ydp)?;
let yb = e.dtoh(&e.qmatvec_gemm_raw(&wd, &xd, tt, in_f, out_f, gt, row_bytes)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("GEMM {tname} [{:?}] T={tt}: rel={rel:.2e} {}", t.ggml_type,
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
if let Some(mirror) = &wgmma_mirror {
let (aq, ad) = e.quantize_q8_1(&xd, tt, in_f)?;
let yw = e.dtoh(&e.qmatvec_gemm_q8_0_wgmma_raw(mirror, &aq, &ad, tt, in_f, out_f)?)?;
let dw = maxdiff(&yb, &yw);
let relw = dw / scale;
println!("GEMM {tname} wgmma T={tt}: rel={relw:.2e} {}",
if relw < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
if let Some(m16) = &f16_mirror {
let yf = e.dtoh(&e.qmatvec_gemm_f16_raw(m16, &xd, tt, in_f, out_f)?)?;
let df = maxdiff(&yb, &yf);
let relf = df / scale;
println!("GEMM {tname} f16 T={tt}: rel={relf:.2e} {}",
if relf < 1e-2 { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
{
fn f16_bits(x: f32) -> u16 {
let b = x.to_bits();
let s = ((b >> 16) & 0x8000) as u16;
if x == 0.0 { return s; }
let he = ((b >> 23) & 0xff) as i32 - 127 + 15; let m = b & 0x7f_ffff;
let mut h = ((he as u32) << 10) | (m >> 13);
let rem = m & 0x1fff;
if rem > 0x1000 || (rem == 0x1000 && (h & 1) == 1) { h += 1; }
s | h as u16
}
let m_sizes: [i32; 8] = [1, 3, 17, 33, 64, 129, 200, 300];
let n_active = m_sizes.len();
let mut ex_off_host = vec![0i32; n_active + 1];
for (g, m) in m_sizes.iter().enumerate() { ex_off_host[g + 1] = ex_off_host[g] + m; }
let n_pairs = *ex_off_host.last().unwrap() as usize;
let snap = |v: f32| (v * 256.0).round() / 256.0;
for (in_f, out_f) in [(512usize, 300usize), (480, 192)] {
let w_f32: Vec<f32> = (0..n_active * out_f * in_f)
.map(|i| snap(pr(i + 101) - 0.5)).collect();
let a_f32: Vec<f32> = (0..n_pairs * in_f)
.map(|i| snap(pr(i + 211) - 0.5)).collect();
let scales: Vec<f32> = (0..n_pairs).map(|p| 1.0 + (p % 5) as f32 * 0.25).collect();
let mut cpu = vec![0f32; n_pairs * out_f];
for g in 0..n_active {
let (lo, hi) = (ex_off_host[g] as usize, ex_off_host[g + 1] as usize);
for p in lo..hi {
let arow = &a_f32[p * in_f..][..in_f];
for o in 0..out_f {
let wrow = &w_f32[(g * out_f + o) * in_f..][..in_f];
let s: f32 = wrow.iter().zip(arow).map(|(w, a)| w * a).sum();
cpu[p * out_f + o] = s * scales[p];
}
}
}
let to_bytes = |v: &[f32]| -> Vec<u8> {
v.iter().flat_map(|&x| f16_bits(x).to_le_bytes()).collect()
};
let wd = e.htod_bytes(&to_bytes(&w_f32))?;
let ad = e.htod_bytes(&to_bytes(&a_f32))?;
let sd = e.htod(&scales)?;
let offd = e.htod_i32(&ex_off_host)?;
let y_legacy = e.dtoh(&e.moe_f16g_gemm_sk_raw(&wd, &ad, &sd, &ex_off_host, &offd,
in_f, out_f, n_pairs, -1, 0, 0)?)?;
let scale = cpu.iter().map(|v| v.abs()).fold(0.0f32, f32::max).max(1e-3);
let rel = maxdiff(&cpu, &y_legacy) / scale;
println!("f16g-sk (in={in_f} out={out_f} skew 1..300) grid-scan vs oracle: \
rel={rel:.2e} {}",
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
let mut y_tail_deep: Option<Vec<f32>> = None;
let mut y_tail_leg: Option<Vec<f32>> = None;
for (name, cross, tail) in [("visitor-hybrid(cross=64,deep-tail)", 64, 1),
("visitor-128", 1, 1),
("visitor-32-deep-tail", i32::MAX, 1),
("visitor-32-legacy-tail", i32::MAX, 0)] {
let yv = e.dtoh(&e.moe_f16g_gemm_sk_raw(&wd, &ad, &sd, &ex_off_host, &offd,
in_f, out_f, n_pairs, 0, cross, tail)?)?;
let d = maxdiff(&y_legacy, &yv);
println!("f16g-sk (in={in_f} out={out_f}) {name} vs grid-scan: maxdiff={d:.2e} {}",
if d == 0.0 { "OK (byte-identical)" } else { fails += 1; "FAIL" });
if cross == i32::MAX {
if tail == 1 { y_tail_deep = Some(yv); } else { y_tail_leg = Some(yv); }
}
}
if let (Some(yd), Some(yl)) = (&y_tail_deep, &y_tail_leg) {
let d = maxdiff(yd, yl);
println!("f16g-sk-tail (in={in_f} out={out_f}) deep(64x3st) vs legacy(32x2st): \
maxdiff={d:.2e} {}",
if d == 0.0 { "OK (byte-identical)" } else { fails += 1; "FAIL" });
}
}
}
{
let m_sizes: [i32; 8] = [1, 3, 17, 33, 64, 129, 200, 300];
let n_active = m_sizes.len();
let mut ex_off_host = vec![0i32; n_active + 1];
for (g, m) in m_sizes.iter().enumerate() { ex_off_host[g + 1] = ex_off_host[g] + m; }
let n_pairs = *ex_off_host.last().unwrap() as usize;
let (in_f, out_f) = (512usize, 300usize); let n_expert = n_active;
for (qname, qtype, sbb) in [("q4_K", memra_engine::QT_Q4_K, 144usize),
("q6_K", memra_engine::QT_Q6_K, 210usize),
("iq4_xs", memra_engine::QT_IQ4_XS, 136usize),
("iq3_s", memra_engine::QT_IQ3_S, 110usize)] {
let row_bytes = in_f / 256 * sbb;
let ex_bytes = out_f * row_bytes;
let mut slab = vec![0u8; n_expert * ex_bytes];
for (i, b) in slab.iter_mut().enumerate() {
*b = (pr(i + 313) * 256.0) as u8;
}
for ex in 0..n_expert {
for r in 0..out_f {
for s in 0..(in_f / 256) {
let off = ex * ex_bytes + r * row_bytes + s * sbb;
let seed = ex * 131 + r * 7 + s;
let h = |k: usize| -> [u8; 2] {
(0x2C00u16 + ((pr(seed + k) * 512.0) as u16)).to_le_bytes()
};
if qtype == memra_engine::QT_Q4_K {
slab[off..off + 2].copy_from_slice(&h(1));
slab[off + 2..off + 4].copy_from_slice(&h(2));
} else if qtype == memra_engine::QT_Q6_K {
slab[off + 208..off + 210].copy_from_slice(&h(1));
} else {
slab[off..off + 2].copy_from_slice(&h(1));
}
}
}
}
let slab_d = e.htod_bytes(&slab)?;
let base = {
use cudarc::driver::DevicePtr;
let s = e.stream();
let (p, _g) = slab_d.device_ptr(&s);
p as u64
};
let tab: Vec<u64> = (0..n_expert).map(|ex| base + (ex * ex_bytes) as u64).collect();
let tab_d = e.htod_u64(&tab)?;
let ex_ids: Vec<i32> = (0..n_active as i32).rev().collect();
let exi_d = e.htod_i32(&ex_ids)?;
let act: Vec<u8> = (0..n_pairs * in_f).flat_map(|i| {
let h = (0x2C00u16 + ((pr(i + 619) * 4096.0) as u16))
| (((i & 1) as u16) << 15);
h.to_le_bytes()
}).collect();
let ad = e.htod_bytes(&act)?;
let scales: Vec<f32> = (0..n_pairs).map(|p| 0.5 + pr(p + 733)).collect();
let sd = e.htod(&scales)?;
let offd = e.htod_i32(&ex_off_host)?;
let ws = e.moe_f16g_dequant_raw(&tab_d, 0, n_expert, &exi_d,
in_f, out_f, n_active, qtype, row_bytes)?;
for (name, cross, tail) in [("hybrid(cross=64,deep-tail)", 64, 1),
("all-128", 1, 1),
("all-32-deep-tail", i32::MAX, 1),
("all-32-legacy-tail", i32::MAX, 0)] {
let y_ws = e.dtoh(&e.moe_f16g_gemm_sk_raw(&ws, &ad, &sd, &ex_off_host, &offd,
in_f, out_f, n_pairs, 0, cross, tail)?)?;
let y_dq = e.dtoh(&e.moe_kq_gemm_sk_raw(&tab_d, 0, n_expert, &exi_d, &ad, &sd,
&ex_off_host, &offd, in_f, out_f,
n_pairs, qtype, row_bytes, cross, tail)?)?;
let d = maxdiff(&y_ws, &y_dq);
println!("f16g-kq-direct [{qname} synth in={in_f} out={out_f}] {name} \
vs workspace: maxdiff={d:.2e} {}",
if d == 0.0 { "OK (byte-identical)" } else { fails += 1; "FAIL" });
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
let o35b = kc_model("f16g-kq-direct", "ornith-1.0-35b-Q4_K_M.gguf",
&["/data/ai-ml/hf-models/ornith-1.0-35b-gguf/ornith-1.0-35b-Q4_K_M.gguf"],
&gguf_arg);
if let Some(path) = o35b.as_deref() {
let g = GgufFile::open(path)?;
let mut cases: Vec<(String, i32, usize)> = Vec::new();
if let Some(t) = g.find("blk.0.ffn_gate_exps.weight")
.filter(|t| t.ggml_type == GgmlType::Q4_K) {
let _ = t; cases.push(("blk.0.ffn_gate_exps.weight".into(),
memra_engine::QT_Q4_K, 144));
}
for l in 0..48 {
let name = format!("blk.{l}.ffn_down_exps.weight");
if g.find(&name).map(|t| t.ggml_type == GgmlType::Q6_K).unwrap_or(false) {
cases.push((name, memra_engine::QT_Q6_K, 210));
break;
}
}
let m_sizes: [i32; 6] = [5, 33, 64, 80, 129, 17];
let n_active = m_sizes.len();
let mut ex_off_host = vec![0i32; n_active + 1];
for (gg, m) in m_sizes.iter().enumerate() { ex_off_host[gg + 1] = ex_off_host[gg] + m; }
let n_pairs = *ex_off_host.last().unwrap() as usize;
for (tname, qtype, sbb) in cases {
let t = g.find(&tname).unwrap();
let (in_f, out_f, ne) = (t.ne[0] as usize, t.ne[1] as usize, t.ne[2] as usize);
if in_f % 256 != 0 || ne < n_active {
println!("f16g-kq-direct [{tname}] SKIP (in_f={in_f} ne={ne})");
continue;
}
let row_bytes = in_f / 256 * sbb;
let ex_bytes = out_f * row_bytes;
let raw = g.tensor_data(t);
let slab_d = e.htod_bytes(&raw[..n_active * ex_bytes])?;
let base = {
use cudarc::driver::DevicePtr;
let s = e.stream();
let (p, _gg) = slab_d.device_ptr(&s);
p as u64
};
let tab: Vec<u64> = (0..n_active).map(|ex| base + (ex * ex_bytes) as u64).collect();
let tab_d = e.htod_u64(&tab)?;
let ex_ids: Vec<i32> = (0..n_active as i32).collect();
let exi_d = e.htod_i32(&ex_ids)?;
let act: Vec<u8> = (0..n_pairs * in_f).flat_map(|i| {
let h = (0x2C00u16 + ((pr(i + 619) * 4096.0) as u16))
| (((i & 1) as u16) << 15);
h.to_le_bytes()
}).collect();
let ad = e.htod_bytes(&act)?;
let scales: Vec<f32> = (0..n_pairs).map(|p| 0.5 + pr(p + 733)).collect();
let sd = e.htod(&scales)?;
let offd = e.htod_i32(&ex_off_host)?;
let ws = e.moe_f16g_dequant_raw(&tab_d, 0, n_active, &exi_d,
in_f, out_f, n_active, qtype, row_bytes)?;
for (name, cross, tail) in [("hybrid(cross=64,deep-tail)", 64, 1),
("all-128", 1, 1),
("all-32-deep-tail", i32::MAX, 1),
("all-32-legacy-tail", i32::MAX, 0)] {
let y_ws = e.dtoh(&e.moe_f16g_gemm_sk_raw(&ws, &ad, &sd, &ex_off_host,
&offd, in_f, out_f, n_pairs, 0, cross, tail)?)?;
let y_dq = e.dtoh(&e.moe_kq_gemm_sk_raw(&tab_d, 0, n_active, &exi_d, &ad,
&sd, &ex_off_host, &offd, in_f, out_f,
n_pairs, qtype, row_bytes, cross, tail)?)?;
let d = maxdiff(&y_ws, &y_dq);
println!("f16g-kq-direct [{tname} in={in_f} out={out_f}] {name} \
vs workspace: maxdiff={d:.2e} {}",
if d == 0.0 { "OK (byte-identical)" } else { fails += 1; "FAIL" });
}
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
let q35 = kc_model("f16g-kq-direct", "Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
&["/data/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf"],
&gguf_arg);
if let Some(path) = q35.as_deref() {
let g = GgufFile::open(path)?;
let mut cases: Vec<(String, &str, i32, usize)> = Vec::new();
for l in 0..48 {
let name = format!("blk.{l}.ffn_gate_exps.weight");
if g.find(&name).map(|t| t.ggml_type == GgmlType::IQ3_S).unwrap_or(false) {
cases.push((name, "iq3_s", memra_engine::QT_IQ3_S, 110));
break;
}
}
for l in 0..48 {
let name = format!("blk.{l}.ffn_down_exps.weight");
if g.find(&name).map(|t| t.ggml_type == GgmlType::IQ4_XS).unwrap_or(false) {
cases.push((name, "iq4_xs", memra_engine::QT_IQ4_XS, 136));
break;
}
}
let m_sizes: [i32; 6] = [5, 33, 64, 80, 129, 17];
let n_active = m_sizes.len();
let mut ex_off_host = vec![0i32; n_active + 1];
for (gg, m) in m_sizes.iter().enumerate() { ex_off_host[gg + 1] = ex_off_host[gg] + m; }
let n_pairs = *ex_off_host.last().unwrap() as usize;
for (tname, qname, qtype, sbb) in cases {
let t = g.find(&tname).unwrap();
let (in_f, out_f, ne) = (t.ne[0] as usize, t.ne[1] as usize, t.ne[2] as usize);
if in_f % 256 != 0 || ne < n_active {
println!("f16g-kq-direct [q35 {tname}] SKIP (in_f={in_f} ne={ne})");
continue;
}
let row_bytes = in_f / 256 * sbb;
let ex_bytes = out_f * row_bytes;
let raw = g.tensor_data(t);
let slab_d = e.htod_bytes(&raw[..n_active * ex_bytes])?;
let base = {
use cudarc::driver::DevicePtr;
let s = e.stream();
let (p, _gg) = slab_d.device_ptr(&s);
p as u64
};
let tab: Vec<u64> = (0..n_active).map(|ex| base + (ex * ex_bytes) as u64).collect();
let tab_d = e.htod_u64(&tab)?;
let ex_ids: Vec<i32> = (0..n_active as i32).collect();
let exi_d = e.htod_i32(&ex_ids)?;
let act: Vec<u8> = (0..n_pairs * in_f).flat_map(|i| {
let h = (0x2C00u16 + ((pr(i + 619) * 4096.0) as u16))
| (((i & 1) as u16) << 15);
h.to_le_bytes()
}).collect();
let ad = e.htod_bytes(&act)?;
let scales: Vec<f32> = (0..n_pairs).map(|p| 0.5 + pr(p + 733)).collect();
let sd = e.htod(&scales)?;
let offd = e.htod_i32(&ex_off_host)?;
let ws = e.moe_f16g_dequant_raw(&tab_d, 0, n_active, &exi_d,
in_f, out_f, n_active, qtype, row_bytes)?;
for (name, cross, tail) in [("hybrid(cross=64,deep-tail)", 64, 1),
("all-128", 1, 1),
("all-32-deep-tail", i32::MAX, 1),
("all-32-legacy-tail", i32::MAX, 0)] {
let y_ws = e.dtoh(&e.moe_f16g_gemm_sk_raw(&ws, &ad, &sd, &ex_off_host,
&offd, in_f, out_f, n_pairs, 0, cross, tail)?)?;
let y_dq = e.dtoh(&e.moe_kq_gemm_sk_raw(&tab_d, 0, n_active, &exi_d, &ad,
&sd, &ex_off_host, &offd, in_f, out_f,
n_pairs, qtype, row_bytes, cross, tail)?)?;
let d = maxdiff(&y_ws, &y_dq);
println!("f16g-kq-direct [q35 {tname} {qname} in={in_f} out={out_f}] {name} \
vs workspace: maxdiff={d:.2e} {}",
if d == 0.0 { "OK (byte-identical)" } else { fails += 1; "FAIL" });
}
}
}
}
{
let iq4xs_gate = |e: &Engine, wd: &_, in_f: usize, out_f: usize, row_bytes: usize,
label: &str, fails: &mut i32|
-> Result<(), Box<dyn std::error::Error>> {
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 47) * 0.1).collect();
let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec_iq4_XS_fast(wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let yb = e.dtoh(&e.qmatvec_mmq_iq4xs_raw(wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("iq4xs-mmq [{label}] T={tt}: rel={rel:.2e} {}",
if rel < 1e-3 { "OK" } else { *fails += 1; "FAIL" });
}
Ok(())
};
{
let (in_f, out_f) = (512usize, 300usize);
let row_bytes = in_f / 256 * 136;
let mut w = vec![0u8; out_f * row_bytes];
for (i, b) in w.iter_mut().enumerate() { *b = (pr(i + 409) * 256.0) as u8; }
for r in 0..out_f {
for s in 0..(in_f / 256) {
let off = r * row_bytes + s * 136;
let h = 0x2C00u16 + ((pr(r * 7 + s + 3) * 512.0) as u16);
w[off..off + 2].copy_from_slice(&h.to_le_bytes());
}
}
let wd = e.htod_bytes(&w)?;
iq4xs_gate(&e, &wd, in_f, out_f, row_bytes, "synth", &mut fails)?;
}
{
use memra_gguf::{GgufFile, GgmlType};
let kat = kc_model("iq4xs-mmq", "Kwaipilot_KAT-Coder-V2.5-Dev-IQ4_XS.gguf",
&["/data/ai-ml/hf-models/kat-coder-v25-dev-gguf/Kwaipilot_KAT-Coder-V2.5-Dev-IQ4_XS.gguf"],
&gguf_arg);
if let Some(path) = kat.as_deref() {
let g = GgufFile::open(path)?;
if let Some(t) = g.tensors.iter().find(|t| {
t.ggml_type == GgmlType::IQ4_XS && t.ne.len() == 2
&& t.ne[0] as usize % 256 == 0 && t.ne[1] >= 128
}) {
let (in_f, out_f) = (t.ne[0] as usize, t.ne[1] as usize);
let raw = g.tensor_data(t);
let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
iq4xs_gate(&e, &wd, in_f, out_f, row_bytes, &t.name, &mut fails)?;
}
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
let gguf_9b_owned = kc_model("nvfp4-gemm", "Qwen3.5-9B-NVFP4-MTP-GGUF.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf",
"/home/ubuntu/memra-bench/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf"],
&gguf_arg)
.or_else(|| kc_model("nvfp4-gemm", "Qwen3.6-27B-NVFP4-Q4_K_M-mtp.gguf",
&["/data/ai-ml/hf-models/qwen36-27b-nvfp4-mtp/Qwen3.6-27B-NVFP4-Q4_K_M-mtp.gguf"],
&gguf_arg));
if let Some(gguf_9b) = gguf_9b_owned.as_deref() {
let g = GgufFile::open(gguf_9b)?;
if let Some(t) = g.find("blk.0.attn_gate.weight").filter(|t| t.ggml_type == GgmlType::Q5_K) {
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 91) * 0.1).collect();
let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec_q5_K_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let yb = e.dtoh(&e.qmatvec_gemm_raw(&wd, &xd, tt, in_f, out_f, memra_engine::QT_Q5_K, row_bytes)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("GEMM blk.0.attn_gate.weight [Q5_K] T={tt}: rel={rel:.2e} {}",
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
}
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 81) * 0.1).collect();
let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec_nvfp4_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let yb = e.dtoh(&e.qmatvec_gemm_raw(&wd, &xd, tt, in_f, out_f, memra_engine::QT_NVFP4, row_bytes)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("GEMM blk.0.ffn_gate.weight [NVFP4] T={tt}: rel={rel:.2e} {}",
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
}
let nvfp4_checks = !cfg!(memra_portable_cuda);
if cfg!(memra_portable_cuda) && !cfg!(memra_hopper_mma) {
println!("portable CUDA: native FP4 and static-MMQ model-backed checks — SKIP");
} else {
if nvfp4_checks {
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let w_f32 = dequant::dequantize(GgmlType::NVFP4, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 83) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let yb = e.dtoh(&e.qmatvec_gemm_nvfp4_fp4_raw(&wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let d = maxdiff(&cpu, &yb);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("FP4-GEMM blk.0.ffn_gate.weight [NVFP4] T={tt}: rel={rel:.2e} (informational; \
authoritative gate = argmax) {}", if rel < 2e-1 { "OK" } else { "HIGH" });
}
}
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let _ = row_bytes;
let w_f32 = dequant::dequantize(GgmlType::NVFP4, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 83) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let yb = e.dtoh(&e.qmatvec_mmq_nvfp4_raw(&wd, &xd, tt, in_f, out_f)?)?;
let rel = maxdiff(&cpu, &yb) / scale;
let y1 = e.dtoh(&e.qmatvec_mmq_nvfp4_raw_v1(&wd, &xd, tt, in_f, out_f)?)?;
let rel_v1 = maxdiff(&cpu, &y1) / scale;
println!("MMQ-GEMM blk.0.ffn_gate.weight [NVFP4] T={tt}: rel={rel:.2e} (informational; \
authoritative gate = argmax) {}", if rel < 2e-1 { "OK" } else { "HIGH" });
println!("MMQ-GEMM-V1 blk.0.ffn_gate.weight [NVFP4] T={tt}: rel={rel_v1:.2e} \
(pre-port oracle; two-level/{}) {}",
if rel > 0.0 { format!("{:.2}x", rel_v1 / rel) } else { "n/a".into() },
if rel <= rel_v1 { "IMPROVED" } else { "WORSE" });
}
for tt in [16usize, 128] {
let x: Vec<f32> = (0..tt * in_f)
.map(|i| {
let token = i / in_f;
let decade = 10.0f32.powi((token % 7) as i32 - 3);
pr(i + 83) * 0.1 * decade
})
.collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let per_token_rel = |y: &[f32]| -> f32 {
(0..tt)
.map(|j| {
let lo = j * out_f;
let hi = lo + out_f;
let s = cpu[lo..hi].iter().map(|v| v.abs()).fold(0.0, f32::max);
if s <= 0.0 { return 0.0; }
maxdiff(&cpu[lo..hi], &y[lo..hi]) / s
})
.fold(0.0, f32::max)
};
let yb = e.dtoh(&e.qmatvec_mmq_nvfp4_raw(&wd, &xd, tt, in_f, out_f)?)?;
let y1 = e.dtoh(&e.qmatvec_mmq_nvfp4_raw_v1(&wd, &xd, tt, in_f, out_f)?)?;
let rel = per_token_rel(&yb);
let rel_v1 = per_token_rel(&y1);
println!("MMQ-GEMM-DYN blk.0.ffn_gate.weight [NVFP4] T={tt}: rel={rel:.2e} \
v1={rel_v1:.2e} ({}) {}",
if rel > 0.0 { format!("{:.2}x", rel_v1 / rel) } else { "n/a".into() },
if rel < rel_v1 { "OK" } else { fails += 1; "FAIL" });
}
{
let tt = 128usize;
let hot: Vec<usize> = (0..8).map(|c| (c * 977 + 13) % in_f).collect();
let is_hot = {
let mut v = vec![false; in_f];
for &c in &hot { v[c] = true; }
v
};
let x: Vec<f32> = (0..tt * in_f)
.map(|i| {
let token = i / in_f;
let chan = i % in_f;
let decade = 10.0f32.powi((token % 7) as i32 - 3);
let boost = if is_hot[chan] { 300.0 } else { 1.0 };
pr(i + 83) * 0.1 * decade * boost
})
.collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let per_token_rel = |y: &[f32]| -> f32 {
(0..tt)
.map(|j| {
let lo = j * out_f;
let hi = lo + out_f;
let s = cpu[lo..hi].iter().map(|v| v.abs()).fold(0.0, f32::max);
if s <= 0.0 { return 0.0; }
maxdiff(&cpu[lo..hi], &y[lo..hi]) / s
})
.fold(0.0, f32::max)
};
let mut rel_k0 = f32::NAN;
for k in [0i32, 4, 8, 16, 32, 64] {
let y = e.dtoh(&e.qmatvec_mmq_nvfp4_raw_res(&wd, &xd, tt, in_f, out_f, k)?)?;
let r = per_token_rel(&y);
if k == 0 {
rel_k0 = r;
println!("MMQ-GEMM-RES blk.0.ffn_gate.weight [NVFP4] T={tt} k=0: rel={r:.2e} \
(baseline, 8 outlier channels @300x)");
continue;
}
let hard = k >= 8;
let better = r < rel_k0;
let verdict = if better { "OK" } else if hard { fails += 1; "FAIL" } else { "FLAT" };
println!("MMQ-GEMM-RES blk.0.ffn_gate.weight [NVFP4] T={tt} k={k}: rel={r:.2e} \
({} vs k=0){} {verdict}",
if r > 0.0 { format!("{:.2}x", rel_k0 / r) } else { "n/a".into() },
if hard { "" } else { " informational" });
}
}
}
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
use memra_gguf::dequant;
use memra_engine::model::repack_nvfp4_split;
use memra_runtime::cpu_linear;
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t);
let w_f32 = dequant::dequantize(GgmlType::NVFP4, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
let wd_rp = e.htod_bytes(&repack_nvfp4_split(raw, out_f))?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 83) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let yb = e.dtoh(&e.qmatvec_mmq_nvfp4_w4a8_raw(&wd, &xd, tt, in_f, out_f)?)?;
let d = maxdiff(&cpu, &yb);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMQ-W4A8 blk.0.ffn_gate.weight [NVFP4] T={tt}: rel={rel:.2e} (int8 band ~1e-3) {}",
if rel < 2e-2 { "OK" } else { fails += 1; "FAIL" });
let yr = e.dtoh(&e.qmatvec_mmq_nvfp4_w4a8_raw_rp(&wd_rp, &xd, tt, in_f, out_f)?)?;
let nbad = yb.iter().zip(yr.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("MMQ-W4A8-RP blk.0.ffn_gate.weight [NVFP4] T={tt}: bit-mismatch {nbad}/{} {}",
yb.len(), if nbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
} for (tname, want, qt) in [("blk.3.attn_q.weight", GgmlType::Q4_K, memra_engine::QT_Q4_K),
("blk.0.attn_gate.weight", GgmlType::Q5_K, memra_engine::QT_Q5_K)] {
let Some(t) = g.find(tname).filter(|t| t.ggml_type == want) else { continue };
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t);
let w_f32 = dequant::dequantize(want, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 87) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let yb = e.dtoh(&e.qmatvec_mmq_q45k_raw(&wd, &xd, tt, in_f, out_f, qt)?)?;
let d = maxdiff(&cpu, &yb);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMQ-GEMM {tname} [{want:?}] T={tt}: rel={rel:.2e} {}",
if rel < 2e-2 { "OK" } else { fails += 1; "FAIL" });
}
}
#[cfg(memra_cutlass)]
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let w_f32 = dequant::dequantize(GgmlType::NVFP4, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
let mut b_packed = e.alloc_u8(out_f * in_f / 2)?;
let mut sfb_lin = e.alloc_u8(out_f * (in_f / 16))?;
e.cutlass_gguf_nvfp4_deinterleave(&wd, row_bytes, &mut b_packed, &mut sfb_lin, out_f, in_f)?;
let mut w_rt_d = e.htod(&vec![0f32; out_f * in_f])?;
e.cutlass_nvfp4_dequant_ref(&b_packed, &sfb_lin, &mut w_rt_d, out_f, in_f)?;
let w_rt = e.dtoh(&w_rt_d)?;
let wmax = w_f32.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-6);
let wrel = maxdiff(&w_f32, &w_rt) / wmax;
println!("CUTLASS-FP4 weight round-trip blk.0.ffn_gate.weight [NVFP4]: rel={wrel:.2e} {}",
if wrel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
let (b_packed_sw, sfb_sw) = e.build_cutlass_weight(&wd, out_f, in_f, row_bytes)?;
for tt in [128usize, 512] { let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 87) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let yhr = e.dtoh(&e.qmatvec_gemm_nvfp4_fp4_raw(&wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let ycl = e.dtoh(&e.cutlass_fp4_gemm(&b_packed_sw, &sfb_sw, &xd, 1.0, tt, out_f, in_f)?)?;
let rel_hr = maxdiff(&cpu, &yhr) / scale;
let rel_cl = maxdiff(&cpu, &ycl) / scale;
let ok = (rel_cl - rel_hr).abs() < 5e-2 && rel_cl < 2e-1;
println!("CUTLASS-FP4 GEMM-band blk.0.ffn_gate.weight [NVFP4] T={tt}: rel_cutlass={rel_cl:.2e} \
rel_handroll={rel_hr:.2e} {}", if ok { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
{
let g35_path = kc_model("q8mmq-gemm", "Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
&["/data/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
"/home/avifenesh/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf"],
&gguf_arg);
if let Some(g35_path) = g35_path {
let g35 = GgufFile::open(&g35_path)?;
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
for tname in ["blk.0.attn_qkv.weight", "blk.0.ffn_gate_shexp.weight"] {
let Some(t) = g35.find(tname).filter(|t| t.ggml_type == GgmlType::Q8_0) else { continue };
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g35.tensor_data(t);
let w_f32 = dequant::dequantize(GgmlType::Q8_0, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 53) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let yb = e.dtoh(&e.qmatvec_mmq_q8_0_raw(&wd, &xd, tt, in_f, out_f)?)?;
let d = maxdiff(&cpu, &yb);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMQ-Q8_0 {tname} [Q8_0 in={in_f} out={out_f}] T={tt}: rel={rel:.2e} {}",
if rel < 2e-2 { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
{
let g12_path = kc_model("q4_0-mmq", "gemma-4-12b-it-qat-q4_0.gguf",
&["/data/ai-ml/models/gemma-4-12b-it-qat/gemma-4-12b-it-qat-q4_0.gguf"],
&gguf_arg);
fn repack_q4_0_split(raw: &[u8], nblocks: usize) -> Vec<u8> {
let mut out = vec![0u8; nblocks * 18];
let dplane = nblocks * 16;
for i in 0..nblocks {
let b = &raw[i * 18..i * 18 + 18];
out[i * 16..i * 16 + 16].copy_from_slice(&b[2..18]);
out[dplane + i * 2] = b[0];
out[dplane + i * 2 + 1] = b[1];
}
out
}
if let Some(g12_path) = g12_path {
let g12 = GgufFile::open(&g12_path)?;
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
for tname in ["blk.0.attn_q.weight", "blk.0.ffn_gate.weight"] {
let Some(t) = g12.find(tname).filter(|t| t.ggml_type == GgmlType::Q4_0) else { continue };
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g12.tensor_data(t);
let w_f32 = dequant::dequantize(GgmlType::Q4_0, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
let wd_rp = e.htod_bytes(&repack_q4_0_split(raw, out_f * in_f / 32))?;
for tt in [16usize, 64, 128, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 59) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let yb = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
let d = maxdiff(&cpu, &yb);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMQ-Q4_0 {tname} [Q4_0 in={in_f} out={out_f}] T={tt}: rel={rel:.2e} {}",
if rel < 2e-2 { "OK" } else { fails += 1; "FAIL" });
let yr = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd_rp, &xd, tt, in_f, out_f, true)?)?;
let nbad = yb.iter().zip(yr.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("MMQ-Q4_0-RP {tname} T={tt}: bit-mismatch {nbad}/{} {}",
yb.len(), if nbad == 0 { "OK" } else { fails += 1; "FAIL" });
memra_engine::MMQ_SK_FORCE.store(0, std::sync::atomic::Ordering::Relaxed);
let clc_avail = unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(0) } == 1;
if clc_avail {
let yt = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(1) };
let yc = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
let nbad = yt.iter().zip(yc.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("MMQ-Q4_0-CLC {tname} T={tt}: bit-mismatch {nbad}/{} {}",
yt.len(), if nbad == 0 { "OK" } else { fails += 1; "FAIL" });
let yc_rp = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd_rp, &xd, tt, in_f, out_f, true)?)?;
let nbad_rp = yt.iter().zip(yc_rp.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("MMQ-Q4_0-CLC-RP {tname} T={tt}: bit-mismatch {nbad_rp}/{} {}",
yt.len(), if nbad_rp == 0 { "OK" } else { fails += 1; "FAIL" });
} else {
println!("MMQ-Q4_0-CLC {tname} T={tt}: SKIP (pre-SM100 build — CLC kernel not compiled)");
}
unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(-1) };
memra_engine::MMQ_SK_FORCE.store(-1, std::sync::atomic::Ordering::Relaxed);
}
}
}
}
{
let g27_path = kc_model("nvfp4-27b-shape", "Qwen3.6-27B-NVFP4-Q4_K_M-mtp.gguf",
&["/data/ai-ml/hf-models/qwen36-27b-nvfp4-mtp/Qwen3.6-27B-NVFP4-Q4_K_M-mtp.gguf"],
&gguf_arg);
if let Some(g27_path) = g27_path {
let g27 = GgufFile::open(&g27_path)?;
for tn in ["blk.0.ffn_down.weight", "blk.0.ffn_gate.weight"] {
if let Some(t) = g27.find(tn).filter(|t| t.ggml_type == GgmlType::NVFP4) {
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g27.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for tt in [16usize, 512] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 71) * 0.1).collect();
let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec_nvfp4_fast(&wd, &xd, tt, in_f, out_f, row_bytes)?)?;
let yb = e.dtoh(&e.qmatvec_mmq_nvfp4_raw(&wd, &xd, tt, in_f, out_f)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMQ-27B {tn} [NVFP4 in={in_f} out={out_f}] T={tt}: rel={rel:.2e} (W4A4-vs-dp4a band ~0.1) {}",
if rel < 2.5e-1 { "OK" } else { "HIGH" });
}
}
}
}
}
{
let g26_path = kc_model("q4_0-sk-arm", "gemma-4-26B_q4_0-it.gguf",
&["/data/ai-ml/hf-models/gemma4-26b-a4b-qat-gguf/gemma-4-26B_q4_0-it.gguf"],
&gguf_arg);
if let Some(g26_path) = g26_path {
let g26 = GgufFile::open(&g26_path)?;
use memra_gguf::dequant;
use memra_runtime::cpu_linear;
{
let (in_f, out_f) = (2112usize, 2560usize);
let nblk = in_f / 32 * out_f;
let mut raw = vec![0u8; nblk * 18];
for (bi, b) in raw.chunks_mut(18).enumerate() {
b[0] = 0x00; b[1] = 0x3C; for k in 0..16 { b[2 + k] = ((bi * 31 + k * 7) % 251) as u8; }
}
let w_f32 = dequant::dequantize(GgmlType::Q4_0, &raw, in_f * out_f);
let wd = e.htod_bytes(&raw)?;
for tt in [103usize, 479] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 83) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let mut y_tile: Vec<f32> = Vec::new();
for (force, label) in [(0i8, "TILE"), (1, "SK")] {
memra_engine::MMQ_SK_FORCE.store(force, std::sync::atomic::Ordering::Relaxed);
let yb = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
let rel = maxdiff(&cpu, &yb) / scale;
println!("MMQ-Q4_0-RAGK {label} [in={in_f} out={out_f} nc=false] T={tt}: rel={rel:.2e} {}",
if rel < 2e-2 { "OK" } else { fails += 1; "FAIL" });
if force == 0 { y_tile = yb; }
}
memra_engine::MMQ_SK_FORCE.store(0, std::sync::atomic::Ordering::Relaxed);
if unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(1) } == 1 {
let yc = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
let nbad = y_tile.iter().zip(yc.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("MMQ-Q4_0-RAGK CLC [in={in_f} out={out_f} nc=false] T={tt}: bit-mismatch {nbad}/{} {}",
yc.len(), if nbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(-1) };
memra_engine::MMQ_SK_FORCE.store(-1, std::sync::atomic::Ordering::Relaxed);
}
}
for tname in ["blk.0.attn_q.weight", "blk.0.attn_k.weight", "blk.0.attn_v.weight",
"blk.0.attn_output.weight", "blk.0.ffn_gate.weight",
"blk.0.ffn_down.weight"] {
let Some(t) = g26.find(tname).filter(|t| t.ggml_type == GgmlType::Q4_0) else { continue };
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g26.tensor_data(t);
let w_f32 = dequant::dequantize(GgmlType::Q4_0, raw, in_f * out_f);
let wd = e.htod_bytes(raw)?;
for tt in [103usize, 229, 479, 1024, 2048, 2151] {
let x: Vec<f32> = (0..tt * in_f).map(|i| pr(i + 83) * 0.1).collect();
let xd = e.htod(&x)?;
let cpu = cpu_linear(&x, &w_f32, tt, in_f, out_f);
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let mut y_tile: Vec<f32> = Vec::new();
for (force, label) in [(0i8, "TILE"), (1, "SK")] {
memra_engine::MMQ_SK_FORCE.store(force, std::sync::atomic::Ordering::Relaxed);
let yb = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
let rel = maxdiff(&cpu, &yb) / scale;
println!("MMQ-Q4_0-NC26 {tname} {label} [in={in_f} out={out_f}] T={tt}: rel={rel:.2e} {}",
if rel < 2e-2 { "OK" } else { fails += 1; "FAIL" });
if force == 0 { y_tile = yb; }
}
memra_engine::MMQ_SK_FORCE.store(0, std::sync::atomic::Ordering::Relaxed);
if unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(1) } == 1 {
let yc = e.dtoh(&e.qmatvec_mmq_q4_0_raw(&wd, &xd, tt, in_f, out_f, false)?)?;
let nbad = y_tile.iter().zip(yc.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("MMQ-Q4_0-NC26 {tname} CLC [in={in_f} out={out_f}] T={tt}: bit-mismatch {nbad}/{} {}",
yc.len(), if nbad == 0 { "OK" } else { fails += 1; "FAIL" });
}
unsafe { memra_engine::mmq_ffi::memra_mmq_q4_0_set_clc(-1) };
memra_engine::MMQ_SK_FORCE.store(-1, std::sync::atomic::Ordering::Relaxed);
}
}
}
}
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType};
let g = GgufFile::open(&path)?;
let mmvq_cases: [(&str, i32, &str); 5] = [
("blk.0.ffn_gate.weight", memra_engine::QT_Q8_0, "q8_0"),
("blk.0.attn_qkv.weight", memra_engine::QT_Q8_0, "q8_0"),
("blk.3.attn_q.weight", memra_engine::QT_Q4_K, "q4_K"),
("blk.0.attn_v.weight", memra_engine::QT_Q6_K, "q6_K"),
("output.weight", memra_engine::QT_Q6_K, "q6_K"),
];
for (tname, want_qt, sel) in mmvq_cases {
let t = match g.find(tname) { Some(t) => t, None => continue };
let gt = match t.ggml_type {
GgmlType::Q8_0 => memra_engine::QT_Q8_0, GgmlType::Q4_K => memra_engine::QT_Q4_K,
GgmlType::Q6_K => memra_engine::QT_Q6_K, GgmlType::NVFP4 => memra_engine::QT_NVFP4,
_ => continue,
};
if gt != want_qt { continue; }
if t.ne.len() > 2 { continue; } let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for mm in [1usize, 2] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 101) * 0.1).collect();
let xd = e.htod(&x)?;
let ydp = match sel {
"q8_0" => e.qmatvec_q8_0_fast(&wd, &xd, mm, in_f, out_f, row_bytes)?,
"q4_K" => e.qmatvec_q4_K_fast(&wd, &xd, mm, in_f, out_f, row_bytes)?,
"q6_K" => e.qmatvec_q6_K_fast(&wd, &xd, mm, in_f, out_f, row_bytes)?,
_ => unreachable!(),
};
let ya = e.dtoh(&ydp)?;
let yb = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, gt, row_bytes, false)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMVQ {tname} [{:?}] m={mm}: rel={rel:.2e} {}", t.ggml_type,
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
}
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType};
let g = GgufFile::open(&path)?;
let pair_sets: [(&str, &str); 3] = [
("blk.0.attn_qkv.weight", "blk.0.attn_gate.weight"), ("blk.0.ffn_gate_shexp.weight", "blk.0.ffn_up_shexp.weight"), ("blk.0.ssm_beta.weight", "blk.0.ssm_alpha.weight"), ];
let grab = |name: &str| -> Option<(usize, usize, usize, Vec<u8>)> {
let t = g.find(name)?;
if t.ggml_type != GgmlType::Q8_0 || t.ne.len() > 2 { return None; }
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t);
Some((in_f, out_f, raw.len() / out_f, raw.to_vec()))
};
for (n0, n1) in pair_sets {
let (Some(t0), Some(t1)) = (grab(n0), grab(n1)) else { continue };
if t0.0 != t1.0 { continue; }
let (in_f, rb) = (t0.0, t0.2);
let w0 = e.htod_bytes(&t0.3)?;
let w1 = e.htod_bytes(&t1.3)?;
let x: Vec<f32> = (0..in_f).map(|i| pr(i + 131) * 0.1).collect();
let xd = e.htod(&x)?;
let r0 = e.dtoh(&e.qmatvec_mmvq_raw(&w0, &xd, 1, in_f, t0.1, memra_engine::QT_Q8_0, rb, false)?)?;
let r1 = e.dtoh(&e.qmatvec_mmvq_raw(&w1, &xd, 1, in_f, t1.1, memra_engine::QT_Q8_0, rb, false)?)?;
let (f0, f1) = e.qmatvec_q8_fused2_raw(&w0, &w1, &xd, in_f, t0.1, t1.1, rb)?;
let (f0, f1) = (e.dtoh(&f0)?, e.dtoh(&f1)?);
let bits_ok = r0.iter().zip(f0.iter()).all(|(a, b)| a.to_bits() == b.to_bits())
&& r1.iter().zip(f1.iter()).all(|(a, b)| a.to_bits() == b.to_bits());
let d = maxdiff(&r0, &f0).max(maxdiff(&r1, &f1));
println!("Q8-FUSED2 {n0}+{n1} [Q8_0] out=({},{}): rel={d:.2e} bits={} {}",
t0.1, t1.1, bits_ok,
if bits_ok { "OK" } else { fails += 1; "FAIL" });
for mm in [2usize, 3, 4, 5, 8] {
let xm: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 151 + mm) * 0.1).collect();
let xmd = e.htod(&xm)?;
let (aq, ad) = e.quantize_q8_1(&xmd, mm, in_f)?;
let mc = memra_engine::Engine::batched_mcols(mm);
let r0 = e.dtoh(&e.qmatvec_mmvq_batched(&w0, &aq, &ad, mm, in_f, t0.1,
memra_engine::QT_Q8_0, rb, mc, 1.0, false)?)?;
let r1 = e.dtoh(&e.qmatvec_mmvq_batched(&w1, &aq, &ad, mm, in_f, t1.1,
memra_engine::QT_Q8_0, rb, mc, 1.0, false)?)?;
let (f0, f1) = e.qmatvec_q8_fused2_t_raw(&w0, &w1, &xmd, mm, in_f, t0.1, t1.1, rb)?;
let (f0, f1) = (e.dtoh(&f0)?, e.dtoh(&f1)?);
let bits_ok = r0.iter().zip(f0.iter()).all(|(a, b)| a.to_bits() == b.to_bits())
&& r1.iter().zip(f1.iter()).all(|(a, b)| a.to_bits() == b.to_bits());
let d = maxdiff(&r0, &f0).max(maxdiff(&r1, &f1));
println!("Q8-FUSED2-B {n0}+{n1} [Q8_0] m={mm} out=({},{}): rel={d:.2e} bits={} {}",
t0.1, t1.1, bits_ok,
if bits_ok { "OK" } else { fails += 1; "FAIL" });
}
}
let tri: [&str; 3] = ["blk.3.attn_q.weight", "blk.3.attn_k.weight", "blk.3.attn_v.weight"];
if let (Some(t0), Some(t1), Some(t2)) = (grab(tri[0]), grab(tri[1]), grab(tri[2])) {
if t0.0 == t1.0 && t1.0 == t2.0 {
let (in_f, rb) = (t0.0, t0.2);
let w0 = e.htod_bytes(&t0.3)?;
let w1 = e.htod_bytes(&t1.3)?;
let w2 = e.htod_bytes(&t2.3)?;
let x: Vec<f32> = (0..in_f).map(|i| pr(i + 137) * 0.1).collect();
let xd = e.htod(&x)?;
let r0 = e.dtoh(&e.qmatvec_mmvq_raw(&w0, &xd, 1, in_f, t0.1, memra_engine::QT_Q8_0, rb, false)?)?;
let r1 = e.dtoh(&e.qmatvec_mmvq_raw(&w1, &xd, 1, in_f, t1.1, memra_engine::QT_Q8_0, rb, false)?)?;
let r2 = e.dtoh(&e.qmatvec_mmvq_raw(&w2, &xd, 1, in_f, t2.1, memra_engine::QT_Q8_0, rb, false)?)?;
let (f0, f1, f2) = e.qmatvec_q8_fused3_raw(&w0, &w1, &w2, &xd, in_f, t0.1, t1.1, t2.1, rb)?;
let (f0, f1, f2) = (e.dtoh(&f0)?, e.dtoh(&f1)?, e.dtoh(&f2)?);
let bits_ok = r0.iter().zip(f0.iter()).all(|(a, b)| a.to_bits() == b.to_bits())
&& r1.iter().zip(f1.iter()).all(|(a, b)| a.to_bits() == b.to_bits())
&& r2.iter().zip(f2.iter()).all(|(a, b)| a.to_bits() == b.to_bits());
let d = maxdiff(&r0, &f0).max(maxdiff(&r1, &f1)).max(maxdiff(&r2, &f2));
println!("Q8-FUSED3 wq+wk+wv [Q8_0] out=({},{},{}): rel={d:.2e} bits={} {}",
t0.1, t1.1, t2.1, bits_ok,
if bits_ok { "OK" } else { fails += 1; "FAIL" });
for mm in [2usize, 3, 4] {
let xm: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 157 + mm) * 0.1).collect();
let xmd = e.htod(&xm)?;
let (aq, ad) = e.quantize_q8_1(&xmd, mm, in_f)?;
let mc = memra_engine::Engine::batched_mcols(mm);
let r0 = e.dtoh(&e.qmatvec_mmvq_batched(&w0, &aq, &ad, mm, in_f, t0.1,
memra_engine::QT_Q8_0, rb, mc, 1.0, false)?)?;
let r1 = e.dtoh(&e.qmatvec_mmvq_batched(&w1, &aq, &ad, mm, in_f, t1.1,
memra_engine::QT_Q8_0, rb, mc, 1.0, false)?)?;
let r2 = e.dtoh(&e.qmatvec_mmvq_batched(&w2, &aq, &ad, mm, in_f, t2.1,
memra_engine::QT_Q8_0, rb, mc, 1.0, false)?)?;
let (f0, f1, f2) = e.qmatvec_q8_fused3_t_raw(&w0, &w1, &w2, &xmd, mm, in_f,
t0.1, t1.1, t2.1, rb)?;
let (f0, f1, f2) = (e.dtoh(&f0)?, e.dtoh(&f1)?, e.dtoh(&f2)?);
let bits_ok = r0.iter().zip(f0.iter()).all(|(a, b)| a.to_bits() == b.to_bits())
&& r1.iter().zip(f1.iter()).all(|(a, b)| a.to_bits() == b.to_bits())
&& r2.iter().zip(f2.iter()).all(|(a, b)| a.to_bits() == b.to_bits());
let d = maxdiff(&r0, &f0).max(maxdiff(&r1, &f1)).max(maxdiff(&r2, &f2));
println!("Q8-FUSED3-B wq+wk+wv [Q8_0] m={mm} out=({},{},{}): rel={d:.2e} bits={} {}",
t0.1, t1.1, t2.1, bits_ok,
if bits_ok { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
let gguf_9b = kc_model("nvfp4-mmvq", "Qwen3.5-9B-NVFP4-MTP-GGUF.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf"],
&gguf_arg);
if let Some(gguf_9b) = gguf_9b {
let g = GgufFile::open(&gguf_9b)?;
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for mm in [1usize, 2] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 111) * 0.1).collect();
let xd = e.htod(&x)?;
let ya = e.dtoh(&e.qmatvec_nvfp4_fast(&wd, &xd, mm, in_f, out_f, row_bytes)?)?;
let yb = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, false)?)?;
let d = maxdiff(&ya, &yb);
let scale = ya.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("MMVQ blk.0.ffn_gate.weight [NVFP4] m={mm}: rel={rel:.2e} {}",
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
{
let e4m3 = |b: u8| -> f32 {
let s = if b & 0x80 != 0 { -1.0f32 } else { 1.0 };
let ex = ((b >> 3) & 0x0F) as i32;
let mn = (b & 0x07) as f32;
if ex == 0 { s * mn * (2f32).powi(-9) }
else if ex == 15 && mn == 7.0 { f32::NAN }
else { s * (1.0 + mn / 8.0) * (2f32).powi(ex - 7) }
};
let qt = memra_engine::QT_F8_E4M3;
for (in_f, out_f) in [(5120usize, 512usize), (2048, 320)] {
let wb: Vec<u8> = (0..in_f * out_f).map(|i| {
let mut b = ((i.wrapping_mul(2654435761) ^ 0x9E3779B9) >> 9) as u8;
if b & 0x7F == 0x7F { b &= 0xF7; }
b
}).collect();
let wd = e.htod_bytes(&wb)?;
let row_bytes = in_f; for mm in [1usize, 2, 5, 9] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 151) * 0.1).collect();
let xd = e.htod(&x)?;
let (aqd, add) = e.quantize_q8_1(&xd, mm, in_f)?;
let y = e.dtoh(&e.qmatvec_mmvq(&wd, &aqd, &add, mm, in_f, out_f, qt, row_bytes,
1.0, false)?)?;
let aq: Vec<i8> = e.stream().clone_dtoh(&aqd)?; e.stream().synchronize()?;
let ad = e.dtoh(&add)?;
let nblk = in_f / 32;
let mut cpu = vec![0f32; mm * out_f];
for t in 0..mm {
for o in 0..out_f {
let mut acc = 0f64;
for blk in 0..nblk {
let mut bs = 0f64;
for j in 0..32 {
let w = e4m3(wb[o * in_f + blk * 32 + j]) as f64;
bs += w * aq[t * in_f + blk * 32 + j] as f64;
}
acc += ad[t * nblk + blk] as f64 * bs;
}
cpu[t * out_f + o] = acc as f32;
}
}
let scale = cpu.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = maxdiff(&cpu, &y) / scale;
let mut ok = rel < 1e-3;
let mut bits_ok = true;
if mm > 1 {
for t in 0..mm {
let xt = &x[t * in_f..(t + 1) * in_f];
let xtd = e.htod(xt)?;
let y1 = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xtd, 1, in_f, out_f, qt,
row_bytes, false)?)?;
bits_ok &= y1.iter().zip(&y[t * out_f..(t + 1) * out_f])
.all(|(a, b)| a.to_bits() == b.to_bits());
}
ok &= bits_ok;
}
println!("E4M3-MMVQ synth [{in_f}x{out_f}] m={mm}: rel={rel:.2e} m1-bits={bits_ok} {}",
if ok { "OK" } else { fails += 1; "FAIL" });
}
for mm in 2..=8usize {
let mcols = memra_engine::Engine::batched_mcols(mm);
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 163) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, qt, row_bytes, false)?)?;
let yb = e.dtoh(&e.qmatvec_batched_raw(&wd, &xd, mm, in_f, out_f, qt, row_bytes,
mcols, false)?)?;
let bits_bad = yref.iter().zip(&yb).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
let d = maxdiff(&yref, &yb);
println!("E4M3-BATCHED synth [{in_f}x{out_f}] m={mm} b{mcols}: rel={d:.2e} bit-bad={bits_bad} {}",
if bits_bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let bitbad = |a: &[f32], b: &[f32]| -> usize {
a.iter().zip(b).filter(|(x, y)| x.to_bits() != y.to_bits()).count()
+ a.len().abs_diff(b.len())
};
let out1 = out_f / 2 + 64; let wb1: Vec<u8> = (0..in_f * out1).map(|i| {
let mut b = ((i.wrapping_mul(2246822519) ^ 0x85EBCA6B) >> 7) as u8;
if b & 0x7F == 0x7F { b &= 0xF7; }
b
}).collect();
let wd1 = e.htod_bytes(&wb1)?;
let out2 = 128usize;
let wb2: Vec<u8> = (0..in_f * out2).map(|i| {
let mut b = ((i.wrapping_mul(3266489917) ^ 0xC2B2AE35) >> 5) as u8;
if b & 0x7F == 0x7F { b &= 0xF7; }
b
}).collect();
let wd2 = e.htod_bytes(&wb2)?;
let (s0, s1, s2) = (0.031_25f32, 0.007_812_5f32, 1.0f32); let x: Vec<f32> = (0..in_f).map(|i| pr(i + 179) * 0.1).collect();
let xd = e.htod(&x)?;
let mut r0 = e.qmatvec_mmvq_raw(&wd, &xd, 1, in_f, out_f, qt, row_bytes, false)?;
let mut r1 = e.qmatvec_mmvq_raw(&wd1, &xd, 1, in_f, out1, qt, row_bytes, false)?;
let mut r2 = e.qmatvec_mmvq_raw(&wd2, &xd, 1, in_f, out2, qt, row_bytes, false)?;
e.scale_inplace(&mut r0, s0, out_f)?;
e.scale_inplace(&mut r1, s1, out1)?;
e.scale_inplace(&mut r2, s2, out2)?;
let (a0, a1) = e.qmatvec_e4m3_fused2_raw(&wd, &wd1, &xd, in_f, out_f, out1,
row_bytes, s0, s1)?;
let bad2 = bitbad(&e.dtoh(&r0)?, &e.dtoh(&a0)?)
+ bitbad(&e.dtoh(&r1)?, &e.dtoh(&a1)?);
println!("E4M3-FUSED2 synth [{in_f}x({out_f}+{out1})] m=1: bit-bad={bad2} {}",
if bad2 == 0 { "OK" } else { fails += 1; "FAIL" });
let (c0, c1, c2) = e.qmatvec_e4m3_fused3_raw(&wd, &wd1, &wd2, &xd, in_f, out_f,
out1, out2, row_bytes, s0, s1, s2)?;
let bad3 = bitbad(&e.dtoh(&r0)?, &e.dtoh(&c0)?)
+ bitbad(&e.dtoh(&r1)?, &e.dtoh(&c1)?)
+ bitbad(&e.dtoh(&r2)?, &e.dtoh(&c2)?);
println!("E4M3-FUSED3 synth [{in_f}x({out_f}+{out1}+{out2})] m=1: bit-bad={bad3} {}",
if bad3 == 0 { "OK" } else { fails += 1; "FAIL" });
for mm in 2..=8usize {
let mcols = memra_engine::Engine::batched_mcols(mm);
let xt: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 191) * 0.1).collect();
let xtd = e.htod(&xt)?;
let mut b0 = e.qmatvec_batched_raw(&wd, &xtd, mm, in_f, out_f, qt, row_bytes,
mcols, false)?;
let mut b1 = e.qmatvec_batched_raw(&wd1, &xtd, mm, in_f, out1, qt, row_bytes,
mcols, false)?;
e.scale_inplace(&mut b0, s0, mm * out_f)?;
e.scale_inplace(&mut b1, s1, mm * out1)?;
let (f0, f1) = e.qmatvec_e4m3_fused2_t_raw(&wd, &wd1, &xtd, mm, in_f, out_f,
out1, row_bytes, s0, s1)?;
let badt = bitbad(&e.dtoh(&b0)?, &e.dtoh(&f0)?)
+ bitbad(&e.dtoh(&b1)?, &e.dtoh(&f1)?);
println!("E4M3-FUSED2-T synth [{in_f}x({out_f}+{out1})] m={mm} b{mcols}: bit-bad={badt} {}",
if badt == 0 { "OK" } else { fails += 1; "FAIL" });
if mm <= 4 {
let mut b2 = e.qmatvec_batched_raw(&wd2, &xtd, mm, in_f, out2, qt,
row_bytes, mcols, false)?;
e.scale_inplace(&mut b2, s2, mm * out2)?;
let (g0, g1, g2) = e.qmatvec_e4m3_fused3_t_raw(&wd, &wd1, &wd2, &xtd, mm,
in_f, out_f, out1, out2,
row_bytes, s0, s1, s2)?;
let bad3t = bitbad(&e.dtoh(&b0)?, &e.dtoh(&g0)?)
+ bitbad(&e.dtoh(&b1)?, &e.dtoh(&g1)?)
+ bitbad(&e.dtoh(&b2)?, &e.dtoh(&g2)?);
println!("E4M3-FUSED3-T synth [{in_f}x({out_f}+{out1}+{out2})] m={mm} b{mcols}: bit-bad={bad3t} {}",
if bad3t == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
}
{
let e4m3_hw = |b: u8| -> f32 {
let s = if b & 0x80 != 0 { -1.0f32 } else { 1.0 };
let ex = ((b >> 3) & 0x0F) as i32;
let mn = (b & 0x07) as f32;
if ex == 0 { s * mn * (2f32).powi(-9) }
else if ex == 15 && mn == 7.0 { f32::NAN }
else { s * (1.0 + mn / 8.0) * (2f32).powi(ex - 7) }
};
const INT_CODES: [u8; 9] = [0x00, 0x38, 0xB8, 0x40, 0xC0, 0x44, 0xC4, 0x48, 0xC8];
let blk_ref = |wb: &[u8], aq: &[i8], ad: &[f32], sc: &[f32],
in_f: usize, out_f: usize, m: usize, scols: usize| -> Vec<f32> {
let nblk = in_f / 32;
let mut y = vec![0f32; m * out_f];
for t in 0..m {
for o in 0..out_f {
let srow = &sc[(o >> 7) * scols..(o >> 7) * scols + scols];
let mut acc = 0f32;
for blk in 0..nblk {
let mut bs = 0f32;
for j in 0..32 {
bs = e4m3_hw(wb[o * in_f + blk * 32 + j])
.mul_add(aq[t * in_f + blk * 32 + j] as f32, bs);
}
acc = (srow[blk >> 2] * ad[t * nblk + blk]).mul_add(bs, acc);
}
y[t * out_f + o] = acc;
}
}
y
};
let shapes: [(usize, usize); 6] = [
(512, 128), (5120, 512), (5120, 320), (2080, 256), (1184, 200), (5120, 1536), ];
let mut codes_seen = [false; 256];
for (in_f, out_f) in shapes {
let srows = out_f.div_ceil(128);
let scols = in_f.div_ceil(128);
for exact in [true, false] {
let arm = if exact { "EXACT" } else { "RAND" };
let wb: Vec<u8> = (0..in_f * out_f).map(|i| {
let h = (i.wrapping_mul(2654435761) ^ 0x9E3779B9) >> 9;
if exact { INT_CODES[h % INT_CODES.len()] }
else {
let c = h as u8;
if c & 0x7F == 0x7F { c & 0xBF } else { c }
}
}).collect();
for b in &wb { codes_seen[*b as usize] = true; }
let sc: Vec<f32> = (0..srows * scols).map(|i| {
let h = (i.wrapping_mul(2246822519) ^ 0x85EBCA6B) >> 7;
if exact { (2f32).powi(((h % 7) as i32) - 3) }
else { 0.002 + 0.5 * (pr(i + 977) * 0.5 + 0.5) }
}).collect();
let wd = e.htod_bytes(&wb)?;
let scd = e.htod(&sc)?;
let nan = e.fp8_blk_nan_count(&wd)?;
for mm in [1usize, 2, 5, 9] {
let mut x: Vec<f32> = if exact {
(0..mm * in_f).map(|i| {
let h = (i.wrapping_mul(3266489917) ^ 0xC2B2AE35) >> 11;
((h % 9) as f32) - 4.0 }).collect()
} else {
(0..mm * in_f).map(|i| pr(i + 211) * 0.1).collect()
};
if exact {
for t in 0..mm {
for b in 0..(in_f / 32) {
let s = if (b + t) & 1 == 0 { 127.0 } else { -127.0 };
x[t * in_f + b * 32 + (b * 11 + t) % 32] = s;
}
}
}
let xd = e.htod(&x)?;
let (aqd, add) = e.quantize_q8_1(&xd, mm, in_f)?;
let aq: Vec<i8> = e.stream().clone_dtoh(&aqd)?; e.stream().synchronize()?;
let ad = e.dtoh(&add)?;
let got = e.dtoh(&e.qmatvec_e4m3_blk_mmvq(&wd, &aqd, &add, &scd, mm, in_f,
out_f, in_f, scols)?)?;
let want = blk_ref(&wb, &aq, &ad, &sc, in_f, out_f, mm, scols);
let q8_lossless = !exact || (0..mm * in_f).all(|i| {
ad[i / 32] == 1.0 && aq[i] as f32 == x[i]
});
let bits_bad = got.iter().zip(&want)
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
let max_abs = maxdiff(&want, &got);
let rms = (want.iter().map(|v| (*v as f64) * (*v as f64)).sum::<f64>()
/ want.len().max(1) as f64).sqrt() as f32;
let rms_rel = if rms > 0.0 { max_abs / rms } else { max_abs };
let mut m1_bits = true;
if mm > 1 {
for t in 0..mm {
let xtd = e.htod(&x[t * in_f..(t + 1) * in_f])?;
let y1 = e.dtoh(&e.qmatvec_e4m3_blk_mmvq_raw(&wd, &xtd, &scd, 1, in_f,
out_f, in_f, scols)?)?;
m1_bits &= y1.iter().zip(&got[t * out_f..(t + 1) * out_f])
.all(|(a, b)| a.to_bits() == b.to_bits());
}
}
let ok = nan == 0 && q8_lossless && m1_bits
&& if exact { bits_bad == 0 } else { rms_rel < 1e-5 };
println!("E4M3-BLK-MMVQ {arm} [{in_f}x{out_f}] s{srows}x{scols} m={mm}: \
rms_rel={rms_rel:.2e} bit-bad={bits_bad}/{} m1-bits={m1_bits} \
q8-lossless={q8_lossless} nan={nan} {}",
want.len(), if ok { "OK" } else { fails += 1; "FAIL" });
}
}
}
let codes = codes_seen.iter().filter(|s| **s).count();
let codes_ok = codes >= 254;
println!("E4M3-BLK-MMVQ code coverage: {codes}/254 legal e4m3 codes exercised \
(0x7F/0xFF excluded — refused by the residency NaN precondition) {}",
if codes_ok { "OK" } else { fails += 1; "FAIL" });
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType};
let g = GgufFile::open(&path)?;
let want: [(GgmlType, i32); 4] = [
(GgmlType::Q8_0, memra_engine::QT_Q8_0), (GgmlType::Q4_K, memra_engine::QT_Q4_K),
(GgmlType::Q5_K, memra_engine::QT_Q5_K), (GgmlType::Q6_K, memra_engine::QT_Q6_K),
];
for (gtype, gt) in want {
let t = match g.tensors.iter().find(|t| t.ggml_type == gtype && t.ne.len() == 2
&& t.ne[0] % 256 == 0 && t.ne[1] >= 4) {
Some(t) => t, None => continue,
};
let tname = t.name.clone();
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for (mm, mcols) in [(2usize, 2usize), (3, 4), (4, 4), (5, 8), (6, 8), (8, 8)] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 131) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, gt, row_bytes, false)?)?;
let ybat = e.dtoh(&e.qmatvec_batched_raw(&wd, &xd, mm, in_f, out_f, gt, row_bytes, mcols, false)?)?;
let d = maxdiff(&yref, &ybat);
let scale = yref.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("BATCHED {tname} [{:?}] m={mm} mcols={mcols}: rel={rel:.2e} {}", t.ggml_type,
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
}
}
if let Some(path) = gguf_arg.clone() {
use memra_gguf::{GgufFile, GgmlType};
let g = GgufFile::open(&path)?;
let want: [(GgmlType, i32); 2] = [
(GgmlType::Q4_K, memra_engine::QT_Q4_K), (GgmlType::Q6_K, memra_engine::QT_Q6_K),
];
for (gtype, gt) in want {
let t = match g.tensors.iter().find(|t| t.ggml_type == gtype && t.ne.len() == 2
&& t.ne[0] % 256 == 0 && t.ne[1] >= 4) {
Some(t) => t, None => continue,
};
let tname = t.name.clone();
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
let mir = e.build_kq_rp4_raw(&wd, in_f, out_f, gt)?;
{
let x: Vec<f32> = (0..in_f).map(|i| pr(i + 151) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, 1, in_f, out_f, gt, row_bytes, false)?)?;
let yrp = e.dtoh(&e.qmatvec_mmvq_raw(&mir, &xd, 1, in_f, out_f, gt, row_bytes, true)?)?;
let bad = yref.iter().zip(&yrp).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("KQRP {tname} [{:?}] m=1 mmvq_rp: bit-bad={bad} {}", t.ggml_type,
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
let tiers: &[(usize, usize)] = if gt == memra_engine::QT_Q6_K {
&[(2, 2), (3, 4), (4, 4), (5, 8), (8, 8), (12, 16)]
} else {
&[(2, 2), (3, 4), (4, 4), (5, 8), (8, 8)]
};
for &(mm, mcols) in tiers {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 161) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_batched_raw(&wd, &xd, mm, in_f, out_f, gt, row_bytes, mcols, false)?)?;
let yrp = e.dtoh(&e.qmatvec_batched_raw(&mir, &xd, mm, in_f, out_f, gt, row_bytes, mcols, true)?)?;
let bad = yref.iter().zip(&yrp).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("KQRP {tname} [{:?}] m={mm} mcols={mcols} batched_rp: bit-bad={bad} {}", t.ggml_type,
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
let gguf_9b = kc_model("nvfp4-batched", "Qwen3.5-9B-NVFP4-MTP-GGUF.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf"],
&gguf_arg);
if let Some(gguf_9b) = gguf_9b {
let g = GgufFile::open(&gguf_9b)?;
if let Some(t) = g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4) {
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let wd = e.htod_bytes(raw)?;
for (mm, mcols) in [(2usize, 2usize), (3, 4), (4, 4), (5, 8), (6, 8), (8, 8)] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 141) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, false)?)?;
let ybat = e.dtoh(&e.qmatvec_batched_raw(&wd, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, mcols, false)?)?;
let d = maxdiff(&yref, &ybat);
let scale = yref.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
println!("BATCHED blk.0.ffn_gate.weight [NVFP4] m={mm} mcols={mcols}: rel={rel:.2e} {}",
if rel < 1e-3 { "OK" } else { fails += 1; "FAIL" });
}
}
if let (Some(tg), Some(tu)) = (
g.find("blk.0.ffn_gate.weight").filter(|t| t.ggml_type == GgmlType::NVFP4),
g.find("blk.0.ffn_up.weight").filter(|t| t.ggml_type == GgmlType::NVFP4),
) {
use memra_engine::model::repack_nvfp4_split;
let in_f = tg.ne[0] as usize; let out_f = tg.ne[1] as usize;
let raw_g = g.tensor_data(tg); let row_bytes = raw_g.len() / out_f;
let raw_u = g.tensor_data(tu);
let wg = e.htod_bytes(raw_g)?;
let wu = e.htod_bytes(raw_u)?;
let wg_rp = e.htod_bytes(&repack_nvfp4_split(raw_g, out_f))?;
let wu_rp = e.htod_bytes(&repack_nvfp4_split(raw_u, out_f))?;
for (mm, mcols) in [(2usize, 2usize), (3, 4), (4, 4), (5, 8), (6, 8), (7, 8)] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 151) * 0.1).collect();
let xd = e.htod(&x)?;
let (aq, ad) = e.quantize_q8_1(&xd, mm, in_f)?;
for (rp, w0, w1) in [(false, &wg, &wu), (true, &wg_rp, &wu_rp)] {
if mm > 4 && !rp { continue; }
let y0ref = e.dtoh(&e.qmatvec_mmvq_batched(w0, &aq, &ad, mm, in_f, out_f,
memra_engine::QT_NVFP4, row_bytes, mcols, 1.0, rp)?)?;
let y1ref = e.dtoh(&e.qmatvec_mmvq_batched(w1, &aq, &ad, mm, in_f, out_f,
memra_engine::QT_NVFP4, row_bytes, mcols, 1.0, rp)?)?;
let (y0d, y1d) = e.qmatvec_batched_dual_raw(w0, w1, &aq, &ad, mm, in_f,
out_f, row_bytes, rp)?;
let (y0, y1) = (e.dtoh(&y0d)?, e.dtoh(&y1d)?);
let bad0 = y0ref.iter().zip(&y0).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
let bad1 = y1ref.iter().zip(&y1).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
println!("DUAL-BATCHED gate+up [NVFP4{}] m={mm} mcols={mcols}: bit-bad={}/{} {}",
if rp { " rp" } else { "" }, bad0, bad1,
if bad0 == 0 && bad1 == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
use memra_engine::model::{repack_nvfp4_split, unpack_nvfp4_split};
let path9 = kc_model("a6-split-plane(9b-fallback)", "Qwen3.5-9B-NVFP4-MTP-GGUF.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen35-9b-nvfp4-gguf/Qwen3.5-9B-NVFP4-MTP-GGUF.gguf"],
&gguf_arg);
let srcs: Vec<String> = gguf_arg.clone().into_iter().chain(path9).collect();
let mut done = false;
for path in srcs {
if done { break; }
let g = match GgufFile::open(&path) { Ok(g) => g, Err(_) => continue };
let mut picks: Vec<_> = g.tensors.iter()
.filter(|t| t.ggml_type == GgmlType::NVFP4 && t.ne.len() == 2 && t.ne[0] % 64 == 0)
.take(2).collect();
if let Some(deep) = g.tensors.iter().find(|t| t.ggml_type == GgmlType::NVFP4
&& t.ne.len() == 2 && t.ne[0] % 512 == 0 && t.ne[0] >= 6144) {
if !picks.iter().any(|p| p.name == deep.name) { picks.push(deep); }
}
for t in picks {
done = true;
let tname = t.name.clone();
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize;
let raw = g.tensor_data(t); let row_bytes = raw.len() / out_f;
let rpb = repack_nvfp4_split(raw, out_f);
let rt_bad = unpack_nvfp4_split(&rpb, out_f).iter().zip(raw.iter())
.filter(|(a, b)| a != b).count();
println!("RP roundtrip {tname}: {} mismatched bytes {}", rt_bad,
if rt_bad == 0 { "OK" } else { fails += 1; "FAIL" });
let wd = e.htod_bytes(raw)?;
let wrp = e.htod_bytes(&rpb)?;
let bit_bad = |a: &[f32], b: &[f32]| a.iter().zip(b)
.filter(|(x, y)| x.to_bits() != y.to_bits()).count();
for mm in [1usize, 2] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 151) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, false)?)?;
let yrp = e.dtoh(&e.qmatvec_mmvq_raw(&wrp, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, true)?)?;
let bad = bit_bad(&yref, &yrp);
println!("RP MMVQ {tname} m={mm}: bit-bad={bad} {}",
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
for (mm, mcols) in [(2usize, 2usize), (3, 4), (4, 4), (5, 8), (6, 8), (7, 8), (8, 8)] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 161) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_mmvq_raw(&wd, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, false)?)?;
let yrp = e.dtoh(&e.qmatvec_batched_raw(&wrp, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, mcols, true)?)?;
let v = e.batched_variant(mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, mcols, true);
if v.starts_with("rpks") {
let d = maxdiff(&yref, &yrp);
let scale = yref.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-3);
let rel = d / scale;
let y2 = e.dtoh(&e.qmatvec_batched_raw(&wrp, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes, mcols, true)?)?;
let det = bit_bad(&yrp, &y2);
println!("RP BATCHED {tname} m={mm} mcols={mcols} [{v}]: rel={rel:.2e} det-bad={det} {}",
if rel < 1e-6 && det == 0 { "OK" } else { fails += 1; "FAIL" });
} else {
let bad = bit_bad(&yref, &yrp);
println!("RP BATCHED {tname} m={mm} mcols={mcols} [{v}]: bit-bad={bad} {}",
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
for mm in [1usize, 5] {
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 171) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_nvfp4_fast(&wd, &xd, mm, in_f, out_f, row_bytes)?)?;
let yrp = e.dtoh(&e.qmatvec_nvfp4_fast_rp(&wrp, &xd, mm, in_f, out_f, row_bytes)?)?;
let bad = bit_bad(&yref, &yrp);
println!("RP DP4A {tname} m={mm}: bit-bad={bad} {}",
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let mm = 128usize;
let x: Vec<f32> = (0..mm * in_f).map(|i| pr(i + 181) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec_gemm_raw(&wd, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4, row_bytes)?)?;
let yrp = e.dtoh(&e.qmatvec_gemm_raw(&wrp, &xd, mm, in_f, out_f, memra_engine::QT_NVFP4_RP, row_bytes)?)?;
let bad = bit_bad(&yref, &yrp);
println!("RP GEMM {tname} T={mm}: bit-bad={bad} {}",
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let x: Vec<f32> = (0..in_f).map(|i| pr(i + 191) * 0.1).collect();
let xd = e.htod(&x)?;
let yref = e.dtoh(&e.qmatvec(&wd, &xd, 1, in_f, out_f, memra_engine::QT_NVFP4, row_bytes)?)?;
let yrp = e.dtoh(&e.qmatvec(&wrp, &xd, 1, in_f, out_f, memra_engine::QT_NVFP4_RP, row_bytes)?)?;
let bad = bit_bad(&yref, &yrp);
println!("RP STAGE-A {tname}: bit-bad={bad} {}",
if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
}
}
{
let (hd, nh, nhkv) = (256usize, 16usize, 4usize);
let scale = 1.0 / (hd as f32).sqrt();
let cpu_sdpa = |q: &[f32], k: &[f32], v: &[f32], t: usize, tkv: usize| -> Vec<f32> {
let mut o = vec![0f32; hd * nh * t];
for head in 0..nh {
let kvh = head / (nh / nhkv);
for qt in 0..t {
let q_pos = (tkv - t) + qt;
let qv = &q[(qt * nh + head) * hd..][..hd];
let mut sc = vec![0f32; tkv];
for tk in 0..tkv {
let kv = &k[(tk * nhkv + kvh) * hd..][..hd];
let mut a = 0.0; for d in 0..hd { a += qv[d] * kv[d]; }
a *= scale; if tk > q_pos { a = -1e30; } sc[tk] = a;
}
let mx = sc.iter().cloned().fold(-1e30f32, f32::max);
let mut sum = 0.0; for s in sc.iter_mut() { *s = (*s - mx).exp(); sum += *s; }
for s in sc.iter_mut() { *s /= sum; }
let ov = &mut o[(qt * nh + head) * hd..][..hd];
for d in 0..hd { let mut a = 0.0; for tk in 0..tkv { a += sc[tk] * v[(tk*nhkv+kvh)*hd+d]; } ov[d] = a; }
}
}
o
};
for (t, tkv) in [(16usize, 16usize), (64, 64), (100, 100), (256, 256)] {
let q: Vec<f32> = (0..hd*nh*t).map(|i| pr(i)*0.2).collect();
let k: Vec<f32> = (0..hd*nhkv*tkv).map(|i| pr(i+7)*0.2).collect();
let v: Vec<f32> = (0..hd*nhkv*tkv).map(|i| pr(i+11)*0.2).collect();
let cpu = cpu_sdpa(&q,&k,&v,t,tkv);
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?; let mut od=e.zeros(hd*nh*t)?;
e.fa_prefill(&qd,&kd,&vd,&mut od,hd,nh,nhkv,t,tkv,scale,true)?;
let g=e.dtoh(&od)?; let d=maxdiff(&cpu,&g);
let sc=cpu.iter().map(|v|v.abs()).fold(0.0,f32::max).max(1e-3); let rel=d/sc;
println!("fa_prefill T={t} Tkv={tkv}: rel={rel:.2e} {}", if rel<2e-2 {"OK"} else {fails+=1;"FAIL"});
}
{
let (hdw, nhw, nkvw, wnd) = (256usize, 4usize, 1usize, 32usize);
let scalew = 1.0f32 / (hdw as f32).sqrt();
let cpu_sdpa_w = |q: &[f32], k: &[f32], v: &[f32], t: usize, tkv: usize| -> Vec<f32> {
let mut o = vec![0.0f32; t * nhw * hdw];
for head in 0..nhw { for qt in 0..t {
let q_pos = (tkv - t) + qt;
let qv = &q[(qt * nhw + head) * hdw..][..hdw];
let mut sc = vec![0.0f32; tkv];
for (tk, s) in sc.iter_mut().enumerate() {
let kv = &k[tk * hdw..][..hdw];
let mut a = 0.0; for d in 0..hdw { a += qv[d] * kv[d]; }
a *= scalew;
if tk > q_pos || (q_pos >= wnd && tk < q_pos - (wnd - 1)) { a = -1e30; }
*s = a;
}
let mx = sc.iter().cloned().fold(-1e30f32, f32::max);
let mut sum = 0.0; for s in sc.iter_mut() { *s = (*s - mx).exp(); sum += *s; }
for s in sc.iter_mut() { *s /= sum; }
let ov = &mut o[(qt * nhw + head) * hdw..][..hdw];
for d in 0..hdw {
let mut a = 0.0; for tk in 0..tkv { a += sc[tk] * v[tk * hdw + d]; }
ov[d] = a;
}
} }
o
};
for (t, tkv) in [(64usize, 64usize), (100, 100)] {
let q: Vec<f32> = (0..hdw*nhw*t).map(|i| pr(i+47)*0.2).collect();
let k: Vec<f32> = (0..hdw*nkvw*tkv).map(|i| pr(i+53)*0.2).collect();
let v: Vec<f32> = (0..hdw*nkvw*tkv).map(|i| pr(i+61)*0.2).collect();
let cpu = cpu_sdpa_w(&q, &k, &v, t, tkv);
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?;
let mut o_f32=e.zeros(hdw*nhw*t)?; let mut o_bf=e.zeros(hdw*nhw*t)?;
e.fa_prefill_w_arm(&qd,&kd,&vd,&mut o_f32,hdw,nhw,nkvw,t,tkv,scalew,true,wnd,true,false)?;
e.fa_prefill_w_arm(&qd,&kd,&vd,&mut o_bf,hdw,nhw,nkvw,t,tkv,scalew,true,wnd,false,false)?;
let gf=e.dtoh(&o_f32)?; let gb=e.dtoh(&o_bf)?;
let d=maxdiff(&cpu,&gf);
let sc=cpu.iter().map(|x|x.abs()).fold(0.0,f32::max).max(1e-3); let rel=d/sc;
println!("fa_prefill_w T={t} Tkv={tkv} w={wnd}: rel={rel:.2e} {}",
if rel<2e-2 {"OK"} else {fails+=1;"FAIL"});
let hp_door = memra_engine::fa_f16pv_on() && memra_engine::faw_hp_on();
if hp_door {
println!("fa_prefill_w bf16-stage T={t}: SKIPPED (hp door numeric class)");
} else {
let nbad = gf.iter().zip(gb.iter()).filter(|(a,b)| a.to_bits()!=b.to_bits()).count();
println!("fa_prefill_w bf16-stage T={t}: bit-mismatch {nbad}/{} {}",
gf.len(), if nbad==0 {"OK"} else {fails+=1;"FAIL"});
}
}
}
{
let (hd5, nh5, nhkv5) = (512usize, 8usize, 1usize);
let scale5 = 1.0f32 / (hd5 as f32).sqrt();
let cpu_sdpa5 = |q: &[f32], k: &[f32], v: &[f32], t: usize, tkv: usize| -> Vec<f32> {
let mut o = vec![0.0f32; t * nh5 * hd5];
for head in 0..nh5 { for qt in 0..t {
let q_pos = (tkv - t) + qt;
let qv = &q[(qt * nh5 + head) * hd5..][..hd5];
let mut sc = vec![0.0f32; tkv];
for (tk, s) in sc.iter_mut().enumerate() {
let kv = &k[tk * hd5..][..hd5];
let mut a = 0.0; for d in 0..hd5 { a += qv[d] * kv[d]; }
a *= scale5; if tk > q_pos { a = -1e30; } *s = a;
}
let mx = sc.iter().cloned().fold(-1e30f32, f32::max);
let mut sum = 0.0; for s in sc.iter_mut() { *s = (*s - mx).exp(); sum += *s; }
for s in sc.iter_mut() { *s /= sum; }
let ov = &mut o[(qt * nh5 + head) * hd5..][..hd5];
for d in 0..hd5 {
let mut a = 0.0; for tk in 0..tkv { a += sc[tk] * v[tk * hd5 + d]; }
ov[d] = a;
}
} }
o
};
for (t, tkv) in [(64usize, 64usize), (100, 100)] {
let q: Vec<f32> = (0..hd5*nh5*t).map(|i| pr(i+13)*0.2).collect();
let k: Vec<f32> = (0..hd5*nhkv5*tkv).map(|i| pr(i+17)*0.2).collect();
let v: Vec<f32> = (0..hd5*nhkv5*tkv).map(|i| pr(i+23)*0.2).collect();
let cpu = cpu_sdpa5(&q, &k, &v, t, tkv);
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?;
let mut o_f32=e.zeros(hd5*nh5*t)?; let mut o_bf=e.zeros(hd5*nh5*t)?;
let mut o_sp=e.zeros(hd5*nh5*t)?; let mut o_sp16=e.zeros(hd5*nh5*t)?;
e.fa_prefill_hd512_arm(&qd,&kd,&vd,&mut o_f32,hd5,nh5,nhkv5,t,tkv,scale5,true,true,false,false)?;
e.fa_prefill_hd512_arm(&qd,&kd,&vd,&mut o_bf,hd5,nh5,nhkv5,t,tkv,scale5,true,false,false,false)?;
e.fa_prefill_hd512_arm(&qd,&kd,&vd,&mut o_sp,hd5,nh5,nhkv5,t,tkv,scale5,true,false,true,false)?;
e.fa_prefill_hd512_arm(&qd,&kd,&vd,&mut o_sp16,hd5,nh5,nhkv5,t,tkv,scale5,true,false,true,true)?;
let gf=e.dtoh(&o_f32)?; let gb=e.dtoh(&o_bf)?; let gs=e.dtoh(&o_sp)?;
let gs16=e.dtoh(&o_sp16)?;
let d=maxdiff(&cpu,&gf);
let sc=cpu.iter().map(|x|x.abs()).fold(0.0,f32::max).max(1e-3); let rel=d/sc;
println!("fa_prefill_hd512 T={t} Tkv={tkv}: rel={rel:.2e} {}",
if rel<2e-2 {"OK"} else {fails+=1;"FAIL"});
let nbad = gf.iter().zip(gb.iter()).filter(|(a,b)| a.to_bits()!=b.to_bits()).count();
println!("fa_prefill_hd512 bf16-stage T={t}: bit-mismatch {nbad}/{} {}",
gf.len(), if nbad==0 {"OK"} else {fails+=1;"FAIL"});
let dsp=maxdiff(&cpu,&gs); let relsp=dsp/sc;
println!("fa_prefill_hd512_sp T={t} Tkv={tkv}: rel={relsp:.2e} {}",
if relsp<2e-2 {"OK"} else {fails+=1;"FAIL"});
let d16=maxdiff(&cpu,&gs16); let rel16=d16/sc;
println!("fa_prefill_hd512_sp16 T={t} Tkv={tkv}: rel={rel16:.2e} {}",
if rel16<2e-2 {"OK"} else {fails+=1;"FAIL"});
}
}
{
let (hd6, nh6, nhkv6) = (512usize, 8usize, 4usize);
let scale6 = 1.0f32 / (hd6 as f32).sqrt();
let grp = nh6 / nhkv6;
let cpu_sdpa6 = |q: &[f32], k: &[f32], v: &[f32], t: usize, tkv: usize| -> Vec<f32> {
let mut o = vec![0.0f32; t * nh6 * hd6];
for head in 0..nh6 { for qt in 0..t {
let kvh = head / grp;
let q_pos = (tkv - t) + qt;
let qv = &q[(qt * nh6 + head) * hd6..][..hd6];
let mut sc = vec![0.0f32; tkv];
for (tk, sv) in sc.iter_mut().enumerate() {
let kv = &k[(tk * nhkv6 + kvh) * hd6..][..hd6];
let mut a = 0.0; for d in 0..hd6 { a += qv[d] * kv[d]; }
a *= scale6; if tk > q_pos { a = -1e30; } *sv = a;
}
let mx = sc.iter().cloned().fold(-1e30f32, f32::max);
let mut sum = 0.0; for sv in sc.iter_mut() { *sv = (*sv - mx).exp(); sum += *sv; }
for sv in sc.iter_mut() { *sv /= sum; }
let ov = &mut o[(qt * nh6 + head) * hd6..][..hd6];
for d in 0..hd6 {
let mut a = 0.0;
for tk in 0..tkv { a += sc[tk] * v[(tk * nhkv6 + kvh) * hd6 + d]; }
ov[d] = a;
}
} }
o
};
for (t, tkv) in [(64usize, 64usize), (100, 100)] {
let q: Vec<f32> = (0..hd6*nh6*t).map(|i| pr(i+29)*0.2).collect();
let k: Vec<f32> = (0..hd6*nhkv6*tkv).map(|i| pr(i+31)*0.2).collect();
let v: Vec<f32> = (0..hd6*nhkv6*tkv).map(|i| pr(i+37)*0.2).collect();
let cpu = cpu_sdpa6(&q, &k, &v, t, tkv);
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?;
let mut o_hp=e.zeros(hd6*nh6*t)?;
e.fa_prefill_hd512_arm(&qd,&kd,&vd,&mut o_hp,hd6,nh6,nhkv6,t,tkv,scale6,true,false,true,true)?;
let gh=e.dtoh(&o_hp)?;
let d=maxdiff(&cpu,&gh);
let sc=cpu.iter().map(|x|x.abs()).fold(0.0,f32::max).max(1e-3); let rel=d/sc;
println!("fa_prefill_hd512 GQA nkv=4 (hp arm) T={t} Tkv={tkv}: rel={rel:.2e} {}",
if rel<2e-2 {"OK"} else {fails+=1;"FAIL"});
}
}
let kv_dim_k = hd * nhkv; let kv_dim_v = hd * nhkv;
let (kbb, vbb) = memra_engine::kv_blk_bytes(); let k_tok_bytes = (kv_dim_k / 32) * kbb;
let v_tok_bytes = (kv_dim_v / 32) * vbb;
let kvq_tol: f32 = 6e-2 * match memra_engine::kv_cache_formats().1 {
"q4_0" => 5.0, "fp8" => 2.5, _ => 1.0,
};
for tkv in [64usize, 128, 257] {
let q: Vec<f32> = (0..hd*nh).map(|i| pr(i+1)*0.2).collect();
let k: Vec<f32> = (0..hd*nhkv*tkv).map(|i| pr(i+7)*0.2).collect();
let v: Vec<f32> = (0..hd*nhkv*tkv).map(|i| pr(i+11)*0.2).collect();
let cpu = cpu_sdpa(&q,&k,&v,1,tkv);
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?;
let mut kc = e.alloc_u8(tkv * k_tok_bytes)?;
let mut vc = e.alloc_u8(tkv * v_tok_bytes)?;
for tok in 0..tkv {
let k_row = kd.slice(tok*kv_dim_k..(tok+1)*kv_dim_k);
let v_row = vd.slice(tok*kv_dim_v..(tok+1)*kv_dim_v);
e.append_kv_quantized_view(&k_row,&v_row,&mut kc,&mut vc,tok,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes, false)?;
}
let kview=e.view_u8(&kc, tkv*k_tok_bytes); let vview=e.view_u8(&vc, tkv*v_tok_bytes);
let sc=cpu.iter().map(|v|v.abs()).fold(0.0,f32::max).max(1e-3);
unsafe { std::env::set_var("MEMRA_NO_FA_VEC", "1"); }
let mut od=e.zeros(hd*nh)?;
e.fa_decode(&qd,&kview,&vview,&mut od,hd,nh,nhkv,tkv,scale,k_tok_bytes,v_tok_bytes)?;
let rel = maxdiff(&cpu,&e.dtoh(&od)?)/sc;
unsafe { std::env::remove_var("MEMRA_NO_FA_VEC"); }
let mut od_v=e.zeros(hd*nh)?;
e.fa_decode(&qd,&kview,&vview,&mut od_v,hd,nh,nhkv,tkv,scale,k_tok_bytes,v_tok_bytes)?;
let rel_v = maxdiff(&cpu,&e.dtoh(&od_v)?)/sc;
println!("fa_decode(KVQ) Tkv={tkv}: rel={rel:.2e} {}", if rel<kvq_tol {"OK"} else {fails+=1;"FAIL"});
let regress = rel_v > rel + 2.5e-3;
println!("fa_decode_vec_q(KVQ) Tkv={tkv}: rel={rel_v:.2e} (scalar {rel:.2e}) {}",
if rel_v<kvq_tol && !regress {"OK"} else {fails+=1;"FAIL"});
if memra_engine::kv_cache_formats() == ("q8_0", "q5_1") && hd % 128 == 0 {
unsafe { std::env::set_var("MEMRA_FA_V3", "1"); }
let mut od_3=e.zeros(hd*nh)?;
e.fa_decode(&qd,&kview,&vview,&mut od_3,hd,nh,nhkv,tkv,scale,k_tok_bytes,v_tok_bytes)?;
unsafe { std::env::remove_var("MEMRA_FA_V3"); }
let rel_3 = maxdiff(&cpu,&e.dtoh(&od_3)?)/sc;
let regress3 = rel_3 > rel + 1e-2;
println!("fa_decode_vec_q_v3(KVQ) Tkv={tkv}: rel={rel_3:.2e} (scalar {rel:.2e}) {}",
if rel_3<kvq_tol && !regress3 {"OK"} else {fails+=1;"FAIL"});
}
}
for (base_len, t) in [(95usize, 5usize), (127, 4), (256, 3), (1000, 5)] {
let tkv_max = base_len + t;
let q: Vec<f32> = (0..hd*nh*t).map(|i| pr(i+3)*0.2).collect();
let k: Vec<f32> = (0..hd*nhkv*tkv_max).map(|i| pr(i+7)*0.2).collect();
let v: Vec<f32> = (0..hd*nhkv*tkv_max).map(|i| pr(i+11)*0.2).collect();
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?;
let mut kc = e.alloc_u8(tkv_max * k_tok_bytes)?;
let mut vc = e.alloc_u8(tkv_max * v_tok_bytes)?;
for tok in 0..tkv_max {
let k_row = kd.slice(tok*kv_dim_k..(tok+1)*kv_dim_k);
let v_row = vd.slice(tok*kv_dim_v..(tok+1)*kv_dim_v);
e.append_kv_quantized_view(&k_row,&v_row,&mut kc,&mut vc,tok,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes, false)?;
}
let mut o_loop = e.zeros(hd*nh*t)?;
for r in 0..t {
let t_kv_r = base_len + r + 1;
let kview=e.view_u8(&kc, t_kv_r*k_tok_bytes);
let vview=e.view_u8(&vc, t_kv_r*v_tok_bytes);
let mut q_row = e.zeros(hd*nh)?;
let q_src = qd.slice(r*nh*hd..(r+1)*nh*hd);
e.copy_view_into(&mut q_row, 0, &q_src, nh*hd)?;
let mut o_row = e.zeros(hd*nh)?;
e.fa_decode(&q_row,&kview,&vview,&mut o_row,hd,nh,nhkv,t_kv_r,scale,k_tok_bytes,v_tok_bytes)?;
e.copy_into(&mut o_loop, r*nh*hd, &o_row, nh*hd)?;
}
let kview=e.view_u8(&kc, tkv_max*k_tok_bytes);
let vview=e.view_u8(&vc, tkv_max*v_tok_bytes);
let mut o_rows = e.zeros(hd*nh*t)?;
e.fa_decode_rows(&qd,&kview,&vview,&mut o_rows,hd,nh,nhkv,base_len,t,scale,k_tok_bytes,v_tok_bytes,None,false, false, None)?;
let a = e.dtoh(&o_loop)?; let b = e.dtoh(&o_rows)?;
let bitdiff = a.iter().zip(&b).filter(|(x,y)| x.to_bits() != y.to_bits()).count();
println!("fa_decode_rows vs per-row loop base={base_len} T={t}: bitdiff={bitdiff} {}",
if bitdiff == 0 {"OK"} else {fails+=1;"FAIL"});
if memra_engine::kv_cache_formats() == ("q8_0", "q5_1") && hd % 128 == 0 {
unsafe { std::env::set_var("MEMRA_FA_V3", "1"); }
let mut o_loop3 = e.zeros(hd*nh*t)?;
for r in 0..t {
let t_kv_r = base_len + r + 1;
let kview=e.view_u8(&kc, t_kv_r*k_tok_bytes);
let vview=e.view_u8(&vc, t_kv_r*v_tok_bytes);
let mut q_row = e.zeros(hd*nh)?;
let q_src = qd.slice(r*nh*hd..(r+1)*nh*hd);
e.copy_view_into(&mut q_row, 0, &q_src, nh*hd)?;
let mut o_row = e.zeros(hd*nh)?;
e.fa_decode(&q_row,&kview,&vview,&mut o_row,hd,nh,nhkv,t_kv_r,scale,k_tok_bytes,v_tok_bytes)?;
e.copy_into(&mut o_loop3, r*nh*hd, &o_row, nh*hd)?;
}
let kview=e.view_u8(&kc, tkv_max*k_tok_bytes);
let vview=e.view_u8(&vc, tkv_max*v_tok_bytes);
let mut o_rows3 = e.zeros(hd*nh*t)?;
e.fa_decode_rows(&qd,&kview,&vview,&mut o_rows3,hd,nh,nhkv,base_len,t,scale,k_tok_bytes,v_tok_bytes,None,false, false, None)?;
unsafe { std::env::remove_var("MEMRA_FA_V3"); }
let a3 = e.dtoh(&o_loop3)?; let b3 = e.dtoh(&o_rows3)?;
let bd3 = a3.iter().zip(&b3).filter(|(x,y)| x.to_bits() != y.to_bits()).count();
println!("fa_decode_rows_v3 vs per-row loop (FA_V3) base={base_len} T={t}: bitdiff={bd3} {}",
if bd3 == 0 {"OK"} else {fails+=1;"FAIL"});
}
}
{
use cudarc::driver::DevicePtr;
for depths in [vec![96usize, 128, 257, 511], vec![200; 8]] {
let b_n = depths.len();
let sp0 = memra_engine::fa_split_keys_pub(depths[0], nhkv);
let eligible = depths.iter().all(|&t| memra_engine::fa_seqs_eligible(t, hd))
&& depths.iter().all(|&t| memra_engine::fa_split_keys_pub(t, nhkv) == sp0);
if !eligible { continue; } let t_kv_max = *depths.iter().max().unwrap();
let kpool: Vec<f32> = (0..kv_dim_k*t_kv_max).map(|i| pr(i+13)*0.2).collect();
let vpool: Vec<f32> = (0..kv_dim_v*t_kv_max).map(|i| pr(i+17)*0.2).collect();
let kpd = e.htod(&kpool)?; let vpd = e.htod(&vpool)?;
let mut kcs: Vec<_> = Vec::new(); let mut vcs: Vec<_> = Vec::new();
let mut kcs2: Vec<_> = Vec::new(); let mut vcs2: Vec<_> = Vec::new();
for &tkv in &depths {
let mut kc = e.alloc_u8(tkv * k_tok_bytes)?;
let mut vc = e.alloc_u8(tkv * v_tok_bytes)?;
for tok in 0..tkv-1 {
let k_row = kpd.slice(tok*kv_dim_k..(tok+1)*kv_dim_k);
let v_row = vpd.slice(tok*kv_dim_v..(tok+1)*kv_dim_v);
e.append_kv_quantized_view(&k_row,&v_row,&mut kc,&mut vc,tok,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes,false)?;
}
let mut kc2 = e.alloc_u8(tkv * k_tok_bytes)?;
let mut vc2 = e.alloc_u8(tkv * v_tok_bytes)?;
for tok in 0..tkv-1 {
let k_row = kpd.slice(tok*kv_dim_k..(tok+1)*kv_dim_k);
let v_row = vpd.slice(tok*kv_dim_v..(tok+1)*kv_dim_v);
e.append_kv_quantized_view(&k_row,&v_row,&mut kc2,&mut vc2,tok,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes,false)?;
}
kcs.push(kc); vcs.push(vc); kcs2.push(kc2); vcs2.push(vc2);
}
let knew: Vec<f32> = (0..kv_dim_k*b_n).map(|i| pr(i+23)*0.2).collect();
let vnew: Vec<f32> = (0..kv_dim_v*b_n).map(|i| pr(i+27)*0.2).collect();
let knd = e.htod(&knew)?; let vnd = e.htod(&vnew)?;
let pos: Vec<i32> = depths.iter().map(|&t| (t - 1) as i32).collect();
let pos_d = e.htod_i32(&pos)?;
for z in 0..b_n {
let k_row = knd.slice(z*kv_dim_k..(z+1)*kv_dim_k);
let v_row = vnd.slice(z*kv_dim_v..(z+1)*kv_dim_v);
e.append_kv_quantized_view(&k_row,&v_row,&mut kcs[z],&mut vcs[z],depths[z]-1,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes,false)?;
}
let es = e.gpu.stream();
let mut ptrs2: Vec<u64> = Vec::new();
for z in 0..b_n {
let (pk, _g) = kcs2[z].device_ptr(&es);
let (pv, _g2) = vcs2[z].device_ptr(&es);
ptrs2.push(pk as u64); ptrs2.push(pv as u64);
}
let table2 = e.htod_u64(&ptrs2)?;
let tv2 = table2.slice(0..2*b_n);
e.append_kv_quantized_seqs(&knd,&vnd,&tv2,&pos_d,b_n,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes)?;
let mut ap_diff = 0usize;
for z in 0..b_n {
let a = e.dtoh_u8(&kcs[z])?; let b = e.dtoh_u8(&kcs2[z])?;
ap_diff += a.iter().zip(&b).filter(|(x,y)| x != y).count();
let a = e.dtoh_u8(&vcs[z])?; let b = e.dtoh_u8(&vcs2[z])?;
ap_diff += a.iter().zip(&b).filter(|(x,y)| x != y).count();
}
println!("append_kv_seqs vs per-seq loop B={b_n}: bytediff={ap_diff} {}",
if ap_diff == 0 {"OK"} else {fails+=1;"FAIL"});
let q: Vec<f32> = (0..hd*nh*b_n).map(|i| pr(i+31)*0.2).collect();
let qd = e.htod(&q)?;
let mut o_loop = e.zeros(hd*nh*b_n)?;
for z in 0..b_n {
let kview = e.view_u8(&kcs[z], depths[z]*k_tok_bytes);
let vview = e.view_u8(&vcs[z], depths[z]*v_tok_bytes);
let mut q_row = e.zeros(hd*nh)?;
let q_src = qd.slice(z*nh*hd..(z+1)*nh*hd);
e.copy_view_into(&mut q_row, 0, &q_src, nh*hd)?;
let mut o_row = e.zeros(hd*nh)?;
e.fa_decode(&q_row,&kview,&vview,&mut o_row,hd,nh,nhkv,depths[z],scale,
k_tok_bytes,v_tok_bytes)?;
e.copy_into(&mut o_loop, z*nh*hd, &o_row, nh*hd)?;
}
let mut ptrs1: Vec<u64> = Vec::new();
for z in 0..b_n {
let (pk, _g) = kcs[z].device_ptr(&es);
let (pv, _g2) = vcs[z].device_ptr(&es);
ptrs1.push(pk as u64); ptrs1.push(pv as u64);
}
let table1 = e.htod_u64(&ptrs1)?;
let tv1 = table1.slice(0..2*b_n);
let mut o_seqs = e.zeros(hd*nh*b_n)?;
e.fa_decode_batch_seqs_v4(&qd,&tv1,&pos_d,&mut o_seqs,hd,nh,nhkv,b_n,
t_kv_max,scale,sp0,k_tok_bytes,v_tok_bytes)?;
let a = e.dtoh(&o_loop)?; let b = e.dtoh(&o_seqs)?;
let bitdiff = a.iter().zip(&b).filter(|(x,y)| x.to_bits() != y.to_bits()).count();
println!("fa_decode_seqs_v4 vs per-seq loop B={b_n} depths={depths:?}: bitdiff={bitdiff} {}",
if bitdiff == 0 {"OK"} else {fails+=1;"FAIL"});
}
}
for (hdd, nhd, nhkvd) in [(256usize, 16usize, 2usize), (256, 32, 8)] {
let kvd = hdd * nhkvd;
let ktb = (kvd / 32) * kbb;
let vtb = (kvd / 32) * vbb;
let t_max = 6272usize;
let kf: Vec<f32> = (0..kvd * t_max).map(|i| pr(i + 7) * 0.2).collect();
let vf: Vec<f32> = (0..kvd * t_max).map(|i| pr(i + 11) * 0.2).collect();
let kfd = e.htod(&kf)?; let vfd = e.htod(&vf)?;
let mut kc = e.alloc_u8(t_max * ktb)?;
let mut vc = e.alloc_u8(t_max * vtb)?;
for tok in 0..t_max {
let k_row = kfd.slice(tok * kvd..(tok + 1) * kvd);
let v_row = vfd.slice(tok * kvd..(tok + 1) * kvd);
e.append_kv_quantized_view(&k_row, &v_row, &mut kc, &mut vc, tok,
kvd, kvd, ktb, vtb, false)?;
}
let q: Vec<f32> = (0..hdd * nhd).map(|i| pr(i + 1) * 0.2).collect();
let qd = e.htod(&q)?;
unsafe { std::env::set_var("MEMRA_FA_DEEP_MIN", "0"); }
for d in [511usize, 512, 513, 3071, 3073, 4096, 6144, 6200] {
let kview = e.view_u8(&kc, d * ktb);
let vview = e.view_u8(&vc, d * vtb);
unsafe { std::env::set_var("MEMRA_FA_DEEP", "0"); }
let mut o_v4 = e.zeros(hdd * nhd)?;
e.fa_decode(&qd, &kview, &vview, &mut o_v4, hdd, nhd, nhkvd, d, scale, ktb, vtb)?;
unsafe { std::env::set_var("MEMRA_FA_DEEP", "1"); }
let mut o_dp = e.zeros(hdd * nhd)?;
e.fa_decode(&qd, &kview, &vview, &mut o_dp, hdd, nhd, nhkvd, d, scale, ktb, vtb)?;
let (a, b) = (e.dtoh(&o_v4)?, e.dtoh(&o_dp)?);
let bd = a.iter().zip(&b).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
println!("fa_decode_v4_deep vs v4 (eager) t_kv={d}: bitdiff={bd} {}",
if bd == 0 { "OK" } else { fails += 1; "FAIL" });
let tdev = e.htod_i32(&[d as i32])?;
for bucket in [d, d + 128] {
unsafe { std::env::set_var("MEMRA_FA_DEEP", "0"); }
let mut o4dc = e.zeros(hdd * nhd)?;
e.fa_decode_dc(&qd, &kview, &vview, &mut o4dc, hdd, nhd, nhkvd, &tdev,
bucket, scale, ktb, vtb, false)?;
unsafe { std::env::set_var("MEMRA_FA_DEEP", "1"); }
let mut odpdc = e.zeros(hdd * nhd)?;
e.fa_decode_dc(&qd, &kview, &vview, &mut odpdc, hdd, nhd, nhkvd, &tdev,
bucket, scale, ktb, vtb, false)?;
let (adc, bdc) = (e.dtoh(&o4dc)?, e.dtoh(&odpdc)?);
let bd2 = adc.iter().zip(&bdc).filter(|(x, y)| x.to_bits() != y.to_bits()).count();
println!("fa_decode_v4_deep vs v4 (dc) t_kv={d} bucket={bucket}: bitdiff={bd2} {}",
if bd2 == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
unsafe { std::env::remove_var("MEMRA_FA_DEEP"); }
unsafe { std::env::remove_var("MEMRA_FA_DEEP_MIN"); }
}
for (t, tkv) in [(64usize, 192usize), (100, 100), (37, 297)] {
let q: Vec<f32> = (0..hd*nh*t).map(|i| pr(i+5)*0.2).collect();
let k: Vec<f32> = (0..hd*nhkv*tkv).map(|i| pr(i+7)*0.2).collect();
let v: Vec<f32> = (0..hd*nhkv*tkv).map(|i| pr(i+11)*0.2).collect();
let qd=e.htod(&q)?; let kd=e.htod(&k)?; let vd=e.htod(&v)?;
let mut kc = e.alloc_u8(tkv * k_tok_bytes)?;
let mut vc = e.alloc_u8(tkv * v_tok_bytes)?;
for tok in 0..tkv {
let k_row = kd.slice(tok*kv_dim_k..(tok+1)*kv_dim_k);
let v_row = vd.slice(tok*kv_dim_v..(tok+1)*kv_dim_v);
e.append_kv_quantized_view(&k_row,&v_row,&mut kc,&mut vc,tok,
kv_dim_k,kv_dim_v,k_tok_bytes,v_tok_bytes, false)?;
}
let kview=e.view_u8(&kc, tkv*k_tok_bytes); let vview=e.view_u8(&vc, tkv*v_tok_bytes);
let mut o_inl = e.zeros(hd*nh*t)?;
e.fa_prefill_view(&qd,&kview,&vview,&mut o_inl,hd,nh,nhkv,t,tkv,scale,true,
k_tok_bytes,v_tok_bytes, false)?;
let mut o_ws = e.zeros(hd*nh*t)?;
e.fa_prefill_view_ws(&qd,&kview,&vview,&mut o_ws,hd,nh,nhkv,t,tkv,scale,true,
k_tok_bytes,v_tok_bytes, false)?;
let a = e.dtoh(&o_inl)?; let b = e.dtoh(&o_ws)?;
let bitdiff = a.iter().zip(&b).filter(|(x,y)| x.to_bits() != y.to_bits()).count();
println!("fa_prefill_view_ws vs inline-dequant T={t} Tkv={tkv}: bitdiff={bitdiff} {}",
if bitdiff == 0 {"OK"} else {fails+=1;"FAIL"});
}
}
{
use memra_gguf::dequant::fp16_to_f32;
let (kfmt, vfmt) = memra_engine::kv_cache_formats();
let (kbb, vbb) = memra_engine::kv_blk_bytes();
let nblk = 4usize; let kv_dim_k = nblk * 32;
let kv_dim_v = nblk * 32;
let k_tok_bytes = (kv_dim_k / 32) * kbb;
let v_tok_bytes = (kv_dim_v / 32) * vbb;
let kin: Vec<f32> = (0..kv_dim_k).map(|i| pr(i + 71) * 1.3).collect();
let mut vin: Vec<f32> = (0..kv_dim_v).map(|i| pr(i + 91) * 0.7 + 0.1).collect();
let step = 0.05f32;
for j in 0..32 { vin[32 + j] = j as f32 * step; }
let kd = e.htod(&kin)?; let vd = e.htod(&vin)?;
let mut kc = e.alloc_u8(k_tok_bytes)?; let mut vc = e.alloc_u8(v_tok_bytes)?;
e.append_kv_quantized(&kd, &vd, &mut kc, &mut vc, 0, kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, false)?;
let kbytes = e.dtoh_u8(&kc)?; let vbytes = e.dtoh_u8(&vc)?;
let f16_to_f32 = |b: &[u8]| -> f32 { fp16_to_f32(u16::from_le_bytes([b[0], b[1]])) };
let e4m3 = |b: u8| -> f32 {
let s = if b & 0x80 != 0 { -1.0f32 } else { 1.0 };
let ex = ((b >> 3) & 0x0F) as i32;
let mn = (b & 0x07) as f32;
if ex == 0 { s * mn * (2f32).powi(-9) } else if ex == 15 && mn == 7.0 { f32::NAN } else { s * (1.0 + mn / 8.0) * (2f32).powi(ex - 7) }
};
let mut k_deq = vec![0f32; kv_dim_k];
for blk in 0..nblk {
let base = blk * kbb;
match kfmt {
"fp8" => for j in 0..32 { k_deq[blk * 32 + j] = e4m3(kbytes[base + j]); },
_ => {
let d = f16_to_f32(&kbytes[base..base + 2]);
for j in 0..32 { k_deq[blk * 32 + j] = d * (kbytes[base + 2 + j] as i8) as f32; }
}
}
}
let kerr = maxdiff(&kin, &k_deq);
let kamax = kin.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-6);
let krel = kerr / kamax;
let ktol = if kfmt == "fp8" { 7e-2 } else { 5e-3 };
println!("kvq {kfmt} K round-trip: rel={krel:.2e} {}", if krel < ktol { "OK" } else { fails += 1; "FAIL" });
let mut v_deq = vec![0f32; kv_dim_v];
for blk in 0..nblk {
let base = blk * vbb;
match vfmt {
"fp8" => for j in 0..32 { v_deq[blk * 32 + j] = e4m3(vbytes[base + j]); },
"q4_0" => {
let d = f16_to_f32(&vbytes[base..base + 2]);
let qs = &vbytes[base + 2..base + 18];
for j in 0..32 {
let q = if j < 16 { (qs[j] & 0x0F) as i32 } else { (qs[j - 16] >> 4) as i32 };
v_deq[blk * 32 + j] = d * (q - 8) as f32;
}
}
_ => {
let d = f16_to_f32(&vbytes[base..base + 2]);
let m = f16_to_f32(&vbytes[base + 2..base + 4]);
let qh = u32::from_le_bytes([vbytes[base + 4], vbytes[base + 5], vbytes[base + 6], vbytes[base + 7]]);
let qs = &vbytes[base + 8..base + 24];
for j in 0..32 {
let lo = if j < 16 { (qs[j] & 0x0F) as i32 } else { (qs[j - 16] >> 4) as i32 };
let hi = (((qh >> j) & 1) << 4) as i32;
v_deq[blk * 32 + j] = d * (lo | hi) as f32 + m;
}
}
}
}
let verr = maxdiff(&vin, &v_deq);
let vamax = vin.iter().map(|v| v.abs()).fold(0.0, f32::max).max(1e-6);
let vrel = verr / vamax;
let vtol = if vfmt == "q5_1" { 3e-2 } else { 7e-2 };
println!("kvq {vfmt} V round-trip: rel={vrel:.2e} {}", if vrel < vtol { "OK" } else { fails += 1; "FAIL" });
if vfmt == "q5_1" {
let bnd_err = (0..32).map(|j| (vin[32 + j] - v_deq[32 + j]).abs()).fold(0.0, f32::max);
let bnd_d = step; println!("kvq q5_1 5th-bit boundary: maxerr={bnd_err:.2e} (d~{bnd_d:.2e}) {}",
if bnd_err < bnd_d { "OK" } else { fails += 1; "FAIL" });
}
}
{
let nblk = 4usize;
let kv_dim_k = nblk * 32;
let kv_dim_v = nblk * 32;
let (kbb, vbb) = memra_engine::kv_blk_bytes();
let k_tok_bytes = (kv_dim_k / 32) * kbb;
let v_tok_bytes = (kv_dim_v / 32) * vbb;
let (t0, t) = (3usize, 7usize);
let cap = t0 + t;
let kin: Vec<f32> = (0..t * kv_dim_k).map(|i| pr(i + 301) * 1.1).collect();
let vin: Vec<f32> = (0..t * kv_dim_v).map(|i| pr(i + 401) * 0.6 - 0.1).collect();
let kd = e.htod(&kin)?; let vd = e.htod(&vin)?;
let mut kc_ref = e.alloc_u8(cap * k_tok_bytes)?; let mut vc_ref = e.alloc_u8(cap * v_tok_bytes)?;
for i in 0..t {
let k_row = kd.slice(i * kv_dim_k..(i + 1) * kv_dim_k);
let v_row = vd.slice(i * kv_dim_v..(i + 1) * kv_dim_v);
e.append_kv_quantized_view(&k_row, &v_row, &mut kc_ref, &mut vc_ref, t0 + i,
kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, false)?;
}
let mut kc_b = e.alloc_u8(cap * k_tok_bytes)?; let mut vc_b = e.alloc_u8(cap * v_tok_bytes)?;
e.append_kv_quantized_rows(&kd, &vd, &mut kc_b, &mut vc_b, t0, t,
kv_dim_k, kv_dim_v, k_tok_bytes, v_tok_bytes, false)?;
let (kr, kb) = (e.dtoh_u8(&kc_ref)?, e.dtoh_u8(&kc_b)?);
let (vr, vb) = (e.dtoh_u8(&vc_ref)?, e.dtoh_u8(&vc_b)?);
let kmis = (t0 * k_tok_bytes..cap * k_tok_bytes).filter(|&i| kr[i] != kb[i]).count();
let vmis = (t0 * v_tok_bytes..cap * v_tok_bytes).filter(|&i| vr[i] != vb[i]).count();
println!("kv append rows-vs-loop bit-identity (T={t}, t0={t0}): k_mismatch={kmis} v_mismatch={vmis} {}",
if kmis == 0 && vmis == 0 { "OK" } else { fails += 1; "FAIL" });
}
{
let (t, n_expert, n_used) = (8usize, 256usize, 8usize);
let mut logits: Vec<f32> = (0..t * n_expert).map(|i| pr(i + 123) * 4.0).collect();
for tok in 0..t { logits[tok * n_expert + 17] = logits[tok * n_expert + 200]; } let host_route = |row: &[f32]| -> (Vec<i32>, Vec<f32>) {
let maxl = row.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut probs = vec![0f32; n_expert];
let mut den = 0f32;
for i in 0..n_expert { let x = (row[i] - maxl).exp(); probs[i] = x; den += x; }
for p in probs.iter_mut() { *p /= den; }
let mut idx: Vec<usize> = (0..n_expert).collect();
idx.sort_by(|&a, &b| probs[b].total_cmp(&probs[a]).then(a.cmp(&b)));
let sel = &idx[..n_used];
let mut w: Vec<f32> = sel.iter().map(|&i| probs[i]).collect();
let mut ws: f32 = w.iter().sum();
ws = ws.max(6.103515625e-5_f32);
for x in w.iter_mut() { *x /= ws; }
(sel.iter().map(|&i| i as i32).collect(), w)
};
let ld = e.htod(&logits)?;
let (sel_d, w_d) = e.moe_router_topk(&ld, t, n_expert, n_used)?;
let sel_g = e.dtoh_i32(&sel_d)?;
let w_g = e.dtoh(&w_d)?;
let mut idx_ok = true;
let mut w_max_rel = 0f32; let mut w_max_ulp = 0i64; for tok in 0..t {
let (sh, wh) = host_route(&logits[tok * n_expert..(tok + 1) * n_expert]);
for j in 0..n_used {
if sel_g[tok * n_used + j] != sh[j] { idx_ok = false; }
let (a, b) = (w_g[tok * n_used + j], wh[j]);
let rel = (a - b).abs() / b.abs().max(1e-12);
if rel > w_max_rel { w_max_rel = rel; }
let ulp = (a.to_bits() as i64 - b.to_bits() as i64).abs();
if ulp > w_max_ulp { w_max_ulp = ulp; }
}
}
println!("moe_router idx-match (incl. tie 17/200): {}", if idx_ok { "OK" } else { fails += 1; "FAIL" });
println!("moe_router weight rel={w_max_rel:.2e} (max {w_max_ulp} ULP, host-exp vs device-expf): {}",
if w_max_rel < 1e-5 { "OK" } else { fails += 1; "FAIL" });
}
{
use memra_gguf::{GgufFile, GgmlType};
use memra_engine::moe_cache::{MoeSlotCache, BlockId, PROJ_GATE};
let gguf_35b = kc_model("d2-cache-bit-identity", "Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
&["/home/avifenesh/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf"],
&gguf_arg);
if let Some(gguf_35b) = gguf_35b {
let g = GgufFile::open(&gguf_35b)?;
let t = g.find("blk.0.ffn_gate_exps.weight").expect("gate_exps");
let in_f = t.ne[0] as usize; let out_f = t.ne[1] as usize; let n_expert = t.ne[2] as usize;
let qt_opt = match t.ggml_type {
GgmlType::IQ3_S => Some(memra_engine::QT_IQ3_S), GgmlType::IQ4_XS => Some(memra_engine::QT_IQ4_XS),
GgmlType::Q6_K => Some(memra_engine::QT_Q6_K), GgmlType::Q8_0 => Some(memra_engine::QT_Q8_0),
other => { println!("D.2 cache: gate_exps {other:?} unhandled — SKIP"); None },
};
if let Some(qt) = qt_opt {
let raw = g.tensor_data(t);
let expert_stride = raw.len() / n_expert;
let row_bytes = raw.len() / (out_f * n_expert);
let ex = 5usize; let host_bytes = &raw[ex * expert_stride..(ex + 1) * expert_stride];
let x: Vec<f32> = (0..in_f).map(|i| pr(i + 999) * 0.1).collect();
let xd = e.htod(&x)?;
let mut scratch = e.alloc_u8(expert_stride)?;
e.stage_expert(host_bytes, &mut scratch, 0)?;
let y_stage = e.dtoh(&e.qmatvec_view(&scratch, 0..expert_stride, &xd.slice(0..in_f), 1,
in_f, out_f, qt, row_bytes)?)?;
let mut cache = MoeSlotCache::new(&e, expert_stride)?;
let id = BlockId::new(0, PROJ_GATE, ex as u16);
let slot = cache.force_admit(id, host_bytes, &e)?;
let y_hit = e.dtoh(&e.qmatvec_view(cache.slot(slot), 0..expert_stride, &xd.slice(0..in_f), 1,
in_f, out_f, qt, row_bytes)?)?;
let _ = cache.dispatch(id, host_bytes, &e)?;
let bitwise = y_stage.iter().zip(&y_hit).all(|(a, b)| a.to_bits() == b.to_bits());
println!("moe cache-HIT bit-identity (stage==cache): {}",
if bitwise { "OK" } else { fails += 1; "FAIL" });
}
}
}
{
use memra_gguf::{GgufFile, GgmlType};
let gguf_q35 = kc_model("fast-router-batch", "Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
&["/data/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf",
"/home/avifenesh/ai-ml/hf-models/qwen36-35b-moe/Qwen3.6-35B-A3B-UD-IQ4_XS.gguf"],
&gguf_arg);
if let Some(p) = gguf_q35 {
let g = GgufFile::open(&p)?;
let tw = g.find("blk.0.ffn_gate_inp.weight").expect("gate_inp");
assert!(matches!(tw.ggml_type, GgmlType::F32), "gate_inp must be F32");
let n_embd = tw.ne[0] as usize;
let n_experts = tw.ne[1] as usize;
let le = |b: &[u8]| f32::from_le_bytes([b[0], b[1], b[2], b[3]]);
let wf: Vec<f32> = g.tensor_data(tw).chunks_exact(4).map(le).collect();
let t_max = 2048usize;
let x: Vec<f32> = (0..t_max * n_embd).map(|i| (pr(i + 7) - 0.5) * 4.0).collect();
let wd = e.htod(&wf)?; let xd = e.htod(&x)?;
let yref = e.dtoh(&e.router_gemv_form(&wd, &xd, n_embd, n_experts, t_max, true, false)?)?;
let ms: [usize; 32] = [1, 2, 3, 4, 5, 7, 8, 9, 15, 16, 17, 31, 32, 33, 63, 64, 65,
75, 127, 128, 129, 255, 256, 257, 511, 512, 513, 1023, 1024,
1025, 2047, 2048];
let (mut r_bits, mut r_minv) = (0usize, 0usize);
for &m in &ms {
let y_p = e.dtoh(&e.router_gemv_form(&wd, &xd, n_embd, n_experts, m, true, false)?)?;
let y_b = e.dtoh(&e.router_gemv_form(&wd, &xd, n_embd, n_experts, m, true, true)?)?;
r_bits += y_p.iter().zip(&y_b).filter(|(a, b)| a.to_bits() != b.to_bits()).count();
r_minv += y_b.iter().zip(&yref[..m * n_experts])
.filter(|(a, b)| a.to_bits() != b.to_bits()).count();
}
println!("router batch-twin bit-identity (real q35 router, {} m-points 1..{t_max}): mism={r_bits} {}",
ms.len(), if r_bits == 0 { "OK" } else { fails += 1; "FAIL" });
println!("router batch-twin m-invariance (rows vs plain m={t_max} prefix): mism={r_minv} {}",
if r_minv == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
{
use memra_engine::moe_cache::{BlockId, DispatchSlot, MoeSlotCache, PROJ_GATE};
let old_slots = std::env::var_os("MEMRA_MOE_SLOTS");
unsafe { std::env::set_var("MEMRA_MOE_SLOTS", "8"); }
let block_len = 4096usize;
let mut cache = MoeSlotCache::new(&e, block_len)?;
let sources: Vec<Vec<u8>> = (0..8).map(|i| vec![0xA0 + i as u8; block_len]).collect();
for (i, src) in sources.iter().enumerate() {
cache.force_admit(BlockId::new(7, PROJ_GATE, i as u16), src, &e)?;
}
let keep = [BlockId::new(7, PROJ_GATE, 0)];
let next_id = BlockId::new(7, PROJ_GATE, 8);
let next = vec![0xF8; block_len];
let queued = cache.prefetch(next_id, &next, &keep, &e)?;
let hidden_while_pending = cache.resident(next_id).is_none();
let DispatchSlot::Resident(next_slot) = cache.dispatch(next_id, &next, &e)?;
let next_got = e.dtoh_u8(cache.slot(next_slot))?[..block_len].to_vec();
let visible_after_wait = cache.resident(next_id) == Some(next_slot);
let _ = cache.dispatch(next_id, &next, &e)?;
let keep_slot = cache.resident(keep[0]);
let keep_got = match keep_slot {
Some(slot) => e.dtoh_u8(cache.slot(slot))?[..block_len].to_vec(),
None => Vec::new(),
};
let counters_ok = cache.hits == 1 && cache.misses == 1
&& cache.staged_bytes == 9 * block_len as u64;
let ok = queued && hidden_while_pending && visible_after_wait
&& next_got == next && keep_got == sources[0] && counters_ok;
if !ok {
eprintln!("[prefetch-check] queued={queued} hidden={hidden_while_pending} \
visible={visible_after_wait} bytes_ok={} keep_ok={} counters: hits={} \
misses={} staged={} (want 1/1/{})",
next_got == next, keep_got == sources[0], cache.hits, cache.misses,
cache.staged_bytes, 9 * block_len);
}
println!("moe async-prefetch ordering + protected victim: {}",
if ok { "OK" } else { fails += 1; "FAIL" });
unsafe {
match old_slots {
Some(v) => std::env::set_var("MEMRA_MOE_SLOTS", v),
None => std::env::remove_var("MEMRA_MOE_SLOTS"),
}
}
}
{
use memra_gguf::nvfp4_repack::{f32_to_q8_0, fp8_e4m3_to_f32};
for &(out_f, in_f) in &[(256usize, 512usize), (136usize, 160usize), (8usize, 32usize),
(5usize, 128usize), (6usize, 160usize)] {
let (rows, cols) = (out_f.div_ceil(128), in_f.div_ceil(128));
let codes: Vec<u8> = (0..out_f * in_f).map(|i| (i % 256) as u8).collect();
let grid: Vec<f32> = (0..rows * cols)
.map(|i| 2f32.powi((i % 10) as i32 - 4) * (1.0 + 0.125 * (i % 3) as f32))
.collect();
let mut cpu: Vec<u8> = Vec::with_capacity(out_f * (in_f / 32) * 34);
for o in 0..out_f {
let row: Vec<f32> = (0..in_f)
.map(|e| fp8_e4m3_to_f32(codes[o * in_f + e]) * grid[(o >> 7) * cols + (e >> 7)])
.collect();
cpu.extend_from_slice(&f32_to_q8_0(&row));
}
let dev = e.fp8_blk_dequant_q8_0(&codes, &grid, out_f, in_f)?;
let gpu = e.dtoh_u8(&dev)?;
let bad = if gpu.len() != cpu.len() {
usize::MAX
} else {
gpu.iter().zip(&cpu).filter(|(a, b)| a != b).count()
};
println!("fp8-blk-gpu Q8_0 bit-parity [{out_f}x{in_f}] bytes={} bad={bad} {}",
cpu.len(), if bad == 0 { "OK" } else { fails += 1; "FAIL" });
}
}
if fails == 0 { println!("\nALL GREEN: kernels match CPU reference."); Ok(()) }
else { Err(format!("{fails} kernel(s) FAILED").into()) }
}