use super::cache::{ForwardScratch, KvCache};
use super::detokenize::{IncrementalDetokenizer, decode_tokens};
use super::model::Qwen35Model;
use super::sampling::sample_token;
use crate::attention::gdn::GatedDeltaNetState;
use crate::error::InferenceError;
use crate::grammar::pda::GrammarState;
use crate::model::qwen35_config::{
GenerateConfig, GenerateOutput, Qwen35Config, decode_cap, force_close_think,
};
use crate::tokenizer::common::Tokenizer;
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,
});
}
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();
prefill_tokens(
self,
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len = prompt_len;
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) {
if !engine.advance(gs, next_id) {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
});
}
}
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,
});
}
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 thinking_closed_seed = Some(next_id) == think_close_id;
if gen_cfg.stop_strings.is_empty() {
let stopped = 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,
)?;
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,
})
} else {
let mut detok = IncrementalDetokenizer::new();
let first_delta = detok.push(&self.tokenizer, next_id);
let mut full = first_delta;
if let Some(hit) = earliest_stop_match(&full, &gen_cfg.stop_strings) {
full.truncate(hit);
return Ok(GenerateOutput {
text: full,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped: true,
});
}
let stopped = 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,
)?;
Ok(GenerateOutput {
text: full,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
})
}
}
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,
});
}
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();
prefill_tokens(
self,
&prompt_ids,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len = prompt_len;
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) {
if !engine.advance(gs, next_id) {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
});
}
}
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,
});
}
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 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 delta = detok.push(&self.tokenizer, next_id);
if !delta.is_empty() {
on_token(&delta);
}
let mut stopped = false;
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) {
if !engine.advance(gs, next_id) {
stopped = true;
break;
}
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
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() {
on_token(&delta);
}
if let Some(end) = reasoning_end_len {
if generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens {
break;
}
}
}
let tail = detok.finish();
if !tail.is_empty() {
on_token(&tail);
}
Ok(GenerateOutput {
text: detok.text(),
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
})
} else {
let mut streamer = StopStreamer::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,
});
}
let mut stopped = false;
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) {
if !engine.advance(gs, next_id) {
stopped = true;
break;
}
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
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;
break;
}
if let Some(end) = reasoning_end_len {
if generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens {
break;
}
}
}
streamer.finish(&detok.finish(), &mut on_token);
stopped |= streamer.stopped;
Ok(GenerateOutput {
text: streamer.into_text(),
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
})
}
}
}
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;
}
}
}
pub(crate) struct StopStreamer<'a> {
full: String,
emitted: usize,
max_stop: usize,
stops: &'a [String],
stopped: bool,
}
impl<'a> StopStreamer<'a> {
pub(crate) fn new(stops: &'a [String]) -> Self {
let max_stop = stops.iter().map(String::len).max().unwrap_or(1);
Self {
full: String::new(),
emitted: 0,
max_stop,
stops,
stopped: false,
}
}
pub(crate) fn push(&mut self, delta: &str, sink: &mut impl FnMut(&str)) -> bool {
if delta.is_empty() {
return false;
}
self.full.push_str(delta);
if let Some(hit) = earliest_stop_match(&self.full, self.stops) {
let slice = &self.full[self.emitted..hit];
if !slice.is_empty() {
sink(slice);
}
self.full.truncate(hit);
self.emitted = self.full.len();
self.stopped = true;
return true;
}
let mut safe = self
.full
.len()
.saturating_sub(self.max_stop.saturating_sub(1));
safe = safe.max(self.emitted);
while safe > self.emitted && !self.full.is_char_boundary(safe) {
safe -= 1;
}
if safe > self.emitted {
sink(&self.full[self.emitted..safe]);
self.emitted = safe;
}
false
}
pub(crate) fn finish(&mut self, tail: &str, sink: &mut impl FnMut(&str)) {
if self.stopped {
return;
}
if !tail.is_empty() {
self.full.push_str(tail);
}
if let Some(hit) = earliest_stop_match(&self.full, self.stops) {
let slice = &self.full[self.emitted..hit];
if !slice.is_empty() {
sink(slice);
}
self.full.truncate(hit);
self.emitted = self.full.len();
self.stopped = true;
return;
}
if self.emitted < self.full.len() {
sink(&self.full[self.emitted..]);
self.emitted = self.full.len();
}
}
pub(crate) fn into_text(self) -> String {
self.full
}
}
#[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,
) -> Result<bool, 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) {
if !engine.advance(gs, next_id) {
return Ok(true);
}
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
return Ok(true);
}
generated_ids.push(next_id);
all_ids.push(next_id);
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
if let Some(end) = reasoning_end_len {
if generated_ids.len().saturating_sub(end) >= gen_cfg.max_new_tokens {
break;
}
}
}
Ok(false)
}
#[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,
) -> Result<bool, InferenceError> {
let cfg = &model.config;
let mut stopped = false;
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) {
if !engine.advance(gs, next_id) {
stopped = true;
break;
}
}
if Some(next_id) == think_close_id {
thinking_closed = true;
}
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
if thinking_closed && reasoning_end_len.is_none() {
reasoning_end_len = Some(generated_ids.len());
}
let delta = detok.push(&model.tokenizer, next_id);
if !delta.is_empty() {
full.push_str(&delta);
}
if let Some(hit) = earliest_stop_match(full, &gen_cfg.stop_strings) {
full.truncate(hit);
stopped = true;
break;
}
if let Some(end) = reasoning_end_len {
if 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);
return Ok(true);
}
}
return Ok(false);
}
Ok(true)
}
pub(crate) fn earliest_stop_match(haystack: &str, stops: &[String]) -> Option<usize> {
stops.iter().filter_map(|s| haystack.find(s.as_str())).min()
}
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(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn earliest_stop_match_single_present() {
assert_eq!(
earliest_stop_match("hello world", &["world".to_string()]),
Some(6)
);
}
#[test]
fn earliest_stop_match_no_match() {
assert_eq!(
earliest_stop_match("hello world", &["foo".to_string()]),
None
);
}
#[test]
fn earliest_stop_match_multiple_returns_earliest() {
assert_eq!(
earliest_stop_match("hello world", &["world".to_string(), "lo".to_string()]),
Some(3)
);
}
#[test]
fn earliest_stop_match_at_index_zero() {
assert_eq!(
earliest_stop_match("stopword rest", &["stop".to_string()]),
Some(0)
);
}
#[test]
fn earliest_stop_match_multibyte_utf8() {
assert_eq!(
earliest_stop_match("世界hello", &["界".to_string()]),
Some(3)
);
}
#[test]
fn earliest_stop_match_empty_stops() {
assert_eq!(earliest_stop_match("hello", &[]), None);
}
#[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 stop_streamer_stop_split_across_deltas_no_double_emit() {
let stops = vec!["World".to_string()];
let mut streamer = StopStreamer::new(&stops);
let mut all_emitted: Vec<String> = Vec::new();
let stopped1 = streamer.push("hel", &mut |s| all_emitted.push(s.to_string()));
assert!(!stopped1);
let stopped2 = streamer.push("lo W", &mut |s| all_emitted.push(s.to_string()));
assert!(!stopped2);
let stopped3 = streamer.push("orld!", &mut |s| all_emitted.push(s.to_string()));
assert!(stopped3, "stop should be detected on third delta");
let concatenated = all_emitted.join("");
assert_eq!(
concatenated, "hello ",
"emitted concat must equal pre-stop text exactly once (BUG 1 regression)"
);
assert_eq!(streamer.into_text(), "hello ");
}
#[test]
fn stop_streamer_stop_at_first_delta() {
let stops = vec!["Stop".to_string()];
let mut streamer = StopStreamer::new(&stops);
let mut emitted: Vec<String> = Vec::new();
let stopped = streamer.push("Stop now", &mut |s| emitted.push(s.to_string()));
assert!(stopped);
assert_eq!(emitted.join(""), "");
assert_eq!(streamer.into_text(), "");
}
#[test]
fn stop_streamer_no_stop_natural_end() {
let stops = vec!["zzz".to_string()];
let mut streamer = StopStreamer::new(&stops);
let mut emitted: Vec<String> = Vec::new();
streamer.push("abc", &mut |s| emitted.push(s.to_string()));
streamer.push("def", &mut |s| emitted.push(s.to_string()));
streamer.finish("", &mut |s| emitted.push(s.to_string()));
assert_eq!(emitted.join(""), "abcdef");
assert_eq!(streamer.into_text(), "abcdef");
}
#[test]
fn stop_streamer_multibyte_no_panic() {
let stops = vec!["STOP".to_string()];
let mut streamer = StopStreamer::new(&stops);
let mut emitted: Vec<String> = Vec::new();
streamer.push("世", &mut |s| emitted.push(s.to_string()));
streamer.push("界x", &mut |s| emitted.push(s.to_string()));
streamer.finish("", &mut |s| emitted.push(s.to_string()));
let concat = emitted.join("");
assert_eq!(concat, streamer.into_text());
}
#[test]
fn stop_streamer_stop_at_delta_boundary() {
let stops = vec!["STOP".to_string()];
let mut streamer = StopStreamer::new(&stops);
let mut emitted: Vec<String> = Vec::new();
let stopped1 = streamer.push("abc", &mut |s| emitted.push(s.to_string()));
assert!(!stopped1);
let stopped2 = streamer.push("STOP", &mut |s| emitted.push(s.to_string()));
assert!(stopped2);
assert_eq!(emitted.join(""), "abc");
assert_eq!(streamer.into_text(), "abc");
}
#[test]
fn stop_streamer_hold_back_emits_safe_prefix_then_finish() {
let stops = vec!["xyz".to_string()]; let mut streamer = StopStreamer::new(&stops);
let mut emitted: Vec<String> = Vec::new();
let stopped = streamer.push("abcde", &mut |s| emitted.push(s.to_string()));
assert!(!stopped);
assert_eq!(emitted.join(""), "abc");
streamer.finish("", &mut |s| emitted.push(s.to_string()));
assert_eq!(emitted.join(""), "abcde");
assert_eq!(streamer.into_text(), "abcde");
}
#[test]
fn stop_streamer_finish_noop_after_stop() {
let stops = vec!["END".to_string()];
let mut streamer = StopStreamer::new(&stops);
let mut emitted: Vec<String> = Vec::new();
let stopped = streamer.push("helloENDextra", &mut |s| emitted.push(s.to_string()));
assert!(stopped);
streamer.finish("should_not_appear", &mut |s| emitted.push(s.to_string()));
assert_eq!(emitted.join(""), "hello");
assert_eq!(streamer.into_text(), "hello");
}
}