use super::cache::{ForwardScratch, KvCache};
use super::detokenize::{IncrementalDetokenizer, decode_tokens};
use super::model::Qwen35Model;
use super::sampling::sample_token;
use super::stop_strings::{
StopStringMatcher, earliest_stop_match, earliest_stop_match_from, stop_scan_search_start,
};
use crate::attention::gdn::GatedDeltaNetState;
use crate::error::InferenceError;
use crate::grammar::pda::GrammarState;
use crate::model::qwen35_config::{
GenerateConfig, GenerateOutput, Qwen35Config, TokenLogprob, decode_cap, force_close_think,
};
use crate::sampling::compute_step_logprobs;
use crate::stop_reason::StopReason;
use crate::tokenizer::common::Tokenizer;
#[cfg(test)]
pub(crate) static FORCE_SERIAL_PREFILL: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
#[cfg(test)]
pub(crate) static SERIAL_PREFILL_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[cfg(test)]
fn force_serial_prefill() -> bool {
FORCE_SERIAL_PREFILL.load(std::sync::atomic::Ordering::SeqCst)
}
#[cfg(not(test))]
fn force_serial_prefill() -> bool {
false
}
impl Qwen35Model {
pub fn generate(
&self,
prompt: &str,
gen_cfg: &GenerateConfig,
) -> Result<GenerateOutput, InferenceError> {
let cfg = &self.config;
let mut rng_state = initial_rng_state(gen_cfg.seed);
let input = self.tokenizer.tokenize(prompt);
let prompt_ids: Vec<u32> = input.input_ids[..input.real_length].to_vec();
let prompt_len = prompt_ids.len();
if prompt_len == 0 {
return Err(InferenceError::Inference("empty prompt".into()));
}
if gen_cfg.max_new_tokens == 0 {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
stop_reason: Some(StopReason::Length),
token_logprobs: vec![],
});
}
let max_context = self.max_context();
let effective_new = decode_cap(gen_cfg.reasoning_budget, gen_cfg.max_new_tokens);
if prompt_len.saturating_add(effective_new) > max_context {
return Err(InferenceError::Inference(format!(
"prompt ({prompt_len} tokens) plus max_new_tokens ({}) exceeds \
model context window ({max_context})",
gen_cfg.max_new_tokens
)));
}
let num_linear = cfg.num_linear_attention_layers();
let num_full = cfg.num_full_attention_layers();
let mut gdn_states: Vec<GatedDeltaNetState> = (0..num_linear)
.map(|_| GatedDeltaNetState::new(cfg))
.collect();
let mut kv_cache = KvCache::new(num_full);
let mut scratch = ForwardScratch::new();
let mut grammar_state: Option<GrammarState> =
gen_cfg.grammar.as_ref().map(|g| g.initial_state());
let mut generated_ids: Vec<u32> = Vec::with_capacity(effective_new);
let mut all_ids = prompt_ids.clone();
let mut token_logprobs: Vec<TokenLogprob> = Vec::new();
let prefill_logits: Vec<f32> = if force_serial_prefill() {
prefill_tokens(
self,
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len = prompt_len;
scratch.logits[..cfg.vocab_size].to_vec()
} else {
match self.prefill_tokens_batched_for_generate(
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
) {
Ok(logits) => logits,
Err(InferenceError::UnsupportedModel(_)) => {
prefill_tokens(
self,
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len = prompt_len;
scratch.logits[..cfg.vocab_size].to_vec()
}
Err(e) => return Err(e),
}
};
scratch.ensure_capacity(cfg, prompt_len);
scratch.logits[..cfg.vocab_size].copy_from_slice(&prefill_logits);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token at step 0; \
no legal first token exists in the current grammar state"
.into(),
));
}
}
let next_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
&all_ids,
&mut rng_state,
);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state)
&& !engine.advance(gs, next_id)
{
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
stop_reason: Some(StopReason::Grammar),
token_logprobs: vec![],
});
}
if should_stop_token(cfg, gen_cfg, next_id) {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: true,
stop_reason: Some(StopReason::Eos),
token_logprobs: vec![],
});
}
generated_ids.push(next_id);
all_ids.push(next_id);
let think_close_id = if gen_cfg.reasoning_budget.is_some() {
self.tokenizer.special_token_id("</think>")
} else {
None
};
let mut policy = DecodePolicy::init(
gen_cfg,
think_close_id,
&mut token_logprobs,
next_id,
&scratch.logits[..cfg.vocab_size],
gen_cfg.temperature,
generated_ids.len(),
false,
);
if gen_cfg.stop_strings.is_empty() {
let (stopped, loop_stop_reason) = decode_loop(
self,
gen_cfg,
&mut all_ids,
&mut generated_ids,
&mut rng_state,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
&mut grammar_state,
&mut policy,
&mut token_logprobs,
)?;
let text = decode_tokens(&self.tokenizer, &generated_ids);
Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
stop_reason: Some(loop_stop_reason),
token_logprobs,
})
} else {
let mut detok = IncrementalDetokenizer::new();
let first_delta = detok.push(&self.tokenizer, next_id);
let mut full = String::new();
let mut token_logprob_end_offsets: Vec<usize> = Vec::new();
if matches!(
policy.check_initial_stop(
&mut token_logprobs,
&mut full,
&mut token_logprob_end_offsets,
&first_delta,
|_| true,
),
StopCheckOutcome::Stopped
) {
return Ok(GenerateOutput {
text: full,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped: true,
stop_reason: Some(StopReason::Eos),
token_logprobs,
});
}
let (stopped, loop_stop_reason) = decode_loop_with_stops(
self,
gen_cfg,
&mut all_ids,
&mut generated_ids,
&mut rng_state,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
&mut detok,
&mut full,
&mut grammar_state,
&mut policy,
&mut token_logprobs,
&mut token_logprob_end_offsets,
)?;
Ok(GenerateOutput {
text: full,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
stop_reason: Some(loop_stop_reason),
token_logprobs,
})
}
}
pub fn generate_streaming(
&self,
prompt: &str,
gen_cfg: &GenerateConfig,
mut on_token: impl FnMut(&str),
) -> Result<GenerateOutput, InferenceError> {
self.generate_streaming_with_cancel(
prompt,
gen_cfg,
|delta| {
on_token(delta);
true
},
|| false,
)
}
pub fn generate_streaming_with_cancel<F, C>(
&self,
prompt: &str,
gen_cfg: &GenerateConfig,
on_token: F,
should_cancel: C,
) -> Result<GenerateOutput, InferenceError>
where
F: FnMut(&str) -> bool,
C: FnMut() -> bool,
{
self.generate_streaming_with_observer(prompt, gen_cfg, on_token, should_cancel, |_| {})
}
pub fn generate_streaming_with_observer<F, C, O>(
&self,
prompt: &str,
gen_cfg: &GenerateConfig,
mut on_token: F,
mut should_cancel: C,
mut on_raw_event: O,
) -> Result<GenerateOutput, InferenceError>
where
F: FnMut(&str) -> bool,
C: FnMut() -> bool,
O: FnMut(RawGenEvent),
{
let cfg = &self.config;
let mut rng_state = initial_rng_state(gen_cfg.seed);
let input = self.tokenizer.tokenize(prompt);
let prompt_ids: Vec<u32> = input.input_ids[..input.real_length].to_vec();
let prompt_len = prompt_ids.len();
if prompt_len == 0 {
return Err(InferenceError::Inference("empty prompt".into()));
}
if gen_cfg.max_new_tokens == 0 {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
stop_reason: Some(StopReason::Length),
token_logprobs: vec![],
});
}
let max_context = self.max_context();
let effective_new = decode_cap(gen_cfg.reasoning_budget, gen_cfg.max_new_tokens);
if prompt_len.saturating_add(effective_new) > max_context {
return Err(InferenceError::Inference(format!(
"prompt ({prompt_len} tokens) plus max_new_tokens ({}) exceeds \
model context window ({max_context})",
gen_cfg.max_new_tokens
)));
}
let num_linear = cfg.num_linear_attention_layers();
let num_full = cfg.num_full_attention_layers();
let mut gdn_states: Vec<GatedDeltaNetState> = (0..num_linear)
.map(|_| GatedDeltaNetState::new(cfg))
.collect();
let mut kv_cache = KvCache::new(num_full);
let mut scratch = ForwardScratch::new();
let mut grammar_state: Option<GrammarState> =
gen_cfg.grammar.as_ref().map(|g| g.initial_state());
let mut generated_ids: Vec<u32> = Vec::with_capacity(effective_new);
let mut all_ids = prompt_ids.clone();
let mut token_logprobs: Vec<TokenLogprob> = Vec::new();
if should_cancel() {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false, stop_reason: Some(StopReason::Interrupt),
token_logprobs: vec![],
});
}
let prefill_logits: Vec<f32> = if force_serial_prefill() {
prefill_tokens(
self,
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len = prompt_len;
scratch.logits[..cfg.vocab_size].to_vec()
} else {
match self.prefill_tokens_batched_for_generate(
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
) {
Ok(logits) => logits,
Err(InferenceError::UnsupportedModel(_)) => {
prefill_tokens(
self,
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len = prompt_len;
scratch.logits[..cfg.vocab_size].to_vec()
}
Err(e) => return Err(e),
}
};
scratch.ensure_capacity(cfg, prompt_len);
scratch.logits[..cfg.vocab_size].copy_from_slice(&prefill_logits);
if should_cancel() {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false, stop_reason: Some(StopReason::Interrupt),
token_logprobs: vec![],
});
}
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token at step 0; \
no legal first token exists in the current grammar state"
.into(),
));
}
}
on_raw_event(RawGenEvent::PrefillEnd);
#[cfg(test)]
test_record_first_sample_entry();
let next_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
&all_ids,
&mut rng_state,
);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state)
&& !engine.advance(gs, next_id)
{
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
stop_reason: Some(StopReason::Grammar),
token_logprobs: vec![],
});
}
if should_stop_token(cfg, gen_cfg, next_id) {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: true,
stop_reason: Some(StopReason::Eos),
token_logprobs: vec![],
});
}
generated_ids.push(next_id);
all_ids.push(next_id);
on_raw_event(RawGenEvent::RawToken {
index: generated_ids.len(),
});
let think_close_id = if gen_cfg.reasoning_budget.is_some() {
self.tokenizer.special_token_id("</think>")
} else {
None
};
let mut policy = DecodePolicy::init(
gen_cfg,
think_close_id,
&mut token_logprobs,
next_id,
&scratch.logits[..cfg.vocab_size],
gen_cfg.temperature,
generated_ids.len(),
true,
);
let mut detok = IncrementalDetokenizer::new();
if gen_cfg.stop_strings.is_empty() {
let mut text = String::new();
let mut throwaway_offsets: Vec<usize> = Vec::new();
let delta = detok.push(&self.tokenizer, next_id);
if matches!(
policy.check_initial_stop(
&mut token_logprobs,
&mut text,
&mut throwaway_offsets,
&delta,
|s| on_token(s),
),
StopCheckOutcome::Interrupted
) {
return Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped: false, stop_reason: Some(StopReason::Interrupt),
token_logprobs,
});
}
let mut stopped = false;
let mut stopped_by_caller = false;
let mut stop_reason = StopReason::Length;
let cap = policy.cap();
for _ in 1..cap {
if should_cancel() {
stopped_by_caller = true;
stop_reason = StopReason::Interrupt;
break;
}
let pos = kv_cache.seq_len;
let Some(&last_token) = all_ids.last() else {
return Err(InferenceError::Inference("empty generation state".into()));
};
self.forward_step(
last_token,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len += 1;
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token; \
no legal continuation exists in the current grammar state"
.into(),
));
}
}
let sampled_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
&all_ids,
&mut rng_state,
);
let generated_len_before = generated_ids.len();
let outcome = policy.transition(
&mut token_logprobs,
sampled_id,
&scratch.logits[..cfg.vocab_size],
gen_cfg.temperature,
generated_len_before,
|next_id| {
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state) {
engine.advance(gs, next_id)
} else {
true
}
},
|next_id| should_stop_token(cfg, gen_cfg, next_id),
|next_id| {
generated_ids.push(next_id);
all_ids.push(next_id);
on_raw_event(RawGenEvent::RawToken {
index: generated_ids.len(),
});
},
|next_id| detok.push(&self.tokenizer, next_id),
&mut text,
&mut throwaway_offsets,
|s, _next_id| on_token(s),
);
let answer_budget_exhausted = match outcome {
StepOutcome::GrammarStop => {
stopped = true;
stop_reason = StopReason::Grammar;
break;
}
StepOutcome::Eos => {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
StepOutcome::Interrupted => {
stopped_by_caller = true;
stop_reason = StopReason::Interrupt;
break;
}
StepOutcome::Stopped => {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
StepOutcome::Emitted {
answer_budget_exhausted,
..
} => answer_budget_exhausted,
};
if answer_budget_exhausted {
break;
}
}
if !stopped_by_caller {
let tail = detok.finish();
if !tail.is_empty() {
text.push_str(&tail);
on_token(&tail);
}
}
Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
stop_reason: Some(stop_reason),
token_logprobs,
})
} else {
let mut text = String::new();
let mut throwaway_offsets: Vec<usize> = Vec::new();
let first_delta = detok.push(&self.tokenizer, next_id);
let initial_outcome = policy.check_initial_stop(
&mut token_logprobs,
&mut text,
&mut throwaway_offsets,
&first_delta,
|s| on_token(s),
);
if matches!(initial_outcome, StopCheckOutcome::Interrupted) {
return Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped: false, stop_reason: Some(StopReason::Interrupt),
token_logprobs,
});
}
if matches!(initial_outcome, StopCheckOutcome::Stopped) {
return Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped: true,
stop_reason: Some(StopReason::Eos),
token_logprobs,
});
}
let mut stopped = false;
let mut stopped_by_caller = false;
let mut stop_reason = StopReason::Length;
let cap = policy.cap();
for _ in 1..cap {
if should_cancel() {
stopped_by_caller = true;
stop_reason = StopReason::Interrupt;
break;
}
let pos = kv_cache.seq_len;
let Some(&last_token) = all_ids.last() else {
return Err(InferenceError::Inference("empty generation state".into()));
};
self.forward_step(
last_token,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len += 1;
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token; \
no legal continuation exists in the current grammar state"
.into(),
));
}
}
let sampled_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
&all_ids,
&mut rng_state,
);
let generated_len_before = generated_ids.len();
let outcome = policy.transition(
&mut token_logprobs,
sampled_id,
&scratch.logits[..cfg.vocab_size],
gen_cfg.temperature,
generated_len_before,
|next_id| {
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state) {
engine.advance(gs, next_id)
} else {
true
}
},
|next_id| should_stop_token(cfg, gen_cfg, next_id),
|next_id| {
generated_ids.push(next_id);
all_ids.push(next_id);
on_raw_event(RawGenEvent::RawToken {
index: generated_ids.len(),
});
},
|next_id| detok.push(&self.tokenizer, next_id),
&mut text,
&mut throwaway_offsets,
|s, _next_id| on_token(s),
);
let answer_budget_exhausted = match outcome {
StepOutcome::GrammarStop => {
stopped = true;
stop_reason = StopReason::Grammar;
break;
}
StepOutcome::Eos => {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
StepOutcome::Interrupted => {
stopped_by_caller = true;
stop_reason = StopReason::Interrupt;
break;
}
StepOutcome::Stopped => {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
StepOutcome::Emitted {
answer_budget_exhausted,
..
} => answer_budget_exhausted,
};
if answer_budget_exhausted {
break;
}
}
if !stopped_by_caller {
let tail_stopped = policy.finish_stop(&mut text, &detok.finish(), |s| on_token(s));
if tail_stopped && !stopped {
stopped = true;
stop_reason = StopReason::Eos;
}
}
Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
stop_reason: Some(stop_reason),
token_logprobs,
})
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RawGenEvent {
PrefillEnd,
RawToken { index: usize },
}
#[cfg(test)]
thread_local! {
static TEST_PREFILL_END_SEEN: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
static TEST_FIRST_SAMPLE_SAW_PREFILL_END: std::cell::Cell<Option<bool>> =
const { std::cell::Cell::new(None) };
}
#[cfg(test)]
fn test_reset_sample_seam() {
TEST_PREFILL_END_SEEN.with(|c| c.set(false));
TEST_FIRST_SAMPLE_SAW_PREFILL_END.with(|c| c.set(None));
}
#[cfg(test)]
fn test_mark_prefill_end_seen() {
TEST_PREFILL_END_SEEN.with(|c| c.set(true));
}
#[cfg(test)]
fn test_record_first_sample_entry() {
TEST_FIRST_SAMPLE_SAW_PREFILL_END.with(|seen_at_sample| {
if seen_at_sample.get().is_none() {
let seen = TEST_PREFILL_END_SEEN.with(std::cell::Cell::get);
seen_at_sample.set(Some(seen));
}
});
}
#[cfg(test)]
fn test_take_first_sample_saw_prefill_end() -> Option<bool> {
TEST_FIRST_SAMPLE_SAW_PREFILL_END.with(std::cell::Cell::get)
}
pub(crate) enum StepOutcome {
GrammarStop,
Eos,
Stopped,
Interrupted,
Emitted {
token_id: u32,
answer_budget_exhausted: bool,
},
}
pub(crate) enum StopCheckOutcome {
Continue,
Stopped,
Interrupted,
}
pub(crate) struct DecodePolicy {
reasoning_budget: Option<usize>,
enable_thinking: bool,
max_new_tokens: usize,
logprobs: Option<usize>,
think_close_id: Option<u32>,
thinking_closed: bool,
reasoning_end_len: Option<usize>,
stop_mode: StopMode,
}
enum StopMode {
Disabled,
Streaming(StopStringMatcher),
FullScan {
stop_strings: Vec<String>,
max_stop: usize,
},
}
impl StopMode {
fn for_config(stop_strings: &[String], streaming: bool) -> Self {
if stop_strings.is_empty() {
StopMode::Disabled
} else if streaming {
StopMode::Streaming(StopStringMatcher::new(stop_strings))
} else {
let max_stop = stop_strings.iter().map(String::len).max().unwrap_or(1);
StopMode::FullScan {
stop_strings: stop_strings.to_vec(),
max_stop,
}
}
}
}
impl DecodePolicy {
#[allow(clippy::too_many_arguments)]
pub(crate) fn init(
gen_cfg: &GenerateConfig,
think_close_id: Option<u32>,
token_logprobs: &mut Vec<TokenLogprob>,
first_emitted_id: u32,
first_logits: &[f32],
temperature: f32,
first_generated_len: usize,
streaming: bool,
) -> Self {
let thinking_closed = Some(first_emitted_id) == think_close_id;
let reasoning_end_len = if thinking_closed {
Some(first_generated_len)
} else {
None
};
let policy = Self {
reasoning_budget: gen_cfg.reasoning_budget,
enable_thinking: gen_cfg.enable_thinking,
max_new_tokens: gen_cfg.max_new_tokens,
logprobs: gen_cfg.logprobs,
think_close_id,
thinking_closed,
reasoning_end_len,
stop_mode: StopMode::for_config(&gen_cfg.stop_strings, streaming),
};
policy.record_logprob(token_logprobs, first_logits, first_emitted_id, temperature);
policy
}
pub(crate) fn cap(&self) -> usize {
decode_cap(self.reasoning_budget, self.max_new_tokens)
}
fn apply_override(&self, generated_len: usize, sampled_id: u32) -> u32 {
force_close_think(
self.reasoning_budget,
self.enable_thinking,
self.thinking_closed,
generated_len,
self.think_close_id,
)
.unwrap_or(sampled_id)
}
fn note_emitted(&mut self, next_id: u32) {
if Some(next_id) == self.think_close_id {
self.thinking_closed = true;
}
}
fn capture_reasoning_end(&mut self, generated_len_after_push: usize) {
if self.thinking_closed && self.reasoning_end_len.is_none() {
self.reasoning_end_len = Some(generated_len_after_push);
}
}
fn answer_budget_exhausted(&self, generated_len: usize) -> bool {
self.reasoning_end_len
.is_some_and(|end| generated_len.saturating_sub(end) >= self.max_new_tokens)
}
fn record_logprob(
&self,
token_logprobs: &mut Vec<TokenLogprob>,
logits: &[f32],
token_id: u32,
temperature: f32,
) {
let Some(top_n) = self.logprobs else {
return;
};
let (logprob, top) = compute_step_logprobs(logits, token_id, temperature, top_n);
token_logprobs.push(TokenLogprob {
token_id,
logprob,
top,
});
}
fn stop_check(
&mut self,
token_logprobs: &mut Vec<TokenLogprob>,
text: &mut String,
token_logprob_end_offsets: &mut Vec<usize>,
delta: &str,
mut emit_confirmed: impl FnMut(&str) -> bool,
) -> StopCheckOutcome {
match &mut self.stop_mode {
StopMode::Disabled => {
if delta.is_empty() {
return StopCheckOutcome::Continue;
}
text.push_str(delta);
if emit_confirmed(delta) {
StopCheckOutcome::Continue
} else {
StopCheckOutcome::Interrupted
}
}
StopMode::Streaming(matcher) => {
let mut interrupted = false;
let stop_matched = matcher.push(delta, &mut |s| {
if !s.is_empty() {
text.push_str(s);
if !interrupted && !emit_confirmed(s) {
interrupted = true;
}
}
});
if interrupted {
StopCheckOutcome::Interrupted
} else if stop_matched {
StopCheckOutcome::Stopped
} else {
StopCheckOutcome::Continue
}
}
StopMode::FullScan {
stop_strings,
max_stop,
} => {
let prev_len = text.len();
if !delta.is_empty() {
text.push_str(delta);
}
if token_logprobs.len() > token_logprob_end_offsets.len() {
token_logprob_end_offsets.push(text.len());
}
let search_start = stop_scan_search_start(text, prev_len, *max_stop);
if let Some(hit) = earliest_stop_match_from(text, stop_strings, search_start) {
text.truncate(hit);
truncate_token_logprobs_to_retained_text(
token_logprobs,
token_logprob_end_offsets,
hit,
);
StopCheckOutcome::Stopped
} else {
StopCheckOutcome::Continue
}
}
}
}
pub(crate) fn check_initial_stop(
&mut self,
token_logprobs: &mut Vec<TokenLogprob>,
text: &mut String,
token_logprob_end_offsets: &mut Vec<usize>,
delta: &str,
emit_confirmed: impl FnMut(&str) -> bool,
) -> StopCheckOutcome {
self.stop_check(
token_logprobs,
text,
token_logprob_end_offsets,
delta,
emit_confirmed,
)
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn transition(
&mut self,
token_logprobs: &mut Vec<TokenLogprob>,
sampled_id: u32,
logits: &[f32],
temperature: f32,
generated_len_before: usize,
mut grammar_advance: impl FnMut(u32) -> bool,
mut is_eos: impl FnMut(u32) -> bool,
mut push: impl FnMut(u32),
mut decode_delta: impl FnMut(u32) -> String,
text: &mut String,
token_logprob_end_offsets: &mut Vec<usize>,
mut emit_confirmed: impl FnMut(&str, u32) -> bool,
) -> StepOutcome {
let next_id = self.apply_override(generated_len_before, sampled_id);
if !grammar_advance(next_id) {
return StepOutcome::GrammarStop;
}
self.note_emitted(next_id);
if is_eos(next_id) {
return StepOutcome::Eos;
}
push(next_id);
let generated_len_after = generated_len_before + 1;
self.record_logprob(token_logprobs, logits, next_id, temperature);
self.capture_reasoning_end(generated_len_after);
let delta = decode_delta(next_id);
let stop_outcome = self.stop_check(
token_logprobs,
text,
token_logprob_end_offsets,
&delta,
|s| emit_confirmed(s, next_id),
);
match stop_outcome {
StopCheckOutcome::Stopped => return StepOutcome::Stopped,
StopCheckOutcome::Interrupted => return StepOutcome::Interrupted,
StopCheckOutcome::Continue => {}
}
StepOutcome::Emitted {
token_id: next_id,
answer_budget_exhausted: self.answer_budget_exhausted(generated_len_after),
}
}
pub(crate) fn finish_stop(
&mut self,
text: &mut String,
tail: &str,
mut emit_confirmed: impl FnMut(&str) -> bool,
) -> bool {
match &mut self.stop_mode {
StopMode::Disabled => {
if !tail.is_empty() {
text.push_str(tail);
emit_confirmed(tail);
}
false
}
StopMode::Streaming(matcher) => {
matcher.finish(tail, &mut |s| {
if !s.is_empty() {
text.push_str(s);
emit_confirmed(s);
}
});
matcher.stopped()
}
StopMode::FullScan { .. } => false,
}
}
}
fn has_finite_logit(logits: &[f32]) -> bool {
logits.iter().any(|&l| l > f32::NEG_INFINITY)
}
fn initial_rng_state(seed: Option<u64>) -> u64 {
match seed {
Some(s) => {
if s == 0 {
1
} else {
s
}
}
None => {
use std::time::SystemTime;
let t = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0x12345678_9abcdef0);
if t == 0 { 1 } else { t }
}
}
}
fn prefill_tokens(
model: &Qwen35Model,
prompt_ids: &[u32],
gdn_states: &mut [GatedDeltaNetState],
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
) {
let prompt_len = prompt_ids.len();
for (pos, &token_id) in prompt_ids.iter().enumerate() {
model.forward_step(token_id, pos, gdn_states, kv_cache, scratch);
if pos < prompt_len - 1 {
kv_cache.seq_len += 1;
}
}
}
#[allow(clippy::too_many_arguments)]
fn decode_loop(
model: &Qwen35Model,
gen_cfg: &GenerateConfig,
all_ids: &mut Vec<u32>,
generated_ids: &mut Vec<u32>,
rng_state: &mut u64,
gdn_states: &mut [GatedDeltaNetState],
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
grammar_state: &mut Option<GrammarState>,
policy: &mut DecodePolicy,
token_logprobs: &mut Vec<TokenLogprob>,
) -> Result<(bool, StopReason), InferenceError> {
let cfg = &model.config;
let cap = policy.cap();
for _ in 1..cap {
let pos = kv_cache.seq_len;
let Some(&last_token) = all_ids.last() else {
return Err(InferenceError::Inference("empty generation state".into()));
};
model.forward_step(last_token, pos, gdn_states, kv_cache, scratch);
kv_cache.seq_len += 1;
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut *grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token; \
no legal continuation exists in the current grammar state"
.into(),
));
}
}
let sampled_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
all_ids,
rng_state,
);
let generated_len_before = generated_ids.len();
let mut throwaway_text = String::new();
let mut throwaway_offsets: Vec<usize> = Vec::new();
let outcome = policy.transition(
token_logprobs,
sampled_id,
&scratch.logits[..cfg.vocab_size],
gen_cfg.temperature,
generated_len_before,
|next_id| {
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut *grammar_state) {
engine.advance(gs, next_id)
} else {
true
}
},
|next_id| should_stop_token(cfg, gen_cfg, next_id),
|next_id| {
generated_ids.push(next_id);
all_ids.push(next_id);
},
|_next_id| String::new(),
&mut throwaway_text,
&mut throwaway_offsets,
|_delta, _next_id| true,
);
match outcome {
StepOutcome::GrammarStop => return Ok((true, StopReason::Grammar)),
StepOutcome::Eos => return Ok((true, StopReason::Eos)),
StepOutcome::Stopped => return Ok((true, StopReason::Eos)),
StepOutcome::Interrupted => return Ok((false, StopReason::Interrupt)),
StepOutcome::Emitted {
answer_budget_exhausted,
..
} => {
if answer_budget_exhausted {
break;
}
}
}
}
Ok((false, StopReason::Length))
}
#[allow(clippy::too_many_arguments)]
fn decode_loop_with_stops(
model: &Qwen35Model,
gen_cfg: &GenerateConfig,
all_ids: &mut Vec<u32>,
generated_ids: &mut Vec<u32>,
rng_state: &mut u64,
gdn_states: &mut [GatedDeltaNetState],
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
detok: &mut IncrementalDetokenizer,
full: &mut String,
grammar_state: &mut Option<GrammarState>,
policy: &mut DecodePolicy,
token_logprobs: &mut Vec<TokenLogprob>,
token_logprob_end_offsets: &mut Vec<usize>,
) -> Result<(bool, StopReason), InferenceError> {
let cfg = &model.config;
let mut stopped = false;
let mut stop_reason = StopReason::Length;
let cap = policy.cap();
for _ in 1..cap {
let pos = kv_cache.seq_len;
let Some(&last_token) = all_ids.last() else {
return Err(InferenceError::Inference("empty generation state".into()));
};
model.forward_step(last_token, pos, gdn_states, kv_cache, scratch);
kv_cache.seq_len += 1;
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut *grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token; \
no legal continuation exists in the current grammar state"
.into(),
));
}
}
let sampled_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
all_ids,
rng_state,
);
let generated_len_before = generated_ids.len();
let outcome = policy.transition(
token_logprobs,
sampled_id,
&scratch.logits[..cfg.vocab_size],
gen_cfg.temperature,
generated_len_before,
|next_id| {
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut *grammar_state) {
engine.advance(gs, next_id)
} else {
true
}
},
|next_id| should_stop_token(cfg, gen_cfg, next_id),
|next_id| {
generated_ids.push(next_id);
all_ids.push(next_id);
},
|next_id| detok.push(&model.tokenizer, next_id),
full,
token_logprob_end_offsets,
|_delta, _next_id| true,
);
let (_next_id, answer_budget_exhausted) = match outcome {
StepOutcome::GrammarStop => {
stopped = true;
stop_reason = StopReason::Grammar;
break;
}
StepOutcome::Eos => {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
StepOutcome::Stopped => {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
StepOutcome::Interrupted => {
stop_reason = StopReason::Interrupt;
break;
}
StepOutcome::Emitted {
token_id,
answer_budget_exhausted,
} => (token_id, answer_budget_exhausted),
};
if answer_budget_exhausted {
break;
}
}
if !stopped {
let tail = detok.finish();
if !tail.is_empty() {
full.push_str(&tail);
if let Some(hit) = earliest_stop_match(full, &gen_cfg.stop_strings) {
full.truncate(hit);
truncate_token_logprobs_to_retained_text(
token_logprobs,
token_logprob_end_offsets,
hit,
);
return Ok((true, StopReason::Eos));
}
}
return Ok((false, StopReason::Length));
}
Ok((stopped, stop_reason))
}
fn truncate_token_logprobs_to_retained_text(
token_logprobs: &mut Vec<TokenLogprob>,
token_logprob_end_offsets: &[usize],
retained_len: usize,
) {
debug_assert_eq!(token_logprobs.len(), token_logprob_end_offsets.len());
let keep = token_logprob_end_offsets.partition_point(|&end| end <= retained_len);
token_logprobs.truncate(keep);
}
pub(crate) fn should_stop_token(
cfg: &Qwen35Config,
gen_cfg: &GenerateConfig,
token_id: u32,
) -> bool {
token_id == cfg.eos_token_id || gen_cfg.stop_token_ids.contains(&token_id)
}
pub(crate) fn check_grammar_not_set(gen_cfg: &GenerateConfig) -> Result<(), InferenceError> {
if gen_cfg.grammar.is_some() {
return Err(InferenceError::InvalidInput(
"grammar-constrained decoding is not yet supported on this path; \
use the Qwen3.5 CPU generate() / generate_streaming() or the generic \
generate() in src/generate.rs, which implement grammar masking"
.into(),
));
}
Ok(())
}
pub(crate) fn check_logprobs_not_set(gen_cfg: &GenerateConfig) -> Result<(), InferenceError> {
if gen_cfg.logprobs.is_some() {
return Err(InferenceError::InvalidInput(
"per-token logprobs are not yet supported on this generation path; \
use the Qwen3.5 CPU generate() / generate_streaming() or the Metal \
generate_streaming(), which implement logprobs capture"
.into(),
));
}
Ok(())
}
pub(crate) fn check_stop_strings_not_set(gen_cfg: &GenerateConfig) -> Result<(), InferenceError> {
if !gen_cfg.stop_strings.is_empty() {
return Err(InferenceError::InvalidInput(
"stop_strings is not yet supported on this generation path; \
use the Qwen3.5 CPU generate() / generate_streaming() or the Metal \
generate() / generate_streaming(), which implement stop-string matching"
.into(),
));
}
Ok(())
}
pub(crate) fn check_reasoning_budget_not_set(
gen_cfg: &GenerateConfig,
) -> Result<(), InferenceError> {
if gen_cfg.reasoning_budget.is_some() {
return Err(InferenceError::InvalidInput(
"reasoning_budget is not yet supported on this generation path; \
use the Qwen3.5 CPU generate() / generate_streaming() or the Metal \
generate_streaming(), which implement reasoning-budget forcing"
.into(),
));
}
Ok(())
}
#[cfg(all(target_os = "macos", feature = "metal-gpu"))]
pub(crate) fn check_mtp_not_requested(gen_cfg: &GenerateConfig) -> Result<(), InferenceError> {
let mtp_enabled = gen_cfg
.enable_mtp
.unwrap_or_else(|| std::env::var("LATTICE_MTP").is_ok());
if mtp_enabled {
return Err(InferenceError::InvalidInput(
"enable_mtp (or LATTICE_MTP) is not supported on the cross-turn \
prefix-cache generation path, which has no MTP draft/verify \
wiring; use the Metal generate() / generate_streaming() paths, \
which implement MTP"
.into(),
));
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn grammar_masking_blocks_argmax_token() {
use crate::grammar::{GrammarEngine, GrammarSpec};
use std::sync::Arc;
let spec = GrammarSpec::Gbnf("root ::= \"t\" | \"f\"\n".to_string());
let vocab = vec![b"t".to_vec(), b"f".to_vec(), b"x".to_vec()];
let engine =
Arc::new(GrammarEngine::new(&spec, vocab).expect("trivial grammar must compile"));
let mut state = engine.initial_state();
let mut logits = vec![1.0_f32, 2.0_f32, 1000.0_f32];
engine.mask_logits(&mut state, &mut logits);
assert_eq!(
logits[2],
f32::NEG_INFINITY,
"mask_logits must set the forbidden token to NEG_INFINITY"
);
assert!(
has_finite_logit(&logits),
"at least one allowed logit must survive the mask"
);
let gen_cfg = GenerateConfig {
temperature: 0.0, ..Default::default()
};
let mut rng = 1u64;
let sampled = sample_token(&logits, &gen_cfg, &[], &mut rng);
assert_ne!(
sampled, 2,
"blocked token must not be selected by the sampler"
);
assert_eq!(
sampled, 1,
"greedy must select the highest remaining allowed logit (index 1)"
);
}
#[test]
fn grammar_all_blocked_mask_detected_by_has_finite_logit() {
let all_blocked = vec![f32::NEG_INFINITY; 8];
assert!(
!has_finite_logit(&all_blocked),
"all-NEG_INFINITY logit buffer must NOT pass the has_finite_logit guard; \
the caller must return a typed error, not silently emit token 0"
);
let mut one_allowed = vec![f32::NEG_INFINITY; 8];
one_allowed[3] = 1.0_f32;
assert!(
has_finite_logit(&one_allowed),
"a single finite logit must pass the guard (grammar still has valid tokens)"
);
}
#[test]
fn grammar_wiring_mask_logits_called_in_generate() {
use crate::attention::gdn::GatedDeltaNetWeights;
use crate::grammar::{GrammarEngine, GrammarSpec};
use crate::lora_hook::NoopLoraHook;
use crate::model::qwen35::{
AttentionWeights, CommonLayerWeights, DenseFfnWeights, FeedForwardWeights,
FullAttentionLayerWeights, ModelWeights,
};
use crate::model::qwen35_config::{LayerType, compute_layer_types};
use crate::rope::RopeTable;
use crate::tokenizer::bpe::BpeTokenizer;
use std::sync::Arc;
const H: usize = 64;
const VOCAB: usize = 97;
const I: usize = 128;
const NUM_LAYERS: usize = 4;
const FULL_INTERVAL: usize = 4;
const HEAD_DIM: usize = 16;
const LINEAR_KH: usize = 4;
const KERNEL: usize = 4;
let cfg = Qwen35Config {
hidden_size: H,
num_hidden_layers: NUM_LAYERS,
vocab_size: VOCAB,
intermediate_size: I,
rms_norm_eps: 1e-6,
num_attention_heads: 4,
num_key_value_heads: 2,
head_dim: HEAD_DIM,
rope_theta: 10_000_000.0,
partial_rotary_factor: 0.25,
rope_parameters: None,
linear_num_key_heads: LINEAR_KH,
linear_num_value_heads: Some(LINEAR_KH),
linear_key_head_dim: HEAD_DIM,
linear_value_head_dim: HEAD_DIM,
linear_conv_kernel_dim: KERNEL,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval: FULL_INTERVAL,
layer_types: compute_layer_types(NUM_LAYERS, FULL_INTERVAL),
layer_mask: vec![true; NUM_LAYERS],
eos_token_id: (VOCAB - 1) as u32,
max_position_embeddings: 1024,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
};
fn rand_vec(rng: &mut u64, len: usize) -> Vec<f32> {
(0..len)
.map(|_| {
*rng ^= *rng << 13;
*rng ^= *rng >> 7;
*rng ^= *rng << 17;
((*rng >> 32) as u32 as f32 / u32::MAX as f32 * 2.0 - 1.0) * 0.02
})
.collect()
}
let mut rng = 0xA55E_u64 | 1;
let qkv_dim = cfg.linear_qkv_dim();
let out_dim = cfg.linear_output_dim();
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let mut layers = Vec::with_capacity(NUM_LAYERS);
for lt in &cfg.layer_types {
let common = CommonLayerWeights {
input_layernorm: rand_vec(&mut rng, H),
post_attention_layernorm: rand_vec(&mut rng, H),
ffn: FeedForwardWeights::Dense(DenseFfnWeights {
gate_proj: rand_vec(&mut rng, I * H),
up_proj: rand_vec(&mut rng, I * H),
down_proj: rand_vec(&mut rng, H * I),
}),
};
let attn = match lt {
LayerType::LinearAttention => AttentionWeights::Linear(GatedDeltaNetWeights {
in_proj_qkv: rand_vec(&mut rng, qkv_dim * H),
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: H,
in_proj_z: rand_vec(&mut rng, out_dim * H),
in_proj_z_rows: out_dim,
in_proj_z_cols: H,
in_proj_b: rand_vec(&mut rng, LINEAR_KH * H),
in_proj_b_rows: LINEAR_KH,
in_proj_b_cols: H,
in_proj_a: rand_vec(&mut rng, LINEAR_KH * H),
in_proj_a_rows: LINEAR_KH,
in_proj_a_cols: H,
a_log: rand_vec(&mut rng, LINEAR_KH),
dt_bias: rand_vec(&mut rng, LINEAR_KH),
conv1d_weight: rand_vec(&mut rng, qkv_dim * KERNEL),
conv_dim: qkv_dim,
kernel_size: KERNEL,
norm_weight: rand_vec(&mut rng, out_dim),
out_proj: rand_vec(&mut rng, H * out_dim),
out_proj_rows: H,
out_proj_cols: out_dim,
}),
LayerType::FullAttention => AttentionWeights::Full(FullAttentionLayerWeights {
q_proj: rand_vec(&mut rng, 2 * q_dim * H),
k_proj: rand_vec(&mut rng, kv_dim * H),
v_proj: rand_vec(&mut rng, kv_dim * H),
o_proj: rand_vec(&mut rng, H * q_dim),
q_norm: rand_vec(&mut rng, HEAD_DIM),
k_norm: rand_vec(&mut rng, HEAD_DIM),
}),
};
layers.push((attn, common));
}
let tok_json = r#"{
"version":"1.0","truncation":null,"padding":null,"added_tokens":[],
"normalizer":null,
"pre_tokenizer":{"type":"ByteLevel","add_prefix_space":false,"trim_offsets":true,"use_regex":true},
"post_processor":null,
"decoder":{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true},
"model":{"type":"BPE","dropout":null,"unk_token":"<unk>","continuing_subword_prefix":null,
"end_of_word_suffix":null,"fuse_unk":false,"byte_fallback":false,"ignore_merges":false,
"vocab":{"<unk>":0,"a":1,"b":2,"c":3,"d":4,"e":5," ":6},"merges":[]}
}"#;
let tokenizer =
BpeTokenizer::from_tokenizer_json_str(tok_json).expect("test tokenizer parses");
let rope = RopeTable::new(
cfg.rope_dim(),
cfg.max_position_embeddings.min(8192),
cfg.rope_theta,
);
let model = Qwen35Model {
config: cfg.clone(),
weights: ModelWeights {
embed_tokens: rand_vec(&mut rng, VOCAB * H),
lm_head: None,
final_norm: rand_vec(&mut rng, H),
layers,
},
tokenizer,
rope,
lora: Box::new(NoopLoraHook),
};
let vocab_bytes: Vec<Vec<u8>> = vec![vec![]; VOCAB];
let spec = GrammarSpec::Gbnf("root ::= \"ok\"\n".to_string());
let engine = Arc::new(
GrammarEngine::new(&spec, vocab_bytes).expect("grammar engine builds with empty vocab"),
);
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
temperature: 0.0,
grammar: Some(engine),
..Default::default()
};
let result = model.generate("a", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"grammar blocking every token must return Err(InvalidInput); got {result:?}"
);
}
#[test]
fn generate_config_default_stop_strings_empty() {
let cfg = GenerateConfig::default();
assert!(cfg.stop_strings.is_empty());
}
#[test]
fn generate_config_stop_strings_field_explicit() {
let cfg = GenerateConfig {
stop_strings: vec!["</s>".to_string(), "\nUser:".to_string()],
..Default::default()
};
assert_eq!(cfg.stop_strings.len(), 2);
assert_eq!(cfg.stop_strings[0], "</s>");
}
#[test]
fn check_stop_strings_not_set_rejects_nonempty() {
let cfg = GenerateConfig {
stop_strings: vec!["</s>".to_string()],
..Default::default()
};
let result = check_stop_strings_not_set(&cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"non-empty stop_strings must be rejected with InvalidInput; got {result:?}"
);
}
#[test]
fn check_stop_strings_not_set_allows_empty() {
assert!(
check_stop_strings_not_set(&GenerateConfig::default()).is_ok(),
"empty stop_strings (the default) must be allowed"
);
}
#[test]
fn check_reasoning_budget_not_set_rejects_some() {
let cfg = GenerateConfig {
reasoning_budget: Some(128),
..Default::default()
};
let result = check_reasoning_budget_not_set(&cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"Some(reasoning_budget) must be rejected with InvalidInput; got {result:?}"
);
}
#[test]
fn check_reasoning_budget_not_set_allows_none() {
assert!(
check_reasoning_budget_not_set(&GenerateConfig::default()).is_ok(),
"reasoning_budget: None (the default) must be allowed"
);
}
const DEFAULT_TINY_TOK_JSON: &str = r#"{
"version":"1.0","truncation":null,"padding":null,"added_tokens":[],
"normalizer":null,
"pre_tokenizer":{"type":"ByteLevel","add_prefix_space":false,"trim_offsets":true,"use_regex":true},
"post_processor":null,
"decoder":{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true},
"model":{"type":"BPE","dropout":null,"unk_token":"<unk>","continuing_subword_prefix":null,
"end_of_word_suffix":null,"fuse_unk":false,"byte_fallback":false,"ignore_merges":false,
"vocab":{"<unk>":0,"a":1,"b":2,"c":3,"d":4,"e":5," ":6},"merges":[]}
}"#;
fn build_tiny_zero_model() -> Qwen35Model {
build_tiny_zero_model_tok(DEFAULT_TINY_TOK_JSON)
}
fn build_tiny_zero_model_tok(tok_json: &str) -> Qwen35Model {
use crate::attention::gdn::GatedDeltaNetWeights;
use crate::lora_hook::NoopLoraHook;
use crate::model::qwen35::{
AttentionWeights, CommonLayerWeights, DenseFfnWeights, FeedForwardWeights,
FullAttentionLayerWeights, ModelWeights,
};
use crate::model::qwen35_config::{LayerType, compute_layer_types};
use crate::rope::RopeTable;
use crate::tokenizer::bpe::BpeTokenizer;
const H: usize = 64;
const VOCAB: usize = 97;
const I: usize = 128;
const NUM_LAYERS: usize = 4;
const FULL_INTERVAL: usize = 4;
const HEAD_DIM: usize = 16;
const LINEAR_KH: usize = 4;
const KERNEL: usize = 4;
let cfg = Qwen35Config {
hidden_size: H,
num_hidden_layers: NUM_LAYERS,
vocab_size: VOCAB,
intermediate_size: I,
rms_norm_eps: 1e-6,
num_attention_heads: 4,
num_key_value_heads: 2,
head_dim: HEAD_DIM,
rope_theta: 10_000_000.0,
partial_rotary_factor: 0.25,
rope_parameters: None,
linear_num_key_heads: LINEAR_KH,
linear_num_value_heads: Some(LINEAR_KH),
linear_key_head_dim: HEAD_DIM,
linear_value_head_dim: HEAD_DIM,
linear_conv_kernel_dim: KERNEL,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval: FULL_INTERVAL,
layer_types: compute_layer_types(NUM_LAYERS, FULL_INTERVAL),
layer_mask: vec![true; NUM_LAYERS],
eos_token_id: (VOCAB - 1) as u32,
max_position_embeddings: 1024,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
};
let z = |len: usize| vec![0.0_f32; len];
let qkv_dim = cfg.linear_qkv_dim();
let out_dim = cfg.linear_output_dim();
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let mut layers = Vec::with_capacity(NUM_LAYERS);
for lt in &cfg.layer_types {
let common = CommonLayerWeights {
input_layernorm: z(H),
post_attention_layernorm: z(H),
ffn: FeedForwardWeights::Dense(DenseFfnWeights {
gate_proj: z(I * H),
up_proj: z(I * H),
down_proj: z(H * I),
}),
};
let attn = match lt {
LayerType::LinearAttention => AttentionWeights::Linear(GatedDeltaNetWeights {
in_proj_qkv: z(qkv_dim * H),
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: H,
in_proj_z: z(out_dim * H),
in_proj_z_rows: out_dim,
in_proj_z_cols: H,
in_proj_b: z(LINEAR_KH * H),
in_proj_b_rows: LINEAR_KH,
in_proj_b_cols: H,
in_proj_a: z(LINEAR_KH * H),
in_proj_a_rows: LINEAR_KH,
in_proj_a_cols: H,
a_log: z(LINEAR_KH),
dt_bias: z(LINEAR_KH),
conv1d_weight: z(qkv_dim * KERNEL),
conv_dim: qkv_dim,
kernel_size: KERNEL,
norm_weight: z(out_dim),
out_proj: z(H * out_dim),
out_proj_rows: H,
out_proj_cols: out_dim,
}),
LayerType::FullAttention => AttentionWeights::Full(FullAttentionLayerWeights {
q_proj: z(2 * q_dim * H),
k_proj: z(kv_dim * H),
v_proj: z(kv_dim * H),
o_proj: z(H * q_dim),
q_norm: z(HEAD_DIM),
k_norm: z(HEAD_DIM),
}),
};
layers.push((attn, common));
}
let tokenizer =
BpeTokenizer::from_tokenizer_json_str(tok_json).expect("test tokenizer parses");
let rope = RopeTable::new(
cfg.rope_dim(),
cfg.max_position_embeddings.min(8192),
cfg.rope_theta,
);
Qwen35Model {
config: cfg.clone(),
weights: ModelWeights {
embed_tokens: z(VOCAB * H),
lm_head: None,
final_norm: z(H),
layers,
},
tokenizer,
rope,
lora: Box::new(NoopLoraHook),
}
}
#[test]
fn stop_reason_length_on_zero_max_tokens() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 0,
temperature: 0.0,
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("zero-token generate must succeed");
assert_eq!(
result.stop_reason,
Some(StopReason::Length),
"max_new_tokens == 0 must return StopReason::Length; got {:?}",
result.stop_reason
);
assert_eq!(
result.generated_tokens, 0,
"zero max_new_tokens must produce no tokens"
);
}
#[test]
fn stop_reason_eos_on_first_stop_token() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
temperature: 0.0,
stop_token_ids: vec![0], ..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("eos-on-first-token generate must succeed");
assert_eq!(
result.stop_reason,
Some(StopReason::Eos),
"stop_token_ids match on first token must return StopReason::Eos; got {:?}",
result.stop_reason
);
}
#[test]
fn eos_token_id_override_forces_continuation_past_matching_stop_token() {
let mut model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
temperature: 0.0,
stop_token_ids: vec![], ..Default::default()
};
model.config_mut().eos_token_id = 0;
let baseline = model
.generate("a", &gen_cfg)
.expect("baseline generate must succeed");
assert_eq!(
baseline.stop_reason,
Some(StopReason::Eos),
"sanity: eos_token_id = 0 must stop generation on the first greedy-sampled \
token (this fixture always samples token 0); got {:?} -- if this fails, \
config_mut() is not reaching the real config should_stop_token reads",
baseline.stop_reason
);
assert_eq!(
baseline.generated_tokens, 0,
"sanity: eos_token_id = 0 must stop before any token is emitted into the \
output (the matching token itself is excluded, per should_stop_token's \
stop-before-append semantics -- see stop_reason_eos_on_first_stop_token \
above), got {}",
baseline.generated_tokens
);
model.config_mut().eos_token_id = u32::MAX;
let result = model
.generate("a", &gen_cfg)
.expect("eos-override generate must succeed");
assert_eq!(
result.stop_reason,
Some(StopReason::Length),
"overriding eos_token_id out of range must force the loop to run to \
max_new_tokens (StopReason::Length), not stop early; got {:?}",
result.stop_reason
);
assert_eq!(
result.generated_tokens, 5,
"eos_token_id override must decode exactly max_new_tokens (5), got {}",
result.generated_tokens
);
}
#[test]
fn eos_token_id_override_suppresses_match_in_should_stop_token() {
let mut model = build_tiny_zero_model();
let base_cfg = GenerateConfig {
stop_token_ids: vec![],
..Default::default()
};
let original_eos_token_id = model.config.eos_token_id;
assert!(
should_stop_token(&model.config, &base_cfg, original_eos_token_id),
"sanity: without the override, the model's own eos_token_id must stop generation"
);
model.config_mut().eos_token_id = u32::MAX;
assert!(
!should_stop_token(&model.config, &base_cfg, original_eos_token_id),
"eos_token_id override must suppress a match against the model's original \
(now-superseded) eos_token_id -- the sentinel u32::MAX itself trivially \
matches u32::MAX, so this must check the ORIGINAL id stays unmatched"
);
}
#[test]
fn stop_reason_grammar_on_advance_false() {
use crate::grammar::{GrammarEngine, GrammarSpec};
use std::sync::Arc;
let model = build_tiny_zero_model();
let spec = GrammarSpec::Gbnf("root ::= \"x\"\n".to_string());
let vocab = vec![b"t".to_vec()];
let engine =
Arc::new(GrammarEngine::new(&spec, vocab).expect("single-token grammar compiles"));
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
temperature: 0.0,
grammar: Some(engine),
stop_token_ids: vec![],
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("grammar-advance-false generate must succeed");
assert_eq!(
result.stop_reason,
Some(StopReason::Grammar),
"grammar advance returning false must return StopReason::Grammar; got {:?}",
result.stop_reason
);
}
fn build_tiny_thinking_model() -> Qwen35Model {
build_tiny_zero_model_tok(
r#"{
"version":"1.0","truncation":null,"padding":null,
"added_tokens":[{"id":7,"content":"</think>","single_word":false,"lstrip":false,"rstrip":false,"normalized":false,"special":false}],
"normalizer":null,
"pre_tokenizer":{"type":"ByteLevel","add_prefix_space":false,"trim_offsets":true,"use_regex":true},
"post_processor":null,
"decoder":{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true},
"model":{"type":"BPE","dropout":null,"unk_token":"<unk>","continuing_subword_prefix":null,
"end_of_word_suffix":null,"fuse_unk":false,"byte_fallback":false,"ignore_merges":false,
"vocab":{"<unk>":0,"a":1,"b":2,"c":3,"d":4,"e":5," ":6},"merges":[]}
}"#,
)
}
#[test]
fn grammar_budget_forced_close_fails_closed() {
use crate::grammar::{GrammarEngine, GrammarSpec};
use std::sync::Arc;
let model = build_tiny_thinking_model();
let close_id = model
.tokenizer
.special_token_id("</think>")
.expect("thinking model tokenizer resolves </think>");
assert!(
close_id >= 7,
"test assumes </think> id ({close_id}) is outside the 7-token grammar vocab"
);
let spec = GrammarSpec::Gbnf("root ::= \"aa\"\n".to_string());
let vocab: Vec<Vec<u8>> = vec![
b"<unk>".to_vec(),
b"a".to_vec(),
b"b".to_vec(),
b"c".to_vec(),
b"d".to_vec(),
b"e".to_vec(),
b" ".to_vec(),
];
let engine = Arc::new(GrammarEngine::new(&spec, vocab).expect("aa grammar compiles"));
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
temperature: 0.0,
enable_thinking: true,
reasoning_budget: Some(1),
grammar: Some(engine),
stop_token_ids: vec![],
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("combined grammar+budget generate must succeed");
assert_eq!(
result.stop_reason,
Some(StopReason::Grammar),
"budget-forced </think> forbidden by grammar must stop with Grammar; got {:?}",
result.stop_reason
);
assert_eq!(
result.token_ids,
vec![1],
"fail-closed: the budget-forced </think> must NOT be emitted; only the \
pre-force 'a' (id 1) survives. token_ids [1, 7] means advance ran on the \
sampled token, not the forced token"
);
assert_eq!(result.generated_tokens, 1);
}
#[test]
fn generate_streaming_with_cancel_true_before_prefill_returns_interrupt() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
temperature: 0.0,
..Default::default()
};
let calls = std::cell::Cell::new(0usize);
let result = model
.generate_streaming_with_cancel(
"a",
&gen_cfg,
|_delta| true,
|| {
let n = calls.get() + 1;
calls.set(n);
n == 1
},
)
.expect("cancelled-before-prefill generate must succeed");
assert!(
!result.stopped,
"a caller cancellation is not an OpenAI stop condition"
);
assert_eq!(result.stop_reason, Some(StopReason::Interrupt));
assert_eq!(result.generated_tokens, 0);
assert!(result.text.is_empty());
}
#[test]
fn generate_streaming_with_cancel_true_after_prefill_returns_interrupt() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
temperature: 0.0,
..Default::default()
};
let calls = std::cell::Cell::new(0usize);
let on_token_calls = std::cell::Cell::new(0usize);
let result = model
.generate_streaming_with_cancel(
"a",
&gen_cfg,
|_delta| {
on_token_calls.set(on_token_calls.get() + 1);
true
},
|| {
let n = calls.get() + 1;
calls.set(n);
n == 2
},
)
.expect("cancelled-after-prefill generate must succeed");
assert!(
!result.stopped,
"a caller cancellation is not an OpenAI stop condition"
);
assert_eq!(result.stop_reason, Some(StopReason::Interrupt));
assert_eq!(
result.generated_tokens, 0,
"post-prefill cancellation must stop before any token is sampled"
);
assert!(result.text.is_empty());
assert_eq!(
on_token_calls.get(),
0,
"post-prefill cancellation must stop before on_token is ever called"
);
}
#[test]
fn generate_streaming_with_cancel_mid_decode_stops_early_fast_path() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 10,
temperature: 0.0,
..Default::default()
};
let calls = std::cell::Cell::new(0usize);
let result = model
.generate_streaming_with_cancel(
"a",
&gen_cfg,
|_delta| true,
|| {
let n = calls.get() + 1;
calls.set(n);
n >= 3
},
)
.expect("mid-decode cancel generate must succeed");
assert!(!result.stopped);
assert_eq!(result.stop_reason, Some(StopReason::Interrupt));
assert_eq!(
result.generated_tokens, 1,
"should_cancel flipping true at the first decode-loop checkpoint must stop \
after exactly the one pre-loop token; got {}",
result.generated_tokens
);
}
#[test]
fn generate_streaming_with_cancel_on_token_false_stops_generation_fast_path() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 10,
temperature: 0.0,
..Default::default()
};
let result = model
.generate_streaming_with_cancel("a", &gen_cfg, |_delta| false, || false)
.expect("on_token-false generate must succeed");
assert!(!result.stopped);
assert_eq!(result.stop_reason, Some(StopReason::Interrupt));
assert_eq!(
result.generated_tokens, 1,
"on_token returning false on the very first delta must stop after exactly \
one token; got {}",
result.generated_tokens
);
}
#[test]
fn generate_streaming_with_cancel_on_token_false_stops_generation_stop_string_path() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 10,
temperature: 0.0,
stop_strings: vec!["ZZZZ".to_string()],
..Default::default()
};
let result = model
.generate_streaming_with_cancel("a", &gen_cfg, |_delta| false, || false)
.expect("on_token-false generate (stop-string path) must succeed");
assert!(!result.stopped);
assert_eq!(result.stop_reason, Some(StopReason::Interrupt));
assert_eq!(
result.generated_tokens, 1,
"on_token returning false on the very first delta must stop after exactly \
one token in the stop-string path too; got {}",
result.generated_tokens
);
}
#[test]
fn generate_streaming_with_cancel_mid_decode_stops_early_stop_string_path() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 10,
temperature: 0.0,
stop_strings: vec!["ZZZZ".to_string()],
..Default::default()
};
let calls = std::cell::Cell::new(0usize);
let result = model
.generate_streaming_with_cancel(
"a",
&gen_cfg,
|_delta| true,
|| {
let n = calls.get() + 1;
calls.set(n);
n >= 3
},
)
.expect("mid-decode cancel generate (stop-string path) must succeed");
assert!(!result.stopped);
assert_eq!(result.stop_reason, Some(StopReason::Interrupt));
assert_eq!(
result.generated_tokens, 1,
"should_cancel flipping true at the first decode-loop checkpoint must stop \
after exactly the one pre-loop token (stop-string path); got {}",
result.generated_tokens
);
}
#[test]
fn raw_observer_prefill_end_precedes_first_raw_token_monotonic_index() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 3,
temperature: 0.0,
..Default::default()
};
let events = std::cell::RefCell::new(Vec::<RawGenEvent>::new());
let result = model
.generate_streaming_with_observer(
"a",
&gen_cfg,
|_delta| true,
|| false,
|evt| events.borrow_mut().push(evt),
)
.expect("generation must succeed");
let events = events.into_inner();
assert_eq!(
events.first(),
Some(&RawGenEvent::PrefillEnd),
"the very first raw event must be PrefillEnd -- prefill completed and logits \
are ready before any token is sampled; got {events:?}"
);
let token_indices: Vec<usize> = events
.iter()
.skip(1)
.map(|e| match e {
RawGenEvent::RawToken { index } => *index,
RawGenEvent::PrefillEnd => {
panic!("PrefillEnd must fire exactly once, at the very start: {events:?}")
}
})
.collect();
assert_eq!(result.generated_tokens, 3);
assert_eq!(
token_indices,
vec![1, 2, 3],
"RawToken events must be one per generated token, in generation order, with a \
monotonically increasing 1-based index equal to generated_tokens so far"
);
}
#[test]
fn raw_observer_prefill_end_precedes_first_sample_not_just_first_raw_token() {
test_reset_sample_seam();
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 3,
temperature: 0.0,
..Default::default()
};
let result = model
.generate_streaming_with_observer(
"a",
&gen_cfg,
|_delta| true,
|| false,
|evt| {
if evt == RawGenEvent::PrefillEnd {
test_mark_prefill_end_seen();
}
},
)
.expect("generation must succeed");
assert_eq!(result.generated_tokens, 3);
assert_eq!(
test_take_first_sample_saw_prefill_end(),
Some(true),
"PrefillEnd must have already fired by the moment sample_token is entered for \
the prefill-derived first token -- not merely before the RawToken callback, \
which a PrefillEnd emitted after sampling (but before the RawToken push) would \
also satisfy while silently including sampling time in the reported prefill \
interval (codex round-2 medium, PR #882)"
);
}
#[test]
fn raw_observer_fires_before_on_token_for_buffered_incomplete_utf8_first_token() {
const UTF8_LEAD_BYTE_TOK_JSON: &str = r#"{
"version":"1.0","truncation":null,"padding":null,"added_tokens":[],
"normalizer":null,
"pre_tokenizer":{"type":"ByteLevel","add_prefix_space":false,"trim_offsets":true,"use_regex":true},
"post_processor":null,
"decoder":{"type":"ByteLevel","add_prefix_space":true,"trim_offsets":true,"use_regex":true},
"model":{"type":"BPE","dropout":null,"unk_token":"<unk>","continuing_subword_prefix":null,
"end_of_word_suffix":null,"fuse_unk":false,"byte_fallback":false,"ignore_merges":false,
"vocab":{"Â":0,"<unk>":1,"a":2},"merges":[]}
}"#;
let model = build_tiny_zero_model_tok(UTF8_LEAD_BYTE_TOK_JSON);
let mut probe = IncrementalDetokenizer::new();
assert_eq!(
probe.push(model.tokenizer(), 0),
"",
"token id 0 (raw byte 0xC2) must be an incomplete UTF-8 lead byte on its own -- \
this test's premise depends on it"
);
let gen_cfg = GenerateConfig {
max_new_tokens: 4,
temperature: 0.0,
..Default::default()
};
let raw_indices = std::cell::RefCell::new(Vec::<usize>::new());
let on_token_calls = std::cell::Cell::new(0usize);
let snapshots_at_raw_event = std::cell::RefCell::new(Vec::<usize>::new());
let result = model
.generate_streaming_with_observer(
"a",
&gen_cfg,
|_delta: &str| {
on_token_calls.set(on_token_calls.get() + 1);
true
},
|| false,
|evt| {
if let RawGenEvent::RawToken { index } = evt {
raw_indices.borrow_mut().push(index);
snapshots_at_raw_event
.borrow_mut()
.push(on_token_calls.get());
}
},
)
.expect("generation over an incomplete-lead-byte vocab must still succeed");
assert_eq!(
result.generated_tokens, 4,
"token id 0 (eos_token_id is 96 on this tiny model, never sampled) never \
satisfies should_stop_token, so all 4 requested tokens must be generated"
);
assert_eq!(
raw_indices.into_inner(),
vec![1, 2, 3, 4],
"one monotonically indexed RawToken event per generated token, regardless of \
detokenizer buffering"
);
assert_eq!(
snapshots_at_raw_event.borrow()[0],
0,
"on_token must not have been called yet at the moment the FIRST RawToken event \
fires -- the prefill-derived first token's delta is buffered (empty) by the \
incomplete-UTF-8 lead byte, so a phase-event trace measured off on_token would \
have missed or mis-timed this token entirely; got {} prior on_token calls",
snapshots_at_raw_event.borrow()[0]
);
}
#[test]
fn stop_string_truncation_drops_stale_token_logprobs() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 10,
temperature: 0.0,
logprobs: Some(0),
stop_strings: vec!["k><unk".to_string()],
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("stop-string generate must succeed");
assert_eq!(
result.text, "<un",
"stop string must truncate at the first match; got {:?}",
result.text
);
assert_eq!(
result.generated_tokens, 2,
"both tokens were sampled before the match completed (can't un-generate); \
got {}",
result.generated_tokens
);
assert!(
result.token_logprobs.is_empty(),
"both tokens' text was only partially retained after truncation, so both \
logprobs entries must be dropped; got {} entries",
result.token_logprobs.len()
);
}
#[test]
fn decode_policy_record_logprob_noop_when_not_requested() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 3,
temperature: 0.0,
logprobs: None,
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("plain generate must succeed");
assert!(
result.token_logprobs.is_empty(),
"logprobs: None must record nothing across the whole generation \
(prefill token via init, decode tokens via transition); got {} entries",
result.token_logprobs.len()
);
}
#[test]
fn init_records_the_prefill_tokens_logprob_before_any_transition_call() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
temperature: 0.0,
logprobs: Some(0),
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("single-token generate must succeed");
assert_eq!(
result.generated_tokens, 1,
"max_new_tokens: 1 must generate exactly the prefill-derived \
first token and never enter decode_loop; got {}",
result.generated_tokens
);
assert_eq!(
result.token_logprobs.len(),
1,
"the sole generated token's logprob must be recorded by \
DecodePolicy::init alone (transition is never called when \
max_new_tokens == 1); got {} entries",
result.token_logprobs.len()
);
assert_eq!(
result.token_logprobs[0].token_id, result.token_ids[0],
"the recorded logprob entry must describe the actual first token"
);
}
#[test]
fn transition_records_one_logprob_per_generated_token() {
let model = build_tiny_zero_model();
let gen_cfg = GenerateConfig {
max_new_tokens: 3,
temperature: 0.0,
logprobs: Some(0),
..Default::default()
};
let result = model
.generate("a", &gen_cfg)
.expect("plain generate must succeed");
assert_eq!(
result.token_logprobs.len(),
result.token_ids.len(),
"every generated token must get exactly one TokenLogprob entry \
when logprobs is requested and nothing truncates the output; \
got {} logprobs for {} tokens",
result.token_logprobs.len(),
result.token_ids.len()
);
for (i, (logprob, &token_id)) in result
.token_logprobs
.iter()
.zip(result.token_ids.iter())
.enumerate()
{
assert_eq!(
logprob.token_id, token_id,
"token_logprobs[{i}] must describe the token actually emitted \
at that position"
);
}
}
fn assert_batched_prefill_matches_serial(
model: &Qwen35Model,
prompt: &str,
max_new_tokens: usize,
) {
let _guard = SERIAL_PREFILL_TEST_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let gen_cfg = crate::model::qwen35_config::GenerateConfig {
max_new_tokens,
temperature: 0.0,
repetition_penalty: 1.0,
..Default::default()
};
FORCE_SERIAL_PREFILL.store(true, std::sync::atomic::Ordering::SeqCst);
let serial = model.generate(prompt, &gen_cfg);
FORCE_SERIAL_PREFILL.store(false, std::sync::atomic::Ordering::SeqCst);
let batched = model.generate(prompt, &gen_cfg);
let serial = serial.expect("serial-prefill generate must succeed");
let batched = batched.expect("batched-prefill generate must succeed");
assert_eq!(
serial.token_ids, batched.token_ids,
"batched-prefill delegation changed generated token ids for prompt {prompt:?}: \
serial={:?} batched={:?}",
serial.token_ids, batched.token_ids
);
assert_eq!(
serial.text, batched.text,
"batched-prefill delegation changed decoded text for prompt {prompt:?}"
);
assert_eq!(serial.stop_reason, batched.stop_reason);
assert_eq!(serial.stopped, batched.stopped);
}
#[test]
#[ignore = "requires local Qwen3.5 checkpoint: set LATTICE_INFERENCE_MODEL_DIR"]
fn generate_batched_prefill_matches_serial_for_seeded_dense_prompt() {
let Ok(model_dir) = std::env::var("LATTICE_INFERENCE_MODEL_DIR") else {
return;
};
let model = Qwen35Model::from_safetensors(std::path::Path::new(&model_dir))
.expect("dense Qwen3.5 model should load successfully");
assert_batched_prefill_matches_serial(
&model,
"The quick brown fox jumps over the lazy dog. In a distant future,",
20,
);
}
#[test]
#[ignore = "requires local Qwen3.5 checkpoint: set LATTICE_INFERENCE_MODEL_DIR"]
fn generate_streaming_batched_prefill_matches_nonstreaming_text() {
let Ok(model_dir) = std::env::var("LATTICE_INFERENCE_MODEL_DIR") else {
return;
};
let model = Qwen35Model::from_safetensors(std::path::Path::new(&model_dir))
.expect("dense Qwen3.5 model should load successfully");
let gen_cfg = crate::model::qwen35_config::GenerateConfig {
max_new_tokens: 20,
temperature: 0.0,
repetition_penalty: 1.0,
..Default::default()
};
let prompt = "The quick brown fox jumps over the lazy dog. In a distant future,";
let non_streaming = model
.generate(prompt, &gen_cfg)
.expect("non-streaming generate must succeed");
let mut streamed_text = String::new();
let streaming = model
.generate_streaming(prompt, &gen_cfg, |delta| streamed_text.push_str(delta))
.expect("streaming generate must succeed");
assert_eq!(
non_streaming.token_ids, streaming.token_ids,
"streaming batched-prefill delegation diverged from non-streaming"
);
assert_eq!(non_streaming.text, streaming.text);
assert_eq!(non_streaming.text, streamed_text);
}
#[test]
#[ignore = "requires local Qwen3.5 checkpoint: set LATTICE_INFERENCE_MODEL_DIR"]
fn generate_batched_prefill_matches_serial_across_prompt_lengths() {
let Ok(model_dir) = std::env::var("LATTICE_INFERENCE_MODEL_DIR") else {
return;
};
let model = Qwen35Model::from_safetensors(std::path::Path::new(&model_dir))
.expect("dense Qwen3.5 model should load successfully");
for words in [8usize, 64, 256] {
let prompt = "hello ".repeat(words);
assert_batched_prefill_matches_serial(&model, prompt.trim_end(), 5);
}
}
#[test]
#[ignore = "requires local Qwen3.5 checkpoint: set LATTICE_INFERENCE_MODEL_DIR; run --release"]
fn public_prefill_ttft_ab_sweep() {
let Ok(model_dir) = std::env::var("LATTICE_INFERENCE_MODEL_DIR") else {
return;
};
let model = Qwen35Model::from_safetensors(std::path::Path::new(&model_dir))
.expect("dense Qwen3.5 model should load successfully");
let _guard = SERIAL_PREFILL_TEST_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let gen_cfg = crate::model::qwen35_config::GenerateConfig {
max_new_tokens: 1,
temperature: 0.0,
repetition_penalty: 1.0,
..Default::default()
};
println!("words\tprompt_tokens\tserial_ms\tbatched_ms\tspeedup");
for words in [64usize, 512, 2000] {
let prompt = "hello ".repeat(words);
let prompt = prompt.trim_end();
FORCE_SERIAL_PREFILL.store(false, std::sync::atomic::Ordering::SeqCst);
let _ = model.generate(prompt, &gen_cfg).unwrap();
FORCE_SERIAL_PREFILL.store(true, std::sync::atomic::Ordering::SeqCst);
let t0 = std::time::Instant::now();
let serial = model.generate(prompt, &gen_cfg).expect("serial generate");
let serial_ms = t0.elapsed().as_secs_f64() * 1000.0;
FORCE_SERIAL_PREFILL.store(false, std::sync::atomic::Ordering::SeqCst);
let t0 = std::time::Instant::now();
let batched = model.generate(prompt, &gen_cfg).expect("batched generate");
let batched_ms = t0.elapsed().as_secs_f64() * 1000.0;
assert_eq!(
serial.token_ids, batched.token_ids,
"TTFT sweep: token mismatch at words={words}"
);
println!(
"{words}\t{}\t{serial_ms:.1}\t{batched_ms:.1}\t{:.3}",
serial.prompt_tokens,
serial_ms / batched_ms.max(1e-6),
);
}
}
}