use memra_engine::Engine;
use memra_engine::qwen4exp_gpu::{LoadOptions, Qwen4ExpGpu, read_checkpoint, read_checkpoint_with};
use memra_gguf::config::{HfConfig, ModelConfig};
use memra_gguf::model_plan::{AttentionPlan, ModelPlan, RopeFactors};
use memra_gguf::tensor_contract::{
CheckpointDialect, ContractOptions, QuantConstraint, TensorMatch,
};
use memra_reference::{ReferenceWeights, deterministic_fixture, execute};
use sha2::{Digest, Sha256};
use std::path::{Path, PathBuf};
type Res<T> = Result<T, Box<dyn std::error::Error>>;
const TINY_CONFIG: &str = r#"{"model_type":"qwen4_exp","text_config":{
"model_type":"qwen4_exp_text","eos_token_id":63,
"num_hidden_layers":4,"hidden_size":16,
"num_attention_heads":2,"num_key_value_heads":1,"head_dim":8,
"intermediate_size":32,"vocab_size":64,"max_position_embeddings":64,
"rms_norm_eps":0.000001,"full_attention_interval":4,
"rope_parameters":{"rope_theta":10000,"partial_rotary_factor":0.25,
"mrope_section":[1,1,2],"mrope_interleaved":true},
"linear_conv_kernel_dim":4,"linear_key_head_dim":4,"linear_value_head_dim":4,
"linear_num_key_heads":1,"linear_num_value_heads":2,
"num_experts":8,"num_experts_per_tok":2,"moe_intermediate_size":8,
"shared_expert_intermediate_size":8,
"indexer_n_heads":1,"indexer_kv_heads":1,"indexer_head_dim":4,
"indexer_compress_ratio":4,"indexer_budget":8,
"hc_count":2,"hc_lowrank":4,
"ngram_size":3,"heads_per_ngram":2,"ngram_vocab_size_base":64,
"make_ngram_vocab_size_divisible_by":128,"split_ngram_parts":2,
"ple_layer_ids":[2],"ple_embed_dim":16,"ple_conv_kernel_size":4,
"output_gate_type":"sigmoid",
"mtp":{"num_hidden_layers":1,"rope_theta":10000},"mtp_num_hidden_layers":1}}"#;
fn argmax(row: &[f32]) -> usize {
let mut best = 0;
for (index, &value) in row.iter().enumerate() {
if value > row[best] {
best = index;
}
}
best
}
struct RowStats {
max_abs: f32,
max_rel: f32,
ref_absmax: f32,
argmax_match: bool,
}
fn compare_row(reference: &[f32], candidate: &[f32]) -> RowStats {
let mut max_abs = 0.0f32;
let mut max_rel = 0.0f32;
let mut ref_absmax = 0.0f32;
for (&r, &c) in reference.iter().zip(candidate) {
let abs = (r - c).abs();
max_abs = max_abs.max(abs);
max_rel = max_rel.max(abs / r.abs().max(1.0));
ref_absmax = ref_absmax.max(r.abs());
}
RowStats {
max_abs,
max_rel,
ref_absmax,
argmax_match: argmax(reference) == argmax(candidate),
}
}
fn hash_name(name: &str) -> u64 {
let mut value = 0xcbf2_9ce4_8422_2325u64;
for byte in name.bytes() {
value ^= byte as u64;
value = value.wrapping_mul(0x0000_0100_0000_01b3);
}
value
}
fn gen_f32(name: &str, elements: usize, center: f32, scale: f32) -> Vec<f32> {
let salt = hash_name(name);
(0..elements)
.map(|index| {
let mut value = index as u64 ^ salt.wrapping_mul(0x9e37_79b9);
value ^= value >> 16;
value = value.wrapping_mul(0x45d9_f3b);
value ^= value >> 16;
let unit = (value as u32) as f32 / u32::MAX as f32;
center + (2.0 * unit - 1.0) * scale
})
.collect()
}
fn gen_bytes(name: &str, elements: usize, lo: u8, hi: u8) -> Vec<u8> {
let salt = hash_name(name);
(0..elements)
.map(|index| {
let mut value = index as u64 ^ salt.wrapping_mul(0x9e37_79b9);
value ^= value >> 13;
value = value.wrapping_mul(0x45d9_f3b);
(lo as u64 + value % (hi as u64 - lo as u64 + 1)) as u8
})
.collect()
}
fn f32_to_bf16(value: f32) -> u16 {
let bits = value.to_bits();
let rounded = bits.wrapping_add(0x7FFF + ((bits >> 16) & 1));
(rounded >> 16) as u16
}
fn bf16_bytes(values: &[f32]) -> Vec<u8> {
values
.iter()
.flat_map(|&v| f32_to_bf16(v).to_le_bytes())
.collect()
}
struct StEntry {
name: String,
dtype: &'static str,
shape: Vec<u64>,
bytes: Vec<u8>,
}
fn write_safetensors(path: &Path, entries: &[StEntry]) -> Res<()> {
let mut header = String::from("{");
let mut offset = 0usize;
for (index, entry) in entries.iter().enumerate() {
if index > 0 {
header.push(',');
}
let end = offset + entry.bytes.len();
let shape = entry
.shape
.iter()
.map(|d| d.to_string())
.collect::<Vec<_>>()
.join(",");
header.push_str(&format!(
"\"{}\":{{\"dtype\":\"{}\",\"shape\":[{}],\"data_offsets\":[{},{}]}}",
entry.name, entry.dtype, shape, offset, end
));
offset = end;
}
header.push('}');
let mut out = Vec::with_capacity(8 + header.len() + offset);
out.extend_from_slice(&(header.len() as u64).to_le_bytes());
out.extend_from_slice(header.as_bytes());
for entry in entries {
out.extend_from_slice(&entry.bytes);
}
std::fs::write(path, out)?;
Ok(())
}
fn first_primes_at_least(base: i64, count: usize) -> Vec<i64> {
let is_prime = |n: i64| {
if n < 2 {
return false;
}
let mut d = 2;
while d * d <= n {
if n % d == 0 {
return false;
}
d += 1;
}
true
};
let mut out = Vec::with_capacity(count);
let mut candidate = base.max(2);
while out.len() < count {
if is_prime(candidate) {
out.push(candidate);
}
candidate += 1;
}
out
}
#[derive(Clone, Copy, PartialEq)]
enum DirKind {
Bf16Fused,
Nvfp4Stacked,
Nvfp4PerExpert,
}
fn synthesize_dir(dir: &Path, cfg: &ModelConfig, plan: &ModelPlan, kind: DirKind) -> Res<()> {
use memra_gguf::model_packs::qwen4_exp::{ExpertDialect, tensor_contract_for};
use memra_gguf::tensor_contract::{LayerTensor, TensorId, TensorOwner};
let pack = memra_gguf::model_packs::for_config(cfg).ok_or("no pack for tiny config")?;
let contract = match kind {
DirKind::Nvfp4PerExpert => {
tensor_contract_for(cfg, plan, ExpertDialect::PerExpertModelopt)?
}
_ => pack.compile_tensor_contract(
cfg,
plan,
CheckpointDialect::HfSafetensors,
ContractOptions::default(),
)?,
};
std::fs::create_dir_all(dir)?;
std::fs::write(dir.join("config.json"), TINY_CONFIG)?;
let n_trunk = plan.layers.len() as u32;
let mut entries: Vec<StEntry> = Vec::new();
for requirement in &contract.requirements {
let elements = |shape: &[u64]| shape.iter().map(|&d| d as usize).product::<usize>();
if requirement.match_mode == TensorMatch::All {
for name in &requirement.names {
let n = elements(&requirement.shape);
entries.push(StEntry {
name: name.clone(),
dtype: "BF16",
shape: requirement.shape.clone(),
bytes: bf16_bytes(&gen_f32(name, n, 0.0, 0.2)),
});
}
continue;
}
let name = requirement.names[0].clone();
if matches!(requirement.quant, QuantConstraint::I64) {
let q = cfg.qwen4exp.as_ref().ok_or("tiny config lost qwen4exp")?;
let heads = q.ngram_heads() as usize;
let ints: Vec<i64> = if name.ends_with("layer_multipliers") {
(0..q.ngram_size as i64)
.map(|i| 1_000_003 + 2 * i * 31)
.collect()
} else {
let sizes = first_primes_at_least(q.ngram_vocab_size_base as i64, heads);
if name.ends_with("ngram_heads_vocab_sizes") {
sizes
} else {
let mut offsets = Vec::with_capacity(heads);
let mut total = 0i64;
for &size in &sizes {
offsets.push(total);
total += size;
}
offsets
}
};
entries.push(StEntry {
name,
dtype: "I64",
shape: requirement.shape.clone(),
bytes: ints.iter().flat_map(|v| v.to_le_bytes()).collect(),
});
continue;
}
if requirement.quant == QuantConstraint::Nvfp4 {
let (out_f, in_f) = (requirement.shape[0], requirement.shape[1]);
let stem = name
.strip_suffix(".weight")
.ok_or("NVFP4 row without .weight")?;
entries.push(StEntry {
name: name.clone(),
dtype: "U8",
shape: vec![out_f, in_f / 2],
bytes: gen_bytes(&name, (out_f * in_f / 2) as usize, 0, 255),
});
entries.push(StEntry {
name: format!("{stem}.weight_scale"),
dtype: "F8_E4M3",
shape: vec![out_f, in_f / 16],
bytes: gen_bytes(
&format!("{name}.weight_scale"),
(out_f * in_f / 16) as usize,
0x28,
0x40,
),
});
entries.push(StEntry {
name: format!("{stem}.weight_scale_2"),
dtype: "F32",
shape: vec![],
bytes: 5.9945243e-5f32.to_le_bytes().to_vec(),
});
entries.push(StEntry {
name: format!("{stem}.input_scale"),
dtype: "F32",
shape: vec![],
bytes: 0.0078125f32.to_le_bytes().to_vec(), });
continue;
}
let trunk_gate_up = matches!(
(&requirement.id, requirement.owner),
(
TensorId::Layer { index, tensor },
TensorOwner::Layer(_)
) if *index < n_trunk && *tensor == LayerTensor::MoeExpertGateUpBank
);
if kind == DirKind::Nvfp4Stacked && trunk_gate_up {
let (n_expert, out_f, in_f) = (
requirement.shape[0],
requirement.shape[1],
requirement.shape[2],
);
let codes = gen_bytes(&name, (n_expert * out_f * in_f / 2) as usize, 0, 255);
let scales = gen_bytes(
&format!("{name}.weight_scale"),
(n_expert * out_f * in_f / 16) as usize,
0x28,
0x40,
);
entries.push(StEntry {
name: name.clone(),
dtype: "U8",
shape: vec![n_expert, out_f, in_f / 2],
bytes: codes,
});
entries.push(StEntry {
name: format!("{name}.weight_scale"),
dtype: "F8_E4M3",
shape: vec![n_expert, out_f, in_f / 16],
bytes: scales,
});
entries.push(StEntry {
name: format!("{name}.weight_scale_2"),
dtype: "F32",
shape: vec![n_expert],
bytes: vec![0.03125f32; n_expert as usize] .iter()
.flat_map(|v| v.to_le_bytes())
.collect(),
});
continue;
}
let n = elements(&requirement.shape);
let (center, scale) = if name.ends_with("linear_attn.norm.weight") {
(1.0, 0.1)
} else if name.contains("norm") && name.ends_with(".weight") {
(0.0, 0.2)
} else if name.ends_with("linear_attn.A_log") {
(-0.7, 0.3)
} else if name.ends_with("linear_attn.dt_bias") {
(0.0, 0.1)
} else {
let in_f = *requirement.shape.last().unwrap() as f32;
(0.0, 1.0 / in_f.sqrt())
};
entries.push(StEntry {
name: name.clone(),
dtype: "BF16",
shape: requirement.shape.clone(),
bytes: bf16_bytes(&gen_f32(&name, n, center, scale)),
});
}
write_safetensors(&dir.join("model.safetensors"), &entries)?;
Ok(())
}
struct TempDir(PathBuf);
impl Drop for TempDir {
fn drop(&mut self) {
let _ = std::fs::remove_dir_all(&self.0);
}
}
struct ArmResult {
prefill_worst: (f32, f32),
decode_worst: (f32, f32),
}
#[allow(clippy::too_many_arguments)]
fn run_arm(
label: &str,
e: &Engine,
plan: &ModelPlan,
weights: &ReferenceWeights,
model: &Qwen4ExpGpu,
prompt: &[u32],
decode_feed: &[u32],
lines: &mut Vec<String>,
failures: &mut usize,
) -> Res<ArmResult> {
const MAX_ABS: f32 = 0.01;
const MAX_REL: f32 = 0.01;
let n = prompt.len();
let vocab = plan.vocab_size as usize;
let all_tokens: Vec<u32> = prompt.iter().chain(decode_feed.iter()).copied().collect();
let reference = execute(plan, weights, &all_tokens)?;
let mut state = model.alloc_state(e, all_tokens.len())?;
let mut record = |phase: &str, row: usize, stats: &RowStats, failures: &mut usize| {
let passed = stats.max_abs <= MAX_ABS && stats.max_rel <= MAX_REL && stats.argmax_match;
if !passed {
*failures += 1;
}
lines.push(format!(
"{label}\t{phase}\trow={row}\tmax_abs={:.3e}\tmax_rel={:.3e}\tref_absmax={:.3e}\targmax_match={}\tpass={passed}",
stats.max_abs, stats.max_rel, stats.ref_absmax, stats.argmax_match
));
};
let prefill = model.prefill(e, prompt, &mut state)?;
assert_eq!(prefill.len(), n * vocab);
let mut prefill_worst = (0.0f32, 0.0f32);
for row in 0..n {
let stats = compare_row(
&reference.logits[row * vocab..(row + 1) * vocab],
&prefill[row * vocab..(row + 1) * vocab],
);
prefill_worst.0 = prefill_worst.0.max(stats.max_abs);
prefill_worst.1 = prefill_worst.1.max(stats.max_rel);
record("prefill", row, &stats, failures);
}
let mut decode_worst = (0.0f32, 0.0f32);
for (step, &token) in decode_feed.iter().enumerate() {
let logits = model.decode_step(e, token, &mut state)?;
assert_eq!(logits.len(), vocab);
let row = n + step;
let stats = compare_row(&reference.logits[row * vocab..(row + 1) * vocab], &logits);
decode_worst.0 = decode_worst.0.max(stats.max_abs);
decode_worst.1 = decode_worst.1.max(stats.max_rel);
record("decode-step", row, &stats, failures);
}
Ok(ArmResult {
prefill_worst,
decode_worst,
})
}
#[allow(clippy::too_many_arguments)]
fn run_mtp_arm(
label: &str,
e: &Engine,
plan: &ModelPlan,
weights: &ReferenceWeights,
model: &Qwen4ExpGpu,
tokens: &[u32],
lines: &mut Vec<String>,
failures: &mut usize,
) -> Res<String> {
const MAX_ABS: f32 = 0.01;
const MAX_REL: f32 = 0.01;
let vocab = plan.vocab_size as usize;
let t = tokens.len();
let reference = execute(plan, weights, tokens)?;
let mtp_ref = reference
.mtp
.first()
.ok_or("reference produced no MTP output")?;
let trunk_wide = reference
.layer_hidden
.last()
.ok_or("reference produced no trunk hidden")?;
let wide = trunk_wide.len() / t;
let wide_dev = e.htod(trunk_wide)?;
let mut worst = (0.0f32, 0.0f32);
let mut record = |phase: &str,
row: usize,
stats: &RowStats,
with_argmax: bool,
worst: &mut (f32, f32),
failures: &mut usize,
lines: &mut Vec<String>| {
let passed = stats.max_abs <= MAX_ABS
&& stats.max_rel <= MAX_REL
&& (!with_argmax || stats.argmax_match);
if !passed {
*failures += 1;
}
worst.0 = worst.0.max(stats.max_abs);
worst.1 = worst.1.max(stats.max_rel);
lines.push(format!(
"{label}\t{phase}\trow={row}\tmax_abs={:.3e}\tmax_rel={:.3e}\tref_absmax={:.3e}\targmax_match={}\tpass={passed}",
stats.max_abs, stats.max_rel, stats.ref_absmax, stats.argmax_match
));
};
let mut batched = model.mtp_state(e, t + 1)?;
let (logits_d, carrier_d) =
model.mtp_draft_forward(e, tokens, &wide_dev, 0, &mut batched, 0, true)?;
let logits = e.dtoh_view(&logits_d.slice(0..t * vocab))?;
let carrier = e.dtoh_view(&carrier_d.slice(0..t * wide))?;
for row in 0..t {
let stats = compare_row(
&mtp_ref.logits[row * vocab..(row + 1) * vocab],
&logits[row * vocab..(row + 1) * vocab],
);
record(
"mtp-batched",
row,
&stats,
true,
&mut worst,
failures,
lines,
);
}
for row in 0..t {
let stats = compare_row(
&mtp_ref.hidden[row * wide..(row + 1) * wide],
&carrier[row * wide..(row + 1) * wide],
);
record(
"mtp-carrier",
row,
&stats,
false,
&mut worst,
failures,
lines,
);
}
let mut chained = model.mtp_state(e, t + 1)?;
for row in 0..t {
let (ld, cd) = model.mtp_draft_forward(
e,
&tokens[row..row + 1],
&wide_dev,
row,
&mut chained,
0,
true,
)?;
let step_logits = e.dtoh_view(&ld.slice(0..vocab))?;
let stats = compare_row(
&mtp_ref.logits[row * vocab..(row + 1) * vocab],
&step_logits,
);
record("mtp-step", row, &stats, true, &mut worst, failures, lines);
model.mtp_recycle(&mut chained, ld, cd);
}
Ok(format!(
"{label}: draft parity worst abs {:.3e} rel {:.3e} over batched+carrier+steps",
worst.0, worst.1
))
}
#[allow(clippy::too_many_arguments)]
fn run_defer_arm(
label: &str,
e: &Engine,
model: &mut Qwen4ExpGpu,
vocab: usize,
prompt: &[u32],
lines: &mut Vec<String>,
failures: &mut usize,
with_trim: bool,
) -> Res<String> {
use memra_engine::qwen4exp_gpu::SpecOpts;
let max_new = 24usize;
let spec_k = 3usize;
let cap = prompt.len() + max_new + spec_k + 4;
let mut plain_state = model.alloc_state(e, cap)?;
let logits = model.prefill(e, prompt, &mut plain_state)?;
let mut next = argmax(&logits[(prompt.len() - 1) * vocab..]) as u32;
let mut plain = vec![next];
for _ in 1..max_new {
let row = model.decode_step(e, next, &mut plain_state)?;
next = argmax(&row) as u32;
plain.push(next);
}
let mut checked = 0usize;
let check_config_k = |model: &Qwen4ExpGpu,
sk: usize,
cfg_name: &str,
arms: &[(&str, SpecOpts)],
lines: &mut Vec<String>,
failures: &mut usize|
-> Res<()> {
let mut base: Option<(Vec<u32>, (usize, u64, u64, usize, usize))> = None;
for (arm_name, opts) in arms {
let scap = prompt.len() + max_new + sk + 4;
let mut ss = model.alloc_state(e, scap)?;
let mut ds = model.mtp_state(e, scap)?;
let report = model.spec_generate_ext(
e, e, prompt, max_new, sk, &mut ss, &mut ds, None, *opts, None,
)?;
let counters = (
report.rounds,
report.drafted,
report.accepted,
report.guard_stops,
report.zero_draft_rounds,
);
let vs_plain = report.tokens == plain;
let vs_base = match base.as_ref() {
Some((toks, ctrs)) => &report.tokens == toks && counters == *ctrs,
None => true,
};
let pass = vs_plain && vs_base;
if !pass {
*failures += 1;
}
lines.push(format!(
"{label}\t{cfg_name}\t{arm_name}\tk={sk}\ttokens={max_new}\trounds={}\tdrafted={}\taccepted={}\tguard_stops={}\tzero_draft={}\tvs_plain={vs_plain}\tvs_host={vs_base}\tpass={pass}",
report.rounds,
report.drafted,
report.accepted,
report.guard_stops,
report.zero_draft_rounds,
));
if base.is_none() {
base = Some((report.tokens, counters));
}
}
Ok(())
};
let mut check_config =
|model: &Qwen4ExpGpu,
cfg_name: &str,
arms: &[(&str, SpecOpts)],
lines: &mut Vec<String>,
failures: &mut usize|
-> Res<()> { check_config_k(model, spec_k, cfg_name, arms, lines, failures) };
let host = SpecOpts::default();
let defer = SpecOpts {
defer: true,
..Default::default()
};
model.arm_spec_devchain(e)?;
check_config(
model,
"pmin0",
&[("host", host), ("defer", defer)],
lines,
failures,
)?;
checked += 2;
let g = |pmin: f32, defer: bool, gsync: bool| SpecOpts {
pmin,
defer,
defer_guard_sync: gsync,
..Default::default()
};
check_config(
model,
"pmin0.5",
&[
("host", g(0.5, false, false)),
("defer", g(0.5, true, false)),
("defer-gsync", g(0.5, true, true)),
],
lines,
failures,
)?;
checked += 3;
let mut mixed = None;
'sweep: for &sk in &[spec_k, 6usize, 2] {
for &p in &[
0.4f32, 0.35, 0.3, 0.25, 0.2, 0.15, 0.1, 0.08, 0.05, 0.02, 0.01, 0.002,
] {
let scap = prompt.len() + max_new + sk + 4;
let mut ss = model.alloc_state(e, scap)?;
let mut ds = model.mtp_state(e, scap)?;
let r = model.spec_generate_ext(
e,
e,
prompt,
max_new,
sk,
&mut ss,
&mut ds,
None,
g(p, false, false),
None,
)?;
if r.guard_stops > r.zero_draft_rounds && r.drafted > 0 {
mixed = Some((sk, p));
break 'sweep;
}
}
}
match mixed {
Some((sk, p)) => {
let name = format!("k{sk}+pmin{p}-midchain");
check_config_k(
model,
sk,
&name,
&[
("host", g(p, false, false)),
("defer", g(p, true, false)),
("defer-gsync", g(p, true, true)),
],
lines,
failures,
)?;
checked += 3;
}
None => {
lines.push(format!(
"{label}\tmidchain-pmin\tNONE\tcovered_by=guard-trunc-pin+box-defer-ab\t(no swept (k,pmin) mixes on the deterministic fixture)"
));
}
}
{
use memra_engine::qwen4exp_gpu::spec_guard_trunc;
let cases: &[(&[f32], f32, usize, &str)] = &[
(
&[0.9, 0.9, 0.1, 0.9],
0.3,
2,
"mid-chain dip, later recovery ignored",
),
(&[0.1, 0.9, 0.9], 0.3, 0, "zero-draft round"),
(&[0.9, 0.8, 0.7], 0.3, 3, "no stop"),
(
&[0.3, 0.29],
0.3,
1,
"boundary: p == pmin passes (strict <)",
),
(&[], 0.3, 0, "empty window"),
(&[0.5, 0.4, 0.3, 0.2], 0.35, 2, "monotone decay"),
];
let mut ok = true;
for (probs, pmin, want, what) in cases {
let got = spec_guard_trunc(probs, *pmin);
if got != *want {
ok = false;
*failures += 1;
lines.push(format!(
"{label}\tguard-trunc-pin\t{what}\tgot={got}\twant={want}\tpass=false"
));
}
}
if ok {
lines.push(format!(
"{label}\tguard-trunc-pin\t{} windows incl. mid-chain dip\tpass=true",
cases.len()
));
checked += 1;
}
}
if with_trim {
let rev: Vec<u32> = (0..vocab as u32).rev().collect();
model.build_draft_trim(e, &rev)?;
model.arm_spec_devchain(e)?;
let p = mixed.map(|(_, p)| p).unwrap_or(0.5);
check_config(
model,
&format!("trim-rev+pmin{p}"),
&[
("host", g(p, false, false)),
("defer", g(p, true, false)),
("defer-gsync", g(p, true, true)),
],
lines,
failures,
)?;
checked += 3;
model.clear_draft_trim();
}
model.clear_spec_devchain();
Ok(format!(
"{label}: deferred-round identity over {checked} arms (byte identity vs plain + counter identity vs host, incl. a mid-chain guard stop)"
))
}
fn main() -> Res<()> {
let receipt_path = std::env::args()
.nth(1)
.ok_or("usage: qwen4exp-gpu-gate <receipt.tsv>")?;
let pack =
memra_gguf::model_packs::by_alias("qwen4_exp").ok_or("qwen4_exp pack is not registered")?;
let plan = pack.compile_tiny_plan()?;
let cfg = ModelConfig::from_hf(&HfConfig::parse(TINY_CONFIG));
assert_eq!(
pack.compile_plan(&cfg)?,
plan,
"gate tiny config drifted from the pack tiny_plan"
);
let prompt: Vec<u32> = vec![
1, 7, 13, 2, 41, 9, 30, 5, 22, 63, 11, 3, 47, 8, 19, 26, 4, 35,
];
let decode_feed: Vec<u32> = vec![12, 55, 63, 6, 28, 40, 17];
memra_engine::qwen4exp_gpu::apply_env_seams();
let seams_env = std::env::var("MEMRA_Q4E_SEAMS").unwrap_or_default();
if !seams_env.contains("kvq") {
memra_engine::qwen4exp_gpu::set_kv_quant(false);
}
if !seams_env.contains("idxq") {
memra_engine::qwen4exp_gpu::set_idxq("f32");
}
let engine = Engine::new(0)?;
let mut lines: Vec<String> = Vec::new();
let mut failures = 0usize;
let mut summaries: Vec<String> = Vec::new();
summaries.push(memra_engine::qwen4exp_gpu::gate_nvfp4_sel_matvec(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_qmatvec_bf16(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_hc_micro_kernels(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_gdn_step_kernels(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_hc_diet_kernels(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_route_kernel(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_kvq_kernels(&engine)?);
let fixture = deterministic_fixture(&plan)?;
let mut model = Qwen4ExpGpu::from_reference_weights(&engine, &plan, &fixture.weights)?;
let result = run_arm(
"fixture",
&engine,
&plan,
&fixture.weights,
&model,
&prompt,
&decode_feed,
&mut lines,
&mut failures,
)?;
summaries.push(format!(
"fixture: prefill worst abs {:.3e} rel {:.3e}; decode worst abs {:.3e} rel {:.3e}",
result.prefill_worst.0,
result.prefill_worst.1,
result.decode_worst.0,
result.decode_worst.1
));
summaries.push(run_mtp_arm(
"mtp-fixture",
&engine,
&plan,
&fixture.weights,
&model,
&prompt,
&mut lines,
&mut failures,
)?);
for (label, kind) in [
("dir-bf16", DirKind::Bf16Fused),
("dir-nvfp4-stacked", DirKind::Nvfp4Stacked),
("dir-nvfp4-perexpert", DirKind::Nvfp4PerExpert),
] {
let dir = TempDir(
std::env::temp_dir().join(format!("qwen4exp-gate-{label}-{}", std::process::id())),
);
synthesize_dir(&dir.0, &cfg, &plan, kind)?;
let checkpoint = read_checkpoint(&dir.0)?;
assert_eq!(
checkpoint.plan, plan,
"dir loader compiled a different plan"
);
let model = Qwen4ExpGpu::load_from_dir(&engine, &dir.0)?;
let reference_weights = checkpoint.into_reference_weights()?;
let mut trunk_plan = plan.clone();
trunk_plan.mtp_blocks.clear();
let result = run_arm(
label,
&engine,
&trunk_plan,
&reference_weights,
&model,
&prompt,
&decode_feed,
&mut lines,
&mut failures,
)?;
summaries.push(format!(
"{label}: prefill worst abs {:.3e} rel {:.3e}; decode worst abs {:.3e} rel {:.3e}",
result.prefill_worst.0,
result.prefill_worst.1,
result.decode_worst.0,
result.decode_worst.1
));
}
{
let dir = TempDir(
std::env::temp_dir().join(format!("qwen4exp-gate-diet-dir-{}", std::process::id())),
);
synthesize_dir(&dir.0, &cfg, &plan, DirKind::Bf16Fused)?;
let mut dmodel = Qwen4ExpGpu::load_from_dir(&engine, &dir.0)?;
let run_rows = |m: &Qwen4ExpGpu| -> Res<Vec<f32>> {
let mut state = m.alloc_state(&engine, prompt.len() + decode_feed.len() + 2)?;
let mut rows = m.prefill(&engine, &prompt, &mut state)?;
for &tok in &decode_feed {
rows.extend(m.decode_step(&engine, tok, &mut state)?);
}
Ok(rows)
};
let before = run_rows(&dmodel)?;
let freed = dmodel.trunk_f32_diet(&engine)?;
let after = run_rows(&dmodel)?;
let identical = before
.iter()
.zip(&after)
.all(|(a, b)| a.to_bits() == b.to_bits());
let pass = identical && freed > 0;
if !pass {
failures += 1;
}
lines.push(format!(
"trunk-diet\tfreed_bytes={freed}\tbits_identical={identical}\tpass={pass}"
));
summaries.push(format!(
"trunk-diet: dir-bf16 pre/post rows {} (freed {} bytes)",
if identical {
"BIT-IDENTICAL"
} else {
"DIVERGED"
},
freed
));
}
{
let dir = TempDir(
std::env::temp_dir().join(format!("qwen4exp-gate-mtp-dir-{}", std::process::id())),
);
synthesize_dir(&dir.0, &cfg, &plan, DirKind::Bf16Fused)?;
let opts = LoadOptions {
load_mtp: true,
..Default::default()
};
let mut model = Qwen4ExpGpu::load_from_dir_with(&engine, &dir.0, opts)?;
assert!(model.has_mtp(), "load_mtp did not materialize the draft");
let reference_weights = read_checkpoint_with(&dir.0, opts)?.into_reference_weights()?;
summaries.push(run_mtp_arm(
"mtp-dir-bf16",
&engine,
&plan,
&reference_weights,
&model,
&prompt,
&mut lines,
&mut failures,
)?);
summaries.push(run_defer_arm(
"mtp-spec-defer-dirbf16",
&engine,
&mut model,
plan.vocab_size as usize,
&prompt,
&mut lines,
&mut failures,
false,
)?);
}
{
let vocab = plan.vocab_size as usize;
let max_new = 24usize;
let spec_k = 3usize;
let cap = prompt.len() + max_new + spec_k + 4;
let mut plain_state = model.alloc_state(&engine, cap)?;
let logits = model.prefill(&engine, &prompt, &mut plain_state)?;
let mut next = argmax(&logits[(prompt.len() - 1) * vocab..]) as u32;
let mut plain = vec![next];
for _ in 1..max_new {
let row = model.decode_step(&engine, next, &mut plain_state)?;
next = argmax(&row) as u32;
plain.push(next);
}
let mut spec_state = model.alloc_state(&engine, cap)?;
let mut draft_state = model.mtp_state(&engine, cap)?;
let report = model.spec_generate(
&engine,
&prompt,
max_new,
spec_k,
&mut spec_state,
&mut draft_state,
None,
)?;
let matched = report.tokens == plain;
if !matched {
failures += 1;
}
lines.push(format!(
"mtp-spec-tiny\tbyte-identity\tk={spec_k}\ttokens={max_new}\trounds={}\taccepted={}/{}\tspec={:?}\tplain={:?}\tpass={matched}",
report.rounds, report.accepted, report.drafted, report.tokens, plain
));
summaries.push(format!(
"mtp-spec-tiny: spec-vs-plain byte-identity {} (k={spec_k}, {max_new} tokens, \
accepted {}/{} over {} rounds)",
if matched { "PASS" } else { "FAIL" },
report.accepted,
report.drafted,
report.rounds
));
}
summaries.push(run_defer_arm(
"mtp-spec-defer",
&engine,
&mut model,
plan.vocab_size as usize,
&prompt,
&mut lines,
&mut failures,
true,
)?);
{
let vocab = plan.vocab_size as usize;
let short = &prompt[0..3];
let mut ua = model.alloc_state(&engine, short.len() + 4)?;
let la = model.prefill(&engine, short, &mut ua)?;
let mut aa = model.alloc_state(&engine, short.len() + 8)?;
model.spec_arm(&engine, &mut aa, 4)?; let lb = model.prefill(&engine, short, &mut aa)?;
let fed = argmax(&la[(short.len() - 1) * vocab..]) as u32;
let ra = model.decode_step(&engine, fed, &mut ua)?;
let rb = model.decode_step(&engine, fed, &mut aa)?;
let pre_same =
la.len() == lb.len() && la.iter().zip(&lb).all(|(a, b)| a.to_bits() == b.to_bits());
let step_same =
ra.len() == rb.len() && ra.iter().zip(&rb).all(|(a, b)| a.to_bits() == b.to_bits());
let pass = pre_same && step_same;
if !pass {
failures += 1;
}
lines.push(format!(
"mtp-armed-prefill-bit\tprompt_len=3\tk_cap=4\tprefill_bits={pre_same}\tstep_bits={step_same}\tpass={pass}"
));
summaries.push(format!(
"mtp-armed-prefill-bit: short-prompt armed vs unarmed prefill {} (prefill rows + 1 step, bitwise)",
if pass { "BIT-IDENTICAL" } else { "DIVERGED" }
));
}
{
let vocab = plan.vocab_size as usize;
let seq: Vec<u32> = prompt.iter().chain(decode_feed.iter()).copied().collect();
let reference = execute(&plan, &fixture.weights, &seq)?;
let n = prompt.len();
let chunk_len = 4usize;
for keep in 1..chunk_len {
let mut state = model.alloc_state(&engine, seq.len() + 2)?;
model.spec_arm(&engine, &mut state, chunk_len + 1)?;
let _ = model.prefill(&engine, &prompt, &mut state)?;
let chunk = &decode_feed[0..chunk_len];
let chunk_logits = model.prefill(&engine, chunk, &mut state)?;
let mut worst = (0.0f32, 0.0f32);
let mut ok = true;
for row in 0..keep {
let stats = compare_row(
&reference.logits[(n + row) * vocab..(n + row + 1) * vocab],
&chunk_logits[row * vocab..(row + 1) * vocab],
);
worst.0 = worst.0.max(stats.max_abs);
worst.1 = worst.1.max(stats.max_rel);
ok &= stats.max_abs <= 0.01 && stats.max_rel <= 0.01 && stats.argmax_match;
}
model.verify_rewind(&engine, &mut state, keep)?;
for (step, &token) in decode_feed[keep..].iter().enumerate() {
let row = n + keep + step;
let logits = model.decode_step(&engine, token, &mut state)?;
let stats = compare_row(&reference.logits[row * vocab..(row + 1) * vocab], &logits);
worst.0 = worst.0.max(stats.max_abs);
worst.1 = worst.1.max(stats.max_rel);
ok &= stats.max_abs <= 0.01 && stats.max_rel <= 0.01 && stats.argmax_match;
}
if !ok {
failures += 1;
}
lines.push(format!(
"mtp-rewind\tkeep={keep}\tmax_abs={:.3e}\tmax_rel={:.3e}\tpass={ok}",
worst.0, worst.1
));
summaries.push(format!(
"mtp-rewind keep={keep}: chunk+rewind+decode vs reference worst abs {:.3e} rel {:.3e} ({})",
worst.0,
worst.1,
if ok { "PASS" } else { "FAIL" }
));
}
}
{
use memra_engine::qwen4exp_gpu::{set_idxq, set_kv_quant};
let vocab = plan.vocab_size as usize;
let run_rows = |model: &Qwen4ExpGpu| -> Res<Vec<f32>> {
let mut state = model.alloc_state(&engine, prompt.len() + decode_feed.len() + 2)?;
let mut rows = model.prefill(&engine, &prompt, &mut state)?;
for &tok in &decode_feed {
rows.extend(model.decode_step(&engine, tok, &mut state)?);
}
Ok(rows)
};
set_kv_quant(false);
let f32_rows = run_rows(&model)?;
set_kv_quant(true);
let qa = run_rows(&model)?;
let qb = run_rows(&model)?;
set_kv_quant(false);
let deterministic = qa.iter().zip(&qb).all(|(a, b)| a.to_bits() == b.to_bits());
let finite = qa.iter().all(|x| x.is_finite());
let mut worst = (0.0f32, 0.0f32);
let mut argmax_flips = 0usize;
let n_rows = qa.len() / vocab;
for r in 0..n_rows {
let s = compare_row(
&f32_rows[r * vocab..(r + 1) * vocab],
&qa[r * vocab..(r + 1) * vocab],
);
worst.0 = worst.0.max(s.max_abs);
worst.1 = worst.1.max(s.max_rel);
if !s.argmax_match {
argmax_flips += 1;
}
}
let pass = deterministic && finite;
if !pass {
failures += 1;
}
lines.push(format!(
"kvq-fixture\tdeterministic={deterministic}\tfinite={finite}\tenvelope_abs={:.3e}\tenvelope_rel={:.3e}\targmax_flips={argmax_flips}/{n_rows}\tpass={pass}",
worst.0, worst.1
));
summaries.push(format!(
"kvq-fixture: determinism {} / finite {}; envelope vs f32 twin abs {:.3e} rel {:.3e}, argmax flips {argmax_flips}/{n_rows} (REPORT — the hard cross-config gate is the real checkpoint)",
if deterministic { "PASS" } else { "FAIL" },
if finite { "PASS" } else { "FAIL" },
worst.0,
worst.1
));
set_kv_quant(true);
{
let max_new = 24usize;
let spec_k = 3usize;
let cap = prompt.len() + max_new + spec_k + 4;
let mut plain_state = model.alloc_state(&engine, cap)?;
let logits = model.prefill(&engine, &prompt, &mut plain_state)?;
let mut next = argmax(&logits[(prompt.len() - 1) * vocab..]) as u32;
let mut plain = vec![next];
for _ in 1..max_new {
let row = model.decode_step(&engine, next, &mut plain_state)?;
next = argmax(&row) as u32;
plain.push(next);
}
let mut spec_state = model.alloc_state(&engine, cap)?;
let mut draft_state = model.mtp_state(&engine, cap)?;
let report = model.spec_generate(
&engine,
&prompt,
max_new,
spec_k,
&mut spec_state,
&mut draft_state,
None,
)?;
let matched = report.tokens == plain;
if !matched {
failures += 1;
}
lines.push(format!(
"kvq-spec-byte-identity\tk={spec_k}\ttokens={max_new}\taccepted={}/{}\tpass={matched}",
report.accepted, report.drafted
));
summaries.push(format!(
"kvq-spec: spec-vs-plain byte identity {} with the quantized cache on BOTH arms \
(k={spec_k}, {max_new} tokens, accepted {}/{})",
if matched { "PASS" } else { "FAIL" },
report.accepted,
report.drafted
));
}
set_kv_quant(false);
for mode in ["q8", "bf16"] {
set_idxq(mode);
memra_engine::qwen4exp_gpu::set_idx_cache(true);
let on_rows = run_rows(&model)?;
memra_engine::qwen4exp_gpu::set_idx_cache(false);
let off_rows = run_rows(&model)?;
memra_engine::qwen4exp_gpu::set_idx_cache(true);
let identical = on_rows
.iter()
.zip(&off_rows)
.all(|(a, b)| a.to_bits() == b.to_bits());
if !identical {
failures += 1;
}
lines.push(format!(
"idxq-{mode}-interleave\tidxcache_on_vs_off_bits={identical}\tpass={identical}"
));
summaries.push(format!(
"idxq-{mode}: idxcache ON vs OFF (host- vs device-quantized rows) {}",
if identical {
"BIT-IDENTICAL"
} else {
"DIVERGED"
}
));
let mut worst = (0.0f32, 0.0f32);
let mut argmax_flips = 0usize;
for r in 0..n_rows {
let s = compare_row(
&f32_rows[r * vocab..(r + 1) * vocab],
&on_rows[r * vocab..(r + 1) * vocab],
);
worst.0 = worst.0.max(s.max_abs);
worst.1 = worst.1.max(s.max_rel);
if !s.argmax_match {
argmax_flips += 1;
}
}
lines.push(format!(
"idxq-{mode}-envelope\tabs={:.3e}\trel={:.3e}\targmax_flips={argmax_flips}/{n_rows}",
worst.0, worst.1
));
}
set_idxq("q8");
{
let max_new = 24usize;
let spec_k = 3usize;
let cap = prompt.len() + max_new + spec_k + 4;
let mut plain_state = model.alloc_state(&engine, cap)?;
let logits = model.prefill(&engine, &prompt, &mut plain_state)?;
let mut next = argmax(&logits[(prompt.len() - 1) * vocab..]) as u32;
let mut plain = vec![next];
for _ in 1..max_new {
let row = model.decode_step(&engine, next, &mut plain_state)?;
next = argmax(&row) as u32;
plain.push(next);
}
let mut spec_state = model.alloc_state(&engine, cap)?;
let mut draft_state = model.mtp_state(&engine, cap)?;
let report = model.spec_generate(
&engine,
&prompt,
max_new,
spec_k,
&mut spec_state,
&mut draft_state,
None,
)?;
let matched = report.tokens == plain;
if !matched {
failures += 1;
}
lines.push(format!(
"idxq-q8-spec-byte-identity\tk={spec_k}\ttokens={max_new}\taccepted={}/{}\tpass={matched}",
report.accepted, report.drafted
));
summaries.push(format!(
"idxq-q8-spec: spec-vs-plain byte identity {} (quantized raw-key cache both arms)",
if matched { "PASS" } else { "FAIL" }
));
}
set_idxq("f32");
}
let set_yarn = |plan: &mut ModelPlan, factors: RopeFactors| {
for layer in plan.layers.iter_mut() {
if let AttentionPlan::Full(attention) = &mut layer.attention {
attention.rope.factors = factors;
}
}
for mtp in plan.mtp_blocks.iter_mut() {
if let AttentionPlan::Full(attention) = &mut mtp.layer.attention {
attention.rope.factors = factors;
}
}
};
{
let mut yarn_plan = plan.clone();
set_yarn(
&mut yarn_plan,
RopeFactors::Yarn {
factor: 2.0,
original_context: 8,
beta_fast: 32.0,
beta_slow: 1.0,
},
);
let yarn_model =
Qwen4ExpGpu::from_reference_weights(&engine, &yarn_plan, &fixture.weights)?;
let result = run_arm(
"fixture-yarn",
&engine,
&yarn_plan,
&fixture.weights,
&yarn_model,
&prompt,
&decode_feed,
&mut lines,
&mut failures,
)?;
summaries.push(format!(
"fixture-yarn (factor 2, mscale 1.0693): prefill worst abs {:.3e} rel {:.3e}; \
decode worst abs {:.3e} rel {:.3e}",
result.prefill_worst.0,
result.prefill_worst.1,
result.decode_worst.0,
result.decode_worst.1
));
}
{
let mut identity_plan = plan.clone();
set_yarn(
&mut identity_plan,
RopeFactors::Yarn {
factor: 1.0,
original_context: 64,
beta_fast: 32.0,
beta_slow: 1.0,
},
);
let identity_model =
Qwen4ExpGpu::from_reference_weights(&engine, &identity_plan, &fixture.weights)?;
let vocab = plan.vocab_size as usize;
let cap = prompt.len() + decode_feed.len() + 2;
let mut rows_base: Vec<Vec<f32>> = Vec::new();
let mut rows_yarn: Vec<Vec<f32>> = Vec::new();
for (m, sink) in [(&model, &mut rows_base), (&identity_model, &mut rows_yarn)] {
let mut state = m.alloc_state(&engine, cap)?;
let prefill = m.prefill(&engine, &prompt, &mut state)?;
for row in 0..prompt.len() {
sink.push(prefill[row * vocab..(row + 1) * vocab].to_vec());
}
for &token in &decode_feed {
sink.push(m.decode_step(&engine, token, &mut state)?);
}
}
let mut identical = true;
for (row, (a, b)) in rows_base.iter().zip(rows_yarn.iter()).enumerate() {
let same = a.len() == b.len()
&& a.iter()
.zip(b.iter())
.all(|(x, y)| x.to_bits() == y.to_bits());
if !same {
identical = false;
failures += 1;
lines.push(format!(
"yarn-identity\trow={row}\tbit_identical=false\tpass=false"
));
}
}
if identical {
lines.push(format!(
"yarn-identity\trows={}\tbit_identical=true\tpass=true",
rows_base.len()
));
}
summaries.push(format!(
"yarn-identity: factor-1.0 yarn vs plain rope {} ({} rows bit-compared)",
if identical { "BIT-IDENTICAL" } else { "FAIL" },
rows_base.len()
));
}
summaries.push(memra_engine::qwen4exp_gpu::gate_sdpa_blocklist(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_qsa_index_score(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_qsa_index_topk(&engine)?);
summaries.push(memra_engine::qwen4exp_gpu::gate_ple_ngram_cache()?);
summaries.push(memra_engine::qwen4exp_gpu::gate_seam_table()?);
{
let vocab = plan.vocab_size as usize;
let max_new = 24usize;
let spec_k = 3usize;
let cap = prompt.len() + max_new + spec_k + 4;
let mut plain_state = model.alloc_state(&engine, cap)?;
let logits = model.prefill(&engine, &prompt, &mut plain_state)?;
let mut next = argmax(&logits[(prompt.len() - 1) * vocab..]) as u32;
let mut plain = vec![next];
for _ in 1..max_new {
let row = model.decode_step(&engine, next, &mut plain_state)?;
next = argmax(&row) as u32;
plain.push(next);
}
let mut spec_state = model.alloc_state(&engine, cap)?;
let mut draft_state = model.mtp_state(&engine, cap)?;
let opts = memra_engine::qwen4exp_gpu::SpecOpts {
prefill_chunk: Some(8),
wide_ring: Some(16),
..Default::default()
};
let report = model.spec_generate_ext(
&engine,
&engine,
&prompt,
max_new,
spec_k,
&mut spec_state,
&mut draft_state,
None,
opts,
None,
)?;
let matched = report.tokens == plain;
if !matched {
failures += 1;
}
lines.push(format!(
"mtp-spec-ring\tbyte-identity\tchunk=8\tring=16\ttokens={max_new}\trounds={}\tspec={:?}\tplain={:?}\tpass={matched}",
report.rounds, report.tokens, plain
));
summaries.push(format!(
"mtp-spec-ring: chunked co-prefill (8) + wide ring (16) spec-vs-plain \
byte-identity {} ({} rounds)",
if matched { "PASS" } else { "FAIL" },
report.rounds
));
}
{
memra_engine::qwen4exp_gpu::set_longatt("force");
let vocab = plan.vocab_size as usize;
let cap = prompt.len() + decode_feed.len() + 2;
let mut rows_long: Vec<Vec<f32>> = Vec::new();
{
let mut state = model.alloc_state(&engine, cap)?;
let prefill = model.prefill(&engine, &prompt, &mut state)?;
for row in 0..prompt.len() {
rows_long.push(prefill[row * vocab..(row + 1) * vocab].to_vec());
}
for &token in &decode_feed {
rows_long.push(model.decode_step(&engine, token, &mut state)?);
}
}
memra_engine::qwen4exp_gpu::set_longatt("auto");
let mut rows_base: Vec<Vec<f32>> = Vec::new();
{
let mut state = model.alloc_state(&engine, cap)?;
let prefill = model.prefill(&engine, &prompt, &mut state)?;
for row in 0..prompt.len() {
rows_base.push(prefill[row * vocab..(row + 1) * vocab].to_vec());
}
for &token in &decode_feed {
rows_base.push(model.decode_step(&engine, token, &mut state)?);
}
}
let mut identical = true;
for (row, (a, b)) in rows_base.iter().zip(rows_long.iter()).enumerate() {
let same = a.len() == b.len()
&& a.iter()
.zip(b.iter())
.all(|(x, y)| x.to_bits() == y.to_bits());
if !same {
identical = false;
failures += 1;
lines.push(format!(
"fixture-longatt\trow={row}\tbit_identical=false\tpass=false"
));
}
}
if identical {
lines.push(format!(
"fixture-longatt\trows={}\tbit_identical=true\tpass=true",
rows_base.len()
));
}
summaries.push(format!(
"fixture-longatt: forced block-list attention vs masked {} ({} rows bit-compared)",
if identical { "BIT-IDENTICAL" } else { "FAIL" },
rows_base.len()
));
}
{
let vocab = plan.vocab_size as usize;
let cap = prompt.len() + decode_feed.len() + 2;
let mut one_state = model.alloc_state(&engine, cap)?;
let one = model.prefill(&engine, &prompt, &mut one_state)?;
let one_last = &one[(prompt.len() - 1) * vocab..prompt.len() * vocab];
let mut chunk_state = model.alloc_state(&engine, cap)?;
let chunked = model.prefill_extend(&engine, &prompt, &mut chunk_state, 5)?;
let mut worst = (0.0f32, 0.0f32);
let mut ok = chunked.len() == vocab;
let mut fold = |s: RowStats, ok: &mut bool, worst: &mut (f32, f32)| {
worst.0 = worst.0.max(s.max_abs);
worst.1 = worst.1.max(s.max_rel);
*ok &= s.max_abs <= 0.01 && s.max_rel <= 0.01 && s.argmax_match;
};
fold(compare_row(one_last, &chunked), &mut ok, &mut worst);
for &token in &decode_feed {
let a = model.decode_step(&engine, token, &mut one_state)?;
let b = model.decode_step(&engine, token, &mut chunk_state)?;
fold(compare_row(&a, &b), &mut ok, &mut worst);
}
if !ok {
failures += 1;
}
lines.push(format!(
"prefill-extend\tchunk=5\tmax_abs={:.3e}\tmax_rel={:.3e}\tpass={ok}",
worst.0, worst.1
));
summaries.push(format!(
"prefill-extend: chunked (5) vs one-shot prefill + decode worst abs {:.3e} rel \
{:.3e} ({})",
worst.0,
worst.1,
if ok { "PASS" } else { "FAIL" }
));
}
let executable = std::fs::read(std::env::current_exe()?)?;
let sha256 = Sha256::digest(&executable)
.iter()
.map(|byte| format!("{byte:02x}"))
.collect::<String>();
let all_tokens: Vec<u32> = prompt.iter().chain(decode_feed.iter()).copied().collect();
let mut receipt = format!(
"# qwen4exp-gpu-gate\tbinary_sha256={sha256}\tpolicy=max_abs<=0.01 max_rel<=0.01 argmax\tprompt_len={}\tdecode_steps={}\tvocab={}\tplan=tiny(qwen4_exp pack)\tarms=fixture,mtp-fixture,dir-bf16,dir-nvfp4-stacked,dir-nvfp4-perexpert,mtp-dir-bf16,mtp-spec-tiny,mtp-spec-defer,mtp-spec-defer-dirbf16,mtp-armed-prefill-bit,fixture-yarn,yarn-identity\ttokens={all_tokens:?}\n",
prompt.len(),
decode_feed.len(),
plan.vocab_size
);
for line in &lines {
receipt.push_str(line);
receipt.push('\n');
}
for summary in &summaries {
receipt.push_str(&format!("# summary\t{summary}\n"));
}
receipt.push_str(&format!("# verdict\tfailures={failures}\n"));
std::fs::write(&receipt_path, &receipt)?;
if failures > 0 {
eprintln!(
"qwen4exp-gpu-gate FAILED: {failures} rows out of tolerance (receipt {receipt_path})"
);
std::process::exit(1);
}
for summary in &summaries {
println!("qwen4exp-gpu-gate PASS [{summary}]");
}
println!("receipt: {receipt_path}");
Ok(())
}