use memra_engine::Engine;
use memra_engine::forward::argmax;
use memra_engine::hybrid::HybridModel;
use memra_gguf::GgufFile;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let path = std::env::args()
.nth(1)
.expect("usage: pp2-gate <model.gguf> [P] [N] [split]");
let p: usize = std::env::args().nth(2).and_then(|s| s.parse().ok()).unwrap_or(16);
let n: usize = std::env::args().nth(3).and_then(|s| s.parse().ok()).unwrap_or(32);
let split_arg: Option<usize> = std::env::args().nth(4).and_then(|s| s.parse().ok());
unsafe {
std::env::remove_var("MEMRA_PP_STAGES");
std::env::remove_var("MEMRA_PP_SPLIT");
}
let knobs = format!(
"streams={} overlap={} devices={}",
if memra_engine::pp::pp2_streams_off() { "OFF(inc1 seam)" } else { "per-stage" },
if memra_engine::pp::pp2_overlap() { "1(double-buffered)" } else { "0" },
std::env::var("MEMRA_PP_DEVICES").unwrap_or_else(|_| "default(primary)".into()),
);
println!("pp2-gate increment-2 config: {knobs}");
let e = Engine::new(0)?;
let g = GgufFile::open(&path)?;
let m = HybridModel::load(&e, &g)?;
let n_layers = m.layers.len();
let split = split_arg.unwrap_or(n_layers / 2);
let prompt: Vec<u32> = (0..p).map(|i| (100 + (i * 7) % 900) as u32).collect();
let mut cache_ref = memra_engine::cache::Cache::new(&e, &m.cfg, p + n + 8)?;
let mut inputs: Vec<u32> = Vec::with_capacity(p + n);
let mut ref_logits: Vec<Vec<f32>> = Vec::with_capacity(p + n);
let mut next = 0u32;
for step in 0..p + n {
let tok = if step < p { prompt[step] } else { next };
inputs.push(tok);
let ll = m.decode_step(&e, tok, &mut cache_ref)?;
next = argmax(&ll) as u32;
ref_logits.push(ll);
}
let n_vocab = ref_logits[0].len();
unsafe {
std::env::set_var("MEMRA_PP_STAGES", "2");
if let Some(s) = split_arg {
std::env::set_var("MEMRA_PP_SPLIT", s.to_string());
}
}
assert_eq!(
memra_engine::pp::pp2_split(n_layers),
Some(split),
"pp2 door failed to open (n_layers={n_layers}, split={split})"
);
let mut cache_pp = memra_engine::pp::new_cache(&e, &m.cfg, p + n + 8)?;
let mut bad_steps = 0usize;
let mut first: Option<(usize, usize, f32, f32)> = None; for (step, &tok) in inputs.iter().enumerate() {
let ll = m.decode_step(&e, tok, &mut cache_pp)?;
let r = &ref_logits[step];
let diffs = ll
.iter()
.zip(r.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
if diffs > 0 {
bad_steps += 1;
let (idx, (a, b)) = ll
.iter()
.zip(r.iter())
.enumerate()
.find(|(_, (a, b))| a.to_bits() != b.to_bits())
.map(|(i, (a, b))| (i, (*b, *a)))
.unwrap();
if first.is_none() {
first = Some((step, idx, a, b));
}
if bad_steps <= 5 {
println!(
"MISMATCH step {step} ({}): {diffs}/{n_vocab} logits differ, first @[{idx}] ref={a:?} pp2={b:?}",
if step < p { "prime" } else { "gen" }
);
}
}
}
let total = p + n;
if bad_steps == 0 {
println!(
"pp2 gate PASS: {total} steps ({p} prime + {n} gen) BIT-IDENTICAL logits \
(n_vocab={n_vocab}, n_layers={n_layers}, stage0=[0,{split}), stage1=[{split},{n_layers}); {knobs})"
);
Ok(())
} else {
let (s, i, a, b) = first.unwrap();
println!(
"pp2 gate FAIL: {bad_steps}/{total} steps mismatched (first @ step {s} idx {i}: \
ref={a:?} pp2={b:?}; split={split}/{n_layers}; {knobs})"
);
std::process::exit(1);
}
}