use memra_engine::Engine;
use memra_engine::hybrid::HybridModel;
use memra_gguf::GgufFile;
use memra_gguf::micro_gguf::write_glm_dsa_micro;
fn gpu_guard() -> std::sync::MutexGuard<'static, ()> {
static GPU: std::sync::Mutex<()> = std::sync::Mutex::new(());
GPU.lock().unwrap_or_else(|poisoned| poisoned.into_inner())
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
#[allow(clippy::unusual_byte_groupings)] fn gpu_glm_dsa_micro_block_forward_runs() {
let _gpu = gpu_guard();
let p = std::env::temp_dir().join(format!("memra-mla-fwd-gpu-{}.gguf", std::process::id()));
write_glm_dsa_micro(&p, 0x4_F0_2026).unwrap();
let g = GgufFile::open(&p).unwrap();
let e = Engine::new(0).expect("CUDA device 0");
let model = HybridModel::load(&e, &g).expect("glm-dsa micro fixture loads");
std::fs::remove_file(&p).ok();
let tokens: Vec<u32> = vec![1, 5, 9, 13, 17, 21];
let logits = model.forward(&e, &tokens).expect("glm-dsa micro prefill");
let n_vocab = model.cfg.n_vocab as usize;
assert_eq!(logits.len(), tokens.len() * n_vocab, "logits shape");
assert!(
logits.iter().all(|v| v.is_finite()),
"MLA forward produced non-finite logits"
);
let spread = logits.iter().copied().fold(f32::NEG_INFINITY, f32::max)
- logits.iter().copied().fold(f32::INFINITY, f32::min);
assert!(
spread > 1e-4,
"logits degenerate (flat) — a projection is dead"
);
}
use memra_engine::mla::{
MlaDims, MlaInputs, mla_attend_absorbed, mla_attend_naive, rope_interleaved,
};
use memra_gguf::micro_gguf::Rng;
fn maxdiff(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len(), "compared tensors differ in length");
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs())
.fold(0.0f32, f32::max)
}
fn maxabs(a: &[f32]) -> f32 {
a.iter().map(|x| x.abs()).fold(0.0f32, f32::max)
}
struct Case {
q_nope: Vec<f32>,
q_pe: Vec<f32>,
c_kv: Vec<f32>,
k_pe: Vec<f32>,
w_uk: Vec<f32>,
w_uv: Vec<f32>,
}
impl Case {
fn new(d: &MlaDims, t_q: usize, t_kv: usize, seed: u64) -> Self {
let mut rng = Rng(seed | 1);
let ws = 1.0 / (d.kv_rank as f32).sqrt();
Case {
q_nope: rng.fill(t_q * d.n_head * d.d_nope, 1.0),
q_pe: rng.fill(t_q * d.n_head * d.d_rope, 1.0),
c_kv: rng.fill(t_kv * d.kv_rank, 1.0),
k_pe: rng.fill(t_kv * d.d_rope, 1.0),
w_uk: rng.fill(d.n_head * d.d_nope * d.kv_rank, ws),
w_uv: rng.fill(d.n_head * d.d_v * d.kv_rank, ws),
}
}
fn inputs<'a>(&'a self, t_q: usize, t_kv: usize) -> MlaInputs<'a> {
MlaInputs {
q_nope: &self.q_nope,
q_pe: &self.q_pe,
c_kv: &self.c_kv,
k_pe: &self.k_pe,
w_uk: &self.w_uk,
w_uv: &self.w_uv,
t_q,
t_kv,
}
}
fn cache_rows(&self, d: &MlaDims, t_kv: usize) -> Vec<f32> {
let mut rows = Vec::with_capacity(t_kv * (d.kv_rank + d.d_rope));
for t in 0..t_kv {
rows.extend_from_slice(&self.c_kv[t * d.kv_rank..(t + 1) * d.kv_rank]);
rows.extend_from_slice(&self.k_pe[t * d.d_rope..(t + 1) * d.d_rope]);
}
rows
}
fn wk_b(&self, d: &MlaDims) -> Vec<f32> {
let (nh, dn, r) = (d.n_head, d.d_nope, d.kv_rank);
let mut w = vec![0.0f32; nh * r * dn];
for h in 0..nh {
for l in 0..r {
for p in 0..dn {
w[h * r * dn + l * dn + p] = self.w_uk[h * dn * r + p * r + l];
}
}
}
w
}
}
fn gpu_core_parity(e: &Engine, d: &MlaDims, t_q: usize, t_kv: usize, seed: u64, tol: f32) {
let (nh, dn, dr, dv, r) = (d.n_head, d.d_nope, d.d_rope, d.d_v, d.kv_rank);
let c = Case::new(d, t_q, t_kv, seed);
let x = c.inputs(t_q, t_kv);
let want_absorbed = mla_attend_absorbed(d, &x);
let want_naive = mla_attend_naive(d, &x);
let q_nope = e.htod(&c.q_nope).unwrap();
let q_pe = e
.htod(if dr == 0 { &[0.0f32][..] } else { &c.q_pe })
.unwrap();
let cache = e.htod(&c.cache_rows(d, t_kv)).unwrap();
let wk_b = e.htod(&c.wk_b(d)).unwrap();
let wv_b = e.htod(&c.w_uv).unwrap();
let mut q_lat = e.uninit(t_q * nh * r).unwrap();
e.mla_absorb_q(&q_nope, &wk_b, &mut q_lat, t_q, nh, dn, r)
.unwrap();
let mut o_lat = e.uninit(t_q * nh * r).unwrap();
e.mla_attn_absorbed(
&q_lat,
&q_pe,
&cache,
&mut o_lat,
nh,
r,
dr,
t_q,
t_kv,
d.scale(),
)
.unwrap();
let mut out = e.uninit(t_q * nh * dv).unwrap();
e.mla_decompress_v(&o_lat, &wv_b, &mut out, t_q, nh, dv, r)
.unwrap();
let got = e.dtoh(&out).unwrap();
assert!(
got.iter().all(|v| v.is_finite()),
"GPU MLA core produced non-finite values (dims {d:?} t_q {t_q} t_kv {t_kv})"
);
let scale = maxabs(&want_absorbed).max(1.0);
for (name, want) in [("absorbed", &want_absorbed), ("naive", &want_naive)] {
let md = maxdiff(&got, want);
assert!(
md <= tol * scale,
"GPU MLA core vs CPU {name} oracle: maxdiff {md:.3e} (scale {scale:.3e}, \
rel {:.3e}, tol {tol:.1e}) dims {d:?} t_q {t_q} t_kv {t_kv} seed {seed}",
md / scale
);
}
assert!(
maxabs(&got) > 1e-6,
"GPU MLA core output is degenerate (all ~zero)"
);
}
fn glm52() -> MlaDims {
MlaDims::GLM52
}
fn glm5_next() -> MlaDims {
MlaDims::GLM5_NEXT
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn gpu_mla_prefill_parity_vs_cpu_oracle() {
let _gpu = gpu_guard();
let e = Engine::new(0).expect("CUDA device 0");
let small = [
MlaDims {
n_head: 4,
d_nope: 24,
d_rope: 8,
d_v: 32,
kv_rank: 64,
},
MlaDims {
n_head: 3,
d_nope: 12,
d_rope: 4,
d_v: 20,
kv_rank: 48,
},
MlaDims {
n_head: 4,
d_nope: 32,
d_rope: 0,
d_v: 32,
kv_rank: 64,
},
];
for (i, d) in small.iter().enumerate() {
for seed in [7u64, 1234, 0xB1E55ED] {
gpu_core_parity(&e, d, 6, 6, seed + i as u64, 1e-5); gpu_core_parity(&e, d, 3, 11, seed + 5 + i as u64, 1e-5); }
}
for d in [glm52(), glm5_next()] {
gpu_core_parity(&e, &d, 4, 4, 20260827, 1e-4);
gpu_core_parity(&e, &d, 3, 9, 20260828, 1e-4);
gpu_core_parity(&e, &d, 17, 17, 20260829, 1e-4); }
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn gpu_mla_decode_stepwise_parity_vs_cpu_oracle() {
let _gpu = gpu_guard();
let e = Engine::new(0).expect("CUDA device 0");
for (d, tol) in [
(
MlaDims {
n_head: 4,
d_nope: 24,
d_rope: 8,
d_v: 32,
kv_rank: 64,
},
1e-5,
),
(
MlaDims {
n_head: 4,
d_nope: 32,
d_rope: 0,
d_v: 32,
kv_rank: 64,
},
1e-5,
),
(glm52(), 1e-4),
(glm5_next(), 1e-4),
] {
let (nh, dn, dr, dv, r) = (d.n_head, d.d_nope, d.d_rope, d.d_v, d.kv_rank);
let width = r + dr;
let prefill = 5usize;
let steps = 4usize;
let total = prefill + steps;
let c = Case::new(&d, total, total, 0x5_7E9 + total as u64);
let rows = c.cache_rows(&d, total);
let mut cache = e.zeros(total * width).unwrap();
let wk_b = e.htod(&c.wk_b(&d)).unwrap();
let wv_b = e.htod(&c.w_uv).unwrap();
let append = |lo: usize, hi: usize, cache: &mut _| {
let n = hi - lo;
let c_kv = e.htod(&c.c_kv[lo * r..hi * r]).unwrap();
let k_pe = e
.htod(if dr == 0 {
&[0.0f32][..]
} else {
&c.k_pe[lo * dr..hi * dr]
})
.unwrap();
e.mla_append_latent(cache, &c_kv, &k_pe, lo, n, r, dr)
.unwrap();
};
append(0, prefill, &mut cache);
let got_rows = e.dtoh(&cache).unwrap();
assert_eq!(
&got_rows[..prefill * width],
&rows[..prefill * width],
"append_latent did not reproduce the [c_kv | k_pe] row layout (dims {d:?})"
);
for step in 0..steps {
let pos = prefill + step;
append(pos, pos + 1, &mut cache);
let t_kv = pos + 1;
let q_nope = e
.htod(&c.q_nope[pos * nh * dn..(pos + 1) * nh * dn])
.unwrap();
let q_pe = e
.htod(if dr == 0 {
&[0.0f32][..]
} else {
&c.q_pe[pos * nh * dr..(pos + 1) * nh * dr]
})
.unwrap();
let mut q_lat = e.uninit(nh * r).unwrap();
e.mla_absorb_q(&q_nope, &wk_b, &mut q_lat, 1, nh, dn, r)
.unwrap();
let mut o_lat = e.uninit(nh * r).unwrap();
e.mla_attn_absorbed(
&q_lat,
&q_pe,
&cache,
&mut o_lat,
nh,
r,
dr,
1,
t_kv,
d.scale(),
)
.unwrap();
let mut out = e.uninit(nh * dv).unwrap();
e.mla_decompress_v(&o_lat, &wv_b, &mut out, 1, nh, dv, r)
.unwrap();
let got = e.dtoh(&out).unwrap();
let want = mla_attend_absorbed(
&d,
&MlaInputs {
q_nope: &c.q_nope[pos * nh * dn..(pos + 1) * nh * dn],
q_pe: &c.q_pe[pos * nh * dr..(pos + 1) * nh * dr],
c_kv: &c.c_kv[..t_kv * r],
k_pe: &c.k_pe[..t_kv * dr],
w_uk: &c.w_uk,
w_uv: &c.w_uv,
t_q: 1,
t_kv,
},
);
let scale = maxabs(&want).max(1.0);
let md = maxdiff(&got, &want);
assert!(
md <= tol * scale,
"decode step {step} (pos {pos}, t_kv {t_kv}): maxdiff {md:.3e} \
(scale {scale:.3e}, rel {:.3e}, tol {tol:.1e}) dims {d:?}",
md / scale
);
}
}
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn gpu_mla_rope_interleaved_matches_cpu() {
let _gpu = gpu_guard();
let e = Engine::new(0).expect("CUDA device 0");
const BASE: f32 = 8_000_000.0; for &d_rope in &[8usize, 64] {
let (n_pos, n_vec) = (7usize, 5usize);
let mut rng = Rng(0x0DE_2026);
let host = rng.fill(n_pos * n_vec * d_rope, 1.0);
let positions: Vec<i32> = (0..n_pos as i32).map(|p| p * 13).collect();
let mut want = host.clone();
#[allow(clippy::needless_range_loop)]
for p in 0..n_pos {
for v in 0..n_vec {
let off = (p * n_vec + v) * d_rope;
rope_interleaved(
&mut want[off..off + d_rope],
d_rope,
positions[p] as f32,
BASE,
);
}
}
let mut x = e.htod(&host).unwrap();
let pos_d = e.htod_i32(&positions).unwrap();
e.mla_rope_interleaved(&mut x, &pos_d, n_pos, n_vec, d_rope, BASE)
.unwrap();
let got = e.dtoh(&x).unwrap();
let md = maxdiff(&got, &want);
assert!(
md <= 1e-4 * maxabs(&want).max(1.0),
"interleaved rope d_rope {d_rope}: maxdiff {md:.3e}"
);
}
let host = vec![1.0f32, -2.0, 3.0];
let mut x = e.htod(&host).unwrap();
let pos_d = e.htod_i32(&[0i32]).unwrap();
e.mla_rope_interleaved(&mut x, &pos_d, 1, 1, 0, BASE)
.unwrap();
assert_eq!(e.dtoh(&x).unwrap(), host, "NoPE rope must not modify data");
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
#[allow(clippy::unusual_byte_groupings)] fn gpu_mla_cached_prime_decode_matches_stateless_forward() {
let _gpu = gpu_guard();
let p = std::env::temp_dir().join(format!("memra-mla-cached-{}.gguf", std::process::id()));
write_glm_dsa_micro(&p, 0xC_ACE_0802).unwrap();
let g = GgufFile::open(&p).unwrap();
let e = Engine::new(0).expect("CUDA device 0");
let model = HybridModel::load(&e, &g).expect("glm-dsa micro fixture loads");
std::fs::remove_file(&p).ok();
let tokens: Vec<u32> = (0..20u32).map(|i| (i * 7 + 3) % 128).collect();
let n_vocab = model.cfg.n_vocab as usize;
let full = model.forward(&e, &tokens).expect("stateless MLA prefill");
let want = &full[(tokens.len() - 1) * n_vocab..];
let mut cache = memra_kv::Cache::new(&e, &model.cfg, 64).expect("latent cache allocates");
for (il, layer) in cache.latent.iter().enumerate() {
assert!(
layer.is_some(),
"layer {il}: StatePlan::LatentKvCache did not allocate a latent plane"
);
}
let (_prime_logits, _, _) = model
.prime_cache(&e, &tokens[..tokens.len() - 1], &mut cache, 0)
.expect("MLA prime through the latent cache");
let observed: Vec<usize> = cache
.latent
.iter()
.map(|l| l.as_ref().unwrap().len)
.collect();
for il in 0..model.layers.len() {
assert_eq!(
observed[il],
tokens.len() - 1,
"trunk layer {il}: latent plane length after prime (all planes: {observed:?})"
);
}
for il in model.layers.len()..observed.len() {
assert_eq!(
observed[il], 0,
"MTP block {il}: the trunk prime must not append to the MTP latent plane \
(all planes: {observed:?})"
);
}
let (got, _) = model
.decode_step_h(&e, tokens[tokens.len() - 1], &mut cache)
.expect("MLA T=1 decode through the latent cache");
assert_eq!(got.len(), n_vocab);
assert!(got.iter().all(|v| v.is_finite()));
let scale = maxabs(want).max(1.0);
let md = maxdiff(&got, want);
assert!(
md <= 0.005 && md <= 0.005 * scale,
"cached prime+decode vs stateless forward: maxdiff {md:.3e} (scale {scale:.3e}, \
rel {:.3e}) exceeds the glm_dsa checkpoint-parity bar 5e-3",
md / scale
);
let am = |v: &[f32]| {
v.iter()
.enumerate()
.fold((0usize, f32::NEG_INFINITY), |b, (i, &x)| {
if x > b.1 { (i, x) } else { b }
})
.0
};
assert_eq!(am(&got), am(want), "cached vs stateless argmax disagree");
}