use anyhow::{Context, Result};
use serde_json::Value as JsonValue;
use sha2::{Digest, Sha256};
use std::fs;
use std::io::Read;
use std::path::Path;
use std::time::SystemTime;
use crate::cli;
use crate::inference::models::gemma4::{DecodeRegime, MlxModelWeights};
use crate::serve::config::Gemma4Config;
use crate::serve::forward_mlx_shared::cosine_pairwise_f32;
use crate::serve::gpu;
use crate::serve::header;
#[derive(Debug, Clone)]
pub struct CosineStats {
pub mean: f32,
pub min: f32,
pub p1: f32,
pub p50: f32,
pub p99: f32,
pub n_pairs: usize,
}
#[derive(Debug, Clone)]
pub struct GateHEnvelope {
pub cosine: CosineStats,
pub argmax_flip_rate: f32,
pub ppl_dense: f64,
pub ppl_tq: f64,
pub ppl_delta_pct: f64, pub n_steps: usize,
}
#[derive(Debug, Clone)]
struct PassCapture {
pre_replay_argmax: Vec<u32>,
final_tokens: Vec<u32>,
nll_per_step: Vec<f32>,
}
pub fn cmd_parity_capture_tq_quality(
model_path: &Path,
output_dir: &Path,
prompt_name: &str,
max_tokens: Option<usize>,
) -> Result<()> {
if prompt_name == "all" {
anyhow::bail!(
"parity capture --tq-quality requires a single prompt; \
`--prompt all` is not supported (Gate H fixtures are per-prompt)"
);
}
let evals_dir = Path::new("tests/evals");
let prompt_file = evals_dir.join("prompts").join(format!("{prompt_name}.txt"));
anyhow::ensure!(
prompt_file.exists(),
"Prompt file not found: {}",
prompt_file.display()
);
let prompt_text = fs::read_to_string(&prompt_file)?.trim().to_string();
let tokens = max_tokens.unwrap_or(1000);
let dump_root =
std::env::temp_dir().join(format!("hf2q_gate_h_capture_{}", std::process::id()));
fs::create_dir_all(&dump_root)?;
let dense_dump_dir = dump_root.join("dense");
let tq_dump_dir = dump_root.join("tq");
fs::create_dir_all(&dense_dump_dir)?;
fs::create_dir_all(&tq_dump_dir)?;
eprintln!("=== Gate H Capture: {} ===", prompt_name);
eprintln!("Model: {}", model_path.display());
eprintln!(
"Prompt: {} ({} chars)",
prompt_name,
prompt_text.len()
);
eprintln!("Tokens: {}", tokens);
eprintln!("Dump root: {}", dump_root.display());
eprintln!();
let (envelope, dense_capture) = run_two_regime_decode(
model_path,
&prompt_text,
tokens,
&dump_root,
&dense_dump_dir,
&tq_dump_dir,
)?;
let model_sha =
sha256_file(model_path).with_context(|| format!("hash model: {}", model_path.display()))?;
let git_head = git_head_sha();
let captured_at = iso8601_utc_now();
let fixture = serde_json::json!({
"git_head": git_head,
"model_sha256": model_sha,
"prompt": prompt_name,
"max_tokens": tokens,
"captured_at": captured_at,
"dense_tokens": dense_capture.final_tokens,
"dense_nll_per_step": dense_capture
.nll_per_step
.iter()
.map(|n| *n as f64)
.collect::<Vec<_>>(),
"envelope": {
"cosine_mean": envelope.cosine.mean as f64,
"cosine_p1": envelope.cosine.p1 as f64,
"argmax_div": envelope.argmax_flip_rate as f64,
"ppl_delta_pct": envelope.ppl_delta_pct,
}
});
fs::create_dir_all(output_dir)?;
let out_path = output_dir.join(format!("{prompt_name}_tq_quality.json"));
let pretty = serde_json::to_string_pretty(&fixture).context("serialize fixture")?;
fs::write(&out_path, pretty)?;
eprintln!();
eprintln!("Wrote fixture: {}", out_path.display());
eprintln!(
" cosine_mean={:.6} cosine_p1={:.6} argmax_div={:.4} ppl_delta={:.4}",
envelope.cosine.mean, envelope.cosine.p1, envelope.argmax_flip_rate, envelope.ppl_delta_pct
);
let _ = fs::remove_dir_all(&dump_root);
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn cmd_parity_check_tq_quality(
model_path: &Path,
prompt_name: &str,
fixture_path: &Path,
cosine_mean_floor: f32,
cosine_p1_floor: f32,
argmax_max: f32,
ppl_delta_max: f32,
max_tokens: Option<usize>,
) -> Result<()> {
anyhow::ensure!(
fixture_path.exists(),
"Gate H fixture not found: {}.\n\
Hint: run `hf2q parity capture --tq-quality --model <gguf> \
--prompt {prompt_name}` first (iter-112 closes this loop).",
fixture_path.display()
);
let fixture_str = fs::read_to_string(fixture_path)
.with_context(|| format!("read fixture: {}", fixture_path.display()))?;
let fixture: JsonValue = serde_json::from_str(&fixture_str).context("parse fixture JSON")?;
let fixture_prompt = fixture["prompt"].as_str().unwrap_or("");
anyhow::ensure!(
fixture_prompt == prompt_name,
"Fixture prompt mismatch: fixture has {:?}, --prompt is {:?}",
fixture_prompt,
prompt_name
);
let fixture_tokens = fixture["max_tokens"].as_u64().unwrap_or(0) as usize;
let tokens = max_tokens.unwrap_or(fixture_tokens.max(1));
anyhow::ensure!(
tokens == fixture_tokens,
"Token count mismatch: fixture max_tokens={fixture_tokens}, \
--max-tokens={tokens}. Re-capture or pass --max-tokens={fixture_tokens}."
);
let cur_sha = sha256_file(model_path).unwrap_or_else(|_| "unknown".into());
if let Some(fixture_sha) = fixture["model_sha256"].as_str() {
if fixture_sha != cur_sha {
eprintln!(
"[gate-h] warning: model_sha256 mismatch (fixture={}, current={})",
fixture_sha, cur_sha
);
}
}
let evals_dir = Path::new("tests/evals");
let prompt_file = evals_dir.join("prompts").join(format!("{prompt_name}.txt"));
anyhow::ensure!(
prompt_file.exists(),
"Prompt file not found: {}",
prompt_file.display()
);
let prompt_text = fs::read_to_string(&prompt_file)?.trim().to_string();
let dump_root = std::env::temp_dir().join(format!("hf2q_gate_h_check_{}", std::process::id()));
fs::create_dir_all(&dump_root)?;
let dense_dump_dir = dump_root.join("dense");
let tq_dump_dir = dump_root.join("tq");
fs::create_dir_all(&dense_dump_dir)?;
fs::create_dir_all(&tq_dump_dir)?;
eprintln!("=== Gate H Check: {} ===", prompt_name);
eprintln!("Model: {}", model_path.display());
eprintln!("Fixture: {}", fixture_path.display());
eprintln!("Tokens: {}", tokens);
eprintln!(
"Floors: cosine_mean>={cosine_mean_floor:.6} \
cosine_p1>={cosine_p1_floor:.6} \
argmax<={argmax_max:.4} \
ppl_delta<={ppl_delta_max:.4}"
);
eprintln!();
let (envelope, _dense_capture) = run_two_regime_decode(
model_path,
&prompt_text,
tokens,
&dump_root,
&dense_dump_dir,
&tq_dump_dir,
)?;
let _ = fs::remove_dir_all(&dump_root);
eprintln!();
eprintln!("=== Gate H Envelope (this run) ===");
eprintln!(
" cosine: mean={:.6} min={:.6} p1={:.6} p50={:.6} p99={:.6} n_pairs={}",
envelope.cosine.mean,
envelope.cosine.min,
envelope.cosine.p1,
envelope.cosine.p50,
envelope.cosine.p99,
envelope.cosine.n_pairs,
);
eprintln!(
" argmax_flip_rate: {:.4} ({} flips / {} steps)",
envelope.argmax_flip_rate,
(envelope.argmax_flip_rate * envelope.n_steps as f32).round() as usize,
envelope.n_steps,
);
eprintln!(
" PPL: dense={:.4} tq={:.4} delta={:.4}",
envelope.ppl_dense, envelope.ppl_tq, envelope.ppl_delta_pct,
);
if let Some(fixture_env) = fixture.get("envelope") {
eprintln!();
eprintln!("=== Fixture Envelope (frozen reference) ===");
eprintln!(
" cosine_mean={:.6} cosine_p1={:.6} argmax_div={:.4} ppl_delta={:.4}",
fixture_env["cosine_mean"].as_f64().unwrap_or(f64::NAN),
fixture_env["cosine_p1"].as_f64().unwrap_or(f64::NAN),
fixture_env["argmax_div"].as_f64().unwrap_or(f64::NAN),
fixture_env["ppl_delta_pct"].as_f64().unwrap_or(f64::NAN),
);
}
let mut failures: Vec<String> = Vec::new();
if envelope.cosine.mean < cosine_mean_floor {
failures.push(format!(
"cosine_mean {:.6} < floor {:.6}",
envelope.cosine.mean, cosine_mean_floor
));
}
if envelope.cosine.p1 < cosine_p1_floor {
failures.push(format!(
"cosine_p1 {:.6} < floor {:.6}",
envelope.cosine.p1, cosine_p1_floor
));
}
if envelope.argmax_flip_rate > argmax_max {
failures.push(format!(
"argmax_flip_rate {:.4} > max {:.4}",
envelope.argmax_flip_rate, argmax_max
));
}
if (envelope.ppl_delta_pct as f32) > ppl_delta_max {
failures.push(format!(
"ppl_delta {:.4} > max {:.4}",
envelope.ppl_delta_pct, ppl_delta_max
));
}
eprintln!();
if failures.is_empty() {
println!("PASS: Gate H envelope holds (ADR-007 §853-866)");
Ok(())
} else {
for f in &failures {
eprintln!("FAIL: {}", f);
}
anyhow::bail!("Gate H failed: {} threshold(s) tripped", failures.len())
}
}
fn run_two_regime_decode(
model_path: &Path,
prompt_text: &str,
tokens: usize,
dump_root: &Path,
dense_dump_dir: &Path,
tq_dump_dir: &Path,
) -> Result<(GateHEnvelope, PassCapture)> {
unsafe {
std::env::set_var("HF2Q_DUMP_SDPA_MAX_POS", tokens.to_string());
}
let tokenizer_path = super::find_tokenizer(model_path, None)?;
let mut ctx = gpu::GpuContext::new().map_err(|e| anyhow::anyhow!("GPU init: {e}"))?;
let gguf = mlx_native::gguf::GgufFile::open(model_path)
.map_err(|e| anyhow::anyhow!("GGUF open: {e}"))?;
let cfg = Gemma4Config::from_gguf(&gguf)?;
let mut progress = header::LoadProgress::new(false, 1, 0);
let mut mlx_w = MlxModelWeights::load_from_gguf(&gguf, &cfg, &mut ctx, &mut progress)?;
let mut tokenizer = tokenizers::Tokenizer::from_file(&tokenizer_path)
.map_err(|e| anyhow::anyhow!("Tokenizer: {e}"))?;
tokenizer
.with_truncation(None)
.map_err(|e| anyhow::anyhow!("Tokenizer truncation: {e}"))?;
let rendered = super::render_chat_template(
&gguf,
&cli::GenerateArgs {
model: model_path.to_path_buf(),
prompt: Some(prompt_text.to_string()),
prompt_file: None,
tokenizer: None,
config: None,
temperature: 0.0,
top_p: 1.0,
top_k: 0,
min_p: 0.0,
repetition_penalty: 1.0,
max_tokens: tokens,
mmproj: None,
image: None,
chat_template: None,
chat_template_file: None,
benchmark: false,
speculative: false,
kv_bits: None,
enable_thinking: false,
no_thinking: false,
ignore_eos: false,
},
Some(&tokenizer),
prompt_text,
)?;
let encoding = tokenizer
.encode(rendered.as_str(), false)
.map_err(|e| anyhow::anyhow!("Tokenize: {e}"))?;
let prompt_tokens: Vec<u32> = encoding.get_ids().to_vec();
let prompt_len = prompt_tokens.len();
eprintln!(
"[gate-h] Pass 1/2: dense forced regime ({} tokens)...",
tokens
);
mlx_w.set_decode_regime(DecodeRegime::ForceDense);
mlx_w.set_dump_overrides(Some(dump_root.to_path_buf()), Some(true));
let dense_capture = run_one_pass(
&mut mlx_w,
&mut ctx,
&prompt_tokens,
tokens,
None,
)?;
move_sdpa_dumps(dump_root, dense_dump_dir, prompt_len, tokens)?;
eprintln!(
"[gate-h] Pass 2/2: TQ forced regime + dense-token replay ({} tokens)...",
tokens
);
mlx_w.set_decode_regime(DecodeRegime::ForceTq);
mlx_w.set_replay_tokens(dense_capture.final_tokens.clone());
mlx_w.set_dump_overrides(Some(dump_root.to_path_buf()), Some(true));
let tq_capture = run_one_pass(
&mut mlx_w,
&mut ctx,
&prompt_tokens,
tokens,
Some(&dense_capture.final_tokens),
)?;
move_sdpa_dumps(dump_root, tq_dump_dir, prompt_len, tokens)?;
let num_layers = mlx_w.layers.len();
let cosine = synthesize_cosine(
dense_dump_dir,
tq_dump_dir,
num_layers,
prompt_len,
tokens,
&mlx_w,
)?;
let ppl_dense = nll_to_ppl(&dense_capture.nll_per_step);
let ppl_tq = nll_to_ppl(&tq_capture.nll_per_step);
let ppl_delta_pct = if ppl_dense > 0.0 && ppl_dense.is_finite() {
((ppl_tq - ppl_dense).abs() / ppl_dense) as f64
} else {
f64::NAN
};
let n = dense_capture
.final_tokens
.len()
.min(tq_capture.pre_replay_argmax.len());
let mut flips = 0usize;
let mut flip_positions: Vec<(usize, u32, u32)> = Vec::new();
for i in 0..n {
if tq_capture.pre_replay_argmax[i] != dense_capture.final_tokens[i] {
flips += 1;
flip_positions.push((
i,
dense_capture.final_tokens[i],
tq_capture.pre_replay_argmax[i],
));
}
}
let argmax_flip_rate = if n == 0 { 0.0 } else { flips as f32 / n as f32 };
if !flip_positions.is_empty() {
eprintln!("[GATE_H_DIAG] argmax flip positions (step, dense_token, tq_token):");
for (step, d, t) in &flip_positions {
eprintln!(" step {}: dense={} tq={}", step, d, t);
}
}
let envelope = GateHEnvelope {
cosine,
argmax_flip_rate,
ppl_dense: ppl_dense as f64,
ppl_tq: ppl_tq as f64,
ppl_delta_pct,
n_steps: n,
};
Ok((envelope, dense_capture))
}
fn run_one_pass(
mlx_w: &mut MlxModelWeights,
ctx: &mut gpu::GpuContext,
prompt_tokens: &[u32],
tokens: usize,
replay: Option<&[u32]>,
) -> Result<PassCapture> {
let first_returned = mlx_w.forward_prefill(prompt_tokens, tokens, ctx)?;
let scored_first = match replay {
Some(r) if !r.is_empty() => r[0],
_ => first_returned,
};
let first_nll = mlx_w.token_nll_from_logits(scored_first)?;
let mut pre_replay_argmax: Vec<u32> = Vec::with_capacity(tokens);
let mut final_tokens: Vec<u32> = Vec::with_capacity(tokens);
let mut nll_per_step: Vec<f32> = Vec::with_capacity(tokens);
pre_replay_argmax.push(first_returned);
final_tokens.push(scored_first);
nll_per_step.push(first_nll);
for step in 1..tokens {
let prev_token = final_tokens[step - 1];
let pos = prompt_tokens.len() + step - 1;
let mut profile = None;
let returned = mlx_w.forward_decode(prev_token, pos, ctx, &mut profile)?;
let live_argmax = argmax_from_logits(mlx_w.logits_view()?);
let scored = match replay {
Some(r) if step < r.len() => r[step],
_ => returned,
};
let nll = mlx_w.token_nll_from_logits(scored)?;
pre_replay_argmax.push(live_argmax);
final_tokens.push(scored);
nll_per_step.push(nll);
}
Ok(PassCapture {
pre_replay_argmax,
final_tokens,
nll_per_step,
})
}
fn argmax_from_logits(logits: &[f32]) -> u32 {
let mut best_idx = 0usize;
let mut best_val = f32::NEG_INFINITY;
for (i, &v) in logits.iter().enumerate() {
if v > best_val {
best_val = v;
best_idx = i;
}
}
best_idx as u32
}
fn move_sdpa_dumps(src: &Path, dst: &Path, prompt_len: usize, tokens: usize) -> Result<()> {
if !src.exists() {
return Ok(());
}
for entry in fs::read_dir(src)? {
let entry = entry?;
let path = entry.path();
if !path.is_file() {
continue;
}
let name = match path.file_name().and_then(|n| n.to_str()) {
Some(n) => n,
None => continue,
};
if !name.starts_with("hf2q_sdpa_out_layer") {
continue;
}
let dest = dst.join(name);
if let Err(e) = fs::rename(&path, &dest) {
if let Err(copy_err) = fs::copy(&path, &dest) {
anyhow::bail!(
"move sdpa dump {} -> {}: rename={e}; copy={copy_err}",
path.display(),
dest.display()
);
}
let _ = fs::remove_file(&path);
}
}
let _ = (prompt_len, tokens); Ok(())
}
fn synthesize_cosine(
dense_dir: &Path,
tq_dir: &Path,
num_layers: usize,
prompt_len: usize,
tokens: usize,
mlx_w: &MlxModelWeights,
) -> Result<CosineStats> {
let mut cosines: Vec<f32> = Vec::with_capacity(num_layers * tokens);
let mut cosines_per_layer: Vec<Vec<f32>> = vec![Vec::new(); num_layers];
let mut cosines_per_layer_with_step: Vec<Vec<(usize, f32)>> = vec![Vec::new(); num_layers];
for layer_idx in 0..num_layers {
let nh = mlx_w.num_attention_heads;
let hd = mlx_w.layers[layer_idx].head_dim;
let n_elems = nh * hd;
for step in 0..tokens {
let seq_pos = if step == 0 {
prompt_len
} else {
prompt_len + step - 1
};
let fname = format!("hf2q_sdpa_out_layer{:02}_pos{}.bin", layer_idx, seq_pos);
let dense_path = dense_dir.join(&fname);
let tq_path = tq_dir.join(&fname);
if !dense_path.exists() || !tq_path.exists() {
continue;
}
let dense_vec = read_f32_bin(&dense_path, n_elems)?;
let tq_vec = read_f32_bin(&tq_path, n_elems)?;
let cs = cosine_pairwise_f32(&dense_vec, &tq_vec);
if cs.is_finite() {
cosines.push(cs);
cosines_per_layer[layer_idx].push(cs);
cosines_per_layer_with_step[layer_idx].push((step, cs));
}
}
}
eprintln!("[GATE_H_DIAG] per-layer cosine_mean (sdpa_out dense vs TQ):");
for (layer_idx, cs) in cosines_per_layer.iter().enumerate() {
if cs.is_empty() {
eprintln!(" layer {:02}: (no pairs)", layer_idx);
continue;
}
let mean = cs.iter().copied().map(|x| x as f64).sum::<f64>() / cs.len() as f64;
let min = cs.iter().copied().fold(f32::INFINITY, f32::min);
let (min_step, _) = cosines_per_layer_with_step[layer_idx]
.iter()
.copied()
.min_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
.unwrap_or((usize::MAX, f32::INFINITY));
eprintln!(
" layer {:02}: mean={:.6} min={:.6} min_step={} n={}",
layer_idx,
mean,
min,
min_step,
cs.len()
);
}
if cosines.is_empty() {
anyhow::bail!(
"Gate H synthesis: no cosine pairs found. Check that \
HF2Q_DUMP_SDPA_MAX_POS plumbing actually fired (decode dump \
path requires HF2Q_DUMP_ALL_CACHE=1 + dense-or-TQ SDPA branch \
reached at sdpa_out write time)."
);
}
cosines.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
let n = cosines.len();
let mean = cosines.iter().map(|x| *x as f64).sum::<f64>() / n as f64;
let min = cosines[0];
let p1 = cosines[((n as f64 * 0.01) as usize).min(n - 1)];
let p50 = cosines[n / 2];
let p99 = cosines[((n as f64 * 0.99) as usize).min(n - 1)];
Ok(CosineStats {
mean: mean as f32,
min,
p1,
p50,
p99,
n_pairs: n,
})
}
fn read_f32_bin(path: &Path, n_elems: usize) -> Result<Vec<f32>> {
let mut f = fs::File::open(path).with_context(|| format!("open dump: {}", path.display()))?;
let mut bytes = Vec::with_capacity(n_elems * 4);
f.read_to_end(&mut bytes)
.with_context(|| format!("read dump: {}", path.display()))?;
let bytes_f32_count = bytes.len() / 4;
anyhow::ensure!(
bytes_f32_count >= n_elems,
"dump {} has {} f32 elems, expected >= {}",
path.display(),
bytes_f32_count,
n_elems
);
let mut out = Vec::with_capacity(n_elems);
for i in 0..n_elems {
let off = i * 4;
let arr = [bytes[off], bytes[off + 1], bytes[off + 2], bytes[off + 3]];
out.push(f32::from_le_bytes(arr));
}
Ok(out)
}
fn nll_to_ppl(nlls: &[f32]) -> f64 {
if nlls.is_empty() {
return f64::NAN;
}
let sum: f64 = nlls.iter().map(|x| *x as f64).sum();
(sum / nlls.len() as f64).exp()
}
fn sha256_file(path: &Path) -> Result<String> {
let mut f = fs::File::open(path).with_context(|| format!("open: {}", path.display()))?;
let mut hasher = Sha256::new();
let mut buf = vec![0u8; 1024 * 1024];
loop {
let n = f.read(&mut buf)?;
if n == 0 {
break;
}
hasher.update(&buf[..n]);
}
let digest = hasher.finalize();
Ok(hex::encode(digest))
}
fn git_head_sha() -> String {
use std::process::Command;
Command::new("git")
.arg("rev-parse")
.arg("HEAD")
.output()
.ok()
.and_then(|o| {
if o.status.success() {
String::from_utf8(o.stdout)
.ok()
.map(|s| s.trim().to_string())
} else {
None
}
})
.unwrap_or_else(|| "unknown".to_string())
}
fn iso8601_utc_now() -> String {
let now = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.unwrap_or_default();
let secs = now.as_secs() as i64;
let days = secs.div_euclid(86_400);
let secs_of_day = secs.rem_euclid(86_400) as u32;
let h = secs_of_day / 3600;
let m = (secs_of_day / 60) % 60;
let s = secs_of_day % 60;
let z = days + 719_468; let era = if z >= 0 { z } else { z - 146_096 } / 146_097;
let doe = (z - era * 146_097) as u32; let yoe = (doe - doe / 1460 + doe / 36_524 - doe / 146_096) / 365;
let y = yoe as i64 + era * 400;
let doy = doe - (365 * yoe + yoe / 4 - yoe / 100);
let mp = (5 * doy + 2) / 153;
let d = doy - (153 * mp + 2) / 5 + 1;
let mo = if mp < 10 { mp + 3 } else { mp - 9 };
let year = if mo <= 2 { y + 1 } else { y };
format!("{:04}-{:02}-{:02}T{:02}:{:02}:{:02}Z", year, mo, d, h, m, s)
}