use memra_engine::Engine;
use memra_engine::forward::argmax;
use memra_engine::glm_spec::Glm5SpecKnobs;
use memra_engine::hybrid::HybridModel;
use memra_gguf::GgmlType;
use memra_gguf::config::{HfConfig, ModelConfig};
use memra_gguf::model_plan::{ModelPlan, StatePlan};
use memra_gguf::source::{TensorSource, TensorView};
use memra_gguf::tensor_contract::{
CheckpointDialect, ContractOptions, LayerTensor, MtpTensor, OutputHead, TensorContract,
TensorId, TensorMatch,
};
use memra_reference::{ReferenceTensor, ReferenceWeights, deterministic_fixture};
use std::borrow::Cow;
use std::collections::BTreeMap;
const HIDDEN: usize = 128;
const VOCAB: u32 = 32;
const K: usize = 7;
fn mini_config_json() -> String {
r#"{
"model_type": "glm5_next_text",
"num_hidden_layers": 4,
"num_nextn_predict_layers": 1,
"hidden_size": 128,
"intermediate_size": 64,
"vocab_size": 32,
"max_position_embeddings": 512,
"rms_norm_eps": 1e-05,
"hidden_act": "silu",
"swiglu_limit": 10.0,
"tie_word_embeddings": true,
"hc_mult": 4,
"hc_eps": 1e-06,
"hc_sinkhorn_iters": 20,
"mhc": true,
"layer_types": ["linear_attention", "deepseek_sparse_attention",
"linear_attention", "deepseek_sparse_attention"],
"mlp_layer_types": ["dense", "sparse", "sparse", "sparse"],
"first_k_dense_replace": 1,
"indexer_types": ["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, 2],
"full_attn_layers": [1, 3]
},
"num_attention_heads": 2,
"num_key_value_heads": 2,
"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": 4,
"num_experts_per_tok": 2,
"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 varied(len: usize, seed: u64, spread: f32) -> Vec<f32> {
(0..len)
.map(|i| {
let x = (i as u64)
.wrapping_mul(6364136223846793005)
.wrapping_add(seed)
.rotate_left(17) as f64
/ u64::MAX as f64;
1.0 + spread * (x as f32 - 0.5)
})
.collect()
}
fn fixture_weights(plan: &ModelPlan) -> ReferenceWeights {
let mut weights = deterministic_fixture(plan)
.expect("deterministic glm5 hc+mtp fixture")
.weights;
for (tensor, seed) in [
(MtpTensor::EmbeddingNorm, 0xE0_12u64),
(MtpTensor::HiddenNorm, 0x40_77),
(MtpTensor::OutputNorm, 0x5EAD),
] {
weights.insert(
TensorId::Mtp { depth: 0, tensor },
ReferenceTensor::new(vec![HIDDEN], varied(HIDDEN, seed, 0.8)).unwrap(),
);
}
weights
}
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 is_expert_bank(id: &TensorId) -> bool {
matches!(
id,
TensorId::Layer {
tensor: LayerTensor::MoeExpertGateBank
| LayerTensor::MoeExpertUpBank
| LayerTensor::MoeExpertDownBank,
..
}
)
}
fn fixture_source(config: &ModelConfig, plan: &ModelPlan) -> FixtureSource {
let weights = fixture_weights(plan);
let contract = TensorContract::for_plan(
plan,
CheckpointDialect::Gguf,
ContractOptions {
output_head: OutputHead::TiedToEmbedding,
},
)
.expect("contract for the mini glm5_next hc+mtp plan");
let mut tensors = BTreeMap::new();
let mut bank_stems: Vec<String> = Vec::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 elements: usize = req.shape.iter().map(|&d| d as usize).product();
assert_eq!(
elements,
tensor.data.len(),
"shape mismatch for {:?}",
req.id
);
let expert_bank = is_expert_bank(&req.id);
let (bytes, ggml_type) = if expert_bank {
(
memra_gguf::nvfp4_repack::f32_to_nvfp4(&tensor.data),
GgmlType::NVFP4,
)
} else {
(
tensor.data.iter().flat_map(|v| v.to_le_bytes()).collect(),
GgmlType::F32,
)
};
let names = match req.match_mode {
TensorMatch::OneOf => &req.names[..1],
TensorMatch::All => req.names.as_slice(),
};
for name in names {
if expert_bank && let Some(stem) = name.strip_suffix(".weight") {
bank_stems.push(stem.to_string());
}
tensors.insert(
name.clone(),
OwnedTensor {
bytes: bytes.clone(),
ne: req.shape.clone(),
ggml_type,
},
);
}
}
let n_expert = config
.moe
.as_ref()
.expect("mini fixture carries an MoE block")
.expert_count as usize;
let macros: Vec<f32> = (0..n_expert)
.map(|e| {
if e < n_expert / 2 {
0.5 + 0.1 * e as f32
} else {
1.2 + 0.1 * (e - n_expert / 2) as f32
}
})
.collect();
assert!(!bank_stems.is_empty(), "no routed expert banks collected");
for stem in &bank_stems {
tensors.insert(
format!("{stem}.scale"),
OwnedTensor {
bytes: macros.iter().flat_map(|v| v.to_le_bytes()).collect(),
ne: vec![n_expert as u64],
ggml_type: GgmlType::F32,
},
);
}
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()
}
fn bit_diffs(a: &[f32], b: &[f32]) -> usize {
assert_eq!(a.len(), b.len());
a.iter()
.zip(b)
.filter(|(x, y)| x.to_bits() != y.to_bits())
.count()
}
fn fresh_primed(
e: &Engine,
m: &HybridModel,
plan: &ModelPlan,
prompt: &[u32],
max_ctx: usize,
) -> (memra_engine::cache::Cache, Vec<f32>) {
let mut cache = memra_engine::pp::new_cache_planned(e, &m.cfg, plan, max_ctx)
.expect("cache for the mini glm5 model");
let (logits, _seed, _hiddens) = m.prime_cache(e, prompt, &mut cache, 0).expect("hc prime");
(cache, logits)
}
fn plain_tape(
e: &Engine,
m: &HybridModel,
plan: &ModelPlan,
prompt: &[u32],
max_new: usize,
) -> Vec<u32> {
let (mut cache, logits) = fresh_primed(e, m, plan, prompt, prompt.len() + max_new + 16);
let mut tape = Vec::with_capacity(max_new);
tape.push(argmax(&logits) as u32);
while tape.len() < max_new {
let ll = m
.decode_step(e, *tape.last().unwrap(), &mut cache)
.expect("plain decode step");
tape.push(argmax(&ll) as u32);
}
tape
}
struct Verdicts {
fails: usize,
}
impl Verdicts {
fn arm(&mut self, name: &str, ok: bool, detail: &str) {
if ok {
println!("glm5-spec-ppn gate PASS [{name}]: {detail}");
} else {
println!("glm5-spec-ppn gate FAIL [{name}]: {detail}");
self.fails += 1;
}
}
}
#[allow(clippy::too_many_lines)]
fn main() -> Result<(), Box<dyn std::error::Error>> {
let stages: usize = std::env::args()
.nth(1)
.and_then(|s| s.parse().ok())
.unwrap_or(2);
let p: usize = std::env::args()
.nth(2)
.and_then(|s| s.parse().ok())
.unwrap_or(24);
let n: usize = std::env::args()
.nth(3)
.and_then(|s| s.parse().ok())
.unwrap_or(20);
if std::env::var("NVIDIA_TF32_OVERRIDE").as_deref() != Ok("0") {
unsafe { std::env::set_var("NVIDIA_TF32_OVERRIDE", "0") };
}
unsafe {
std::env::set_var("MEMRA_PP_STAGES", stages.to_string());
std::env::set_var("MEMRA_GLM5_MTP", "1");
}
let devices_env = std::env::var("MEMRA_PP_DEVICES").unwrap_or_default();
let primary_dev: usize = devices_env
.split(',')
.next()
.and_then(|s| s.trim().parse().ok())
.unwrap_or(0);
let knobs = format!(
"stages={stages} streams={} overlap={} devices={} splits={} shard={}",
if memra_engine::pp::pp2_streams_off() {
"OFF(same-stream seam)"
} else {
"per-stage"
},
if memra_engine::pp::pp2_overlap() {
"1(double-buffered)"
} else {
"0"
},
if devices_env.is_empty() {
"default(primary)"
} else {
&devices_env
},
std::env::var("MEMRA_PP_SPLITS").unwrap_or_else(|_| "default(even)".into()),
if memra_engine::pp::pp_shard_off() {
"OFF(bring-up placement)"
} else {
"per-stage"
},
);
println!("glm5-spec-ppn-gate config: {knobs}");
let config = ModelConfig::from_hf(&HfConfig::parse(&mini_config_json()));
let plan = 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");
let e = Engine::new(primary_dev)?;
let m = HybridModel::load_from_source(&e, &fixture_source(&config, &plan))?;
let topology = m
.hyper
.as_ref()
.expect("the fixture must load as a HyperConnections trunk");
assert!(
m.mtp.is_some(),
"the NextN draft head must load (MEMRA_GLM5_MTP=1)"
);
println!(
"hc topology: streams={} collapse={:?}",
topology.streams, topology.collapse
);
let n_layers = m.layers.len();
let fence = memra_engine::pp::pp_cuts(n_layers).unwrap_or_else(|| {
panic!("ppn door failed to open (n_layers={n_layers}, stages={stages})")
});
assert_eq!(
fence.len() - 1,
stages,
"fence {fence:?} != stages {stages}"
);
for w in fence.windows(2) {
assert!(w[1] > w[0], "fence {fence:?} leaves an empty stage");
}
let mut recur_stages: Vec<usize> = Vec::new();
let mut latent_stages: Vec<usize> = Vec::new();
for layer in &plan.layers {
let s = memra_engine::pp::stage_of(&fence, layer.index as usize);
match layer.state {
StatePlan::Recurrent { .. } => recur_stages.push(s),
StatePlan::LatentKvCache { .. } => latent_stages.push(s),
_ => {}
}
}
println!(
"stage fence: {fence:?}; Recurrent(KDA) on stages {recur_stages:?}, \
LatentKvCache(MLA+kpool) on stages {latent_stages:?}"
);
let spread = recur_stages
.iter()
.chain(latent_stages.iter())
.collect::<std::collections::BTreeSet<_>>();
assert!(
spread.len() >= 2,
"fence {fence:?} puts every stateful layer on one stage — the per-stage rollback \
contract would go untested; choose a split that separates them"
);
let prompt = tokens(p, 0xA11CE);
let vt = tokens(K + 1, 0xBEEF);
let cc = tokens(12, 0xC0FFEE);
let max_ctx = p + K + cc.len() + 16;
eprintln!("[phase] references: door OFF");
unsafe { std::env::remove_var("MEMRA_PP_STAGES") };
assert!(
memra_engine::pp::pp_cuts(n_layers).is_none(),
"the reference phase must run with the door SHUT"
);
let mut ref_rows: Vec<Vec<f32>> = Vec::with_capacity(vt.len());
{
let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
for &tok in &vt {
ref_rows.push(m.decode_step(&e, tok, &mut cache)?);
}
}
let n_vocab = ref_rows[0].len();
let mut ref_cont: Vec<Vec<Vec<f32>>> = Vec::with_capacity(K + 1);
for j in 0..=K {
let keep = j + 1;
let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
for &tok in &vt[..keep] {
let _ = m.decode_step(&e, tok, &mut cache)?;
}
let mut rows = Vec::with_capacity(cc.len());
for &tok in &cc {
rows.push(m.decode_step(&e, tok, &mut cache)?);
}
ref_cont.push(rows);
}
let tape = plain_tape(&e, &m, &plan, &prompt, n);
unsafe { std::env::set_var("MEMRA_PP_STAGES", stages.to_string()) };
let mut v = Verdicts { fails: 0 };
eprintln!("[phase] arm W0: door ON, plain ppN decode re-pin");
{
let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
let mut bad = 0usize;
for (r, &tok) in vt.iter().enumerate() {
let got = m.decode_step(&e, tok, &mut cache)?;
let diffs = bit_diffs(&got, &ref_rows[r]);
if diffs > 0 {
bad += 1;
println!("[W0] row {r}: {diffs}/{n_vocab} logits differ");
}
}
v.arm(
"W0 plain-ppn-decode",
bad == 0,
&format!(
"{} rows bit-identical vs door-OFF (n_vocab={n_vocab})",
vt.len()
),
);
}
eprintln!("[phase] arm W1: door ON, verify walk (ppN twin)");
{
let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
let (vlogits, _collapsed, _ckpt) = m.glm5_verify_rows(&e, &vt, &mut cache)?;
let host = e.dtoh(&vlogits)?;
let mut bad = 0usize;
for (r, plain) in ref_rows.iter().enumerate() {
let diffs = bit_diffs(&host[r * n_vocab..(r + 1) * n_vocab], plain);
if diffs > 0 {
bad += 1;
println!("[W1] row {r}: {diffs}/{n_vocab} logits differ");
}
}
v.arm(
"W1 verify-walk",
bad == 0,
&format!(
"{} verify rows bit-identical to plain decode under the split (and to \
plain ppN decode via W0)",
vt.len()
),
);
}
eprintln!("[phase] arm A: door ON, accept-j rollback battery");
#[allow(clippy::needless_range_loop)]
for j in 0..=K {
let keep = j + 1;
let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
let pos0 = cache.pos;
let (_vl, _coll, ckpt) = m.glm5_verify_rows(&e, &vt, &mut cache)?;
m.glm5_verify_rollback(&e, &mut cache, &ckpt, keep)?;
assert_eq!(
cache.pos,
pos0 + keep,
"rollback must land pos at snap+keep"
);
let mut bad = 0usize;
for (step, &tok) in cc.iter().enumerate() {
let got = m.decode_step(&e, tok, &mut cache)?;
let diffs = bit_diffs(&got, &ref_cont[j][step]);
if diffs > 0 {
bad += 1;
println!("[A] j={j} continue step {step}: {diffs}/{n_vocab} logits differ");
}
}
v.arm(
&format!("A accept-j={j}"),
bad == 0,
&format!("{} continue steps bit-identical", cc.len()),
);
}
eprintln!("[phase] arm E: door ON, e2e spec-vs-plain tapes");
for k in 1..=K {
let (out, drafted, accepted) = m.generate_spec_glm5(&e, &prompt, n, k)?;
v.arm(
&format!("E natural K={k}"),
out == tape,
&format!("tape identical to door-OFF plain greedy ({accepted}/{drafted} accepted)"),
);
}
for k in [3usize, K] {
let tape_for_override = tape.clone();
let mut over = move |round: usize, ki: usize, _greedy: u32| -> u32 {
let cursor = 1 + round * (k + 1);
let pos = cursor + ki;
if pos < tape_for_override.len() {
tape_for_override[pos]
} else {
0
}
};
let (out, drafted, accepted) = m.generate_spec_glm5_gated(
&e,
&prompt,
n,
k,
Glm5SpecKnobs {
draft_override: Some(&mut over),
..Default::default()
},
)?;
let plumbing_live = accepted * 2 >= drafted;
v.arm(
&format!("E forced-accept K={k}"),
out == tape && plumbing_live,
&format!("tape identical, {accepted}/{drafted} accepted (full-accept path exercised)"),
);
}
{
let k = K;
let tape_for_override = tape.clone();
let committed_before = move |round: usize| -> usize {
let mut c = 1usize;
for r in 0..round {
c += (r % k) + 1;
}
c
};
let mut over = move |round: usize, ki: usize, _greedy: u32| -> u32 {
let j_target = round % k;
let cursor = committed_before(round);
let pos = cursor + ki;
let correct = if pos < tape_for_override.len() {
tape_for_override[pos]
} else {
0
};
if ki < j_target {
correct
} else {
(correct + 1) % VOCAB
}
};
let (out, drafted, accepted) = m.generate_spec_glm5_gated(
&e,
&prompt,
n,
k,
Glm5SpecKnobs {
draft_override: Some(&mut over),
..Default::default()
},
)?;
v.arm(
&format!("E forced-rejection sweep K={k}"),
out == tape,
&format!("every partial-keep rollback cycled, tape identical ({accepted}/{drafted})"),
);
}
eprintln!("[phase] arm R1: RED stale-KDA state");
{
let j = 2usize;
let keep = j + 1;
let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
let (_vl, _coll, ckpt) = m.glm5_verify_rows(&e, &vt, &mut cache)?;
let streams_off = memra_engine::pp::pp2_streams_off();
let rt = if streams_off {
None
} else {
Some(memra_engine::pp::PpNRt::get(&e)?)
};
#[allow(clippy::type_complexity)]
let with_stage_engine =
|il: usize,
f: &mut dyn FnMut(&Engine) -> Result<(), Box<dyn std::error::Error>>|
-> Result<(), Box<dyn std::error::Error>> {
match rt {
Some(rt) => {
let s = memra_engine::pp::stage_of(&fence, il);
let _g = rt.enter(s);
f(rt.engine(s, &e))
}
None => f(&e),
}
};
let mut stale: Vec<
Option<(
cudarc::driver::CudaSlice<f32>,
cudarc::driver::CudaSlice<f32>,
)>,
> = Vec::new();
for il in 0..n_layers {
let mut cloned = None;
if cache.recur[il].is_some() {
with_stage_engine(il, &mut |es| {
let rl = cache.recur[il].as_ref().unwrap();
cloned = Some((
es.clone_dtod(&rl.conv_state)?,
es.clone_dtod(&rl.ssm_state)?,
));
Ok(())
})?;
}
stale.push(cloned);
}
m.glm5_verify_rollback(&e, &mut cache, &ckpt, keep)?;
let mut mutated = 0usize;
for (il, s) in stale.into_iter().enumerate() {
if let Some((conv, ssm)) = s {
with_stage_engine(il, &mut |es| {
let rl = cache.recur[il].as_mut().unwrap();
es.copy_into(&mut rl.conv_state, 0, &conv, conv.len())?;
es.copy_into(&mut rl.ssm_state, 0, &ssm, ssm.len())?;
Ok(())
})?;
mutated += 1;
}
}
assert!(
mutated > 0,
"the mutation must touch at least one KDA layer"
);
let mut diffs_total = 0usize;
for (step, &tok) in cc.iter().enumerate() {
let got = m.decode_step(&e, tok, &mut cache)?;
diffs_total += bit_diffs(&got, &ref_cont[j][step]);
}
v.arm(
"R1 stale-KDA RED",
diffs_total > 0,
&format!("bites — {diffs_total} differing logits across the continuation"),
);
}
eprintln!("[phase] arm R2: RED pool-key clamp");
{
let keep = 1usize; let (mut cache, _l) = fresh_primed(&e, &m, &plan, &prompt, max_ctx);
let (_vl, _coll, ckpt) = m.glm5_verify_rows(&e, &vt, &mut cache)?;
let pre: Vec<Option<usize>> = cache
.latent
.iter()
.take(n_layers)
.map(|p| p.as_ref().map(|p| p.index_pools_ready))
.collect();
m.glm5_verify_rollback(&e, &mut cache, &ckpt, keep)?;
let mut mutated = 0usize;
#[allow(clippy::needless_range_loop)]
for il in 0..n_layers {
if let (Some(plane), Some(ready)) = (cache.latent[il].as_mut(), pre[il])
&& ready > plane.index_pools_ready
{
plane.index_pools_ready = ready;
mutated += 1;
}
}
assert!(
mutated > 0,
"the walk+rollback did not move index_pools_ready — this red arm is vacuous"
);
match m.decode_step(&e, tokens(1, 0xD00D)[0], &mut cache) {
Err(err) => {
let msg = err.to_string();
v.arm(
"R2 pool-key RED",
msg.contains("index_pools_ready"),
&format!("bites by name — {msg}"),
);
}
Ok(_) => v.arm(
"R2 pool-key RED",
false,
"continuing over keys finalized past j did NOT fail — the tripwire is dead \
under the split",
),
}
}
eprintln!("[phase] arm R3: RED rollback disabled");
{
let mut over = |_round: usize, ki: usize, greedy: u32| -> u32 {
if ki == 0 {
(greedy + 1) % VOCAB
} else {
greedy
}
};
match m.generate_spec_glm5_gated(
&e,
&prompt,
n,
K,
Glm5SpecKnobs {
draft_override: Some(&mut over),
disable_rollback: true,
..Default::default()
},
) {
Ok((out, drafted, accepted)) => v.arm(
"R3 rollback-disabled RED",
out != tape,
&format!("bites — tape diverged with rollback disabled ({accepted}/{drafted})"),
),
Err(err) => v.arm(
"R3 rollback-disabled RED",
true,
&format!("bites — loop failed loudly: {err}"),
),
}
}
println!("==========================================================");
if v.fails == 0 {
println!("glm5-spec-ppn gate: ALL ARMS PASS (P={p} N={n} K={K}, fence={fence:?}; {knobs})");
Ok(())
} else {
println!(
"glm5-spec-ppn gate: {} ARM(S) FAILED (fence={fence:?}; {knobs})",
v.fails
);
std::process::exit(1);
}
}