use memra_engine::dsv4_gpu::Dsv4Gpu;
use memra_gguf::dsv4_forward::ActQuantVariant;
use std::io::Write;
use std::path::Path;
fn argmax(v: &[f32]) -> u32 {
let mut best = 0usize;
for i in 1..v.len() {
if v[i] > v[best] {
best = i;
}
}
best as u32
}
fn sha256_hex(bytes: &[u8]) -> String {
use sha2::{Digest, Sha256};
let mut h = Sha256::new();
h.update(bytes);
h.finalize().iter().map(|b| format!("{b:02x}")).collect()
}
fn load_prompts(path: &Path) -> Vec<(String, Vec<u32>)> {
let s = std::fs::read_to_string(path).expect("read prompts");
let mut out = Vec::new();
let mut rest = s.as_str();
while let Some(i) = rest.find("\"pool\"") {
rest = &rest[i + 6..];
let q0 = rest.find('"').expect("pool open quote");
let after = &rest[q0 + 1..];
let q1 = after.find('"').expect("pool close quote");
let pool = after[..q1].to_string();
rest = &after[q1 + 1..];
let j = rest.find("\"ids\"").expect("ids key");
rest = &rest[j + 5..];
let b0 = rest.find('[').expect("ids open");
let b1 = rest.find(']').expect("ids close");
let ids: Vec<u32> = rest[b0 + 1..b1]
.split(',')
.filter_map(|t| t.trim().parse::<u32>().ok())
.collect();
rest = &rest[b1 + 1..];
assert!(!ids.is_empty(), "empty prompt ids for pool {pool}");
out.push((pool, ids));
}
assert!(!out.is_empty(), "no prompts parsed");
out
}
struct ArmAgg {
tokens: usize,
decode_us: u64,
prefill_us: u64,
}
impl ArmAgg {
fn new() -> Self {
ArmAgg {
tokens: 0,
decode_us: 0,
prefill_us: 0,
}
}
fn ms_per_token(&self) -> f64 {
self.decode_us as f64 / 1e3 / self.tokens as f64
}
fn tok_s(&self) -> f64 {
self.tokens as f64 * 1e6 / self.decode_us as f64
}
}
fn main() {
let args: Vec<String> = std::env::args().collect();
if args.len() < 4 {
eprintln!(
"usage: dsv4-drafted-corpus <model-dir> <prompts.json> <out-dir> [n_new] [reps] \
[dev0,dev1]"
);
std::process::exit(2);
}
let t0 = std::time::Instant::now();
let dir = Path::new(&args[1]);
let prompts = load_prompts(Path::new(&args[2]));
let out_dir = Path::new(&args[3]);
std::fs::create_dir_all(out_dir).expect("mkdir out");
let n_new: usize = args
.get(4)
.map(|x| x.parse().expect("n_new"))
.unwrap_or(128);
let reps: usize = args.get(5).map(|x| x.parse().expect("reps")).unwrap_or(3);
let devices: Vec<usize> = args
.get(6)
.map(|s| s.split(',').map(|x| x.parse().expect("device")).collect())
.unwrap_or_else(|| vec![0, 1]);
if std::env::var("MEMRA_DSV4_DRAFTER").as_deref() != Ok("dspark") {
eprintln!("REFUSE: this cell requires MEMRA_DSV4_DRAFTER=dspark");
std::process::exit(2);
}
if std::env::var("MEMRA_DSV4_DECODE_PATH").as_deref() != Ok("device") {
eprintln!("REFUSE: this cell requires MEMRA_DSV4_DECODE_PATH=device");
std::process::exit(2);
}
let variant = ActQuantVariant::ClampOnly;
let max_p = prompts.iter().map(|(_, i)| i.len()).max().unwrap();
let max_seq = (max_p + n_new + 96).max(256);
println!(
"dsv4-drafted-corpus | model {} | prompts {} (max {max_p} ids) | n_new {n_new} | reps \
{reps} | devices {devices:?} | GREEDY ONLY (this lane's drafted path; sampling \
verify is a separate rung)",
dir.display(),
prompts.len()
);
let gpu = Dsv4Gpu::load(dir, &devices, variant, max_seq).expect("load");
println!(
"loaded: split at layer {}, verify tmax {}, t={:.0}s",
gpu.split_at,
gpu.verify_tmax(),
t0.elapsed().as_secs_f64()
);
if let Ok(spec) = std::env::var("MEMRA_DSV4_SPEC_DEPTH_SWEEP") {
let depths: Vec<usize> = spec
.split(',')
.filter_map(|t| t.trim().parse::<usize>().ok())
.filter(|d| *d >= 1)
.collect();
assert!(
!depths.is_empty(),
"MEMRA_DSV4_SPEC_DEPTH_SWEEP parsed empty"
);
run_depth_sweep(&gpu, &prompts, n_new, reps, out_dir, &depths, t0);
return;
}
let mut plain = ArmAgg::new();
let mut draft = ArmAgg::new();
let mut rounds_total = 0usize;
let mut accepted_total = 0usize;
let mut t_batch_sum = 0usize;
let mut t_batch_n = 0usize;
let mut identity_fails: Vec<String> = Vec::new();
let mut per_pool: std::collections::BTreeMap<String, (usize, u64, u64, usize, usize)> =
Default::default();
let mut rows: Vec<String> = Vec::new();
let mut all_rounds: Vec<(usize, usize, usize, usize, u64, Vec<f32>)> = Vec::new();
for rep in 0..reps {
for (pi, (pool, ids)) in prompts.iter().enumerate() {
let plain_first = (rep + pi) % 2 == 0;
let run_plain = || {
let mut state = gpu.alloc_decode_state().expect("alloc state");
let pt0 = std::time::Instant::now();
let pre = gpu.prefill_with_cache(ids, &mut state).expect("prefill");
let pre_us = pt0.elapsed().as_micros() as u64;
let mut t = argmax(&pre.logits);
let mut toks = Vec::with_capacity(n_new);
let dt0 = std::time::Instant::now();
for step in 0..n_new {
toks.push(t);
if step + 1 == n_new {
break;
}
t = gpu.decode_step_greedy(t, &mut state).expect("plain step");
}
(toks, dt0.elapsed().as_micros() as u64, pre_us)
};
let run_draft = || {
let mut state = gpu.alloc_decode_state().expect("alloc state");
let mut dstate = gpu.dspark_alloc_state().expect("alloc dspark");
let mut vstate = gpu.alloc_verify_state().expect("alloc verify");
let dt0 = std::time::Instant::now();
let out = gpu
.spec_greedy_batched_with(ids, n_new, &mut state, &mut dstate, &mut vstate)
.expect("drafted run");
let total_us = dt0.elapsed().as_micros() as u64;
let rounds_us: u64 = out.rounds.iter().map(|r| r.round_us).sum();
let pre_us = total_us.saturating_sub(rounds_us);
(out, rounds_us, pre_us)
};
let ((plain_toks, p_us, ppre), (draft_out, d_us, dpre)) = if plain_first {
let a = run_plain();
let b = run_draft();
(a, b)
} else {
let b = run_draft();
let a = run_plain();
(a, b)
};
plain.prefill_us += ppre;
draft.prefill_us += dpre;
let d_rounds = draft_out.rounds.len();
let d_acc: usize = draft_out.rounds.iter().map(|r| r.accepts).sum();
for r in &draft_out.rounds {
if r.t_batch > 0 {
t_batch_sum += r.t_batch;
t_batch_n += 1;
}
all_rounds.push((
r.t_batch,
r.t_cap,
r.accepts,
r.emitted,
r.round_us,
r.confidence.clone(),
));
}
let draft_toks = draft_out.tokens;
if draft_toks != plain_toks {
let first = draft_toks
.iter()
.zip(&plain_toks)
.position(|(a, b)| a != b)
.unwrap_or(plain_toks.len().min(draft_toks.len()));
identity_fails.push(format!(
"IDENTITY FAIL rep{rep} prompt{pi} ({pool}): drafted != plain at generated \
index {first} (drafted {:?} vs plain {:?})",
draft_toks.get(first),
plain_toks.get(first)
));
}
plain.tokens += plain_toks.len();
plain.decode_us += p_us;
draft.tokens += draft_toks.len();
draft.decode_us += d_us;
rounds_total += d_rounds;
accepted_total += d_acc;
let e = per_pool.entry(pool.clone()).or_insert((0, 0, 0, 0, 0));
e.0 += plain_toks.len();
e.1 += p_us;
e.2 += d_us;
e.3 += d_rounds;
e.4 += d_acc;
let row = format!(
"rep{rep} p{pi:02} [{pool}] ids {} | plain {:.2} ms/tok | drafted {:.2} ms/tok \
| {:.3}x | rounds {d_rounds} accepted {d_acc} ({:.3}/round, {:.3} tok/round)",
ids.len(),
p_us as f64 / 1e3 / plain_toks.len() as f64,
d_us as f64 / 1e3 / draft_toks.len() as f64,
p_us as f64 / d_us as f64,
d_acc as f64 / d_rounds.max(1) as f64,
draft_toks.len() as f64 / d_rounds.max(1) as f64
);
println!("{row}");
rows.push(row);
}
println!(
"--- rep{rep} cumulative: plain {:.2} ms/tok ({:.1} tok/s) | drafted {:.2} ms/tok \
({:.1} tok/s) | {:.3}x | t={:.0}s",
plain.ms_per_token(),
plain.tok_s(),
draft.ms_per_token(),
draft.tok_s(),
plain.ms_per_token() / draft.ms_per_token(),
t0.elapsed().as_secs_f64()
);
}
println!("\n=== OWNER-CORPORA DRAFTED CELL (measured, bench not serving) ===");
println!(
"prompts {} x reps {reps} x n_new {n_new} = {} tokens per arm",
prompts.len(),
plain.tokens
);
println!(
"plain : {:.2} ms/token = {:.1} tok/s bs=1 (prefill excluded: {:.1} ms total)",
plain.ms_per_token(),
plain.tok_s(),
plain.prefill_us as f64 / 1e3
);
println!(
"drafted : {:.2} ms/token = {:.1} tok/s bs=1 (prefill+prime excluded: {:.1} ms total)",
draft.ms_per_token(),
draft.tok_s(),
draft.prefill_us as f64 / 1e3
);
println!(
"SPEEDUP : {:.3}x ({:.1} -> {:.1} tok/s)",
plain.ms_per_token() / draft.ms_per_token(),
plain.tok_s(),
draft.tok_s()
);
memra_engine::dsv4_gpu::dsv4_phase_report(
"drafted round",
rounds_total as u64,
plain.ms_per_token() * 1000.0,
);
println!(
"acceptance (CORRECTNESS observable, never a speed claim — dspark-q38 law): rounds \
{rounds_total} | accepted {accepted_total} | {:.4}/round | {:.4} tokens/round | mean T \
forwarded {:.4}",
accepted_total as f64 / rounds_total.max(1) as f64,
draft.tokens as f64 / rounds_total.max(1) as f64,
if t_batch_n > 0 {
t_batch_sum as f64 / t_batch_n as f64
} else {
0.0
}
);
{
let steady: Vec<&(usize, usize, usize, usize, u64, Vec<f32>)> = all_rounds
.iter()
.filter(|r| r.0 > 0 && r.0 == r.1)
.collect();
if !steady.is_empty() {
let n = steady.len() as f64;
let mean_us = steady.iter().map(|r| r.4 as f64).sum::<f64>() / n;
let mean_t = steady.iter().map(|r| r.0 as f64).sum::<f64>() / n;
let mean_acc = steady.iter().map(|r| r.2 as f64).sum::<f64>() / n;
let tau = steady.iter().map(|r| r.3 as f64).sum::<f64>() / n;
println!(
"[sps] depth-pinned rounds {} | mean T forwarded {:.4} | mean round {:.3} ms \
=> SPS(B) {:.2} rounds/s | accepts {:.4}/round | tau* {:.4} tok/round | \
Theta = tau*.SPS = {:.2} tok/s",
steady.len(),
mean_t,
mean_us / 1e3,
1e6 / mean_us,
mean_acc,
tau,
tau * 1e6 / mean_us
);
let kmax = steady.iter().map(|r| r.0 - 1).max().unwrap_or(0);
let mut pos: Vec<String> = Vec::new();
for j in 0..kmax {
let reached = steady.iter().filter(|r| r.2 >= j).count();
let took = steady.iter().filter(|r| r.2 > j).count();
if reached > 0 {
pos.push(format!("p{}={:.4}", j + 1, took as f64 / reached as f64));
}
}
println!(
"[sps] per-position CONDITIONAL acceptance (this model, this box): {}",
pos.join(" ")
);
}
let side = out_dir.join("spec_rounds.json");
let mut f = std::fs::File::create(&side).expect("rounds sidecar");
writeln!(f, "{{\"schema\": \"dsv4-spec-rounds-v1\", \"rounds\": [").expect("w");
for (i, r) in all_rounds.iter().enumerate() {
let conf =
r.5.iter()
.map(|c| format!("{c:.6}"))
.collect::<Vec<_>>()
.join(",");
writeln!(
f,
" {{\"t_batch\": {}, \"t_cap\": {}, \"accepts\": {}, \"emitted\": {}, \
\"round_us\": {}, \"conf\": [{conf}]}}{}",
r.0,
r.1,
r.2,
r.3,
r.4,
if i + 1 == all_rounds.len() { "" } else { "," }
)
.expect("w");
}
writeln!(f, "]}}").expect("w");
println!(
"[sps] {} rounds banked {}",
all_rounds.len(),
side.display()
);
}
println!("per pool:");
for (pool, (tk, pu, du, rd, ac)) in &per_pool {
println!(
" {pool:<8} tokens {tk:<6} plain {:.2} ms/tok drafted {:.2} ms/tok {:.3}x \
accepted/round {:.3}",
*pu as f64 / 1e3 / *tk as f64,
*du as f64 / 1e3 / *tk as f64,
*pu as f64 / *du as f64,
*ac as f64 / (*rd).max(1) as f64
);
}
println!(
"greedy spec==plain identity on real prompts: {} / {} prompt-runs",
prompts.len() * reps - identity_fails.len(),
prompts.len() * reps
);
let mut f = std::fs::File::create(out_dir.join("drafted_corpus.json")).expect("json");
write!(
f,
"{{\n \"prompts\": {},\n \"reps\": {reps},\n \"n_new\": {n_new},\n \
\"tokens_per_arm\": {},\n \"plain_ms_per_token\": {:.4},\n \
\"plain_tok_s\": {:.3},\n \"drafted_ms_per_token\": {:.4},\n \
\"drafted_tok_s\": {:.3},\n \"speedup\": {:.4},\n \"rounds\": {rounds_total},\n \
\"accepted\": {accepted_total},\n \"accepted_per_round\": {:.4},\n \
\"tokens_per_round\": {:.4},\n \"identity_fails\": {},\n \"rows_sha256\": \"{}\"\n}}\n",
prompts.len(),
plain.tokens,
plain.ms_per_token(),
plain.tok_s(),
draft.ms_per_token(),
draft.tok_s(),
plain.ms_per_token() / draft.ms_per_token(),
accepted_total as f64 / rounds_total.max(1) as f64,
draft.tokens as f64 / rounds_total.max(1) as f64,
identity_fails.len(),
sha256_hex(rows.join("\n").as_bytes())
)
.expect("write json");
if identity_fails.is_empty() {
println!(
"\nCORPORA CELL: identity [PASS] | banked {} | total elapsed {:.0}s",
out_dir.join("drafted_corpus.json").display(),
t0.elapsed().as_secs_f64()
);
} else {
println!(
"\nCORPORA CELL: identity [FAIL] {} finding(s)",
identity_fails.len()
);
for l in &identity_fails {
println!(" {l}");
}
std::process::exit(1);
}
}
struct DepthArm {
depth: usize,
tokens: usize,
decode_us: u64,
rounds: usize,
accepts: usize,
t_batch_sum: usize,
recs: Vec<(usize, usize, usize, usize, u64, Vec<f32>)>,
}
impl DepthArm {
fn new(depth: usize) -> Self {
DepthArm {
depth,
tokens: 0,
decode_us: 0,
rounds: 0,
accepts: 0,
t_batch_sum: 0,
recs: Vec::new(),
}
}
}
#[allow(clippy::too_many_arguments)]
fn run_depth_sweep(
gpu: &Dsv4Gpu,
prompts: &[(String, Vec<u32>)],
n_new: usize,
reps: usize,
out_dir: &Path,
depths: &[usize],
t0: std::time::Instant,
) {
println!(
"\n=== DRAFT-DEPTH SWEEP | depths {depths:?} | plain baseline shared | {} prompts x {reps} \
reps x {n_new} tok | ONE load, ONE thermal window, arm order rotated ===",
prompts.len()
);
let mut plain = ArmAgg::new();
let mut arms: Vec<DepthArm> = depths.iter().map(|d| DepthArm::new(*d)).collect();
let mut identity_fails: Vec<String> = Vec::new();
for rep in 0..reps {
for (pi, (pool, ids)) in prompts.iter().enumerate() {
let n_arms = depths.len() + 1;
let rot = (rep + pi) % n_arms;
let mut plain_toks: Option<Vec<u32>> = None;
for step in 0..n_arms {
let which = (rot + step) % n_arms;
if which == 0 {
let mut state = gpu.alloc_decode_state().expect("alloc state");
let pre = gpu.prefill_with_cache(ids, &mut state).expect("prefill");
let mut t = argmax(&pre.logits);
let mut toks = Vec::with_capacity(n_new);
let dt0 = std::time::Instant::now();
for st in 0..n_new {
toks.push(t);
if st + 1 == n_new {
break;
}
t = gpu.decode_step_greedy(t, &mut state).expect("plain step");
}
plain.decode_us += dt0.elapsed().as_micros() as u64;
plain.tokens += toks.len();
plain_toks = Some(toks);
} else {
let ai = which - 1;
let depth = arms[ai].depth;
let mut state = gpu.alloc_decode_state().expect("alloc state");
let mut dstate = gpu.dspark_alloc_state().expect("alloc dspark");
let mut vstate = gpu.alloc_verify_state().expect("alloc verify");
let dt0 = std::time::Instant::now();
let out = gpu
.spec_greedy_batched_depth(
ids,
n_new,
&mut state,
&mut dstate,
&mut vstate,
depth,
)
.expect("drafted run");
let total_us = dt0.elapsed().as_micros() as u64;
let rounds_us: u64 = out.rounds.iter().map(|r| r.round_us).sum();
let _prefill_us = total_us.saturating_sub(rounds_us);
let a = &mut arms[ai];
a.decode_us += rounds_us;
a.tokens += out.tokens.len();
a.rounds += out.rounds.len();
for r in &out.rounds {
a.accepts += r.accepts;
a.t_batch_sum += r.t_batch;
a.recs.push((
r.t_batch,
r.t_cap,
r.accepts,
r.emitted,
r.round_us,
r.confidence.clone(),
));
}
if let Some(pt) = &plain_toks {
if &out.tokens != pt {
let first = out
.tokens
.iter()
.zip(pt)
.position(|(a, b)| a != b)
.unwrap_or(pt.len().min(out.tokens.len()));
identity_fails.push(format!(
"IDENTITY FAIL rep{rep} prompt{pi} ({pool}) T={depth}: drafted \
!= plain at generated index {first}"
));
}
}
}
}
}
println!(
" rep {} / {reps} done, t={:.0}s",
rep + 1,
t0.elapsed().as_secs_f64()
);
}
let p_ms = plain.ms_per_token();
println!("\n=== DEPTH-SWEEP RESULT (owner corpora, measured, bench not serving) ===");
println!(
"PLAIN : {:.3} ms/tok = {:.2} tok/s ({} tokens)",
p_ms,
plain.tok_s(),
plain.tokens
);
println!(
"\n{:<4} {:>10} {:>9} {:>8} {:>9} {:>9} {:>9} {:>11} {:>9}",
"T", "ms/tok", "tok/s", "vs plain", "meanT", "acc/rnd", "tau*", "round ms", "SPS/s"
);
let mut table_json: Vec<String> = Vec::new();
for a in &arms {
let steady: Vec<&(usize, usize, usize, usize, u64, Vec<f32>)> =
a.recs.iter().filter(|r| r.0 > 0 && r.0 == r.1).collect();
let (mean_round_us, tau, sps) = if steady.is_empty() {
(0.0, 0.0, 0.0)
} else {
let n = steady.len() as f64;
let m = steady.iter().map(|r| r.4 as f64).sum::<f64>() / n;
let e = steady.iter().map(|r| r.3 as f64).sum::<f64>() / n;
(m, e, 1e6 / m)
};
let ms = a.decode_us as f64 / 1e3 / a.tokens as f64;
println!(
"{:<4} {:>10.3} {:>9.2} {:>8.3}x {:>9.4} {:>9.4} {:>9.4} {:>11.3} {:>9.2}",
a.depth,
ms,
a.tokens as f64 * 1e6 / a.decode_us as f64,
p_ms / ms,
a.t_batch_sum as f64 / a.rounds.max(1) as f64,
a.accepts as f64 / a.rounds.max(1) as f64,
tau,
mean_round_us / 1e3,
sps
);
let kmax = steady.iter().map(|r| r.0 - 1).max().unwrap_or(0);
let mut pos: Vec<String> = Vec::new();
for j in 0..kmax {
let reached = steady.iter().filter(|r| r.2 >= j).count();
let took = steady.iter().filter(|r| r.2 > j).count();
if reached > 0 {
pos.push(format!("p{}={:.4}", j + 1, took as f64 / reached as f64));
}
}
println!(
" per-position conditional acceptance: {}",
pos.join(" ")
);
table_json.push(format!(
"{{\"T\": {}, \"ms_per_tok\": {ms:.4}, \"tok_s\": {:.4}, \"speedup\": {:.4}, \
\"mean_T\": {:.4}, \"accepts_per_round\": {:.4}, \"tau_star\": {tau:.4}, \
\"mean_round_us\": {mean_round_us:.1}, \"sps\": {sps:.3}, \"rounds\": {}, \
\"steady_rounds\": {}}}",
a.depth,
a.tokens as f64 * 1e6 / a.decode_us as f64,
p_ms / ms,
a.t_batch_sum as f64 / a.rounds.max(1) as f64,
a.accepts as f64 / a.rounds.max(1) as f64,
a.rounds,
steady.len()
));
let side = out_dir.join(format!("spec_rounds_T{}.json", a.depth));
let mut f = std::fs::File::create(&side).expect("sidecar");
writeln!(
f,
"{{\"schema\": \"dsv4-spec-rounds-v1\", \"T\": {}, \"rounds\": [",
a.depth
)
.expect("w");
for (i, r) in a.recs.iter().enumerate() {
let conf =
r.5.iter()
.map(|c| format!("{c:.6}"))
.collect::<Vec<_>>()
.join(",");
writeln!(
f,
" {{\"t_batch\": {}, \"t_cap\": {}, \"accepts\": {}, \"emitted\": {}, \
\"round_us\": {}, \"conf\": [{conf}]}}{}",
r.0,
r.1,
r.2,
r.3,
r.4,
if i + 1 == a.recs.len() { "" } else { "," }
)
.expect("w");
}
writeln!(f, "]}}").expect("w");
}
println!(
"\nidentity across ALL depths: {}",
if identity_fails.is_empty() {
"PASS (every drafted arm reproduced plain token-for-token)".to_string()
} else {
format!("FAIL ({} cases)", identity_fails.len())
}
);
for l in identity_fails.iter().take(10) {
println!(" {l}");
}
println!(
"*** acceptance is a CORRECTNESS observable, never a speed claim (dspark-q38 law); \
the speed statement is the measured ms/tok column ***"
);
let jf = out_dir.join("depth_sweep.json");
let mut f = std::fs::File::create(&jf).expect("json");
write!(
f,
"{{\n \"cell\": \"dsv4-depth-sweep\",\n \"corpora\": \"owner-sxc\",\n \
\"n_new\": {n_new},\n \"reps\": {reps},\n \"prompts\": {},\n \
\"plain_ms_per_tok\": {p_ms:.4},\n \"plain_tok_s\": {:.4},\n \
\"identity_fails\": {},\n \"arms\": [{}]\n}}\n",
prompts.len(),
plain.tok_s(),
identity_fails.len(),
table_json.join(", ")
)
.expect("w");
println!(
"banked {} | total elapsed {:.0}s",
jf.display(),
t0.elapsed().as_secs_f64()
);
}