use memra_engine::dsv4_gpu::{Dsv4Gpu, GpuCapture};
use memra_gguf::dsv4_forward::{
FixtureSpec, drift_coeff, expert_arm_native, quant_depth_of, read_npz,
};
use std::collections::BTreeSet;
use std::path::Path;
const U_BF16: f64 = 1.0 / 256.0;
fn class_coeff(name: &str, d_b: f64) -> f64 {
let d_q = if expert_arm_native() {
quant_depth_of(name)
} else {
0.0
};
drift_coeff(d_b, d_q)
}
struct ArrayCheck {
name: String,
shape: Vec<usize>,
max_abs: f64,
max_rel: f64,
threshold: f64,
n_over: usize,
verdict: &'static str,
note: String,
}
enum Policy {
BitExact,
Derived {
coeff: f64,
flip_budget_frac: f64,
flip_rel: f32,
flip_abs_of_amax: f32,
},
Logits {
coeff: f64,
top_ids: Vec<u32>,
native: bool,
},
}
fn depth_of(name: &str) -> u32 {
let core = name.strip_prefix("c160_").unwrap_or(name);
let n: u32 = core
.strip_prefix("layer")
.and_then(|r| r.split('_').next())
.and_then(|x| x.parse().ok())
.unwrap_or(0);
if core.contains("_out") && !core.contains("attn_out") {
2 * (n + 1)
} else if core.contains("attn_out") {
2 * n + 1
} else {
2 * n
}
}
fn policy_for(name: &str, spec: &FixtureSpec) -> Option<Policy> {
let native = expert_arm_native();
match name {
"embed_out" => Some(Policy::BitExact),
"final_logits_last" => Some(Policy::Logits {
coeff: class_coeff(name, 86.0),
top_ids: spec.top20_ids.clone(),
native,
}),
"mtp_logits_last" => Some(Policy::Logits {
coeff: class_coeff(name, 88.0),
top_ids: spec
.mtp_top20_ids
.clone()
.expect("fixture banks mtp_logits_last but json has no mtp_top20"),
native,
}),
n if n.contains("indexer_kv") => Some(Policy::Derived {
coeff: class_coeff(n, depth_of(n) as f64),
flip_budget_frac: if native { 0.75 } else { 0.05 },
flip_rel: 1.0, flip_abs_of_amax: 0.0,
}),
n if n.contains("index_score") => Some(Policy::Derived {
coeff: class_coeff(n, depth_of(n) as f64),
flip_budget_frac: if native { 0.5 } else { 0.05 },
flip_rel: 0.0,
flip_abs_of_amax: if native { 0.25 } else { 0.05 },
}),
n => Some(Policy::Derived {
coeff: class_coeff(n, depth_of(n) as f64),
flip_budget_frac: 0.0,
flip_rel: 0.0,
flip_abs_of_amax: 0.0,
}),
}
}
fn top_ids(v: &[f32], k: usize) -> Vec<u32> {
let mut order: Vec<usize> = (0..v.len()).collect();
order.sort_by(|&a, &b| {
v[b].partial_cmp(&v[a])
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.cmp(&b))
});
order.into_iter().take(k).map(|x| x as u32).collect()
}
#[allow(clippy::too_many_arguments)]
fn check_array(
name: &str,
got: &[f32],
got_shape: &[usize],
ref_shape: &[usize],
ref_vals: &[f32],
policy: &Policy,
) -> ArrayCheck {
let mut c = ArrayCheck {
name: name.to_string(),
shape: got_shape.to_vec(),
max_abs: 0.0,
max_rel: 0.0,
threshold: 0.0,
n_over: 0,
verdict: "PASS",
note: String::new(),
};
if got_shape != ref_shape {
c.verdict = "FAIL";
c.note = format!("shape mismatch: got {got_shape:?}, fixture {ref_shape:?}");
return c;
}
if got.iter().any(|x| x.is_nan()) {
c.verdict = "FAIL";
c.note = "NaN in computed array".into();
return c;
}
let mut inf_mismatch = 0usize;
let mut absmax_ref = 0f32;
for (&g, &r) in got.iter().zip(ref_vals) {
if g.is_infinite() != r.is_infinite() || (g.is_infinite() && g != r) {
inf_mismatch += 1;
}
if r.is_finite() {
absmax_ref = absmax_ref.max(r.abs());
}
}
if inf_mismatch > 0 {
c.verdict = "FAIL";
c.note = format!("{inf_mismatch} ±inf-mask position mismatches (causality)");
return c;
}
let mut diffs: Vec<(usize, f64)> = Vec::new();
for (i, (&g, &r)) in got.iter().zip(ref_vals).enumerate() {
if !r.is_finite() {
continue;
}
let d = (g as f64 - r as f64).abs();
if d > c.max_abs {
c.max_abs = d;
}
let rel = d / (r.abs() as f64).max(1e-6);
if rel > c.max_rel {
c.max_rel = rel;
}
if d > 0.0 {
diffs.push((i, d));
}
}
match policy {
Policy::BitExact => {
if c.max_abs > 0.0 {
c.verdict = "FAIL";
c.note = format!(
"{} non-identical elements (bit-exact required)",
diffs.len()
);
}
}
Policy::Derived {
coeff,
flip_budget_frac,
flip_rel,
flip_abs_of_amax,
} => {
let ev = (2.0 * (got.len().max(2) as f64).ln()).sqrt(); c.threshold = coeff.max(U_BF16) * ev * absmax_ref as f64;
let over: Vec<&(usize, f64)> = diffs.iter().filter(|(_, d)| *d > c.threshold).collect();
c.n_over = over.len();
if !over.is_empty() {
let budget = (*flip_budget_frac * got.len() as f64).floor() as usize;
let bound_ok = |i: usize, d: f64| -> bool {
if *flip_abs_of_amax > 0.0 {
d <= *flip_abs_of_amax as f64 * absmax_ref as f64
} else if *flip_rel > 0.0 {
d <= *flip_rel as f64 * (got[i].abs() as f64).max(ref_vals[i].abs() as f64)
} else {
false
}
};
if over.len() <= budget && over.iter().all(|(i, d)| bound_ok(*i, *d)) {
c.note = format!(
"{} quantizer-flip exceeder(s) within the documented bound (budget {})",
over.len(),
budget
);
} else {
c.verdict = "FAIL";
c.note = format!(
"{} elements over threshold (first at flat index {})",
over.len(),
over.first().map(|(i, _)| *i).unwrap_or(0)
);
}
}
}
Policy::Logits {
coeff,
top_ids: want,
native,
} => {
let ev = (2.0 * (got.len().max(2) as f64).ln()).sqrt();
c.threshold = coeff * ev * absmax_ref as f64;
let got_top = top_ids(got, 21);
let ref_top = top_ids(ref_vals, 21);
if want[..5] != ref_top[..5] {
c.verdict = "FAIL";
c.note = format!(
"fixture self-inconsistency: json top5 {:?} vs npz {:?}",
&want[..5],
&ref_top[..5]
);
return c;
}
let top1_ok = got_top[0] == ref_top[0];
let s5g: BTreeSet<u32> = got_top[..5].iter().cloned().collect();
let s5r: BTreeSet<u32> = ref_top[..5].iter().cloned().collect();
let s20g: BTreeSet<u32> = got_top[..20].iter().cloned().collect();
let s20r: BTreeSet<u32> = ref_top[..20].iter().cloned().collect();
let overlap20 = s20g.intersection(&s20r).count();
let ref20 = ref_vals[ref_top[19] as usize];
let gap_ref = ref20 - ref_vals[ref_top[20] as usize];
let gap_got = got[got_top[19] as usize] - got[got_top[20] as usize];
let band = 3.0 * std::f64::consts::SQRT_2 * coeff * (ref20.abs() as f64);
let mut out_of_band: Vec<(u32, f32)> = Vec::new();
for id in s20r.symmetric_difference(&s20g) {
let dist = (ref_vals[*id as usize] - ref20).abs() as f64;
if dist > band {
out_of_band.push((*id, ref_vals[*id as usize]));
}
}
c.note = format!(
"top1 {}=={} | top5 set {} | top20 overlap {}/20 | 20/21 gap ref {:.4} gpu {:.4} | boundary band ±{:.3}: {} set-diff id(s), {} out-of-band",
got_top[0],
ref_top[0],
if s5g == s5r { "EXACT" } else { "MISMATCH" },
overlap20,
gap_ref,
gap_got,
band,
s20r.symmetric_difference(&s20g).count(),
out_of_band.len()
);
if !out_of_band.is_empty() {
c.note.push_str(&format!(" | OUT-OF-BAND: {out_of_band:?}"));
}
let top1_band_ok = top1_ok || {
let m = (ref_vals[ref_top[0] as usize] as f64
- ref_vals[got_top[0] as usize] as f64)
.abs();
m <= 3.0
* std::f64::consts::SQRT_2
* coeff
* (ref_vals[ref_top[0] as usize].abs() as f64)
};
let ref5_boundary = ref_vals[ref_top[4] as usize];
let band5 = 3.0 * std::f64::consts::SQRT_2 * coeff * (ref5_boundary.abs() as f64);
let top5_band_ok = s5r
.symmetric_difference(&s5g)
.all(|id| (ref_vals[*id as usize] - ref5_boundary).abs() as f64 <= band5);
if !top1_band_ok {
c.note.push_str(" | top1 OUT-OF-BAND");
}
if *native && s5g != s5r && !top5_band_ok {
c.note.push_str(" | top5 OUT-OF-BAND");
}
let strict_ok = if *native {
top1_band_ok && top5_band_ok
} else {
top1_ok && s5g == s5r
};
if !strict_ok || !out_of_band.is_empty() || c.max_abs > c.threshold {
c.verdict = "FAIL";
}
}
}
c
}
fn vram_table(gpu: &Dsv4Gpu, tag: &str) {
match gpu.vram_report() {
Ok(rows) => {
for (dev, free, total, resident) in rows {
println!(
"[vram {tag}] dev{dev}: used {:.2} GiB / {:.2} GiB (free {:.2}), loader-resident {:.2} GiB",
(total - free) as f64 / 2f64.powi(30),
total as f64 / 2f64.powi(30),
free as f64 / 2f64.powi(30),
resident as f64 / 2f64.powi(30),
);
}
}
Err(err) => println!("[vram {tag}] report failed: {err}"),
}
}
fn expert_dequant_subgate(gpu: &Dsv4Gpu, samples: &[(&str, usize, &str)]) -> bool {
use cudarc::driver::DevicePtr;
use memra_engine::dsv4_gpu::ExpertKind;
let mut ok = true;
for &(prefix, ex, proj) in samples {
let name = format!("{prefix}.ffn.experts.{ex}.{proj}");
let (shape, host) = gpu.model.tensor_f32(&name);
let (rows, cols) = (shape[0], shape[1]);
let (stage, layer) = if prefix == "mtp.0" {
let m = gpu.mtp.as_ref().expect("mtp loaded");
(&gpu.stages[gpu.stages.len() - 1], &m.layer)
} else {
let il: u32 = prefix.strip_prefix("layers.").unwrap().parse().unwrap();
let stage = &gpu.stages[gpu.layer_stage[il as usize]];
(
stage,
stage.layers.iter().find(|l| l.il == il).expect("layer"),
)
};
let pi = match proj {
"w1" => 0usize,
"w2" => 1,
_ => 2,
};
stage.gpu.ctx.bind_to_thread().expect("bind ctx");
let stream = stage.gpu.stream();
let wbytes = rows * cols / 2;
let sbytes = match layer.expert_kind {
ExpertKind::Nvfp4 => rows * cols / 16,
ExpertKind::Mxfp4 => rows * cols / 32,
};
let wp = (layer.experts_w.device_ptr(&stream).0 as usize + (ex * 3 + pi) * wbytes)
as *const std::os::raw::c_void;
let scp = (layer.experts_sc.device_ptr(&stream).0 as usize + (ex * 3 + pi) * sbytes)
as *const std::os::raw::c_void;
let dst = stage.deq[pi].device_ptr(&stream).0 as *mut std::os::raw::c_void;
let sv = stream.cu_stream() as *mut std::os::raw::c_void;
let rc = unsafe {
match layer.expert_kind {
ExpertKind::Nvfp4 => memra_engine::dsv4_ffi::memra_dsv4_nvfp4_deq_bf16(
wp,
scp,
layer.experts_s2[ex * 3 + pi],
rows as i32,
cols as i32,
dst,
sv,
),
ExpertKind::Mxfp4 => memra_engine::dsv4_ffi::memra_dsv4_mxfp4_deq_bf16(
wp,
scp,
rows as i32,
cols as i32,
dst,
sv,
),
}
};
assert_eq!(rc, 0, "dequant kernel rc {rc}");
let mut raw = vec![0u8; rows * cols * 2];
stream
.memcpy_dtoh(&stage.deq[pi].slice(0..rows * cols * 2), &mut raw[..])
.expect("dtoh");
stream.synchronize().expect("sync");
let mut mismatches = 0usize;
for i in 0..rows * cols {
let b = u16::from_le_bytes([raw[2 * i], raw[2 * i + 1]]);
let g = f32::from_bits((b as u32) << 16);
if g.to_bits() != host[i].to_bits() {
if mismatches == 0 {
println!(" [FAIL] {name}[{i}]: gpu {g} vs host {}", host[i]);
}
mismatches += 1;
}
}
let verdict = if mismatches == 0 { "PASS" } else { "FAIL" };
println!(
" [{verdict}] expert-dequant {name} [{rows}x{cols}]: {mismatches} mismatches (bit-exact required)"
);
ok &= mismatches == 0;
}
ok
}
fn main() {
let args: Vec<String> = std::env::args().collect();
if args.len() < 3 {
eprintln!("usage: dsv4-gpu-gate <model-dir> <fixtures.json> [dev0,dev1]");
std::process::exit(2);
}
let t0 = std::time::Instant::now();
let dir = Path::new(&args[1]);
let spec = FixtureSpec::load(Path::new(&args[2]));
let devices: Vec<usize> = args
.get(3)
.map(|s| {
s.split(',')
.map(|x| x.parse().expect("device ordinal"))
.collect()
})
.unwrap_or_else(|| vec![0, 1]);
println!(
"dsv4-gpu-gate | model {} | fixtures {} | variant {} ({:?}) | devices {devices:?}",
dir.display(),
args[2],
spec.variant_tag,
spec.variant
);
let npz = read_npz(&spec.npz_path);
let mut failures: Vec<String> = Vec::new();
for (name, (shape, sha)) in &spec.arrays {
match npz.get(name) {
None => failures.push(format!("fixture npz missing array {name}")),
Some((nshape, _, nsha)) => {
if nshape != shape || nsha != sha {
failures.push(format!("fixture integrity mismatch on {name}"));
}
}
}
}
if !failures.is_empty() {
for f in &failures {
println!(" FAIL: {f}");
}
std::process::exit(1);
}
println!(
"fixture integrity: {} arrays, payload sha256 all match",
npz.len()
);
let max_seq = 4096.min(
spec.tokens_160
.as_ref()
.map(|t| t.len())
.unwrap_or(0)
.max(512),
);
let gpu = Dsv4Gpu::load(dir, &devices, spec.variant, max_seq).expect("load");
println!(
"loaded: split at layer {} (stage0 {} layers, stage1 {} layers), t={:.0}s",
gpu.split_at,
gpu.stages[0].layers.len(),
gpu.stages[1].layers.len(),
t0.elapsed().as_secs_f64()
);
vram_table(&gpu, "post-load");
let n_trunk = gpu.model.mc.n_layer - gpu.model.mc.nextn_predict_layers;
let stage1_prefix = format!("layers.{}", gpu.split_at);
let last_prefix = format!("layers.{}", n_trunk - 1);
let mut samples: Vec<(&str, usize, &str)> = vec![
("layers.0", 0, "w1"),
("layers.2", 100, "w3"),
("layers.20", 7, "w1"), (&stage1_prefix, 31, "w2"),
(&last_prefix, 255, "w2"),
];
if gpu.mtp.is_some() {
samples.push(("mtp.0", 7, "w1")); samples.push(("mtp.0", 200, "w2"));
} else {
println!(
" (MXFP4 mtp.* samples skipped: no NextN block resident — DSpark drafter lane owns that gate)"
);
}
println!("\n== expert-dequant sub-gate (GPU kernels vs lane-1 host decoders) ==");
let deq_ok = expert_dequant_subgate(&gpu, &samples);
let mut cap32: BTreeSet<u32> = BTreeSet::new();
let mut cap160: BTreeSet<u32> = BTreeSet::new();
for name in spec.arrays.keys() {
if let Some(rest) = name.strip_prefix("c160_layer") {
if let Ok(n) = rest.split('_').next().unwrap_or("").parse::<u32>() {
cap160.insert(n);
}
} else if let Some(rest) = name.strip_prefix("layer") {
if let Ok(n) = rest.split('_').next().unwrap_or("").parse::<u32>() {
cap32.insert(n);
}
}
}
let hc = gpu.model.cfg().hc_mult as usize;
let hidden = gpu.model.mc.n_embd as usize;
let mut table: Vec<ArrayCheck> = Vec::new();
let mut checked: BTreeSet<String> = BTreeSet::new();
let mut skipped: Vec<String> = Vec::new();
let mut check = |name: &str, got: &[f32], got_shape: &[usize], table: &mut Vec<ArrayCheck>| {
let Some((ref_shape, ref_vals, _)) = npz.get(name) else {
return;
};
let Some(policy) = policy_for(name, &spec) else {
return;
};
let c = check_array(name, got, got_shape, ref_shape, ref_vals, &policy);
println!(
" [{}] {} {:?}: max-abs {:.3e} max-rel {:.3e} thr {:.3e} over {}{}",
c.verdict,
c.name,
c.shape,
c.max_abs,
c.max_rel,
c.threshold,
c.n_over,
if c.note.is_empty() {
String::new()
} else {
format!(" | {}", c.note)
}
);
checked.insert(name.to_string());
table.push(c);
};
let ids = &spec.tokens_32;
let s = ids.len();
let mut cap = GpuCapture {
want: cap32.clone(),
..Default::default()
};
let fwd = gpu
.forward(ids, Some(&mut cap), None)
.expect("forward 32")
.expect("logits expected");
let logits = fwd.logits;
println!("32-token forward done t={:.0}s", t0.elapsed().as_secs_f64());
vram_table(&gpu, "post-warmup");
if let Some(e0) = &cap.embed_out {
check("embed_out", e0, &[1, s, hidden], &mut table);
}
for lid in &cap32 {
if let Some(h) = cap.layer_out.get(lid) {
check(
&format!("layer{lid}_out"),
h,
&[1, s, hc, hidden],
&mut table,
);
}
if let Some(a) = cap.attn_out.get(lid) {
check(
&format!("layer{lid}_attn_out"),
a,
&[1, s, hidden],
&mut table,
);
}
if let Some((kv, nb)) = cap.compressor_kv.get(lid) {
check(
&format!("layer{lid}_compressor_kv"),
kv,
&[1, *nb, kv.len() / nb],
&mut table,
);
}
if let Some((kv, nb)) = cap.indexer_kv.get(lid) {
check(
&format!("layer{lid}_indexer_kv"),
kv,
&[1, *nb, kv.len() / nb],
&mut table,
);
}
if let Some((sc, nb)) = cap.index_score.get(lid) {
check(
&format!("layer{lid}_index_score"),
sc,
&[1, s, *nb],
&mut table,
);
}
}
check("final_logits_last", &logits, &[logits.len()], &mut table);
if spec.arrays.contains_key("mtp_logits_last") {
if gpu.mtp.is_some() {
let ml = gpu.mtp_logits_last(&fwd.h_last, ids).expect("mtp forward");
check("mtp_logits_last", &ml, &[ml.len()], &mut table);
println!("mtp forward done t={:.0}s", t0.elapsed().as_secs_f64());
} else {
skipped.push(
"mtp_logits_last (MTP block absent from artifact — SKIPPED, not PASS)".into(),
);
}
}
if let (Some(ids160), Some(&max_l)) = (&spec.tokens_160, cap160.iter().max()) {
let s2 = ids160.len();
let mut cap2 = GpuCapture {
want: cap160.clone(),
..Default::default()
};
let r = gpu
.forward(ids160, Some(&mut cap2), Some(max_l))
.expect("forward 160");
assert!(r.is_none(), "early exit expected");
println!(
"160-token partial forward done t={:.0}s",
t0.elapsed().as_secs_f64()
);
for lid in &cap160 {
if let Some(h) = cap2.layer_out.get(lid) {
check(
&format!("c160_layer{lid}_out_last"),
&h[(s2 - 1) * hc * hidden..],
&[1, hc, hidden],
&mut table,
);
}
if let Some(a) = cap2.attn_out.get(lid) {
check(
&format!("c160_layer{lid}_attn_out_last"),
&a[(s2 - 1) * hidden..],
&[1, hidden],
&mut table,
);
}
if let Some((kv, nb)) = cap2.compressor_kv.get(lid) {
check(
&format!("c160_layer{lid}_compressor_kv"),
kv,
&[1, *nb, kv.len() / nb],
&mut table,
);
}
}
}
vram_table(&gpu, "post-gate");
for name in spec.arrays.keys() {
if !checked.contains(name) && name != "mtp_logits_last" {
failures.push(format!(
"banked array {name} was never computed by the gate"
));
}
}
if !deq_ok {
failures.push("expert-dequant sub-gate failed".into());
}
for c in &table {
if c.verdict == "FAIL" {
failures.push(format!("{}: {}", c.name, c.note));
}
}
println!(
"\n== GPU gate table (variant {}, {} class) ==",
spec.variant_tag,
if expert_arm_native() {
"NATIVE expert arm: C = sqrt(d_b*u_b^2 + d_q*u_q^2), u_q = 2^-4"
} else {
"bf16-dequant arm: u = 2^-8 doctrine"
}
);
println!("| array | shape | max-abs | max-rel | threshold | verdict |");
println!("|---|---|---|---|---|---|");
for c in &table {
println!(
"| {} | {:?} | {:.3e} | {:.3e} | {:.3e} | {}{} |",
c.name,
c.shape,
c.max_abs,
c.max_rel,
c.threshold,
c.verdict,
if c.note.is_empty() {
String::new()
} else {
format!(" ({})", c.note)
}
);
}
for skip in &skipped {
println!("| SKIPPED: {skip} |");
}
println!(
"\nelapsed: {:.1}s (informational, single-run, not a perf claim)",
t0.elapsed().as_secs_f64()
);
if failures.is_empty() {
println!(
"GPU OUTPUT-SAMPLE GATE [{}]: PASS ({} arrays compared, {} skipped, 0 failures)",
spec.variant_tag,
table.len(),
skipped.len()
);
} else {
println!(
"GPU OUTPUT-SAMPLE GATE [{}]: FAIL ({} failures)",
spec.variant_tag,
failures.len()
);
for f in &failures {
println!(" FAIL: {f}");
}
std::process::exit(1);
}
}