use anyhow::Context;
use std::path::Path;
use std::process::Command;
const N_TOKENS: usize = 24;
const PROMPT: &str = "The capital of France is";
const PREFILL_MIN_TOKENS: usize = 8;
pub struct VerifyArgs {
pub model: String,
pub backend: String,
pub emit: bool,
pub prompt_tokens: Option<usize>,
pub prompt: Option<String>,
pub allow_multiple_instances: bool,
}
const TAG: &str = "FRINK_VERIFY_TOKENS ";
const LEN_TAG: &str = "FRINK_VERIFY_PROMPT_LEN ";
pub fn run(args: VerifyArgs) -> anyhow::Result<()> {
let prompt = args.prompt.clone().unwrap_or_else(|| PROMPT.to_string());
if args.emit {
return emit_tokens(&args.model, &prompt, args.prompt_tokens);
}
let (reference, prompt_len) = child_tokens(&args, "cpu", &prompt)?;
let (candidate, _) = child_tokens(&args, &args.backend, &prompt)?;
if prompt_len < PREFILL_MIN_TOKENS {
eprintln!(
"verify: prompt is {prompt_len} tokens, under the {PREFILL_MIN_TOKENS} at which the \
batched-prefill attention kernels turn on — this run checks decode only. \
Pass --prompt-tokens {PREFILL_MIN_TOKENS} or more (or a longer --prompt) to \
cover prefill."
);
}
if reference.is_empty() {
anyhow::bail!("CPU reference produced no tokens; cannot verify");
}
let diverge = reference
.iter()
.zip(&candidate)
.position(|(a, b)| a != b)
.or_else(|| {
(reference.len() != candidate.len()).then_some(reference.len().min(candidate.len()))
});
match diverge {
None => {
println!(
"verify {}: OK — {} tokens identical on cpu and {} ({})",
short(&args.model),
reference.len(),
args.backend,
prompt_desc(prompt_len)
);
Ok(())
}
Some(i) => {
println!(
"verify {}: DIVERGED at token {i} ({}) — cpu={:?} {}={:?}",
short(&args.model),
prompt_desc(prompt_len),
&reference[i.saturating_sub(2)..reference.len().min(i + 3)],
args.backend,
&candidate[i.saturating_sub(2)..candidate.len().min(i + 3)],
);
anyhow::bail!(
"{} disagrees with the CPU reference from token {i}",
args.backend
)
}
}
}
fn prompt_desc(prompt_len: usize) -> String {
if prompt_len >= PREFILL_MIN_TOKENS {
format!("{prompt_len}-token prompt, prefill covered")
} else {
format!("{prompt_len}-token prompt, decode only")
}
}
fn short(model: &str) -> String {
Path::new(model)
.file_stem()
.map(|s| s.to_string_lossy().into_owned())
.unwrap_or_else(|| model.to_string())
}
fn child_tokens(
args: &VerifyArgs,
backend: &str,
prompt: &str,
) -> anyhow::Result<(Vec<u32>, usize)> {
let exe = std::env::current_exe()?;
let mut cmd = Command::new(&exe);
cmd.arg("verify")
.args(["-m", &args.model])
.args(["--backend", backend])
.args(["--prompt", prompt])
.arg("--emit");
if let Some(n) = args.prompt_tokens {
cmd.args(["--prompt-tokens", &n.to_string()]);
}
if args.allow_multiple_instances {
cmd.env("FRINK_ALLOW_MULTIPLE_INSTANCES", "1");
}
let metal_attn = std::env::var("FRINK_METAL_ATTN").unwrap_or_else(|_| "1".to_string());
let out = cmd
.env("FRINK_METAL", if backend == "cpu" { "0" } else { "1" })
.env(
"FRINK_METAL_ATTN",
if backend == "cpu" { "0" } else { &metal_attn },
)
.env("FRINK_CUDA", if backend == "cuda" { "1" } else { "0" })
.output()
.with_context(|| format!("spawning verify child for {backend}"))?;
if !out.status.success() {
anyhow::bail!(
"verify child for {backend} failed: {}",
String::from_utf8_lossy(&out.stderr)
.lines()
.last()
.unwrap_or("(no stderr)")
);
}
let text = String::from_utf8_lossy(&out.stdout);
let line = text
.lines()
.find_map(|l| l.strip_prefix(TAG))
.with_context(|| format!("verify child for {backend} printed no token line"))?;
let prompt_len = text
.lines()
.find_map(|l| l.strip_prefix(LEN_TAG))
.and_then(|l| l.trim().parse::<usize>().ok())
.with_context(|| format!("verify child for {backend} printed no prompt-length line"))?;
Ok((
line.split_whitespace()
.filter_map(|t| t.parse::<u32>().ok())
.collect(),
prompt_len,
))
}
fn emit_tokens(model: &str, prompt: &str, prompt_tokens: Option<usize>) -> anyhow::Result<()> {
let path = crate::pull::resolve_model_path(model)?;
let (ids, prompt_len) = crate::verify_engine::greedy_token_ids(
Path::new(&path),
prompt,
frink_models::tokenizer::SpecialTokens::Parse,
N_TOKENS,
prompt_tokens,
)?;
let joined: Vec<String> = ids.iter().map(|t| t.to_string()).collect();
println!("{LEN_TAG}{prompt_len}");
println!("{TAG}{}", joined.join(" "));
Ok(())
}