use memra_engine::{Engine, SCRATCH_ALLOC_CALLS};
use memra_gguf::GgmlType;
use memra_gguf::config::{HfConfig, ModelConfig};
use memra_gguf::model_plan::ModelPlan;
use memra_gguf::source::{TensorSource, TensorView};
use memra_gguf::tensor_contract::{
CheckpointDialect, ContractOptions, OutputHead, TensorContract, TensorId, TensorMatch,
};
use memra_reference::{ReferenceTensor, deterministic_fixture};
use std::borrow::Cow;
use std::collections::BTreeMap;
use std::sync::atomic::Ordering;
const VOCAB: u32 = 32;
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())
}
fn force_true_f32() {
static ONCE: std::sync::Once = std::sync::Once::new();
ONCE.call_once(|| {
if std::env::var("NVIDIA_TF32_OVERRIDE").as_deref() != Ok("0") {
unsafe { std::env::set_var("NVIDIA_TF32_OVERRIDE", "0") };
}
});
}
fn mini_config_json() -> String {
r#"{
"model_type": "glm5_next_text",
"num_hidden_layers": 8,
"num_nextn_predict_layers": 0,
"hidden_size": 128,
"intermediate_size": 64,
"vocab_size": 32,
"max_position_embeddings": 512,
"rms_norm_eps": 1e-05,
"hidden_act": "silu",
"swiglu_limit": 1e30,
"tie_word_embeddings": true,
"hc_mult": 4,
"hc_eps": 1e-06,
"hc_sinkhorn_iters": 20,
"mhc": true,
"layer_types": ["linear_attention", "linear_attention", "linear_attention", "deepseek_sparse_attention", "linear_attention", "linear_attention", "linear_attention", "deepseek_sparse_attention"],
"mlp_layer_types": ["dense", "sparse", "sparse", "sparse", "sparse", "sparse", "sparse", "sparse"],
"first_k_dense_replace": 1,
"indexer_types": ["full", "full", "full", "full", "full", "full", "full", "full"],
"linear_attn_config": {
"num_heads": 1,
"head_dim": 128,
"short_conv_kernel_size": 4,
"gate_lower_bound": -5.0,
"kda_layers": [0, 1, 2, 4, 5, 6],
"full_attn_layers": [3, 7]
},
"num_attention_heads": 1,
"num_key_value_heads": 1,
"q_lora_rank": 16,
"kv_lora_rank": 16,
"qk_head_dim": 16,
"qk_nope_head_dim": 16,
"qk_rope_head_dim": 0,
"v_head_dim": 16,
"mla_use_nope": true,
"index_n_heads": 1,
"index_head_dim": 8,
"index_topk": 8,
"index_kpool": 4,
"index_kpool_always_select_tail": true,
"index_kpool_compress": true,
"indexer_rope_interleave": true,
"index_share_for_mtp_iteration": true,
"n_routed_experts": 64,
"num_experts_per_tok": 8,
"moe_intermediate_size": 64,
"n_shared_experts": 1,
"scoring_func": "sigmoid",
"topk_method": "noaux_tc",
"routed_scaling_factor": 2.5,
"norm_topk_prob": true,
"n_group": 1,
"topk_group": 1,
"head_dim": 0,
"attention_bias": false,
"moe_router_dtype": "float32",
"dtype": "bfloat16"
}"#
.to_string()
}
fn mini_config() -> ModelConfig {
ModelConfig::from_hf(&HfConfig::parse(&mini_config_json()))
}
fn mini_plan(config: &ModelConfig) -> ModelPlan {
memra_gguf::model_packs::for_config(config)
.expect("glm5_next model pack matches the mini config")
.compile_plan(config)
.expect("mini glm5_next plan compiles")
}
struct OwnedTensor {
bytes: Vec<u8>,
ne: Vec<u64>,
ggml_type: GgmlType,
}
struct FixtureSource {
config: ModelConfig,
tensors: BTreeMap<String, OwnedTensor>,
}
impl TensorSource for FixtureSource {
fn config(&self) -> ModelConfig {
self.config.clone()
}
fn find(&self, name: &str) -> Option<TensorView<'_>> {
let t = self.tensors.get(name)?;
Some(TensorView {
bytes: Cow::Borrowed(&t.bytes),
ggml_type: t.ggml_type,
ne: t.ne.clone(),
})
}
}
fn fixture_source(
config: &ModelConfig,
plan: &ModelPlan,
weights: &BTreeMap<TensorId, ReferenceTensor>,
) -> FixtureSource {
let contract = TensorContract::for_plan(
plan,
CheckpointDialect::Gguf,
ContractOptions {
output_head: OutputHead::TiedToEmbedding,
},
)
.expect("contract for the mini hyper-connections plan");
let mut tensors = BTreeMap::new();
for req in contract
.requirements
.iter()
.filter(|r| r.required || weights.contains_key(&r.id))
{
let tensor = weights
.get(&req.id)
.unwrap_or_else(|| panic!("reference fixture is missing {:?}", req.id));
let names = match req.match_mode {
TensorMatch::OneOf => &req.names[..1],
TensorMatch::All => req.names.as_slice(),
};
let is_expert = names.iter().any(|n| {
n.contains("ffn_gate_exps") || n.contains("ffn_up_exps") || n.contains("ffn_down_exps")
});
let (bytes, ggml_type) = if is_expert {
let row = req.shape[0] as usize;
assert!(
row.is_multiple_of(64),
"expert row {row} is not a multiple of 64; NVFP4 blocks by 64"
);
let mut out = Vec::new();
for chunk in tensor.data.chunks(row) {
out.extend_from_slice(&memra_gguf::nvfp4_repack::f32_to_nvfp4(chunk));
}
(out, GgmlType::NVFP4)
} else {
(
tensor.data.iter().flat_map(|v| v.to_le_bytes()).collect(),
GgmlType::F32,
)
};
for name in names {
tensors.insert(
name.clone(),
OwnedTensor {
bytes: bytes.clone(),
ne: req.shape.clone(),
ggml_type,
},
);
}
}
FixtureSource {
config: config.clone(),
tensors,
}
}
fn tokens(n: usize, seed: u64) -> Vec<u32> {
let mut s = seed | 1;
(0..n)
.map(|_| {
s = s
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((s >> 33) as u32) % VOCAB
})
.collect()
}
struct Harness {
engine: Engine,
model: memra_engine::hybrid::HybridModel,
plan: ModelPlan,
}
impl Harness {
fn new() -> Self {
force_true_f32();
let config = mini_config();
let plan = mini_plan(&config);
let fixture = deterministic_fixture(&plan).expect("deterministic hc fixture");
let source = fixture_source(&config, &plan, &fixture.weights);
let engine = Engine::new(0).expect("CUDA engine on device 0");
let model =
memra_engine::hybrid::HybridModel::load_from_source_without_mtp(&engine, &source)
.expect("mini hyper-connections model loads");
Self {
engine,
model,
plan,
}
}
fn reseat_first_recurrent_layer(
&self,
cache: &mut memra_engine::cache::Cache,
) -> Option<usize> {
let (il, rl) = cache
.recur
.iter_mut()
.enumerate()
.find_map(|(i, r)| r.as_mut().map(|r| (i, r)))?;
use cudarc::driver::DevicePtr;
let n = rl.ssm_state.len();
let mut fresh = self.engine.uninit(n).expect("fresh ssm_state");
self.engine
.copy_into(&mut fresh, 0, &rl.ssm_state, n)
.expect("copy the state into the fresh buffer");
{
let st = self.engine.stream();
let (old_p, _g0) = rl.ssm_state.device_ptr(&st);
let (alt_p, _g1) = rl.ssm_state_alt.device_ptr(&st);
let (new_p, _g2) = fresh.device_ptr(&st);
eprintln!("[gate] reseat il={il} ssm 0x{old_p:x} -> 0x{new_p:x} (alt 0x{alt_p:x})");
}
rl.ssm_state = fresh;
Some(il)
}
fn decode_bits(&self, ids: &[u32], prompt: usize, steps: usize) -> (Vec<Vec<u32>>, u64) {
self.decode_bits_reseat(ids, prompt, steps, None)
}
fn decode_bits_reseat(
&self,
ids: &[u32],
prompt: usize,
steps: usize,
reseat_at: Option<usize>,
) -> (Vec<Vec<u32>>, u64) {
let mut cache =
memra_engine::cache::Cache::new_planned(&self.engine, &self.model.cfg, &self.plan, 64)
.expect("cache for the mini hc model");
let (_primed, _seed, _hiddens) = self
.model
.prime_cache(&self.engine, &ids[..prompt], &mut cache, 0)
.expect("GPU hc prime");
let alloc0 = SCRATCH_ALLOC_CALLS.load(Ordering::Relaxed);
let mut out = Vec::with_capacity(steps);
for step in 0..steps {
if reseat_at == Some(step) {
let il = self.reseat_first_recurrent_layer(&mut cache);
eprintln!("[gate] forced re-seat at step {step} (layer {il:?})");
}
let logits = self
.model
.decode_step(&self.engine, ids[prompt + step], &mut cache)
.expect("GPU hc decode step");
assert!(
logits.iter().all(|v| v.is_finite()),
"step {step}: non-finite logits"
);
out.push(logits.iter().map(|v| v.to_bits()).collect());
}
let allocs = SCRATCH_ALLOC_CALLS.load(Ordering::Relaxed) - alloc0;
(out, allocs)
}
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn graph_door_decode_matches_eager_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(40, 0x5EED);
let (prompt, steps) = (8usize, 16usize);
unsafe {
std::env::set_var("MEMRA_HTOD_DIET", "1");
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
}
let cap_before_eager = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed);
let (eager, _) = h.decode_bits(&ids, prompt, steps);
assert_eq!(
memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed),
cap_before_eager,
"MEMRA_GLM5_DECODE_GRAPH=0 did not disarm the door: the eager arm captured a graph"
);
let cap0 = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed);
let rep0 = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "1");
}
let (graphed, _) = h.decode_bits(&ids, prompt, steps);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
}
let captures = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed) - cap0;
let replays = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed) - rep0;
let layers = memra_engine::GLM5_DECODE_GRAPH_LAYERS.load(Ordering::Relaxed);
println!("door: captures={captures} replays={replays} captured_layers={layers}");
assert!(
captures > 0 && replays > 0,
"VACUOUS: the door never captured ({captures}) or never replayed ({replays}); its \
refusal reason is on stderr as `[glm5-decode-graph] eager: ...`. This gate says nothing \
about capture unless capture happened"
);
let mut bad = 0usize;
for (step, (a, b)) in eager.iter().zip(&graphed).enumerate() {
let d = a.iter().zip(b).filter(|(x, y)| x != y).count();
if d > 0 {
if bad < 4 {
let i = a.iter().zip(b).position(|(x, y)| x != y).unwrap();
println!(
" step {step}: {d}/{} logits differ, first at {i}: eager {} graph {}",
a.len(),
f32::from_bits(a[i]),
f32::from_bits(b[i])
);
}
bad += 1;
}
}
assert_eq!(
bad, 0,
"{bad}/{steps} decode steps diverge between the eager walk and the captured/replayed one"
);
println!("graph door: {steps} steps bit-identical to eager ({replays} replays)");
let recaptures_before = memra_engine::GLM5_DECODE_GRAPH_RECAPTURES.load(Ordering::Relaxed);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "1");
std::env::set_var("MEMRA_GLM5_GRAPH_RECAPTURE", "1");
}
let (reseated, _) = h.decode_bits_reseat(&ids, prompt, steps, Some(steps / 2));
let recaptures =
memra_engine::GLM5_DECODE_GRAPH_RECAPTURES.load(Ordering::Relaxed) - recaptures_before;
unsafe {
std::env::set_var("MEMRA_GLM5_GRAPH_RECAPTURE", "0");
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
}
assert!(
recaptures > 0,
"VACUOUS RE-CAPTURE ARM: the forced re-seat at step {} did not rebuild any stage \
(recaptures={recaptures}). Either the re-seated layer is not inside a captured run, or \
the pointer signature no longer sees a re-seat.",
steps / 2
);
let mut bad_rs = 0usize;
for (step, (a, b)) in eager.iter().zip(&reseated).enumerate() {
if a != b {
if bad_rs < 4 {
let i = a.iter().zip(b).position(|(x, y)| x != y).unwrap_or(0);
println!(
" RE-SEAT step {step}: first differing logit {i}: eager {} graph {}",
f32::from_bits(a[i]),
f32::from_bits(b[i])
);
}
bad_rs += 1;
}
}
assert_eq!(
bad_rs, 0,
"{bad_rs}/{steps} steps diverge across a forced re-capture ({recaptures} rebuilds)"
);
println!("re-capture arm: {steps} steps bit-identical across {recaptures} forced rebuilds");
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn graph_door_mla_halves_match_eager_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(40, 0x5EED);
let (prompt, steps) = (8usize, 16usize);
unsafe {
std::env::set_var("MEMRA_HTOD_DIET", "1");
std::env::set_var("MEMRA_MLA_SEG_WS", "1");
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA", "0");
}
let (eager, _) = h.decode_bits(&ids, prompt, steps);
let cap0 = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed);
let rep0 = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed);
let halves0 = memra_engine::GLM5_DECODE_GRAPH_MLA_HALVES.load(Ordering::Relaxed);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "1");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA", "1");
}
let (graphed, _) = h.decode_bits(&ids, prompt, steps);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA", "0");
std::env::set_var("MEMRA_MLA_SEG_WS", "0");
}
let captures = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed) - cap0;
let replays = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed) - rep0;
let halves = memra_engine::GLM5_DECODE_GRAPH_MLA_HALVES.load(Ordering::Relaxed) - halves0;
println!("mla-halves door: captures={captures} replays={replays} mla_halves={halves}");
assert!(
captures > 0 && replays > 0,
"VACUOUS: the door never captured ({captures}) or never replayed ({replays})"
);
assert!(
halves > 0,
"VACUOUS: no MLA layer was captured in halves (the plan fell back to KDA-only runs); \
the fixture has MLA layers at 3 and 7"
);
let mut bad = 0usize;
for (step, (a, b)) in eager.iter().zip(&graphed).enumerate() {
let d = a.iter().zip(b).filter(|(x, y)| x != y).count();
if d > 0 {
if bad < 4 {
let i = a.iter().zip(b).position(|(x, y)| x != y).unwrap();
println!(
" step {step}: {d}/{} logits differ, first at {i}: eager {} graph {}",
a.len(),
f32::from_bits(a[i]),
f32::from_bits(b[i])
);
}
bad += 1;
}
}
assert_eq!(
bad, 0,
"{bad}/{steps} decode steps diverge between the eager walk and the MLA-halves captured one"
);
println!(
"mla-halves door: {steps} steps bit-identical to eager ({replays} replays, {halves} halves)"
);
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn graph_door_mla_mid_match_eager_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(40, 0x5EED);
let (prompt, steps) = (8usize, 16usize);
unsafe {
std::env::set_var("MEMRA_HTOD_DIET", "1");
std::env::set_var("MEMRA_MLA_SEG_WS", "1");
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA", "0");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA_MID", "0");
}
let (eager, _) = h.decode_bits(&ids, prompt, steps);
let mids0 = memra_engine::GLM5_DECODE_GRAPH_MLA_MIDS.load(Ordering::Relaxed);
let cap0 = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed);
let rep0 = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed);
let halves0 = memra_engine::GLM5_DECODE_GRAPH_MLA_HALVES.load(Ordering::Relaxed);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "1");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA", "1");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA_MID", "1");
}
let (graphed, _) = h.decode_bits(&ids, prompt, steps);
unsafe {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "0");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA", "0");
std::env::set_var("MEMRA_GLM5_GRAPH_MLA_MID", "0");
std::env::set_var("MEMRA_MLA_SEG_WS", "0");
}
let mids = memra_engine::GLM5_DECODE_GRAPH_MLA_MIDS.load(Ordering::Relaxed) - mids0;
let captures = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed) - cap0;
let replays = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed) - rep0;
let halves = memra_engine::GLM5_DECODE_GRAPH_MLA_HALVES.load(Ordering::Relaxed) - halves0;
println!(
"mla-mid door: captures={captures} replays={replays} mla_halves={halves} mla_mids={mids}"
);
assert!(
captures > 0 && replays > 0,
"VACUOUS: the door never captured ({captures}) or never replayed ({replays})"
);
assert!(
mids > 0,
"VACUOUS: no MLA middle was captured (mla_mids={mids}, mla_halves={halves}); the plan \
refused the live middle on the fixture's MLA layers 3 and 7"
);
let mut bad = 0usize;
for (step, (a, b)) in eager.iter().zip(&graphed).enumerate() {
let d = a.iter().zip(b).filter(|(x, y)| x != y).count();
if d > 0 {
if bad < 4 {
let i = a.iter().zip(b).position(|(x, y)| x != y).unwrap();
println!(
" step {step}: {d}/{} logits differ, first at {i}: eager {} graph {}",
a.len(),
f32::from_bits(a[i]),
f32::from_bits(b[i])
);
}
bad += 1;
}
}
assert_eq!(
bad, 0,
"{bad}/{steps} decode steps diverge between the eager walk and the whole-MLA captured one"
);
println!(
"mla-mid door: {steps} steps bit-identical to eager ({replays} replays, {mids} live \
middles, {halves} halves)"
);
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn mla_rows_live_matches_rows_exact_bitwise() {
use memra_engine::hybrid::Mixer;
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(40, 0x5EED);
let (prompt, t, rounds) = (8usize, 3usize, 4usize);
unsafe {
std::env::set_var("MEMRA_MLA_SEG_WS", "1");
}
let e = &h.engine;
let prime = |_: usize| {
let mut cache =
memra_engine::cache::Cache::new_planned(e, &h.model.cfg, &h.plan, 64).expect("cache");
h.model
.prime_cache(e, &ids[..prompt], &mut cache, 0)
.expect("prime");
cache
};
let mut ca = prime(0);
let mut cb = prime(1);
let n_embd = h.model.cfg.n_embd as usize;
let il = (0..h.model.layers.len())
.find(|&i| matches!(h.model.layers[i].mixer, Mixer::Mla(_)))
.expect("the fixture carries an MLA layer");
let Mixer::Mla(mla) = &h.model.layers[il].mixer else {
unreachable!()
};
let pool = mla.index.as_ref().expect("fixture indexer").geom.pool;
let mut pos = ca.pos;
for round in 0..rounds {
let x: Vec<f32> = (0..t * n_embd)
.map(|i| ((i * 7919 + round * 104729) % 1000) as f32 / 500.0 - 1.0)
.collect();
let hin = e.htod(&x).unwrap();
let pos_all: Vec<i32> = (0..t).map(|r| (pos + r) as i32).collect();
let pos_d = e.htod_i32(&pos_all).unwrap();
let ya = h
.model
.mla_attn_cached_rows_exact(e, mla, &hin, &pos_d, t, il, &mut ca)
.expect("rows-exact");
let yb = h
.model
.mla_attn_cached_rows_live(e, mla, &hin, &pos_d, t, il, &mut cb)
.expect("rows-live");
e.stream().synchronize().unwrap();
{
let lb = cb.latent[il].as_mut().unwrap();
lb.len += t;
lb.index_pools_ready = lb.len / pool;
}
let (va, vb) = (e.dtoh(&ya).unwrap(), e.dtoh(&yb).unwrap());
let d = va
.iter()
.zip(&vb)
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
assert_eq!(
d,
0,
"round {round} (pos {pos}): {d}/{} output words differ",
va.len()
);
let (la, lb) = (
ca.latent[il].as_ref().unwrap(),
cb.latent[il].as_ref().unwrap(),
);
assert_eq!(la.len, lb.len, "round {round}: latent len");
assert_eq!(
la.index_pools_ready, lb.index_pools_ready,
"round {round}: pools_ready"
);
assert_eq!(
e.dtoh_i32(&la.len_d).unwrap(),
e.dtoh_i32(&lb.len_d).unwrap(),
"round {round}: len_d"
);
let n = la.len * la.width;
let (ra, rb) = (e.dtoh(&la.rows).unwrap(), e.dtoh(&lb.rows).unwrap());
assert!(
ra[..n]
.iter()
.zip(&rb[..n])
.all(|(a, b)| a.to_bits() == b.to_bits()),
"round {round}: latent rows differ"
);
if let (Some(ka), Some(kb)) = (la.index_pool_keys.as_ref(), lb.index_pool_keys.as_ref()) {
let m = la.index_pools_ready * mla.index.as_ref().unwrap().geom.head_dim;
let (pa, pb) = (e.dtoh(ka).unwrap(), e.dtoh(kb).unwrap());
assert!(
pa[..m]
.iter()
.zip(&pb[..m])
.all(|(a, b)| a.to_bits() == b.to_bits()),
"round {round}: pool keys differ"
);
}
pos += t;
}
unsafe {
std::env::set_var("MEMRA_MLA_SEG_WS", "0");
}
println!(
"mla rows-live: {rounds} rounds of t={t} rows bit-identical to rows-exact (layer {il}, pool {pool}, positions {prompt}..{pos})"
);
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn spec_dev_io_verify_rows_match_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(40, 0x5EED);
let (prompt, t) = (8usize, 4usize);
let e = &h.engine;
let prime = || {
let mut cache =
memra_engine::cache::Cache::new_planned(e, &h.model.cfg, &h.plan, 64).expect("cache");
h.model
.prime_cache(e, &ids[..prompt], &mut cache, 0)
.expect("prime");
cache
};
let rows = &ids[prompt..prompt + t];
unsafe {
std::env::set_var("MEMRA_GLM5_SPEC_DEV_IO", "0");
}
let mut ca = prime();
let (la, _, _) = h
.model
.glm5_verify_rows(e, rows, &mut ca)
.expect("verify rows (copy arm)");
let c0 = memra_engine::SPEC_DEV_IO_AVOIDED.load(Ordering::Relaxed);
unsafe {
std::env::set_var("MEMRA_GLM5_SPEC_DEV_IO", "1");
}
let mut cb = prime();
let (lb, _, _) = h
.model
.glm5_verify_rows(e, rows, &mut cb)
.expect("verify rows (launch arm)");
unsafe {
std::env::set_var("MEMRA_GLM5_SPEC_DEV_IO", "0");
}
e.stream().synchronize().unwrap();
let avoided = memra_engine::SPEC_DEV_IO_AVOIDED.load(Ordering::Relaxed) - c0;
assert!(
avoided >= 2,
"VACUOUS: the door replaced {avoided} copies (expected positions + rows)"
);
let (va, vb) = (e.dtoh(&la).unwrap(), e.dtoh(&lb).unwrap());
assert_eq!(va.len(), vb.len());
let d = va
.iter()
.zip(&vb)
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
assert_eq!(
d,
0,
"{d}/{} verify logits differ between the copy and launch arms",
va.len()
);
println!(
"spec-dev-io: verify rows t={t} bitwise across arms ({avoided} pageable copies replaced)"
);
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn verify_snapshot_reuse_matches_fresh_clones_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(48, 0x5EED);
let (prompt, t) = (8usize, 3usize);
let e = &h.engine;
let prime = || {
let mut cache =
memra_engine::cache::Cache::new_planned(e, &h.model.cfg, &h.plan, 64).expect("cache");
h.model
.prime_cache(e, &ids[..prompt], &mut cache, 0)
.expect("prime");
cache
};
let (mut ca, mut cb) = (prime(), prime());
let mut pool: Vec<Option<cudarc::driver::CudaSlice<f32>>> = Vec::new();
for round in 0..3 {
let rows = &ids[prompt + round * t..prompt + (round + 1) * t];
let (la, _, ca_ck) = h.model.glm5_verify_rows(e, rows, &mut ca).expect("fresh");
let (lb, _, mut cb_ck) = h
.model
.glm5_verify_rows_reusing(e, rows, &mut cb, std::mem::take(&mut pool))
.expect("reusing");
e.stream().synchronize().unwrap();
let (va, vb) = (e.dtoh(&la).unwrap(), e.dtoh(&lb).unwrap());
let d = va
.iter()
.zip(&vb)
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
assert_eq!(d, 0, "round {round}: {d}/{} verify logits differ", va.len());
h.model
.glm5_verify_rollback(e, &mut ca, &ca_ck, 2)
.expect("rollback A");
h.model
.glm5_verify_rollback(e, &mut cb, &cb_ck, 2)
.expect("rollback B");
pool = std::mem::take(&mut cb_ck.kda_ssm_snap_buffers());
e.stream().synchronize().unwrap();
for il in 0..h.model.layers.len() {
if let (Some(ra), Some(rb)) = (ca.recur[il].as_ref(), cb.recur[il].as_ref()) {
let (sa, sb) = (
e.dtoh(&ra.ssm_state).unwrap(),
e.dtoh(&rb.ssm_state).unwrap(),
);
assert!(
sa.iter().zip(&sb).all(|(a, b)| a.to_bits() == b.to_bits()),
"round {round} layer {il}: ssm state differs after rollback"
);
}
}
}
assert!(
pool.iter().any(|b| b.is_some()),
"VACUOUS: no snapshot buffer was carried"
);
println!(
"verify snapshot reuse: 3 rounds of t={t} bitwise (logits + rolled-back ssm state) vs fresh clones"
);
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn verify_graph_live_arm_matches_rows_exact_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(48, 0x5EED);
let (prompt, t) = (8usize, 3usize);
let e = &h.engine;
let prime = || {
let mut cache =
memra_engine::cache::Cache::new_planned(e, &h.model.cfg, &h.plan, 64).expect("cache");
h.model
.prime_cache(e, &ids[..prompt], &mut cache, 0)
.expect("prime");
cache
};
unsafe {
std::env::set_var("MEMRA_MLA_SEG_WS", "1");
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "0");
}
let (mut ca, mut cb) = (prime(), prime());
let c0 = memra_engine::GLM5_VERIFY_LIVE_MLA_CALLS.load(Ordering::Relaxed);
for round in 0..3 {
let rows = &ids[prompt + round * t..prompt + (round + 1) * t];
unsafe {
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "0");
}
let (la, _, ca_ck) = h
.model
.glm5_verify_rows(e, rows, &mut ca)
.expect("rows-exact arm");
unsafe {
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "1");
}
let (lb, _, cb_ck) = h
.model
.glm5_verify_rows(e, rows, &mut cb)
.expect("live arm");
unsafe {
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "0");
}
e.stream().synchronize().unwrap();
let (va, vb) = (e.dtoh(&la).unwrap(), e.dtoh(&lb).unwrap());
let d = va
.iter()
.zip(&vb)
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
assert_eq!(d, 0, "round {round}: {d}/{} verify logits differ", va.len());
for il in 0..h.model.layers.len() {
if let (Some(pa), Some(pb)) = (ca.latent[il].as_ref(), cb.latent[il].as_ref()) {
assert_eq!(pa.len, pb.len, "round {round} layer {il}: latent len");
assert_eq!(
pa.index_pools_ready, pb.index_pools_ready,
"round {round} layer {il}: pools_ready"
);
assert_eq!(
e.dtoh_i32(&pa.len_d).unwrap(),
e.dtoh_i32(&pb.len_d).unwrap(),
"round {round} layer {il}: len_d"
);
let n = pa.len * pa.width;
let (ra, rb) = (e.dtoh(&pa.rows).unwrap(), e.dtoh(&pb.rows).unwrap());
assert!(
ra[..n]
.iter()
.zip(&rb[..n])
.all(|(a, b)| a.to_bits() == b.to_bits()),
"round {round} layer {il}: latent rows differ"
);
}
}
h.model
.glm5_verify_rollback(e, &mut ca, &ca_ck, 2)
.expect("rollback A");
h.model
.glm5_verify_rollback(e, &mut cb, &cb_ck, 2)
.expect("rollback B");
}
unsafe {
std::env::set_var("MEMRA_MLA_SEG_WS", "0");
}
let live = memra_engine::GLM5_VERIFY_LIVE_MLA_CALLS.load(Ordering::Relaxed) - c0;
assert!(
live >= 3,
"VACUOUS: {live} live MLA calls (expected one per MLA layer per round)"
);
println!("verify-graph arm 1: 3 rounds of t={t} bitwise vs rows-exact ({live} live MLA calls)");
}
#[test]
#[ignore = "needs a CUDA device — run under flock /tmp/memra-5090.lock"]
fn verify_graph_replays_match_eager_bitwise() {
let _gpu = gpu_guard();
let h = Harness::new();
let ids = tokens(64, 0x5EED);
let (prompt, t, rounds) = (8usize, 3usize, 6usize);
let e = &h.engine;
let prime = || {
let mut cache =
memra_engine::cache::Cache::new_planned(e, &h.model.cfg, &h.plan, 64).expect("cache");
h.model
.prime_cache(e, &ids[..prompt], &mut cache, 0)
.expect("prime");
cache
};
unsafe {
std::env::set_var("MEMRA_MLA_SEG_WS", "1");
std::env::set_var("MEMRA_MOE_VROWS_DEV_TABLES", "1");
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "0");
}
let (mut ca, mut cb) = (prime(), prime());
let mut pool = memra_engine::glm_spec::VerifyGraphPool::default();
let mut snaps: Vec<Option<cudarc::driver::CudaSlice<f32>>> = Vec::new();
let cap0 = memra_engine::GLM5_VERIFY_GRAPH_CAPTURES.load(Ordering::Relaxed);
let rep0 = memra_engine::GLM5_VERIFY_GRAPH_REPLAYS.load(Ordering::Relaxed);
let sc0 = memra_engine::GLM5_VERIFY_GRAPH_SELFCHECK_FAILS.load(Ordering::Relaxed);
for round in 0..rounds {
let rows = &ids[prompt + round * 2..prompt + round * 2 + t];
unsafe {
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "0");
}
let (la, _, ca_ck) = h
.model
.glm5_verify_rows(e, rows, &mut ca)
.expect("eager arm");
unsafe {
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "1");
}
let (lb, _, mut cb_ck) = h
.model
.glm5_verify_rows_graphed(
e,
rows,
&mut cb,
std::mem::take(&mut snaps),
Some(&mut pool),
)
.expect("graphed arm");
e.stream().synchronize().unwrap();
let (va, vb) = (e.dtoh(&la).unwrap(), e.dtoh(&lb).unwrap());
let d = va
.iter()
.zip(&vb)
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
assert_eq!(
d,
0,
"round {round}: {d}/{} verify logits differ (captures so far {})",
va.len(),
memra_engine::GLM5_VERIFY_GRAPH_CAPTURES.load(Ordering::Relaxed) - cap0
);
for il in 0..h.model.layers.len() {
if let (Some(pa), Some(pb)) = (ca.latent[il].as_ref(), cb.latent[il].as_ref()) {
assert_eq!(pa.len, pb.len, "round {round} layer {il}: latent len");
assert_eq!(
pa.index_pools_ready, pb.index_pools_ready,
"round {round} layer {il}: pools_ready"
);
assert_eq!(
e.dtoh_i32(&pa.len_d).unwrap(),
e.dtoh_i32(&pb.len_d).unwrap(),
"round {round} layer {il}: len_d"
);
}
}
h.model
.glm5_verify_rollback(e, &mut ca, &ca_ck, 2)
.expect("rollback A");
h.model
.glm5_verify_rollback(e, &mut cb, &cb_ck, 2)
.expect("rollback B");
snaps = cb_ck.kda_ssm_snap_buffers();
pool.reclaim_rows(e, &mut cb_ck);
println!(
" round {round}: pool after reclaim {:?}",
e.vws_pool_state()
);
e.stream().synchronize().unwrap();
for il in 0..h.model.layers.len() {
if let (Some(ra), Some(rb)) = (ca.recur[il].as_ref(), cb.recur[il].as_ref()) {
let (sa, sb) = (
e.dtoh(&ra.ssm_state).unwrap(),
e.dtoh(&rb.ssm_state).unwrap(),
);
assert!(
sa.iter().zip(&sb).all(|(a, b)| a.to_bits() == b.to_bits()),
"round {round} layer {il}: ssm state differs after rollback"
);
}
}
}
unsafe {
std::env::set_var("MEMRA_GLM5_VERIFY_GRAPH", "0");
std::env::set_var("MEMRA_MLA_SEG_WS", "0");
std::env::set_var("MEMRA_MOE_VROWS_DEV_TABLES", "0");
}
let captures = memra_engine::GLM5_VERIFY_GRAPH_CAPTURES.load(Ordering::Relaxed) - cap0;
let replays = memra_engine::GLM5_VERIFY_GRAPH_REPLAYS.load(Ordering::Relaxed) - rep0;
let scf = memra_engine::GLM5_VERIFY_GRAPH_SELFCHECK_FAILS.load(Ordering::Relaxed) - sc0;
println!(
"verify-graph arm 2: captures={captures} replays={replays} selfcheck_fails={scf} pool_replays={}",
pool.replays()
);
assert_eq!(scf, 0, "self-check failures");
assert!(
captures >= 1,
"VACUOUS: the pool never captured (refused? see the announce line)"
);
assert!(
replays >= rounds as u64 - 2,
"VACUOUS: {replays} replays for {rounds} rounds"
);
println!(
"verify-graph arm 2: {rounds} rounds of t={t} bitwise vs eager ({captures} captures, {replays} replays)"
);
}