use memra_engine::Engine;
use memra_engine::forward::argmax;
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, OutputHead, TensorContract, TensorId,
TensorMatch,
};
use memra_reference::{ReferenceTensor, deterministic_fixture};
use std::borrow::Cow;
use std::collections::BTreeMap;
const HIDDEN: usize = 128;
const VOCAB: u32 = 32;
const LAYERS: usize = 4;
fn mini_config_json() -> String {
r#"{
"model_type": "glm5_next_text",
"num_hidden_layers": 4,
"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": 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()
}
struct OwnedTensor {
bytes: Vec<u8>,
ne: Vec<u64>,
ggml_type: GgmlType,
}
fn is_expert_bank(id: &TensorId) -> bool {
matches!(
id,
TensorId::Layer {
tensor: LayerTensor::MoeExpertGateBank
| LayerTensor::MoeExpertUpBank
| LayerTensor::MoeExpertDownBank,
..
}
)
}
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 glm5_next hc 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 elements: usize = req.shape.iter().map(|&d| d as usize).product();
assert_eq!(
elements,
tensor.data.len(),
"fixture {:?} has {} elements, contract requires {elements}",
req.id,
tensor.data.len()
);
let (bytes, ggml_type) = if is_expert_bank(&req.id) {
(
memra_gguf::nvfp4_repack::f32_to_q8_0(&tensor.data),
GgmlType::Q8_0,
)
} 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 {
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()
}
fn tape_hash(tape: &[u32]) -> u64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for t in tape {
for b in t.to_le_bytes() {
h ^= b as u64;
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
}
h
}
struct ArmCheck {
name: &'static str,
bad_steps: usize,
checked_steps: usize,
tape_bad: bool,
tape_checked: bool,
first: Option<(usize, usize, f32, f32)>, }
impl ArmCheck {
fn new(name: &'static str) -> Self {
ArmCheck {
name,
bad_steps: 0,
checked_steps: 0,
tape_bad: false,
tape_checked: false,
first: None,
}
}
fn check(&mut self, step: usize, phase: &str, got: &[f32], r: &[f32]) {
self.checked_steps += 1;
assert_eq!(
got.len(),
r.len(),
"[{}] step {step} ({phase}): logit row length {} != reference {}",
self.name,
got.len(),
r.len()
);
let diffs = got
.iter()
.zip(r.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
if diffs > 0 {
self.bad_steps += 1;
let (idx, (a, b)) = got
.iter()
.zip(r.iter())
.enumerate()
.find(|(_, (a, b))| a.to_bits() != b.to_bits())
.map(|(i, (a, b))| (i, (*b, *a)))
.unwrap();
if self.first.is_none() {
self.first = Some((step, idx, a, b));
}
if self.bad_steps <= 5 {
println!(
"[{}] MISMATCH step {step} ({phase}): {diffs}/{} logits differ, first \
@[{idx}] ref={a:?} pp={b:?}",
self.name,
r.len()
);
}
}
}
fn check_tape(&mut self, got: &[u32], want: &[u32]) {
self.tape_checked = true;
let hg = tape_hash(got);
let hw = tape_hash(want);
if got == want {
println!(
"[{}] greedy tape MATCH: {} tokens, fnv1a={hg:#018x}",
self.name,
got.len()
);
} else {
self.tape_bad = true;
let at = got
.iter()
.zip(want)
.position(|(a, b)| a != b)
.unwrap_or_else(|| got.len().min(want.len()));
println!(
"[{}] greedy tape DIVERGED at index {at}: ref={hw:#018x} pp={hg:#018x}",
self.name
);
}
}
}
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(6);
let n: usize = std::env::args()
.nth(3)
.and_then(|s| s.parse().ok())
.unwrap_or(8);
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()) };
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-hyper-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");
assert_eq!(
plan.layers.len(),
LAYERS,
"the fixture plan must carry {LAYERS} trunk layers"
);
assert_eq!(plan.hidden_size as usize, HIDDEN);
let fixture = deterministic_fixture(&plan).expect("deterministic glm5_next hc fixture");
let source = fixture_source(&config, &plan, &fixture.weights);
let e = Engine::new(primary_dev)?;
let m = HybridModel::load_from_source_without_mtp(&e, &source)?;
let topology = m.hyper.as_ref().expect(
"the fixture must load as a HyperConnections trunk — otherwise this gate is \
measuring the generic ppN arm that `ppn-gate` already covers",
);
println!(
"hc topology: streams={} collapse={:?} sinkhorn_iters={}",
topology.streams, topology.collapse, topology.sinkhorn_iterations
);
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; a stage with no layers proves nothing"
);
}
println!("stage fence: {fence:?} over {n_layers} layers");
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!(
"state placement: Recurrent(KDA) layers on stages {recur_stages:?}, \
LatentKvCache(MLA+kpool) layers on stages {latent_stages:?}"
);
assert!(
!recur_stages.is_empty() && !latent_stages.is_empty(),
"the fixture must carry BOTH a Recurrent and a LatentKvCache layer; got \
recur={recur_stages:?} latent={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 ({spread:?}); this arm cannot \
see a per-stage cache placement bug — choose a split that separates them"
);
let ids = tokens(p + n, 0x0915_5EED);
let max_ctx = p + n + 8;
eprintln!("[phase] reference A: door OFF, step-by-step decode over the whole sequence");
unsafe { std::env::remove_var("MEMRA_PP_STAGES") };
assert!(
memra_engine::pp::pp_cuts(n_layers).is_none(),
"the reference arm must run with the door SHUT"
);
let mut ref_step_logits: Vec<Vec<f32>> = Vec::with_capacity(p + n);
let mut ref_step_tape: Vec<u32> = Vec::with_capacity(p + n);
{
let mut cache = memra_engine::cache::Cache::new_planned(&e, &m.cfg, &plan, max_ctx)?;
for &tok in &ids {
let ll = m.decode_step(&e, tok, &mut cache)?;
ref_step_tape.push(argmax(&ll) as u32);
ref_step_logits.push(ll);
}
}
let n_vocab = ref_step_logits[0].len();
eprintln!("[phase] reference C: door OFF, stateless prefill (forward / forward_last)");
let ref_forward = m.forward(&e, &ids)?;
let ref_forward_last = m.forward_last(&e, &ids)?;
eprintln!("[phase] reference B: door OFF, prime + decode continuation");
let mut ref_cont_logits: Vec<Vec<f32>> = Vec::with_capacity(n);
let mut ref_cont_tape: Vec<u32> = Vec::with_capacity(n);
let ref_prime_logits: Vec<f32>;
let ref_prime_hiddens: Vec<f32>;
{
let mut cache = memra_engine::cache::Cache::new_planned(&e, &m.cfg, &plan, max_ctx)?;
let (primed, _seed, hiddens) = m.prime_cache(&e, &ids[..p], &mut cache, 0)?;
ref_prime_logits = primed;
ref_prime_hiddens = e.dtoh(&hiddens)?;
for &tok in &ids[p..] {
let ll = m.decode_step(&e, tok, &mut cache)?;
ref_cont_tape.push(argmax(&ll) as u32);
ref_cont_logits.push(ll);
}
}
let prime_calls =
memra_engine::hybrid_forward::hyper_prime_ranges(p, n_layers, m.gdn_prime_grid_on()).len();
println!(
"prime schedule: {prime_calls} call(s) for P={p} (MEMRA_PRIME_CHUNK={:?})",
std::env::var("MEMRA_PRIME_CHUNK").ok()
);
if let Ok(chunk) = std::env::var("MEMRA_PRIME_CHUNK")
&& chunk.parse::<usize>().is_ok_and(|c| c > 0 && c < p)
{
assert!(
prime_calls > 1,
"MEMRA_PRIME_CHUNK={chunk} < P={p} but the prime schedule is one call — the \
chunked arm would be vacuous"
);
}
assert!(p >= 6, "overlay arm needs P >= 6 (spans at 1..3 and p-2)");
let subs = tokens(3, 0x5EED_0BE1);
let overlay_spans: Vec<(usize, usize, usize)> = vec![(1, 0, 2), (p - 2, 2, 1)];
let placeholder: u32 = 7 % VOCAB;
let mut ids_ph = ids.clone();
let mut ids_sub = ids.clone();
for &(pos, row_off, n_rows) in &overlay_spans {
for r in 0..n_rows {
ids_ph[pos + r] = placeholder;
ids_sub[pos + r] = subs[row_off + r];
}
}
let overlay = memra_engine::vision::EmbedOverlay {
rows: m.embed(&e, &subs)?,
spans: overlay_spans.clone(),
};
let red_shift = std::env::var("MEMRA_GLM5V_GATE_RED").as_deref() == Ok("span-shift");
let arm_overlay = if red_shift {
println!(
"glm5-hyper-ppn-gate RED ARM: MEMRA_GLM5V_GATE_RED=span-shift — every overlay \
span moved +1; this run MUST fail"
);
memra_engine::vision::EmbedOverlay {
rows: overlay.rows.clone(),
spans: overlay_spans
.iter()
.map(|&(pos, row_off, n_rows)| (pos + 1, row_off, n_rows))
.collect(),
}
} else {
memra_engine::vision::EmbedOverlay {
rows: overlay.rows.clone(),
spans: overlay.spans.clone(),
}
};
eprintln!("[phase] reference D: door OFF, substituted-token prime + decode (overlay truth)");
let mut ref_ov_cont: Vec<Vec<f32>> = Vec::with_capacity(n);
let ref_ov_prime: Vec<f32>;
{
let mut cache = memra_engine::cache::Cache::new_planned(&e, &m.cfg, &plan, max_ctx)?;
let (primed, _seed, _hiddens) = m.prime_cache(&e, &ids_sub[..p], &mut cache, 0)?;
ref_ov_prime = primed;
for &tok in &ids_sub[p..] {
ref_ov_cont.push(m.decode_step(&e, tok, &mut cache)?);
}
}
let split_at = p / 2;
let mut ref_ov2_cont: Vec<Vec<f32>> = Vec::with_capacity(n);
let ref_ov2_prime: Vec<f32>;
{
let mut cache = memra_engine::cache::Cache::new_planned(&e, &m.cfg, &plan, max_ctx)?;
let _ = m.prime_cache(&e, &ids_sub[..split_at], &mut cache, p - split_at)?;
let (primed, _seed, _hiddens) = m.prime_cache(&e, &ids_sub[split_at..p], &mut cache, 0)?;
ref_ov2_prime = primed;
for &tok in &ids_sub[p..] {
ref_ov2_cont.push(m.decode_step(&e, tok, &mut cache)?);
}
}
eprintln!("[phase] arm 5a: door OFF, serial overlay splice");
let mut overlay_serial_arm = ArmCheck::new("overlay-serial");
{
let mut cache = memra_engine::cache::Cache::new_planned(&e, &m.cfg, &plan, max_ctx)?;
let (primed, _seed, _hiddens) =
m.prime_cache_overlaid(&e, &ids_ph[..p], &mut cache, 0, Some(&arm_overlay))?;
overlay_serial_arm.check(0, "overlay prime last row", &primed, &ref_ov_prime);
for (k, &tok) in ids_sub[p..].iter().enumerate() {
let ll = m.decode_step(&e, tok, &mut cache)?;
overlay_serial_arm.check(k + 1, "decode after overlay prime", &ll, &ref_ov_cont[k]);
}
}
unsafe { std::env::set_var("MEMRA_PP_STAGES", stages.to_string()) };
eprintln!("[phase] arm 1: door ON, decode-serial split walk");
let mut decode_arm = ArmCheck::new("decode-serial");
{
let mut cache = memra_engine::pp::new_cache_planned(&e, &m.cfg, &plan, max_ctx)?;
let mut tape: Vec<u32> = Vec::with_capacity(p + n);
for (step, &tok) in ids.iter().enumerate() {
let ll = m.decode_step(&e, tok, &mut cache)?;
tape.push(argmax(&ll) as u32);
decode_arm.check(step, "decode", &ll, &ref_step_logits[step]);
}
decode_arm.check_tape(&tape, &ref_step_tape);
}
eprintln!("[phase] arm 2: door ON, prime-twin split walk");
let mut prime_arm = ArmCheck::new("prime-twin");
{
let mut cache = memra_engine::pp::new_cache_planned(&e, &m.cfg, &plan, max_ctx)?;
let (primed, _seed, hiddens) = m.prime_cache(&e, &ids[..p], &mut cache, 0)?;
prime_arm.check(0, "prime last row", &primed, &ref_prime_logits);
let got_hiddens = e.dtoh(&hiddens)?;
prime_arm.check(0, "prime hiddens stack", &got_hiddens, &ref_prime_hiddens);
let mut tape: Vec<u32> = Vec::with_capacity(n);
for (k, &tok) in ids[p..].iter().enumerate() {
let ll = m.decode_step(&e, tok, &mut cache)?;
tape.push(argmax(&ll) as u32);
prime_arm.check(k + 1, "decode after prime", &ll, &ref_cont_logits[k]);
}
prime_arm.check_tape(&tape, &ref_cont_tape);
}
eprintln!("[phase] arm 3: door ON, prefill-twin split walk");
let mut prefill_arm = ArmCheck::new("prefill-twin");
{
let got = m.forward(&e, &ids)?;
prefill_arm.check(0, "forward (all rows)", &got, &ref_forward);
let got_last = m.forward_last(&e, &ids)?;
prefill_arm.check(1, "forward_last", &got_last, &ref_forward_last);
}
println!(
"glm5-hyper-ppn-gate NOTE: pipelined arm skipped — `decode_step_h_ppn_deferred` calls \
refuse_hyper(), so deferred readback is not wired for the mHC residual. This gate \
covers the serial split walk only; do not cite it for the pipelined arm."
);
eprintln!("[phase] arm 5b: door ON, monolithic overlay splice");
let mut overlay_ppn_arm = ArmCheck::new("overlay-ppn");
{
let mut cache = memra_engine::pp::new_cache_planned(&e, &m.cfg, &plan, max_ctx)?;
let (primed, _seed, _hiddens) =
m.prime_cache_overlaid(&e, &ids_ph[..p], &mut cache, 0, Some(&arm_overlay))?;
overlay_ppn_arm.check(0, "overlay prime last row", &primed, &ref_ov_prime);
for (k, &tok) in ids_sub[p..].iter().enumerate() {
let ll = m.decode_step(&e, tok, &mut cache)?;
overlay_ppn_arm.check(k + 1, "decode after overlay prime", &ll, &ref_ov_cont[k]);
}
}
eprintln!("[phase] arm 5c: door ON, two-call windowed overlay splice");
let mut overlay_win_arm = ArmCheck::new("overlay-ppn-windowed");
{
let mut cache = memra_engine::pp::new_cache_planned(&e, &m.cfg, &plan, max_ctx)?;
let w0 = arm_overlay.window(0, split_at);
let _ = m.prime_cache_overlaid(
&e,
&ids_ph[..split_at],
&mut cache,
p - split_at,
w0.as_ref(),
)?;
let w1 = arm_overlay.window(split_at, p - split_at);
let (primed, _seed, _hiddens) =
m.prime_cache_overlaid(&e, &ids_ph[split_at..p], &mut cache, 0, w1.as_ref())?;
overlay_win_arm.check(
0,
"windowed overlay prime last row",
&primed,
&ref_ov2_prime,
);
for (k, &tok) in ids_sub[p..].iter().enumerate() {
let ll = m.decode_step(&e, tok, &mut cache)?;
overlay_win_arm.check(k + 1, "decode after windowed prime", &ll, &ref_ov2_cont[k]);
}
}
let mut fail = false;
for arm in [
&decode_arm,
&prime_arm,
&prefill_arm,
&overlay_serial_arm,
&overlay_ppn_arm,
&overlay_win_arm,
] {
assert!(
arm.checked_steps > 0,
"[{}] compared ZERO steps — a vacuous arm never prints PASS",
arm.name
);
if arm.bad_steps == 0 && !arm.tape_bad {
println!(
"glm5-hyper-ppn gate PASS [{}]: {} comparisons BIT-IDENTICAL vs the unsplit hc \
walk (n_vocab={n_vocab}, P={p} N={n}, fence={fence:?}; {knobs})",
arm.name, arm.checked_steps
);
} else {
let detail = match arm.first {
Some((s, i, a, b)) => format!(
"{}/{} comparisons mismatched (first @ step {s} idx {i}: ref={a:?} \
pp={b:?})",
arm.bad_steps, arm.checked_steps
),
None => format!(
"{}/{} comparisons mismatched",
arm.bad_steps, arm.checked_steps
),
};
let tape = match (arm.tape_checked, arm.tape_bad) {
(false, _) => "; this arm carries no greedy tape (stateless walk)",
(true, true) => "; greedy tape DIVERGED",
(true, false) => {
"; greedy tape MATCHED even so — the logit compare is the \
load-bearing bar"
}
};
println!(
"glm5-hyper-ppn gate FAIL [{}]: {detail}{tape} (fence={fence:?}; {knobs})",
arm.name
);
fail = true;
}
}
if fail {
std::process::exit(1);
}
Ok(())
}