use memra_engine::Engine;
use memra_engine::forward::argmax;
use memra_engine::glm5_sel_ledger::{self, SelRow};
use memra_engine::hybrid::HybridModel;
use memra_gguf::source::SafetensorsSource;
use memra_kv::Cache;
use std::sync::atomic::Ordering;
use std::time::Instant;
fn arg_val(rest: &[String], key: &str) -> Option<String> {
rest.iter()
.position(|a| a == key)
.and_then(|i| rest.get(i + 1))
.cloned()
}
fn median(v: &mut [f64]) -> f64 {
v.sort_by(|a, b| a.total_cmp(b));
if v.is_empty() {
return f64::NAN;
}
v[v.len() / 2]
}
type ArmOut = Result<(Vec<u32>, Vec<Vec<SelRow>>, f64, bool), Box<dyn std::error::Error>>;
fn reseat_first_recurrent_layer(
e: &Engine,
cache: &mut Cache,
) -> Result<Option<usize>, Box<dyn std::error::Error>> {
let Some((il, rl)) = cache
.recur
.iter_mut()
.enumerate()
.find_map(|(i, r)| r.as_mut().map(|r| (i, r)))
else {
return Ok(None);
};
let n = rl.ssm_state.len();
let mut fresh = e.uninit(n)?;
e.copy_into(&mut fresh, 0, &rl.ssm_state, n)?;
rl.ssm_state = fresh;
Ok(Some(il))
}
fn merge_ledger_rows(device: Vec<SelRow>, host: Vec<SelRow>, dev: usize) -> Vec<SelRow> {
let mut by_layer: std::collections::BTreeMap<(usize, u16), SelRow> =
std::collections::BTreeMap::new();
for r in device {
if r.dev != dev || r.sel.iter().any(|&x| x < 0) {
continue;
}
by_layer.insert((r.dev, r.layer), r);
}
for r in host {
if r.dev != dev {
continue;
}
by_layer.insert((r.dev, r.layer), r);
}
by_layer.into_values().collect()
}
#[allow(clippy::too_many_arguments)] fn run_arm(
e: &Engine,
m: &HybridModel,
prompt: &[u32],
steps: usize,
graph_door: bool,
ledger: bool,
reseat_at: Option<usize>,
trace: bool,
) -> ArmOut {
unsafe {
if graph_door {
std::env::set_var("MEMRA_GLM5_DECODE_GRAPH", "1");
} else {
std::env::remove_var("MEMRA_GLM5_DECODE_GRAPH");
}
if ledger {
std::env::set_var("MEMRA_GLM5_GRAPH_SEL_LEDGER", "1");
} else {
std::env::remove_var("MEMRA_GLM5_GRAPH_SEL_LEDGER");
}
}
unsafe {
if trace {
std::env::set_var("MEMRA_GLM5_GRAPH_TRACE", "1");
} else {
std::env::remove_var("MEMRA_GLM5_GRAPH_TRACE");
}
if reseat_at.is_some() {
std::env::set_var("MEMRA_GLM5_GRAPH_RECAPTURE", "1");
} else {
std::env::remove_var("MEMRA_GLM5_GRAPH_RECAPTURE");
}
}
glm5_sel_ledger::reset_host();
memra_engine::glm5_trace_reset();
let mut cache = Cache::new(e, &m.cfg, prompt.len() + steps + 8)?;
let (logits, _h_seed, _hiddens) = m.prime_cache(e, prompt, &mut cache, 0)?;
let mut tok = argmax(&logits) as u32;
let mut tape = vec![tok];
let mut rows: Vec<Vec<SelRow>> = Vec::with_capacity(steps);
if ledger {
glm5_sel_ledger::reset_host();
}
e.stream().synchronize()?;
let t0 = Instant::now();
let mut recaptured = false;
for step in 1..steps {
if trace {
eprintln!(
"[gate] step {step} door={graph_door} replays={} captures={}",
memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed),
memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed),
);
}
if reseat_at == Some(step) && graph_door {
let before = memra_engine::GLM5_DECODE_GRAPH_RECAPTURES.load(Ordering::Relaxed);
let il = reseat_first_recurrent_layer(e, &mut cache)?;
let l = m.decode_step(e, tok, &mut cache)?;
let after = memra_engine::GLM5_DECODE_GRAPH_RECAPTURES.load(Ordering::Relaxed);
recaptured = after > before;
eprintln!(
"[gate] forced re-seat at step {step} (layer {il:?}): captures {before} -> \
{after}, recaptured={recaptured}"
);
tok = argmax(&l) as u32;
tape.push(tok);
if ledger {
let dev = e.ctx().ordinal();
let device = glm5_sel_ledger::drain_device(e)?;
let host = glm5_sel_ledger::take_host();
rows.push(merge_ledger_rows(device, host, dev));
}
continue;
}
let l = m.decode_step(e, tok, &mut cache)?;
if step == 1 {
let (idx, val) =
l.iter()
.enumerate()
.fold((0usize, f32::NEG_INFINITY), |acc, (i, &v)| {
if v > acc.1 { (i, v) } else { acc }
});
let nz = l.iter().filter(|v| **v != 0.0).count();
eprintln!(
"[gate] step 1 door={graph_door} logits: top1 idx={idx} val={val:.6e} \
nonzero={nz}/{}",
l.len()
);
}
tok = argmax(&l) as u32;
tape.push(tok);
if ledger {
let dev = e.ctx().ordinal();
let device = glm5_sel_ledger::drain_device(e)?;
let host = glm5_sel_ledger::take_host();
rows.push(merge_ledger_rows(device, host, dev));
}
}
e.stream().synchronize()?;
let ms_per_token =
t0.elapsed().as_secs_f64() * 1000.0 / (steps.saturating_sub(1)).max(1) as f64;
Ok((tape, rows, ms_per_token, recaptured))
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let rest: Vec<String> = std::env::args().skip(1).collect();
let steps: usize = arg_val(&rest, "--steps")
.and_then(|v| v.parse().ok())
.unwrap_or(64);
let reps: usize = arg_val(&rest, "--reps")
.and_then(|v| v.parse().ok())
.unwrap_or(5);
let prompt_len: usize = arg_val(&rest, "--prompt-len")
.and_then(|v| v.parse().ok())
.unwrap_or(64);
let trace = rest.iter().any(|a| a == "--trace");
let artifact = std::env::var("GLM5_ARTIFACT").map_err(
|_| "GLM5_ARTIFACT must name the glm5_next artifact directory (safetensors checkpoint)",
)?;
println!(
"[glm5-decode-graph-gate] artifact={artifact} steps={steps} reps={reps} \
prompt_len={prompt_len} trace={trace} MEMRA_PP_DEVICES={:?} MEMRA_PP_STAGES={:?} \
MEMRA_HTOD_DIET={:?} MEMRA_HC_DECODE_WS={:?}",
std::env::var("MEMRA_PP_DEVICES").ok(),
std::env::var("MEMRA_PP_STAGES").ok(),
std::env::var("MEMRA_HTOD_DIET").ok(),
std::env::var("MEMRA_HC_DECODE_WS").ok(),
);
let e = Engine::new(0)?;
println!("[glm5-decode-graph-gate] GPU0: {}", e.ctx().name()?);
let src = SafetensorsSource::open(std::path::Path::new(&artifact))?;
let m = HybridModel::load_from_source(&e, &src)?;
println!(
"[glm5-decode-graph-gate] loaded: n_layer={} n_embd={} hyper={}",
m.cfg.n_layer,
m.cfg.n_embd,
m.hyper.is_some(),
);
if m.hyper.is_none() {
return Err(
"this artifact carries no HyperConnections trunk; the door has nothing to capture"
.into(),
);
}
let prompt: Vec<u32> = (0..prompt_len)
.map(|i| (101 + (i * 7) % 900) as u32)
.collect();
let reseat_at = Some((steps / 2).max(2));
let (eager_tape, eager_rows, _, _) = run_arm(&e, &m, &prompt, steps, false, true, None, trace)?;
let replays_before = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed);
let (graph_tape, graph_rows, _, recaptured) =
run_arm(&e, &m, &prompt, steps, true, true, reseat_at, trace)?;
let replays = memra_engine::GLM5_DECODE_GRAPH_REPLAYS.load(Ordering::Relaxed) - replays_before;
let captured_layers = memra_engine::GLM5_DECODE_GRAPH_LAYERS.load(Ordering::Relaxed);
let captures = memra_engine::GLM5_DECODE_GRAPH_CAPTURES.load(Ordering::Relaxed);
println!(
"door: replays={replays} captures={captures} captured_layers={captured_layers} \
forced_recapture={recaptured}"
);
use std::io::Write;
let _ = std::io::stdout().flush();
let mut fail = Vec::new();
if replays == 0 || captured_layers == 0 {
fail.push(format!(
"VACUOUS: the door never replayed (replays={replays}, captured_layers={captured_layers}, \
captures={captures}). The refusal reason is on stderr as [glm5-decode-graph] eager: ..."
));
}
if !recaptured {
fail.push(format!(
"VACUOUS RE-CAPTURE ARM: the forced re-seat at step {reseat_at:?} did not rebuild a \
stage, so this run says nothing about the invalidation path. The gate arms \
MEMRA_GLM5_GRAPH_RECAPTURE for this arm, so the knob is not the reason. Either the \
re-seated layer is not inside a captured run, or the pointer signature no longer \
sees a re-seat: the engine prints `sig_diff=` on the decision line either way."
));
}
if eager_rows.iter().all(|r| r.is_empty()) || graph_rows.iter().all(|r| r.is_empty()) {
fail.push("VACUOUS: the selection ledger recorded no rows on one of the arms".to_string());
}
if eager_tape.len() != graph_tape.len() {
fail.push(format!(
"token tape length {} (eager) != {} (graph)",
eager_tape.len(),
graph_tape.len()
));
}
let mut tok_mismatch = 0usize;
for (i, (a, b)) in eager_tape.iter().zip(graph_tape.iter()).enumerate() {
if a != b {
if tok_mismatch < 5 {
println!(" TOKEN MISMATCH step {i}: eager={a} graph={b}");
}
tok_mismatch += 1;
}
}
if tok_mismatch > 0 {
fail.push(format!(
"{tok_mismatch}/{} token ids differ",
eager_tape.len()
));
}
let mut sel_mismatch = 0usize;
let mut compared = 0usize;
for (step, (ea, ga)) in eager_rows.iter().zip(graph_rows.iter()).enumerate() {
if ea.len() != ga.len() {
fail.push(format!(
"step {step}: ledger row count {} (eager) != {} (graph)",
ea.len(),
ga.len()
));
break;
}
for (er, gr) in ea.iter().zip(ga.iter()) {
compared += 1;
let idx_same = er.layer == gr.layer && er.sel == gr.sel;
let w_same = er.w.len() == gr.w.len()
&& er
.w
.iter()
.zip(gr.w.iter())
.all(|(a, b)| a.to_bits() == b.to_bits());
if !(idx_same && w_same) {
if sel_mismatch < 5 {
println!(
" SELECTION MISMATCH step {step} layer {}/{}: eager sel={:?} w={:?} | \
graph sel={:?} w={:?}",
er.layer, gr.layer, er.sel, er.w, gr.sel, gr.w
);
}
sel_mismatch += 1;
}
}
}
if sel_mismatch > 0 {
fail.push(format!(
"{sel_mismatch}/{compared} (token, layer) expert selections differ"
));
}
println!(
"identity: tokens {}/{} match; selections {}/{} match (head device only under a pp \
split)",
eager_tape.len() - tok_mismatch,
eager_tape.len(),
compared - sel_mismatch,
compared,
);
let mut eager_ms = Vec::with_capacity(reps);
let mut graph_ms = Vec::with_capacity(reps);
for r in 0..reps {
let (_, _, a, _) = run_arm(&e, &m, &prompt, steps, false, false, None, false)?;
let (_, _, b, _) = run_arm(&e, &m, &prompt, steps, true, false, None, false)?;
println!(" rep {r}: eager {a:.3} ms/token graph {b:.3} ms/token");
eager_ms.push(a);
graph_ms.push(b);
}
let me = median(&mut eager_ms.clone());
let mg = median(&mut graph_ms.clone());
println!(
"per-token ms (N={reps}, interleaved A/B): eager median {me:.3} graph median {mg:.3} \
delta {:+.2}%",
100.0 * (me - mg) / me
);
if fail.is_empty() {
println!("ALL GREEN: glm5-decode-graph gate ({steps} steps, {compared} selection rows)");
Ok(())
} else {
for f in &fail {
println!("FAIL: {f}");
}
Err("glm5-decode-graph-gate FAILED".into())
}
}