use lattice_inference::tokenizer::common::Tokenizer as _;
use std::io::Write as _;
use std::path::PathBuf;
use std::time::Instant;
fn parse_arg(args: &[String], flag: &str) -> Option<String> {
args.iter()
.position(|a| a == flag)
.and_then(|i| args.get(i + 1))
.cloned()
}
fn has_flag(args: &[String], flag: &str) -> bool {
args.iter().any(|a| a == flag)
}
fn resolve_model_dir(args: &[String]) -> PathBuf {
if let Some(dir) = parse_arg(args, "--model-dir") {
PathBuf::from(dir)
} else {
let model_name = parse_arg(args, "--model").unwrap_or_else(|| "qwen3.5-0.8b".to_string());
std::env::var("LATTICE_MODEL_CACHE")
.map(PathBuf::from)
.unwrap_or_else(|_| {
let home = std::env::var("HOME").expect("HOME not set");
PathBuf::from(home).join(".lattice").join("models")
})
.join(model_name)
}
}
fn main() {
let args: Vec<String> = std::env::args().collect();
if has_flag(&args, "--emit-phase-events") {
std::process::exit(run_emit_phase_events(&args));
}
let prompt =
parse_arg(&args, "--prompt").unwrap_or_else(|| "What is the meaning of life?".to_string());
let max_tokens: usize = parse_arg(&args, "--max-tokens")
.and_then(|s| s.parse().ok())
.unwrap_or(64);
let seed: Option<u64> = parse_arg(&args, "--seed").and_then(|s| s.parse().ok());
let temperature: Option<f32> = parse_arg(&args, "--temperature").and_then(|s| s.parse().ok());
let model_dir = resolve_model_dir(&args);
println!("Loading model from {model_dir:?}...");
let t0 = Instant::now();
let model = match lattice_inference::model::qwen35::Qwen35Model::from_safetensors(&model_dir) {
Ok(m) => m,
Err(e) => {
eprintln!("Failed to load model: {e}");
std::process::exit(1);
}
};
let load_ms = t0.elapsed().as_millis();
println!("Model loaded in {load_ms}ms\n");
let mut gen_cfg = lattice_inference::model::qwen35_config::GenerateConfig {
max_new_tokens: max_tokens,
seed,
..Default::default()
};
if let Some(t) = temperature {
gen_cfg.temperature = t;
}
if let Some(rp) = parse_arg(&args, "--repetition-penalty").and_then(|s| s.parse().ok()) {
gen_cfg.repetition_penalty = rp;
}
println!("Prompt: {prompt}");
println!(
"Config: temp={}, top_k={}, top_p={}, rep_penalty={}, seed={:?}",
gen_cfg.temperature, gen_cfg.top_k, gen_cfg.top_p, gen_cfg.repetition_penalty, gen_cfg.seed
);
println!("Generating up to {max_tokens} tokens...\n");
let t1 = Instant::now();
match model.generate(&prompt, &gen_cfg) {
Ok(output) => {
let gen_ms = t1.elapsed().as_millis();
let tokens_per_sec = if gen_ms > 0 {
output.generated_tokens as f64 / (gen_ms as f64 / 1000.0)
} else {
0.0
};
println!("--- Generated Text ---");
println!("{}", output.text);
println!("--- Stats ---");
println!("Token IDs: {:?}", output.token_ids);
println!("Prompt tokens: {}", output.prompt_tokens);
println!("Generated tokens: {}", output.generated_tokens);
println!("Time: {gen_ms}ms");
println!("Speed: {tokens_per_sec:.1} tok/s");
}
Err(e) => {
eprintln!("Generation failed: {e}");
std::process::exit(1);
}
}
}
fn emit_phase(t0: Instant, name: &str) {
println!(
"@@bench {{\"ev\":\"phase\",\"name\":\"{name}\",\"monotonic_ns\":{}}}",
t0.elapsed().as_nanos()
);
let _ = std::io::stdout().flush();
}
fn emit_token_available(t0: Instant, token_index: usize) {
println!(
"@@bench {{\"ev\":\"phase\",\"name\":\"token_available\",\"monotonic_ns\":{},\"token_index\":{token_index}}}",
t0.elapsed().as_nanos()
);
let _ = std::io::stdout().flush();
}
const FILLER_TEXT: &str = "The quick brown fox jumps over the lazy dog while the \
autumn wind carries fallen leaves across the quiet stone courtyard and a \
distant bell rings twice before the evening market finally closes its \
wooden stalls for the night and the old lighthouse keeper climbs the \
spiral stairs to trim the lamp before the fog rolls in from the cold \
northern sea and every ship still out on the water turns slowly toward \
the safety of the sheltered harbor lights";
fn build_prompt_of_exact_length(
model: &lattice_inference::model::qwen35::Qwen35Model,
target: usize,
) -> Result<String, String> {
if target == 0 {
return Err("--context must be a positive token count".to_string());
}
let words: Vec<&str> = FILLER_TEXT.split_whitespace().collect();
let mut text = String::new();
let max_word_steps = target * 2 + 64;
for word_idx in 0..max_word_steps {
let real_len = model.tokenizer().tokenize(&text).real_length;
if real_len == target {
return Ok(text);
}
if real_len > target {
break;
}
if !text.is_empty() {
text.push(' ');
}
text.push_str(words[word_idx % words.len()]);
}
if let Some(last_space) = text.rfind(' ') {
text.truncate(last_space);
} else {
text.clear();
}
let fill_chars: Vec<char> = FILLER_TEXT.chars().filter(|c| !c.is_whitespace()).collect();
if fill_chars.is_empty() {
return Err("internal error: empty filler alphabet".to_string());
}
let max_char_steps = target * 4 + 256;
for char_idx in 0..max_char_steps {
let real_len = model.tokenizer().tokenize(&text).real_length;
if real_len == target {
return Ok(text);
}
if real_len > target {
return Err(format!(
"overshot exact token target {target} during character-level backfill \
(reached {real_len} tokens) -- tokenizer merge behavior at this boundary \
is not monotonic for the current filler text; this context point cannot \
be built exactly and the cell must be marked unsupported, not approximated"
));
}
text.push(' ');
text.push(fill_chars[char_idx % fill_chars.len()]);
}
Err(format!(
"could not reach exact token target {target} after word- and character-level \
growth (stuck below target) -- this context point cannot be built exactly"
))
}
fn run_emit_phase_events(args: &[String]) -> i32 {
let t0 = Instant::now();
emit_phase(t0, "load_start");
if parse_arg(args, "--prompt").is_some() {
eprintln!(
"FAIL: --emit-phase-events builds an exact-token-length prompt internally \
from --context; pass --context <N>, not --prompt"
);
return 1;
}
let context: usize = match parse_arg(args, "--context").and_then(|s| s.parse().ok()) {
Some(v) if v >= 1 => v,
_ => {
eprintln!("FAIL: --emit-phase-events requires --context <positive integer>");
return 1;
}
};
let max_tokens: usize = parse_arg(args, "--max-tokens")
.and_then(|s| s.parse().ok())
.unwrap_or(128);
let seed: u64 = parse_arg(args, "--seed")
.and_then(|s| s.parse().ok())
.unwrap_or(42);
let warmup_tokens: usize = parse_arg(args, "--warmup-tokens")
.and_then(|s| s.parse().ok())
.unwrap_or(16);
let model_dir = resolve_model_dir(args);
let mut model =
match lattice_inference::model::qwen35::Qwen35Model::from_safetensors(&model_dir) {
Ok(m) => m,
Err(e) => {
eprintln!("FAIL: could not load model from {model_dir:?}: {e}");
return 1;
}
};
emit_phase(t0, "backend_ready");
let prompt = match build_prompt_of_exact_length(&model, context) {
Ok(p) => p,
Err(e) => {
eprintln!("FAIL: {e}");
return 1;
}
};
let actual_prompt_tokens = model.tokenizer().tokenize(&prompt).real_length;
if actual_prompt_tokens != context {
eprintln!(
"FAIL: internal error -- built prompt tokenizes to {actual_prompt_tokens} \
tokens, expected exactly {context}"
);
return 1;
}
model.set_eos_token_id(u32::MAX);
let base_cfg = lattice_inference::model::qwen35_config::GenerateConfig {
max_new_tokens: max_tokens,
seed: Some(seed),
temperature: 0.0,
top_k: 1,
top_p: 1.0,
repetition_penalty: 1.0,
stop_token_ids: vec![],
enable_thinking: false,
enable_mtp: Some(false),
..Default::default()
};
if warmup_tokens > 0 {
let warmup_cfg = lattice_inference::model::qwen35_config::GenerateConfig {
max_new_tokens: warmup_tokens,
..base_cfg.clone()
};
if let Err(e) =
model.generate_streaming_with_cancel(&prompt, &warmup_cfg, |_delta| true, || false)
{
eprintln!("FAIL: untimed warmup failed: {e}");
return 1;
}
}
emit_phase(t0, "prefill_start");
let mut raw_token_count = 0usize;
let mut delta_call_count = 0usize;
let result = model.generate_streaming_with_observer(
&prompt,
&base_cfg,
|_delta: &str| {
delta_call_count += 1;
true
},
|| false,
|evt| match evt {
lattice_inference::model::qwen35::RawGenEvent::PrefillEnd => {
emit_phase(t0, "prefill_end");
}
lattice_inference::model::qwen35::RawGenEvent::RawToken { index } => {
raw_token_count += 1;
emit_token_available(t0, index);
}
},
);
let output = match result {
Ok(o) => o,
Err(e) => {
eprintln!("FAIL: measured generation failed: {e}");
return 1;
}
};
if raw_token_count != output.generated_tokens {
eprintln!(
"FAIL: internal error -- raw phase-event token count ({raw_token_count}) \
does not match GenerateOutput.generated_tokens ({}); the phase-event \
trace for this trial cannot be trusted",
output.generated_tokens
);
return 1;
}
let delta_matches_generated_tokens = delta_call_count == output.generated_tokens;
let model_dir_json = json_escape(&model_dir.display().to_string());
println!(
"@@bench {{\"ev\":\"summary\",\"prompt_tokens\":{},\"generated_tokens\":{},\
\"requested_max_tokens\":{max_tokens},\"delta_call_count\":{delta_call_count},\
\"delta_matches_generated_tokens\":{delta_matches_generated_tokens},\
\"stopped\":{},\"model_dir\":\"{model_dir_json}\",\"seed\":{seed},\"disable_eos\":true}}",
output.prompt_tokens, output.generated_tokens, output.stopped,
);
let _ = std::io::stdout().flush();
0
}
fn json_escape(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for c in s.chars() {
match c {
'"' => out.push_str("\\\""),
'\\' => out.push_str("\\\\"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
'\t' => out.push_str("\\t"),
c if (c as u32) < 0x20 => out.push_str(&format!("\\u{:04x}", c as u32)),
c => out.push(c),
}
}
out
}