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::record_logprob;
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);
record_logprob(
&mut token_logprobs,
&scratch.logits[..cfg.vocab_size],
next_id,
gen_cfg.temperature,
gen_cfg.logprobs,
);
let think_close_id = if gen_cfg.reasoning_budget.is_some() {
self.tokenizer.special_token_id("</think>")
} else {
None
};
let thinking_closed_seed = Some(next_id) == think_close_id;
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,
think_close_id,
thinking_closed_seed,
&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 = first_delta;
let mut token_logprob_end_offsets: Vec<usize> = if token_logprobs.is_empty() {
Vec::new()
} else {
vec![full.len()]
};
if let Some(hit) = earliest_stop_match(&full, &gen_cfg.stop_strings) {
full.truncate(hit);
truncate_token_logprobs_to_retained_text(
&mut token_logprobs,
&token_logprob_end_offsets,
hit,
);
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,
think_close_id,
thinking_closed_seed,
&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> {
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);
record_logprob(
&mut token_logprobs,
&scratch.logits[..cfg.vocab_size],
next_id,
gen_cfg.temperature,
gen_cfg.logprobs,
);
let think_close_id = if gen_cfg.reasoning_budget.is_some() {
self.tokenizer.special_token_id("</think>")
} else {
None
};
let mut thinking_closed = Some(next_id) == think_close_id;
let mut reasoning_end_len: Option<usize> = None;
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
let mut detok = IncrementalDetokenizer::new();
if gen_cfg.stop_strings.is_empty() {
let mut text = String::new();
let delta = detok.push(&self.tokenizer, next_id);
if !delta.is_empty() {
text.push_str(&delta);
on_token(&delta);
}
let mut stopped = false;
let mut stop_reason = StopReason::Length;
let cap = decode_cap(gen_cfg.reasoning_budget, gen_cfg.max_new_tokens);
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()));
};
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 next_id = force_close_think(
gen_cfg.reasoning_budget,
gen_cfg.enable_thinking,
thinking_closed,
generated_ids.len(),
think_close_id,
)
.unwrap_or(sampled_id);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state)
&& !engine.advance(gs, next_id)
{
stopped = true;
stop_reason = StopReason::Grammar;
break;
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
record_logprob(
&mut token_logprobs,
&scratch.logits[..cfg.vocab_size],
next_id,
gen_cfg.temperature,
gen_cfg.logprobs,
);
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
let delta = detok.push(&self.tokenizer, next_id);
if !delta.is_empty() {
text.push_str(&delta);
on_token(&delta);
}
if let Some(end) = reasoning_end_len
&& generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens
{
break;
}
}
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 streamer = StopStringMatcher::new(&gen_cfg.stop_strings);
let first_delta = detok.push(&self.tokenizer, next_id);
if streamer.push(&first_delta, &mut on_token) {
return Ok(GenerateOutput {
text: streamer.into_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 stop_reason = StopReason::Length;
let cap = decode_cap(gen_cfg.reasoning_budget, gen_cfg.max_new_tokens);
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()));
};
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 next_id = force_close_think(
gen_cfg.reasoning_budget,
gen_cfg.enable_thinking,
thinking_closed,
generated_ids.len(),
think_close_id,
)
.unwrap_or(sampled_id);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut grammar_state)
&& !engine.advance(gs, next_id)
{
stopped = true;
stop_reason = StopReason::Grammar;
break;
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
record_logprob(
&mut token_logprobs,
&scratch.logits[..cfg.vocab_size],
next_id,
gen_cfg.temperature,
gen_cfg.logprobs,
);
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
let delta = detok.push(&self.tokenizer, next_id);
if streamer.push(&delta, &mut on_token) {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
if let Some(end) = reasoning_end_len
&& generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens
{
break;
}
}
streamer.finish(&detok.finish(), &mut on_token);
if streamer.stopped() && !stopped {
stopped = true;
stop_reason = StopReason::Eos;
}
Ok(GenerateOutput {
text: streamer.into_text(),
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
stop_reason: Some(stop_reason),
token_logprobs,
})
}
}
}
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>,
think_close_id: Option<u32>,
thinking_closed_seed: bool,
token_logprobs: &mut Vec<TokenLogprob>,
) -> Result<(bool, StopReason), InferenceError> {
let cfg = &model.config;
let mut thinking_closed = thinking_closed_seed;
let mut reasoning_end_len: Option<usize> = if thinking_closed {
Some(generated_ids.len())
} else {
None
};
let cap = decode_cap(gen_cfg.reasoning_budget, gen_cfg.max_new_tokens);
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 next_id = force_close_think(
gen_cfg.reasoning_budget,
gen_cfg.enable_thinking,
thinking_closed,
generated_ids.len(),
think_close_id,
)
.unwrap_or(sampled_id);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut *grammar_state)
&& !engine.advance(gs, next_id)
{
return Ok((true, StopReason::Grammar));
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
return Ok((true, StopReason::Eos));
}
generated_ids.push(next_id);
all_ids.push(next_id);
record_logprob(
token_logprobs,
&scratch.logits[..cfg.vocab_size],
next_id,
gen_cfg.temperature,
gen_cfg.logprobs,
);
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
if let Some(end) = reasoning_end_len
&& generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens
{
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>,
think_close_id: Option<u32>,
thinking_closed_seed: bool,
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 mut thinking_closed = thinking_closed_seed;
let mut reasoning_end_len: Option<usize> = if thinking_closed {
Some(generated_ids.len())
} else {
None
};
let cap = decode_cap(gen_cfg.reasoning_budget, gen_cfg.max_new_tokens);
let max_stop = gen_cfg
.stop_strings
.iter()
.map(String::len)
.max()
.unwrap_or(1);
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 next_id = force_close_think(
gen_cfg.reasoning_budget,
gen_cfg.enable_thinking,
thinking_closed,
generated_ids.len(),
think_close_id,
)
.unwrap_or(sampled_id);
if let (Some(engine), Some(gs)) = (&gen_cfg.grammar, &mut *grammar_state)
&& !engine.advance(gs, next_id)
{
stopped = true;
stop_reason = StopReason::Grammar;
break;
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
record_logprob(
token_logprobs,
&scratch.logits[..cfg.vocab_size],
next_id,
gen_cfg.temperature,
gen_cfg.logprobs,
);
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
let prev_len = full.len();
let delta = detok.push(&model.tokenizer, next_id);
if !delta.is_empty() {
full.push_str(&delta);
}
if token_logprobs.len() > token_logprob_end_offsets.len() {
token_logprob_end_offsets.push(full.len());
}
let search_start = stop_scan_search_start(full, prev_len, max_stop);
if let Some(hit) = earliest_stop_match_from(full, &gen_cfg.stop_strings, search_start) {
full.truncate(hit);
truncate_token_logprobs_to_retained_text(
token_logprobs,
token_logprob_end_offsets,
hit,
);
stopped = true;
stop_reason = StopReason::Eos;
break;
}
if let Some(end) = reasoning_end_len
&& generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens
{
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(())
}
#[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>");
}
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 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 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()
);
}
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),
);
}
}
}