use memra_engine::Engine;
use memra_engine::cache::Cache;
use memra_engine::decode_batch::DevSamp;
use memra_engine::forward::argmax;
use memra_engine::hybrid::HybridModel;
use memra_gguf::GgufFile;
fn main() -> Result<(), Box<dyn std::error::Error>> {
let mut args = std::env::args().skip(1);
let path = args
.next()
.expect("usage: decode-batch-gate <model.gguf> [--steps N] [--batch B]");
let rest: Vec<String> = args.collect();
let steps: usize = rest
.iter()
.position(|a| a == "--steps")
.and_then(|i| rest.get(i + 1))
.and_then(|v| v.parse().ok())
.unwrap_or(32);
let batches: Vec<usize> = rest
.iter()
.position(|a| a == "--batch")
.and_then(|i| rest.get(i + 1))
.map(|v| {
v.split(',')
.filter_map(|p| p.trim().parse().ok())
.collect::<Vec<usize>>()
})
.filter(|v: &Vec<usize>| !v.is_empty())
.unwrap_or_else(|| vec![4]);
let b_n: usize = batches[0];
let mode: String = rest
.iter()
.position(|a| a == "--mode")
.and_then(|i| rest.get(i + 1))
.cloned()
.unwrap_or_else(|| "config".into());
let strict: bool = mode == "strict";
let ppspec_mode: bool = mode == "ppspec";
let pp_mode: bool = mode == "pp" || ppspec_mode;
let ts: Vec<usize> = rest
.iter()
.position(|a| a == "--ts")
.and_then(|i| rest.get(i + 1))
.map(|v| {
v.split(',')
.filter_map(|p| p.trim().parse().ok())
.collect::<Vec<usize>>()
})
.filter(|v: &Vec<usize>| !v.is_empty())
.unwrap_or_else(|| vec![2, 5, 9]);
let stages: usize = rest
.iter()
.position(|a| a == "--stages")
.and_then(|i| rest.get(i + 1))
.and_then(|v| v.parse().ok())
.unwrap_or(2);
let reps: usize = rest
.iter()
.position(|a| a == "--reps")
.and_then(|i| rest.get(i + 1))
.and_then(|v| v.parse().ok())
.unwrap_or(2);
let plen: u32 = rest
.iter()
.position(|a| a == "--plen")
.and_then(|i| rest.get(i + 1))
.and_then(|v| v.parse().ok())
.unwrap_or(20);
unsafe {
std::env::set_var("MEMRA_GDN_MMA", "0");
}
unsafe {
std::env::set_var("MEMRA_L2_V2", "0");
}
unsafe {
std::env::set_var("MEMRA_FA3", "0");
}
unsafe {
std::env::set_var("MEMRA_GDN_WGMMA", "0");
}
unsafe {
std::env::set_var("MEMRA_MOE_F16G", "0");
}
let primary_dev: usize = if pp_mode {
unsafe {
std::env::set_var("MEMRA_PP_STAGES", stages.to_string());
}
std::env::var("MEMRA_PP_DEVICES")
.ok()
.and_then(|v| v.split(',').next().and_then(|s| s.trim().parse().ok()))
.unwrap_or(0)
} else {
0
};
let e = Engine::new(primary_dev)?;
let (model, arch) = if std::path::Path::new(&path).is_dir() {
let dir = std::path::Path::new(&path);
let src: Box<dyn memra_gguf::source::TensorSource> = if dir.join("manifest.json").exists() {
Box::new(memra_gguf::source::Hy3RepackSource::open(dir)?)
} else {
Box::new(memra_gguf::source::SafetensorsSource::open(dir)?)
};
let m = HybridModel::load_from_source_without_mtp(&e, src.as_ref())?;
(m, "safetensors".to_string())
} else {
let g = GgufFile::open(&path)?;
let a = g.arch().unwrap_or("?").to_string();
(HybridModel::load_without_mtp(&e, &g)?, a)
};
println!(
"loaded {} ({} layers); steps={steps} batch={b_n}",
arch,
model.layers.len()
);
if pp_mode {
let seed: u32 = std::env::var("MEMRA_GATE_SEED")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
if ppspec_mode {
let fails = ppspec_battery(&e, &model, stages, steps, &ts, reps, seed)?;
if fails == 0 {
println!("ALL GREEN: spec-verify PP-{stages} stage-split exactness battery");
return Ok(());
}
return Err("decode-batch-gate --mode ppspec FAILED".into());
}
let fails = pp_battery(&e, &model, stages, steps, &batches, reps, seed, plen)?;
if fails == 0 {
println!("ALL GREEN: batched PP-{stages} stage-split exactness battery");
return Ok(());
}
return Err("decode-batch-gate --mode pp FAILED".into());
}
let seed: u32 = std::env::var("MEMRA_GATE_SEED")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let prompts: Vec<Vec<u32>> = (0..b_n.max(2))
.map(|i| {
(0..plen + i as u32 * 5)
.map(|j| 55 + seed * 13 + i as u32 * 97 + j * 31)
.collect()
})
.collect();
let ctx = 512 + steps + 64;
const G1_EARLY_STEP: usize = 3; const G1_EARLY_K: usize = 4; let canary = std::env::var("MEMRA_GATE_CANARY")
.map(|v| v == "1")
.unwrap_or(false);
let b1_fast_configured = HybridModel::b1_fast_on();
let g1_live_eager_inapplicable =
!strict && !canary && !(b1_fast_configured && model.b1_fast_arch_eligible());
let mut g1_fail = 0usize;
let mut g1_early = 0usize;
let g1_seeds: u32 = if g1_live_eager_inapplicable {
0
} else if strict {
1
} else {
6
};
for gs in 0..g1_seeds {
let p0: Vec<u32> = (0..20).map(|j| 55 + (seed + gs) * 13 + j * 31).collect();
let mut c_ref = Cache::new(&e, &model.cfg, ctx)?;
let mut c_bat = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, &p0, &mut c_ref, 0)?;
let _ = model.prime_cache(&e, &p0, &mut c_bat, 0)?;
let mut t_ref = *p0.last().unwrap();
let mut t_bat = t_ref;
let mut diverged: Option<usize> = None;
for s in 0..steps {
if canary && s == 1 {
t_bat = if t_bat == 0 { 1 } else { t_bat - 1 };
}
let (l_ref, _) = model.decode_step_h(&e, t_ref, &mut c_ref)?;
let l_bat = {
let mut caches = [&mut c_bat];
model
.decode_step_batch(&e, &[t_bat], &mut caches)?
.remove(0)
};
if strict {
let bits_equal = l_ref.len() == l_bat.len()
&& l_ref
.iter()
.zip(l_bat.iter())
.all(|(a, b)| a.to_bits() == b.to_bits());
if !bits_equal {
let md = l_ref
.iter()
.zip(l_bat.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
println!("gate1 step {s}: BIT-DIFF (maxdiff {md:.3e}) FAIL");
g1_fail += 1;
if g1_fail > 3 {
break;
}
}
}
t_ref = argmax(&l_ref) as u32;
t_bat = argmax(&l_bat) as u32;
if t_ref != t_bat {
diverged = Some(s);
break;
}
}
match diverged {
Some(s) if strict => {
println!("gate1 seed {gs} step {s}: token diverged FAIL");
g1_fail += 1;
}
Some(s) if s < G1_EARLY_STEP => {
g1_early += 1;
println!(
"gate1 seed {gs} step {s}: token diverged EARLY \
(step < {G1_EARLY_STEP}; plumbing iff every draw)"
);
}
Some(s) => println!(
"gate1 seed {gs} step {s}: token diverged — accepted \
cross-config drift (WARN)"
),
None => println!("gate1 seed {gs}: agreement all {steps} steps"),
}
}
if g1_live_eager_inapplicable {
println!(
"gate1 (B=1 vs decode_step_h): N/A for the live default; B=1 uses the \
batched numeric class and gate2 checks B=1 vs B={b_n}"
);
} else if !strict {
println!(
"gate1 early draws (step < {G1_EARLY_STEP}): {g1_early}/{g1_seeds} \
(FAIL threshold >= {G1_EARLY_K}; plumbing class = every draw)"
);
if g1_early >= G1_EARLY_K {
g1_fail += 1;
}
println!(
"gate1 (B=1 argmax agreement vs decode_step_h, {steps} steps, \
{g1_seeds} seed(s)): {}",
if g1_fail == 0 { "PASS" } else { "FAIL" }
);
} else {
println!(
"gate1 (B=1 bit-identity vs decode_step_h, {steps} steps, \
{g1_seeds} seed(s)): {}",
if g1_fail == 0 { "PASS" } else { "FAIL" }
);
}
let b1_fast_setting = HybridModel::b1_fast_on();
let b1_fast_live = b1_fast_setting && model.b1_fast_arch_eligible();
HybridModel::set_b1_fast(false);
println!(
"gate2/gate3 B=1 reference arm: batched body (B=1 fast path pinned OFF; \
global setting = {}; effective for this architecture = {})",
if b1_fast_setting { "ON" } else { "OFF" },
if b1_fast_live { "ON" } else { "OFF" }
);
let mut ref_streams: Vec<Vec<u32>> = Vec::new();
let mut ref_logits: Vec<Vec<Vec<f32>>> = Vec::new();
for p in prompts.iter().take(b_n) {
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, p, &mut c, 0)?;
let mut t = *p.last().unwrap();
let mut out = Vec::with_capacity(steps);
let mut ls = Vec::with_capacity(steps);
for _ in 0..steps {
let l = if strict {
model.decode_step_h(&e, t, &mut c)?.0
} else {
let mut caches = [&mut c];
model.decode_step_batch(&e, &[t], &mut caches)?.remove(0)
};
t = argmax(&l) as u32;
out.push(t);
ls.push(l);
}
ref_streams.push(out);
ref_logits.push(ls);
}
let mut caches: Vec<Cache> = Vec::new();
for p in prompts.iter().take(b_n) {
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, p, &mut c, 0)?;
caches.push(c);
}
let mut toks: Vec<u32> = prompts
.iter()
.take(b_n)
.map(|p| *p.last().unwrap())
.collect();
let mut g2_fail = 0usize;
for s in 0..steps {
let mut cache_refs: Vec<&mut Cache> = caches.iter_mut().collect();
let logits = model.decode_step_batch(&e, &toks, &mut cache_refs)?;
for (bi, l) in logits.iter().enumerate() {
toks[bi] = argmax(l) as u32;
if toks[bi] != ref_streams[bi][s] {
println!("gate2 seq {bi}: token DIVERGED from isolated at step {s} FAIL");
g2_fail += 1;
} else if !strict {
let r = &ref_logits[bi][s];
if !(r.len() == l.len()
&& r.iter()
.zip(l.iter())
.all(|(a, b)| a.to_bits() == b.to_bits()))
{
let md = r
.iter()
.zip(l.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
println!(
"gate2 seq {bi} step {s}: LOGIT BIT-DIFF vs isolated \
(maxdiff {md:.3e}) FAIL"
);
g2_fail += 1;
}
}
}
if g2_fail > 6 {
break;
}
}
println!(
"gate2 (B={b_n} vs isolated {}, {steps} steps): {}",
if strict {
"decode_step_h"
} else {
"batched-B=1, bit-checked"
},
if g2_fail == 0 { "PASS" } else { "FAIL" }
);
drop(caches);
drop(ref_logits);
drop(ref_streams);
let mut g3_fail = 0usize;
{
{
let mut caches: Vec<Cache> = Vec::new();
for p in prompts.iter().take(b_n) {
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, p, &mut c, 0)?;
caches.push(c);
}
let mut toks: Vec<u32> = prompts
.iter()
.take(b_n)
.map(|p| *p.last().unwrap())
.collect();
let samp_g: Vec<Option<DevSamp>> = vec![Some((0.0, 0, 0, 0i32, 1.0f32, 0.0f32)); b_n];
for _s in 0..steps.min(16) {
let mut cache_refs: Vec<&mut Cache> = caches.iter_mut().collect();
let (rows, next) =
model.decode_step_batch_sampled(&e, &toks, &mut cache_refs, &samp_g)?;
for (bi, l) in rows.iter().enumerate() {
let host_am = argmax(l) as u32;
let dev = next[bi].expect("greedy device row missing token");
if dev != host_am {
println!(
"gate3a seq {bi}: device argmax {dev} != host argmax {host_am} FAIL"
);
g3_fail += 1;
}
toks[bi] = host_am;
}
if g3_fail > 4 {
break;
}
}
}
let n_s = steps.min(16);
let mut iso: Vec<Vec<u32>> = Vec::with_capacity(b_n);
for bi in 0..b_n {
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, &prompts[bi], &mut c, 0)?;
let mut t = *prompts[bi].last().unwrap();
let mut out = Vec::with_capacity(n_s);
for s in 0..n_s {
let mut refs = [&mut c];
let samp = [Some((
0.7f32,
bi as u64 + 1,
s as u32,
0i32,
1.0f32,
0.0f32,
))];
let (_, nx) = model.decode_step_batch_sampled(&e, &[t], &mut refs, &samp)?;
t = nx[0].expect("sampled row missing token");
out.push(t);
}
iso.push(out);
}
let mut bat: Vec<Vec<u32>> = vec![Vec::with_capacity(n_s); b_n];
{
let mut caches: Vec<Cache> = Vec::new();
for p in prompts.iter().take(b_n) {
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, p, &mut c, 0)?;
caches.push(c);
}
let mut toks: Vec<u32> = prompts
.iter()
.take(b_n)
.map(|p| *p.last().unwrap())
.collect();
for s in 0..n_s {
let samp: Vec<Option<DevSamp>> = (0..b_n)
.map(|bi| Some((0.7f32, bi as u64 + 1, s as u32, 0i32, 1.0f32, 0.0f32)))
.collect();
let mut cache_refs: Vec<&mut Cache> = caches.iter_mut().collect();
let (_, nx) = model.decode_step_batch_sampled(&e, &toks, &mut cache_refs, &samp)?;
for bi in 0..b_n {
toks[bi] = nx[bi].expect("sampled row missing token");
bat[bi].push(toks[bi]);
}
}
}
for bi in 0..b_n {
if iso[bi] != bat[bi] {
let d = iso[bi].iter().zip(&bat[bi]).position(|(a, b)| a != b);
println!(
"gate3b seq {bi}: sampled stream DIVERGED batched-vs-isolated at \
step {d:?} FAIL"
);
g3_fail += 1;
}
}
{
let n_s = steps.min(8);
let mut caches_f: Vec<Cache> = Vec::new();
let mut caches_l: Vec<Cache> = Vec::new();
for p in prompts.iter().take(b_n) {
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, p, &mut c, 0)?;
caches_f.push(c);
let mut c = Cache::new(&e, &model.cfg, ctx)?;
let _ = model.prime_cache(&e, p, &mut c, 0)?;
caches_l.push(c);
}
let mut toks: Vec<u32> = prompts
.iter()
.take(b_n)
.map(|p| *p.last().unwrap())
.collect();
for _s in 0..n_s {
let samp: Vec<Option<DevSamp>> = (0..b_n)
.map(|bi| {
if bi % 2 == 0 {
Some((0.0, 0, 0, 0i32, 1.0f32, 0.0f32))
} else {
None
}
})
.collect();
let (rows_f, next_f) = {
let mut refs: Vec<&mut Cache> = caches_f.iter_mut().collect();
model.decode_step_batch_sampled_lean(&e, &toks, &mut refs, &samp, false)?
};
let (rows_l, next_l) = {
let mut refs: Vec<&mut Cache> = caches_l.iter_mut().collect();
model.decode_step_batch_sampled_lean(&e, &toks, &mut refs, &samp, true)?
};
for bi in 0..b_n {
if samp[bi].is_some() {
if next_f[bi] != next_l[bi] {
println!(
"gate3c seq {bi}: lean token {:?} != full token {:?} FAIL",
next_l[bi], next_f[bi]
);
g3_fail += 1;
}
if !rows_l[bi].is_empty() {
println!("gate3c seq {bi}: lean sampled row NOT empty FAIL");
g3_fail += 1;
}
let parked = e.dtoh(
caches_l[bi]
.last_logits_dev
.as_ref()
.expect("lean row missing device park"),
)?;
let r = &rows_f[bi];
if !(parked.len() == r.len()
&& parked
.iter()
.zip(r.iter())
.all(|(a, b)| a.to_bits() == b.to_bits()))
{
println!("gate3c seq {bi}: parked device logits != full host row FAIL");
g3_fail += 1;
}
toks[bi] = next_f[bi].unwrap();
} else {
let (r, l) = (&rows_f[bi], &rows_l[bi]);
if !(r.len() == l.len()
&& r.iter()
.zip(l.iter())
.all(|(a, b)| a.to_bits() == b.to_bits()))
{
println!("gate3c seq {bi}: unsampled row lean != full FAIL");
g3_fail += 1;
}
toks[bi] = argmax(r) as u32;
}
}
if g3_fail > 8 {
break;
}
}
}
}
println!(
"gate3 (device sampling: greedy==host-argmax + sampled B={b_n} vs isolated \
+ lean-logits identity): {}",
if g3_fail == 0 { "PASS" } else { "FAIL" }
);
if g1_fail + g2_fail + g3_fail == 0 {
println!("ALL GREEN: decode_step_batch exactness battery");
Ok(())
} else {
Err("decode-batch-gate FAILED".into())
}
}
struct BitCheck {
name: String,
bad: usize,
first: Option<(usize, usize, usize, f32, f32)>, compared: usize,
}
fn decode_batch_serial_waves(
e: &Engine,
model: &HybridModel,
toks: &[u32],
caches: &mut [Cache],
mid: Option<usize>,
) -> Result<Vec<Vec<f32>>, Box<dyn std::error::Error>> {
let Some(mid) = mid else {
let mut refs: Vec<&mut Cache> = caches.iter_mut().collect();
return model.decode_step_batch(e, toks, &mut refs);
};
let (toks_a, toks_b) = toks.split_at(mid);
let (caches_a, caches_b) = caches.split_at_mut(mid);
let mut refs_a: Vec<&mut Cache> = caches_a.iter_mut().collect();
let mut rows = model.decode_step_batch(e, toks_a, &mut refs_a)?;
let mut refs_b: Vec<&mut Cache> = caches_b.iter_mut().collect();
rows.extend(model.decode_step_batch(e, toks_b, &mut refs_b)?);
Ok(rows)
}
impl BitCheck {
fn new(name: String) -> Self {
BitCheck {
name,
bad: 0,
first: None,
compared: 0,
}
}
fn check(&mut self, step: usize, row: usize, got: &[f32], r: &[f32]) {
assert_eq!(
got.len(),
r.len(),
"row length mismatch (ref {} vs got {})",
r.len(),
got.len()
);
self.compared += got.len();
let diffs = got
.iter()
.zip(r.iter())
.filter(|(a, b)| a.to_bits() != b.to_bits())
.count();
if diffs > 0 {
self.bad += 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, row, idx, a, b));
}
if self.bad <= 5 {
println!(
"[{}] MISMATCH step {step} row {row}: {diffs}/{} logits differ, \
first @[{idx}] ref={a:?} pp={b:?}",
self.name,
r.len()
);
}
}
}
fn verdict(&self) -> usize {
if self.bad == 0 {
println!(
"pp gate PASS [{}]: {} f32 logits BIT-IDENTICAL (0 differing bits)",
self.name, self.compared
);
0
} else {
let (s, row, i, a, b) = self.first.unwrap();
println!(
"pp gate FAIL [{}]: {} rows mismatched of {} f32 compared (first @ step \
{s} row {row} idx {i}: ref={a:?} pp={b:?})",
self.name, self.bad, self.compared
);
1
}
}
}
#[allow(clippy::too_many_arguments)]
fn pp_battery(
e: &Engine,
model: &HybridModel,
stages: usize,
steps: usize,
batches: &[usize],
reps: usize,
seed: u32,
plen: u32,
) -> Result<usize, Box<dyn std::error::Error>> {
let n_layers = model.layers.len();
let fence = memra_engine::pp::pp_cuts(n_layers).unwrap_or_else(|| {
panic!("pp mode: door failed to open (n_layers={n_layers}, stages={stages})")
});
assert_eq!(
fence.len() - 1,
stages,
"fence {fence:?} != stages {stages}"
);
let devices = std::env::var("MEMRA_PP_DEVICES").unwrap_or_default();
let knobs = format!(
"stages={stages} fence={fence:?} devices={} splits={} shard={} streams={}",
if devices.is_empty() {
"default(primary)".into()
} else {
devices.clone()
},
std::env::var("MEMRA_PP_SPLITS").unwrap_or_else(|_| "default(even)".into()),
if memra_engine::pp::pp_shard_off() {
"OFF(all-primary)"
} else {
"per-stage"
},
if memra_engine::pp::pp2_streams_off() {
"OFF(same-stream)"
} else {
"per-stage"
},
);
println!("pp mode: batched stage-split exactness battery over {n_layers} layers; {knobs}");
println!("pp mode: batches={batches:?} steps={steps} reps={reps} (split arm)");
let b1_live = HybridModel::b1_fast_on();
HybridModel::set_b1_fast(false);
if model.cfg.step35.is_some() {
println!(
"pp mode: B=1 fast path inapplicable for Step35 \
(live correctness default = batched; batched reference pinned)"
);
} else if !model.b1_fast_arch_eligible() {
println!(
"pp mode: B=1 fast path inapplicable for Qwen35-MoE \
(live load-stable default = batched; batched reference pinned)"
);
} else {
println!(
"pp mode: B=1 fast path pinned OFF (live default = {})",
if b1_live { "ON" } else { "OFF" }
);
}
let max_b = *batches.iter().max().unwrap_or(&1);
let ctx = (plen as usize + 5 * max_b) + 512 + steps + 64;
let mut fails = 0usize;
let dual = memra_engine::pp::dual_pp_on();
let exact_wave_cap = if model.cfg.step35.is_some() {
8
} else if model.decode_batch_exact16_ok() {
16
} else {
8
};
let overlap_env = std::env::var_os("MEMRA_PP_OVERLAP");
let host_bounce_env = std::env::var_os("MEMRA_PP_HOST_BOUNCE");
let dual_env = std::env::var_os("MEMRA_DUAL_PP");
if dual {
unsafe {
std::env::set_var("MEMRA_DUAL_PP", "1");
}
unsafe {
std::env::set_var("MEMRA_PP_OVERLAP", "0");
}
let prime_pipe_env = std::env::var_os("MEMRA_PRIME_PIPE");
unsafe {
std::env::set_var("MEMRA_PRIME_PIPE", "0");
}
let prompts: Vec<Vec<u32>> = (0..2)
.map(|i| {
(0..plen + i * 5)
.map(|j| 55 + seed * 13 + i * 97 + j * 31)
.collect()
})
.collect();
let mut caches: Vec<Cache> = Vec::with_capacity(2);
for p in &prompts {
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, p, &mut c, 0)?;
caches.push(c);
}
match prime_pipe_env {
Some(value) => unsafe {
std::env::set_var("MEMRA_PRIME_PIPE", value);
},
None => unsafe {
std::env::remove_var("MEMRA_PRIME_PIPE");
},
}
let before: Vec<usize> = caches.iter().map(|c| c.pos).collect();
let toks: Vec<u32> = prompts.iter().map(|p| *p.last().unwrap()).collect();
let mut refs: Vec<&mut Cache> = caches.iter_mut().collect();
let result = model.decode_step_batch(e, &toks, &mut refs);
drop(refs);
match result {
Err(err)
if err.to_string() == memra_engine::pp::DUAL_PP_SINGLE_SLOT_REFUSAL
&& caches.iter().map(|c| c.pos).eq(before.iter().copied()) =>
{
println!("dual pp negative PASS: {}", err);
}
Err(err) => {
println!("dual pp negative FAIL: wrong refusal or cache mutation: {err}");
fails += 1;
}
Ok(rows) => {
println!(
"dual pp negative FAIL: single-slot cell produced {} token row(s)",
rows.len()
);
fails += 1;
}
}
unsafe {
std::env::set_var("MEMRA_PP_OVERLAP", "1");
}
unsafe {
std::env::set_var("MEMRA_PP_HOST_BOUNCE", "1");
}
let before: Vec<usize> = caches.iter().map(|c| c.pos).collect();
let mut refs: Vec<&mut Cache> = caches.iter_mut().collect();
let result = model.decode_step_batch(e, &toks, &mut refs);
drop(refs);
match result {
Err(err)
if err.to_string() == memra_engine::pp::DUAL_PP_HOST_BOUNCE_REFUSAL
&& caches.iter().map(|c| c.pos).eq(before.iter().copied()) =>
{
println!("dual pp host-bounce negative PASS: {}", err);
}
Err(err) => {
println!(
"dual pp host-bounce negative FAIL: wrong refusal or cache mutation: {err}"
);
fails += 1;
}
Ok(rows) => {
println!(
"dual pp host-bounce negative FAIL: cell produced {} token row(s)",
rows.len()
);
fails += 1;
}
}
match &host_bounce_env {
Some(value) => unsafe {
std::env::set_var("MEMRA_PP_HOST_BOUNCE", value);
},
None => unsafe {
std::env::remove_var("MEMRA_PP_HOST_BOUNCE");
},
}
}
for &b in batches {
let oracle_mid = if dual && b > exact_wave_cap {
let mid = memra_engine::pp::dual_pp_wave_mid(b)
.expect("a width above the exact wave cap must have two waves");
assert!(
mid <= exact_wave_cap && b - mid <= exact_wave_cap,
"B={b} cannot fit two exact oracle waves capped at {exact_wave_cap}"
);
Some(mid)
} else {
None
};
let prompts: Vec<Vec<u32>> = (0..b)
.map(|i| {
(0..plen + i as u32 * 5)
.map(|j| 55 + seed * 13 + i as u32 * 97 + j * 31)
.collect()
})
.collect();
unsafe {
std::env::remove_var("MEMRA_PP_STAGES");
}
let mut inputs: Vec<Vec<u32>> = Vec::with_capacity(steps);
let mut ref_logits: Vec<Vec<Vec<f32>>> = Vec::with_capacity(steps);
{
let mut caches: Vec<Cache> = Vec::with_capacity(b);
for p in prompts.iter() {
let mut c = Cache::new(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, p, &mut c, 0)?;
caches.push(c);
}
let mut toks: Vec<u32> = prompts.iter().map(|p| *p.last().unwrap()).collect();
for _ in 0..steps {
inputs.push(toks.clone());
let rows = decode_batch_serial_waves(e, model, &toks, &mut caches, oracle_mid)?;
for (bi, l) in rows.iter().enumerate() {
toks[bi] = argmax(l) as u32;
}
ref_logits.push(rows);
}
}
let n_vocab = ref_logits[0][0].len();
println!(
"-- B={b}: reference recorded ({steps} steps x {b} rows x {n_vocab} f32, \
door OFF over the sharded placement{})",
oracle_mid
.map(|mid| format!(", serial waves {mid}+{}", b - mid))
.unwrap_or_default()
);
unsafe {
std::env::set_var("MEMRA_PP_STAGES", stages.to_string());
}
for rep in 0..reps.max(1) {
let overlaps0 = memra_engine::pp::dual_pp_overlaps();
let mut chk = BitCheck::new(format!("split B={b} rep{rep}"));
let mut caches: Vec<Cache> = Vec::with_capacity(b);
for p in prompts.iter() {
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, p, &mut c, 0)?;
caches.push(c);
}
for (s, toks) in inputs.iter().enumerate() {
let mut refs: Vec<&mut Cache> = caches.iter_mut().collect();
let rows = model.decode_step_batch(e, toks, &mut refs)?;
for (bi, l) in rows.iter().enumerate() {
chk.check(s, bi, l, &ref_logits[s][bi]);
}
}
fails += chk.verdict();
if dual && b >= 2 {
let overlaps = memra_engine::pp::dual_pp_overlaps() - overlaps0;
if overlaps == 0 {
println!(
"dual pp liveness FAIL [B={b} rep{rep}]: DUAL_PP_OVERLAPS did not advance"
);
fails += 1;
} else {
println!(
"dual pp liveness PASS [B={b} rep{rep}]: DUAL_PP_OVERLAPS +{overlaps}"
);
}
}
}
{
let mut chk = BitCheck::new(format!("unsplit@ppncache B={b}"));
let mut caches: Vec<Cache> = Vec::with_capacity(b);
for p in prompts.iter() {
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, p, &mut c, 0)?;
caches.push(c);
}
unsafe {
std::env::set_var("MEMRA_BATCH_PP", "0");
std::env::set_var("MEMRA_PP_ALLOW_UNSPLIT_BATCH", "1");
}
let r = (|| -> Result<(), Box<dyn std::error::Error>> {
for (s, toks) in inputs.iter().enumerate() {
let rows = decode_batch_serial_waves(e, model, toks, &mut caches, oracle_mid)?;
for (bi, l) in rows.iter().enumerate() {
chk.check(s, bi, l, &ref_logits[s][bi]);
}
}
Ok(())
})();
unsafe {
std::env::remove_var("MEMRA_BATCH_PP");
std::env::remove_var("MEMRA_PP_ALLOW_UNSPLIT_BATCH");
}
r?;
fails += chk.verdict();
}
}
if model.cfg.step35.is_none() && model.b1_fast_arch_eligible() {
let mut chk = BitCheck::new("b1-stagefast vs eager-ppn B=1".to_string());
let prompt: Vec<u32> = (0..24u32).map(|j| 55 + seed * 13 + j * 31).collect();
let n_s = steps.min(16);
HybridModel::set_b1_fast(true);
let mut c_eager = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, &prompt, &mut c_eager, 0)?;
let mut c_batch = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, &prompt, &mut c_batch, 0)?;
let mut tok = *prompt.last().unwrap();
for s in 0..n_s {
let (ref_row, _) = model.decode_step_h(e, tok, &mut c_eager)?;
let got = {
let mut refs: Vec<&mut Cache> = vec![&mut c_batch];
model.decode_step_batch(e, &[tok], &mut refs)?
};
chk.check(s, 0, &got[0], &ref_row);
tok = argmax(&ref_row) as u32;
assert_eq!(
c_eager.pos, c_batch.pos,
"b1-stagefast pos {} != eager pos {} at step {s} — one arm advanced \
the cache differently",
c_batch.pos, c_eager.pos
);
}
fails += chk.verdict();
HybridModel::set_b1_fast(false);
} else if model.cfg.step35.is_some() {
println!(
"pp gate: b1-stagefast arm N/A for Step35 (correctness default is batched at every width)"
);
} else {
println!(
"pp gate: b1-stagefast arm N/A for Qwen35-MoE (load-stable correctness default is batched at every width)"
);
}
{
let b = *batches.iter().max().unwrap();
let prompts: Vec<Vec<u32>> = (0..b)
.map(|i| {
(0..plen + i as u32 * 5)
.map(|j| 55 + seed * 13 + i as u32 * 97 + j * 31)
.collect()
})
.collect();
let n_s = steps.min(8);
let mut ep_fail = 0usize;
let mut caches_f: Vec<Cache> = Vec::with_capacity(b);
let mut caches_l: Vec<Cache> = Vec::with_capacity(b);
for p in prompts.iter() {
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, p, &mut c, 0)?;
caches_f.push(c);
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, p, &mut c, 0)?;
caches_l.push(c);
}
let mut toks: Vec<u32> = prompts.iter().map(|p| *p.last().unwrap()).collect();
for _ in 0..n_s {
let samp: Vec<Option<DevSamp>> = (0..b)
.map(|bi| {
if bi % 2 == 0 {
Some((0.0, 0, 0, 0i32, 1.0f32, 0.0f32))
} else {
None
}
})
.collect();
let (rows_f, next_f) = {
let mut refs: Vec<&mut Cache> = caches_f.iter_mut().collect();
model.decode_step_batch_sampled_lean(e, &toks, &mut refs, &samp, false)?
};
let (rows_l, next_l) = {
let mut refs: Vec<&mut Cache> = caches_l.iter_mut().collect();
model.decode_step_batch_sampled_lean(e, &toks, &mut refs, &samp, true)?
};
for bi in 0..b {
if samp[bi].is_some() {
let host_am = argmax(&rows_f[bi]) as u32;
let dev = next_f[bi].expect("split greedy row missing device token");
if dev != host_am {
println!(
"pp gate epilogue row {bi}: device argmax {dev} != host \
argmax {host_am} FAIL"
);
ep_fail += 1;
}
if next_l[bi] != next_f[bi] {
println!(
"pp gate epilogue row {bi}: lean token {:?} != full token \
{:?} FAIL",
next_l[bi], next_f[bi]
);
ep_fail += 1;
}
if !rows_l[bi].is_empty() {
println!("pp gate epilogue row {bi}: lean sampled row NOT empty FAIL");
ep_fail += 1;
}
let parked = e.dtoh(
caches_l[bi]
.last_logits_dev
.as_ref()
.expect("lean row missing device park"),
)?;
let r = &rows_f[bi];
if !(parked.len() == r.len()
&& parked
.iter()
.zip(r.iter())
.all(|(a, b)| a.to_bits() == b.to_bits()))
{
println!(
"pp gate epilogue row {bi}: parked device logits != full \
host row FAIL"
);
ep_fail += 1;
}
toks[bi] = host_am;
} else {
let (r, l) = (&rows_f[bi], &rows_l[bi]);
if !(r.len() == l.len()
&& r.iter()
.zip(l.iter())
.all(|(a, b)| a.to_bits() == b.to_bits()))
{
println!("pp gate epilogue row {bi}: unsampled row lean != full FAIL");
ep_fail += 1;
}
toks[bi] = argmax(r) as u32;
}
}
if ep_fail > 8 {
break;
}
}
println!(
"pp gate {} [epilogue B={b}]: last-stage device sampling + lean park",
if ep_fail == 0 { "PASS" } else { "FAIL" }
);
fails += usize::from(ep_fail > 0);
}
HybridModel::set_b1_fast(b1_live);
if dual {
match overlap_env {
Some(value) => unsafe {
std::env::set_var("MEMRA_PP_OVERLAP", value);
},
None => unsafe {
std::env::remove_var("MEMRA_PP_OVERLAP");
},
}
match dual_env {
Some(value) => unsafe {
std::env::set_var("MEMRA_DUAL_PP", value);
},
None => unsafe {
std::env::remove_var("MEMRA_DUAL_PP");
},
}
}
println!("pp mode verdict: {fails} failing arm(s); {knobs}");
Ok(fails)
}
fn ppspec_battery(
e: &Engine,
model: &HybridModel,
stages: usize,
rounds: usize,
ts: &[usize],
reps: usize,
seed: u32,
) -> Result<usize, Box<dyn std::error::Error>> {
let n_layers = model.layers.len();
let fence = memra_engine::pp::pp_cuts(n_layers).unwrap_or_else(|| {
panic!("ppspec mode: door failed to open (n_layers={n_layers}, stages={stages})")
});
assert_eq!(
fence.len() - 1,
stages,
"fence {fence:?} != stages {stages}"
);
let devices = std::env::var("MEMRA_PP_DEVICES").unwrap_or_default();
let knobs = format!(
"stages={stages} fence={fence:?} devices={} splits={} shard={} streams={}",
if devices.is_empty() {
"default(primary)".into()
} else {
devices.clone()
},
std::env::var("MEMRA_PP_SPLITS").unwrap_or_else(|_| "default(even)".into()),
if memra_engine::pp::pp_shard_off() {
"OFF(all-primary)"
} else {
"per-stage"
},
if memra_engine::pp::pp2_streams_off() {
"OFF(same-stream)"
} else {
"per-stage"
},
);
println!("ppspec mode: verify stage-split exactness battery over {n_layers} layers; {knobs}");
println!("ppspec mode: T={ts:?} rounds={rounds} reps={reps} (split arm)");
let mut fails = 0usize;
for &t in ts {
assert!(t >= 1, "verify width T must be >= 1");
let prompt: Vec<u32> = (0..24u32).map(|j| 55 + seed * 13 + j * 31).collect();
let ctx = 512 + rounds * t + 64;
unsafe {
std::env::remove_var("MEMRA_PP_STAGES");
}
let mut inputs: Vec<(usize, Vec<u32>)> = Vec::with_capacity(rounds);
let mut ref_logits: Vec<Vec<f32>> = Vec::with_capacity(rounds);
let mut ref_seed: Vec<Vec<f32>> = Vec::with_capacity(rounds);
let mut ref_pos: Vec<usize> = Vec::with_capacity(rounds);
{
let mut c = Cache::new(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, &prompt, &mut c, 0)?;
let mut chunk: Vec<u32> = vec![*prompt.last().unwrap(); t];
for _ in 0..rounds {
let pos0 = c.pos;
inputs.push((pos0, chunk.clone()));
let (l, hs) = model.decode_step_t_h(e, &chunk, pos0, &mut c)?;
let n_vocab = l.len() / t;
chunk = (0..t)
.map(|j| argmax(&l[j * n_vocab..(j + 1) * n_vocab]) as u32)
.collect();
ref_logits.push(l);
ref_seed.push(e.dtoh(&hs)?);
ref_pos.push(c.pos);
}
}
let n_vocab = ref_logits[0].len() / t;
println!(
"-- T={t}: reference recorded ({rounds} rounds x {t} cols x {n_vocab} f32 \
+ h_seed, door OFF over the sharded placement)"
);
unsafe {
std::env::set_var("MEMRA_PP_STAGES", stages.to_string());
}
for rep in 0..reps.max(1) {
let mut chk = BitCheck::new(format!("verify-split T={t} rep{rep}"));
let mut hchk = BitCheck::new(format!("verify-split h_seed T={t} rep{rep}"));
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, &prompt, &mut c, 0)?;
for (r, (pos0, chunk)) in inputs.iter().enumerate() {
assert_eq!(
c.pos, *pos0,
"verify-split pos {} != reference pos {pos0} at round {r} — one arm \
advanced the cache differently",
c.pos
);
let (l, hs) = model.decode_step_t_h(e, chunk, *pos0, &mut c)?;
for j in 0..t {
chk.check(
r,
j,
&l[j * n_vocab..(j + 1) * n_vocab],
&ref_logits[r][j * n_vocab..(j + 1) * n_vocab],
);
}
hchk.check(r, 0, &e.dtoh(&hs)?, &ref_seed[r]);
assert_eq!(
c.pos, ref_pos[r],
"verify-split advanced pos to {} vs reference {} \
at round {r}",
c.pos, ref_pos[r]
);
}
fails += chk.verdict();
fails += hchk.verdict();
}
{
let mut chk = BitCheck::new(format!("verify-unsplit@ppncache T={t}"));
let mut c = memra_engine::pp::new_cache(e, &model.cfg, ctx)?;
let _ = model.prime_cache(e, &prompt, &mut c, 0)?;
unsafe {
std::env::set_var("MEMRA_SPEC_PP", "0");
std::env::set_var("MEMRA_PP_ALLOW_UNSPLIT_BATCH", "1");
}
let r = (|| -> Result<(), Box<dyn std::error::Error>> {
for (r, (pos0, chunk)) in inputs.iter().enumerate() {
let (l, _) = model.decode_step_t_h(e, chunk, *pos0, &mut c)?;
for j in 0..t {
chk.check(
r,
j,
&l[j * n_vocab..(j + 1) * n_vocab],
&ref_logits[r][j * n_vocab..(j + 1) * n_vocab],
);
}
}
Ok(())
})();
unsafe {
std::env::remove_var("MEMRA_SPEC_PP");
std::env::remove_var("MEMRA_PP_ALLOW_UNSPLIT_BATCH");
}
r?;
fails += chk.verdict();
}
}
println!("ppspec mode verdict: {fails} failing arm(s); {knobs}");
Ok(fails)
}