use memra_engine::Engine;
use memra_engine::dflash::DflashDraft;
use memra_engine::hybrid::HybridModel;
use memra_engine::spec::{SpecSampling, sample_boundary_token};
use memra_gguf::GgufFile;
fn load_model(path: &str, e: &Engine) -> Result<HybridModel, Box<dyn std::error::Error>> {
let path = memra_gguf::hf::resolve_arg(path)?;
let is_dir = std::path::Path::new(&path).is_dir();
if is_dir {
let dir = std::path::Path::new(&path);
if dir.join("manifest.json").exists() {
let repack = memra_gguf::source::Hy3RepackSource::open(dir)?;
Ok(HybridModel::load_from_source(e, &repack)?)
} else {
let src = memra_gguf::source::SafetensorsSource::open(dir)?;
Ok(HybridModel::load_from_source(e, &src)?)
}
} else {
let g = GgufFile::open(&path)?;
Ok(HybridModel::load(e, &g)?)
}
}
fn sp(temp: f32, seed: u64, top_k: i32, top_p: f32) -> SpecSampling {
SpecSampling {
temp,
seed,
top_k,
top_p,
min_p: 0.0,
penalty_last_n: 0,
penalty_repeat: 1.0,
penalty_freq: 0.0,
penalty_present: 0.0,
}
}
fn plain_sampled(
model: &HybridModel,
e: &Engine,
prompt: &[u32],
max_new: usize,
cfg: &SpecSampling,
) -> Result<Vec<u32>, Box<dyn std::error::Error>> {
let mut cache = memra_engine::pp::new_cache(e, &model.cfg, prompt.len() + max_new + 8)?;
let mut logits = if prompt.len() >= memra_engine::hybrid_forward::PRIME_MIN_T {
model.prime_cache(e, prompt, &mut cache, 0)?.0
} else {
let mut row = Vec::new();
for &token in prompt {
let mut caches = [&mut cache];
row = model.decode_step_batch(e, &[token], &mut caches)?.remove(0);
}
row
};
let mut sctr = 0u32;
let mut out = Vec::with_capacity(max_new);
for _ in 0..max_new {
let token = sample_boundary_token(e, &logits, cfg, &[], &mut sctr, "trunk-ref")?;
out.push(token);
if out.len() >= max_new {
break;
}
let mut caches = [&mut cache];
logits = model.decode_step_batch(e, &[token], &mut caches)?.remove(0);
}
Ok(out)
}
fn first_divergence(a: &[u32], b: &[u32]) -> Option<usize> {
let n = a.len().min(b.len());
(0..n)
.find(|&i| a[i] != b[i])
.or(if a.len() != b.len() { Some(n) } else { None })
}
fn chi2_q999(df: f64) -> f64 {
let z = 3.0902; df * (1.0 - 2.0 / (9.0 * df) + z * (2.0 / (9.0 * df)).sqrt()).powi(3)
}
fn env_usize(k: &str, d: usize) -> usize {
std::env::var(k)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(d)
}
fn env_f32(k: &str, d: f32) -> f32 {
std::env::var(k)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(d)
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let target = std::env::args()
.nth(1)
.expect("usage: dspark_sample_gate <target> <draft_dir> [t0|tiny|hist|all]");
let draft_dir = std::env::args().nth(2).expect("draft export dir");
let mode = std::env::args().nth(3).unwrap_or_else(|| "all".into());
let e = Engine::new(0)?;
let model = load_model(&target, &e)?;
let draft = DflashDraft::load(&e, std::path::Path::new(&draft_dir))?;
println!(
"target loaded ({} layers); draft block {}, dflash2 {}, markov {}",
model.cfg.n_layer,
draft.cfg.block_size,
draft.dflash2.is_some(),
draft.markov.is_some()
);
let prompt: Vec<u32> = std::env::var("MEMRA_SG_PROMPT")
.ok()
.map(|s| s.split(',').filter_map(|v| v.trim().parse().ok()).collect())
.unwrap_or_else(|| {
vec![
84270, 279, 2701, 7355, 25, 220, 16, 13, 220, 100, 101, 102, 103, 104, 105, 106,
107, 108, 109, 110, 111, 112, 113, 114,
]
});
assert!(
prompt.len() >= 16,
"MEMRA_SG_PROMPT must carry >= 16 token ids (prime_cache floor)"
);
let eos: Vec<u32> = Vec::new();
let mut fails = 0usize;
if mode == "t0" || mode == "all" {
let a = model.generate_spec_dspark(&e, &draft, &prompt, 64, &eos, None)?;
let cfg0 = sp(0.0, 42, 0, 1.0);
let b = model.generate_spec_dspark(&e, &draft, &prompt, 64, &eos, Some(&cfg0))?;
let ok = a == b && !a.is_empty();
println!(
"t0 kill-switch (None == Some(temp=0)): {}",
if ok {
"EXACT"
} else {
fails += 1;
"DIVERGED"
}
);
}
if mode == "tiny" || mode == "all" {
for seed in [42u64, 7, 1234] {
let a = model.generate_spec_dspark(&e, &draft, &prompt, 64, &eos, None)?;
let cfgt = sp(1e-6, seed, 0, 1.0);
let b = model.generate_spec_dspark(&e, &draft, &prompt, 64, &eos, Some(&cfgt))?;
let div = first_divergence(&a, &b);
let ok = div.is_none() && !a.is_empty();
println!(
"tiny-T continuity seed {seed}: {}",
if ok {
"EXACT".into()
} else {
fails += 1;
format!("DIVERGED at {div:?}")
}
);
}
}
if mode == "hist" || mode == "all" {
let n_seeds = env_usize("MEMRA_SG_SEEDS", 2000);
let m_tok = env_usize("MEMRA_SG_TOKENS", 4);
let temp = env_f32("MEMRA_SG_TEMP", 0.8);
let top_k = env_usize("MEMRA_SG_TOPK", 0) as i32;
let top_p = env_f32("MEMRA_SG_TOPP", 1.0);
println!(
"hist: {n_seeds} seeds/arm, {m_tok} positions, temp {temp}, top_k {top_k}, \
top_p {top_p}, prompt len {}",
prompt.len()
);
use std::collections::HashMap;
let mut hist_a: Vec<HashMap<u32, u64>> = vec![HashMap::new(); m_tok];
let mut hist_b: Vec<HashMap<u32, u64>> = vec![HashMap::new(); m_tok];
let t0 = std::time::Instant::now();
for i in 0..n_seeds {
let cfg_a = sp(temp, 1_000_000 + i as u64, top_k, top_p);
let a = plain_sampled(&model, &e, &prompt, m_tok, &cfg_a)?;
for (j, &t) in a.iter().take(m_tok).enumerate() {
*hist_a[j].entry(t).or_insert(0) += 1;
}
let cfg_b = sp(temp, 9_000_000 + i as u64, top_k, top_p);
let b = model.generate_spec_dspark(&e, &draft, &prompt, m_tok, &eos, Some(&cfg_b))?;
for (j, &t) in b.iter().take(m_tok).enumerate() {
*hist_b[j].entry(t).or_insert(0) += 1;
}
if (i + 1) % 500 == 0 {
println!(
" ... {} / {n_seeds} seeds ({:.0}s)",
i + 1,
t0.elapsed().as_secs_f64()
);
}
}
for j in 0..m_tok {
let mut tokens: std::collections::HashSet<u32> = hist_a[j].keys().copied().collect();
tokens.extend(hist_b[j].keys().copied());
let (mut x2, mut buckets) = (0f64, 0usize);
let (mut tail_a, mut tail_b) = (0f64, 0f64);
let mut tv = 0f64;
for &t in &tokens {
let a = *hist_a[j].get(&t).unwrap_or(&0) as f64;
let b = *hist_b[j].get(&t).unwrap_or(&0) as f64;
tv += (a - b).abs();
if a + b >= 10.0 {
x2 += (a - b) * (a - b) / (a + b);
buckets += 1;
} else {
tail_a += a;
tail_b += b;
}
}
if tail_a + tail_b > 0.0 {
x2 += (tail_a - tail_b) * (tail_a - tail_b) / (tail_a + tail_b);
buckets += 1;
}
let df = (buckets.max(2) - 1) as f64;
let bound = chi2_q999(df);
let tvn = tv / (2.0 * n_seeds as f64);
let ok = x2 < bound;
println!(
"pos {j}: X2={x2:.1} df={df:.0} bound(q=.999)={bound:.1} TV={tvn:.4} \
support={} {}",
tokens.len(),
if ok {
"OK"
} else {
fails += 1;
"FAIL"
}
);
}
}
println!(
"== dspark_sample_gate: {} ==",
if fails == 0 { "ALL PASS" } else { "FAIL" }
);
std::process::exit(if fails == 0 { 0 } else { 1 });
}