From 6ecdfb0cc236b1c7f3ef52865c03094748775197 Mon Sep 17 00:00:00 2001 From: OceanLi Date: Fri, 25 Sep 2026 15:15:24 -0400 Subject: [PATCH] fix(inference): keep long prompts whole in lattice chat `lattice chat` tokenized each line with the model's default 4096-token tokenizer cap, so on the CPU backend a prompt between 4097 tokens and the model's context window was generated from its first 4096 tokens, and the printed prompt count showed the shortened length. On the Metal backend an over-long prompt was rejected with a count taken from the already-shortened prompt, and with a zero decode budget it was shortened silently. The CPU backend now raises the tokenizer cap to the model's context window when it loads, as `lattice serve` does. Before generating, both backends count the prompt's full length and refuse a prompt longer than the context window with an error naming its real token count and the limit; the REPL reports the error and reads the next line. --- crates/inference/src/bin/lattice/chat.rs | 203 ++++++++++++++---- .../tests/data/pipeline_boundary_baseline.txt | 1 + 2 files changed, 168 insertions(+), 36 deletions(-) diff --git a/crates/inference/src/bin/lattice/chat.rs b/crates/inference/src/bin/lattice/chat.rs index a758100b4e..c9cf2d446e 100644 --- a/crates/inference/src/bin/lattice/chat.rs +++ b/crates/inference/src/bin/lattice/chat.rs @@ -65,7 +65,10 @@ impl MetalChatBackend { prompt: &str, gen_cfg: &lattice_inference::GenerateConfig, ) -> Result { - self.state.generate(prompt, &self.tokenizer, gen_cfg) + let Self { state, tokenizer } = self; + generate_checked(tokenizer, state.max_context(), prompt, || { + state.generate(prompt, tokenizer, gen_cfg) + }) } } @@ -77,6 +80,55 @@ pub(crate) fn chat_max_cache_len() -> usize { MetalChatBackend::MAX_CACHE_LEN } +enum Backend { + Cpu(Box), + #[cfg(feature = "metal-gpu")] + Metal(Box), +} + +impl Backend { + fn cpu(mut model: lattice_inference::model::qwen35::Qwen35Model) -> Self { + model.ensure_tokenizer_max_seq_len(model.max_context()); + Self::Cpu(Box::new(model)) + } + + fn generate_chat_line( + &mut self, + prompt: &str, + gen_cfg: &lattice_inference::GenerateConfig, + ) -> Result { + match self { + Self::Cpu(model) => { + generate_checked(model.tokenizer(), model.max_context(), prompt, || { + model.generate(prompt, gen_cfg) + }) + } + #[cfg(feature = "metal-gpu")] + Self::Metal(model) => model.generate(prompt, gen_cfg), + } + } +} + +fn generate_checked( + tokenizer: &lattice_inference::BpeTokenizer, + limit: usize, + prompt: &str, + generate: impl FnOnce() -> Result< + lattice_inference::GenerateOutput, + lattice_inference::InferenceError, + >, +) -> Result { + use lattice_inference::Tokenizer; + + let prompt_tokens = tokenizer.tokenize(prompt).pre_truncation_len; + if prompt_tokens > limit { + return Err(lattice_inference::InferenceError::InvalidInput(format!( + "prompt ({prompt_tokens} tokens) exceeds model context window ({limit})" + ))); + } + generate() +} + #[allow(clippy::field_reassign_with_default)] pub(crate) fn run_chat( model_path: &str, @@ -96,16 +148,10 @@ pub(crate) fn run_chat( eprintln!("Loading model from {model_path}..."); - enum Backend { - Cpu(Box), - #[cfg(feature = "metal-gpu")] - Metal(Box), - } - let mut model = match format { backend::ModelFormat::Safetensors => { match lattice_inference::model::qwen35::Qwen35Model::from_safetensors(path) { - Ok(m) => Backend::Cpu(Box::new(m)), + Ok(m) => Backend::cpu(m), Err(e) => { eprintln!("Error: failed to load model: {e}"); std::process::exit(1); @@ -166,34 +212,119 @@ pub(crate) fn run_chat( break; } - match &mut model { - Backend::Cpu(m) => match m.generate(trimmed, &gen_cfg) { - Ok(output) => { - let _ = writeln!(stdout, "{}", output.text); - let _ = writeln!( - stdout, - "[{} prompt tokens, {} generated]", - output.prompt_tokens, output.generated_tokens - ); - } - Err(e) => { - eprintln!("Generation error: {e}"); - } - }, - #[cfg(feature = "metal-gpu")] - Backend::Metal(m) => match m.generate(trimmed, &gen_cfg) { - Ok(output) => { - let _ = writeln!(stdout, "{}", output.text); - let _ = writeln!( - stdout, - "[{} prompt tokens, {} generated]", - output.prompt_tokens, output.generated_tokens - ); - } - Err(e) => { - eprintln!("Generation error: {e}"); - } - }, + match model.generate_chat_line(trimmed, &gen_cfg) { + Ok(output) => { + let _ = writeln!(stdout, "{}", output.text); + let _ = writeln!( + stdout, + "[{} prompt tokens, {} generated]", + output.prompt_tokens, output.generated_tokens + ); + } + Err(e) => { + eprintln!("Generation error: {e}"); + } + } + } +} + +#[cfg(all(test, feature = "test-utils"))] +mod tests { + use super::*; + use lattice_inference::model::qwen35::test_support::tiny_zero_model_with_context; + use lattice_inference::{GenerateConfig, InferenceError}; + + fn count_only() -> GenerateConfig { + let mut cfg = GenerateConfig::default(); + cfg.max_new_tokens = 0; + cfg + } + + #[test] + fn repl_uses_checked_generation_and_cpu_initialization() { + // The stdin-driven entry point must use the same seams as these tests. + let source = include_str!("chat.rs") + .split("#[cfg(all(test,") + .next() + .unwrap(); + assert!(source.contains("Ok(m) => Backend::cpu(m),")); + assert!(source.contains("model.generate_chat_line(trimmed, &gen_cfg)")); + } + + #[test] + fn cpu_chat_generation_keeps_long_prompts() { + let model = tiny_zero_model_with_context(8192); + assert_eq!(model.tokenizer().max_seq_len(), 4096); + let mut backend = Backend::cpu(model); + for n in [4097, 8192] { + let output = backend + .generate_chat_line(&"a".repeat(n), &count_only()) + .unwrap(); + assert_eq!(output.prompt_tokens, n); + assert_eq!(output.generated_tokens, 0); } } + + fn assert_refused(backend: &mut Backend, n: usize, limit: usize) { + let error = backend + .generate_chat_line(&"a".repeat(n), &count_only()) + .unwrap_err(); + assert!( + matches!(error, InferenceError::InvalidInput(ref message) + if message == &format!("prompt ({n} tokens) exceeds model context window ({limit})")), + "unexpected error: {error}" + ); + } + + #[test] + fn cpu_chat_refuses_full_count_and_accepts_next_line() { + for limit in [32, 8192] { + let mut backend = Backend::cpu(tiny_zero_model_with_context(limit)); + for n in [limit + 1, limit + 137] { + assert_refused(&mut backend, n, limit); + } + let output = backend.generate_chat_line("a", &count_only()).unwrap(); + assert_eq!(output.prompt_tokens, 1); + } + } + + #[test] + fn chat_guard_refuses_before_generation_with_metal_tokenizer_cap() { + let model = tiny_zero_model_with_context(8192); + let tokenizer = model.tokenizer(); + let limit = 4096; + assert_eq!(tokenizer.max_seq_len(), limit); + for n in [limit + 1, limit + 137] { + let prompt = "a".repeat(n); + let mut called = false; + let result = generate_checked(tokenizer, limit, &prompt, || { + called = true; + model.generate(&prompt, &count_only()) + }); + assert!(!called, "an overlong prompt must not reach generation"); + assert!(matches!(result, Err(InferenceError::InvalidInput(message)) + if message == format!("prompt ({n} tokens) exceeds model context window ({limit})"))); + } + for n in [1, limit] { + let prompt = "a".repeat(n); + let output = generate_checked(tokenizer, limit, &prompt, || { + model.generate(&prompt, &count_only()) + }) + .unwrap(); + assert_eq!(output.prompt_tokens, n); + } + } + + #[cfg(feature = "metal-gpu")] + #[test] + fn metal_chat_uses_checked_generation() { + // Metal's compatible fixtures are private to its library tests. + let source = include_str!("chat.rs") + .split("enum Backend {") + .next() + .unwrap(); + assert!(source.contains("generate_checked(tokenizer, state.max_context(), prompt, ||")); + assert!(source.contains("state.generate(prompt, tokenizer, gen_cfg)")); + assert_eq!(MetalChatBackend::MAX_CACHE_LEN, 4096); + } } diff --git a/crates/inference/tests/data/pipeline_boundary_baseline.txt b/crates/inference/tests/data/pipeline_boundary_baseline.txt index 1aa806ed57..cf3bac5b33 100644 --- a/crates/inference/tests/data/pipeline_boundary_baseline.txt +++ b/crates/inference/tests/data/pipeline_boundary_baseline.txt @@ -53,6 +53,7 @@ src/bin/lattice_serve.rs crate::model::qwen35::test_support::tiny_zero_model src/bin/lattice/chat.rs crate::forward::metal_qwen35::MetalQwen35State src/bin/lattice/chat.rs crate::model::qwen35_config::Qwen35Config src/bin/lattice/chat.rs crate::model::qwen35::model::Qwen35Model +src/bin/lattice/chat.rs crate::model::qwen35::test_support::tiny_zero_model_with_context src/bin/lattice/doctor.rs crate::model::qwen35_config::Qwen35Config src/bin/lattice/doctor.rs crate::model::qwen35::loading::qwen_required_tensor_names src/bin/lattice/main.rs crate::model::qwen35::model::Qwen35Model