use memra_engine::Engine;
use std::sync::atomic::Ordering;
fn vecf(n: usize, seed: u64) -> Vec<f32> {
let mut s = seed.wrapping_mul(0x9E37_79B9_7F4A_7C15) | 1;
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
((s >> 40) as f32 / (1u64 << 24) as f32 - 0.5) * 0.8
})
.collect()
}
fn run(
e: &Engine,
w: &[f32],
x: &[f32],
m: usize,
in_f: usize,
out_f: usize,
native: bool,
) -> Vec<f32> {
let xd = e.htod(x).unwrap();
let wd = e.htod(w).unwrap();
let mut y = e.uninit(m * out_f).unwrap();
if native {
assert!(
e.gemv_f32_rows_into(&xd, &wd, &mut y, m, in_f, out_f)
.unwrap()
);
} else {
let yy = e.linear(&xd, &wd, m, in_f, out_f).unwrap();
e.copy_into(&mut y, 0, &yy, m * out_f).unwrap();
}
e.stream().synchronize().unwrap();
e.dtoh(&y).unwrap()
}
fn bits(v: &[f32]) -> Vec<u32> {
v.iter().map(|x| x.to_bits()).collect()
}
#[test]
fn gemv_f32_rows_matches_cublas_deterministic_and_m_identical() {
let Ok(e) = Engine::new(0) else {
eprintln!("no CUDA device; skipping");
return;
};
for (in_f, out_f) in [(4096usize, 128usize), (4096, 32), (8192, 24)] {
let w = vecf(out_f * in_f, 31);
let x = vecf(4 * in_f, 32);
let a = run(&e, &w, &x[..in_f], 1, in_f, out_f, false);
let b = run(&e, &w, &x[..in_f], 1, in_f, out_f, true);
let c = run(&e, &w, &x[..in_f], 1, in_f, out_f, true);
assert_eq!(
bits(&b),
bits(&c),
"two native launches differ ({in_f}->{out_f})"
);
let mut worst = 0.0f32;
for (p, q) in a.iter().zip(&b) {
let tol = 1e-5 * (1.0 + q.abs());
worst = worst.max((p - q).abs() / tol);
assert!(
(p - q).abs() <= tol,
"cublas {p} vs native {q} ({in_f}->{out_f})"
);
}
let m4 = run(&e, &w, &x, 4, in_f, out_f, true);
for j in 0..4 {
let one = run(&e, &w, &x[j * in_f..(j + 1) * in_f], 1, in_f, out_f, true);
assert_eq!(
bits(&m4[j * out_f..(j + 1) * out_f]),
bits(&one),
"row {j} of the m=4 launch is not the m=1 launch ({in_f}->{out_f})"
);
}
let mut w2 = w.clone();
w2[in_f / 2] += 0.25;
let r = run(&e, &w2, &x[..in_f], 1, in_f, out_f, true);
assert_ne!(bits(&r), bits(&b), "perturbed W did not move row 0");
println!("gemv_f32_rows {in_f}->{out_f}: worst |cublas-native|/tol = {worst:.3}");
}
let w = vecf(8 * 1000, 33);
let x = vecf(1000, 34);
let xd = e.htod(&x).unwrap();
let wd = e.htod(&w).unwrap();
let mut y = e.uninit(8).unwrap();
assert!(!e.gemv_f32_rows_into(&xd, &wd, &mut y, 1, 1000, 8).unwrap());
let w = vecf(8 * 1024, 35);
let x = vecf(17 * 1024, 36);
let xd = e.htod(&x).unwrap();
let wd = e.htod(&w).unwrap();
let mut y = e.uninit(17 * 8).unwrap();
assert!(!e.gemv_f32_rows_into(&xd, &wd, &mut y, 17, 1024, 8).unwrap());
{
let (in_f, out_f) = (4096usize, 128usize);
let w = vecf(out_f * in_f, 51);
let x = vecf(4 * in_f, 52);
let xd = e.htod(&x).unwrap();
let wd = e.htod(&w).unwrap();
unsafe {
std::env::set_var("MEMRA_F32_GEMV_KERNEL", "1");
}
let c_before = memra_engine::F32_GEMV_KERNEL_DISPATCHES.load(Ordering::Relaxed);
let y4 = e.linear_decode_exact(&xd, &wd, 4, in_f, out_f).unwrap();
let c_after = memra_engine::F32_GEMV_KERNEL_DISPATCHES.load(Ordering::Relaxed);
e.stream().synchronize().unwrap();
let h4 = e.dtoh(&y4).unwrap();
assert_eq!(
c_after - c_before,
1,
"decode-exact rows took {} launches, not one",
c_after - c_before
);
for j in 0..4 {
let xj = e.htod(&x[j * in_f..(j + 1) * in_f]).unwrap();
let yj = e.linear(&xj, &wd, 1, in_f, out_f).unwrap();
e.stream().synchronize().unwrap();
assert_eq!(
bits(&h4[j * out_f..(j + 1) * out_f]),
bits(&e.dtoh(&yj).unwrap()),
"decode-exact row {j} is not the m=1 launch"
);
}
unsafe {
std::env::set_var("MEMRA_F32_GEMV_KERNEL", "0");
}
}
let c0 = memra_engine::F32_GEMV_KERNEL_DISPATCHES.load(Ordering::Relaxed);
unsafe {
std::env::set_var("MEMRA_F32_GEMV_KERNEL", "1");
}
let w = vecf(32 * 4096, 37);
let x = vecf(4096, 38);
let xd = e.htod(&x).unwrap();
let wd = e.htod(&w).unwrap();
let _ = e.linear(&xd, &wd, 1, 4096, 32).unwrap();
unsafe {
std::env::set_var("MEMRA_F32_GEMV_KERNEL", "0");
}
let c1 = memra_engine::F32_GEMV_KERNEL_DISPATCHES.load(Ordering::Relaxed);
assert!(
c1 > c0,
"VACUOUS: `linear` under the door never took the kernel ({c0} -> {c1})"
);
}