use ferrox_models::grammar_sampler::{ConstraintError, GrammarSampler, MaskOutcome};
use ferrox_models::sampling::Sampler;
use ferrox_models::tokenizer::StopTokens;
use crate::generate::{DecodeError, GenerationParams};
use crate::json_mode::mask_logits_for_json;
pub(crate) struct SampleState {
sampler: Sampler,
grammar: Option<GrammarSampler>,
}
impl SampleState {
pub(crate) fn new(seed: u64) -> Self {
Self {
sampler: Sampler::new(seed),
grammar: None,
}
}
#[cfg(test)]
pub(crate) fn grammar_started(&self) -> bool {
self.grammar.is_some()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Step {
Token(usize),
GrammarComplete,
}
pub(crate) fn sample_next(
state: &mut SampleState,
logits: &[f32],
params: &GenerationParams,
history: &[usize],
stop_tokens: &StopTokens,
decode_token: &dyn Fn(usize) -> String,
) -> Result<Step, DecodeError> {
if !params.needs_vocab_logits() {
return Ok(Step::Token(state.sampler.sample(
logits,
¶ms.sampling,
history,
)));
}
debug_assert!(
logits.len() > 1,
"a request needing vocabulary-shaped logits was handed {} of them; \
greedy_gpu_fold_allowed and needs_vocab_logits have drifted apart",
logits.len()
);
if state.grammar.is_none() {
if let Some(grammar) = ¶ms.grammar {
state.grammar = Some(GrammarSampler::new(
grammar.as_ref().clone(),
logits.len(),
|id| decode_token(id).into_bytes(),
|id| stop_tokens.contains(id),
));
}
}
let SampleState { sampler, grammar } = state;
let json_object = params.json_object;
let grammar_ref = grammar.as_ref();
let mut refusal: Option<ConstraintError> = None;
let mut outcome = MaskOutcome::Allowed;
let next = {
let mut mask = |scores: &mut [f32]| {
if json_object {
mask_logits_for_json(scores, decode_token);
}
if let Some(g) = grammar_ref {
match g.mask_logits(scores) {
Ok(o) => outcome = o,
Err(e) => refusal = Some(e),
}
}
};
sampler.sample_with_mask(logits, ¶ms.sampling, history, Some(&mut mask))
};
if let Some(e) = refusal {
return Err(DecodeError::GrammarConstraint {
detail: e.to_string(),
});
}
if outcome == MaskOutcome::Complete {
return Ok(Step::GrammarComplete);
}
if let Some(g) = grammar.as_mut() {
g.accept(next).map_err(|e| DecodeError::GrammarConstraint {
detail: e.to_string(),
})?;
}
Ok(Step::Token(next))
}
#[cfg(test)]
mod tests {
use super::*;
use ferrox_models::grammar::Grammar;
use ferrox_models::sampling::SamplingParams;
use std::sync::Arc;
fn decode_token(id: usize) -> String {
match id {
0 => "<html>".to_string(),
1 => "{".to_string(),
_ => "0".to_string(),
}
}
fn params(json_object: bool, temperature: f32) -> GenerationParams {
GenerationParams {
max_tokens: 8,
sampling: SamplingParams {
temperature,
..SamplingParams::default()
},
seed: 7,
stop: Vec::new(),
stop_token_ids: Vec::new(),
json_object,
grammar: None,
cancel: None,
ignore_eos: false,
}
}
fn letter(id: usize) -> String {
match id {
0 => "a".to_string(),
1 => "b".to_string(),
2 => "c".to_string(),
3 => "d".to_string(),
_ => "</s>".to_string(),
}
}
const LETTER_EOG: usize = 4;
fn letter_stops() -> StopTokens {
StopTokens::from_eos(Some(LETTER_EOG))
}
fn with_grammar(mut p: GenerationParams, src: &str) -> GenerationParams {
p.grammar = Some(Arc::new(
Grammar::from_str_with_root(src, "root").expect("test grammar parses"),
));
p
}
fn token(step: Step) -> usize {
match step {
Step::Token(id) => id,
Step::GrammarComplete => {
panic!("the grammar ended the generation where a token was expected")
}
}
}
fn step(
state: &mut SampleState,
logits: &[f32],
params: &GenerationParams,
history: &[usize],
stops: &StopTokens,
decode: &dyn Fn(usize) -> String,
) -> Result<Step, DecodeError> {
sample_next(state, logits, params, history, stops, decode)
}
#[test]
fn json_mode_masks_at_temperature_zero() {
let logits = vec![9.0, 1.0, 0.5];
let mut state = SampleState::new(7);
let chosen = token(
step(
&mut state,
&logits,
¶ms(true, 0.0),
&[],
&StopTokens::default(),
&decode_token,
)
.expect("json mode has a legal token here"),
);
assert_ne!(chosen, 0, "a non-JSON-safe token was sampled in json mode");
assert_eq!(chosen, 1);
}
#[test]
fn json_mode_masks_when_sampling_too() {
let logits = vec![9.0, 1.0, 0.5];
for seed in 0..16u64 {
let mut state = SampleState::new(seed);
let chosen = token(
step(
&mut state,
&logits,
¶ms(true, 1.0),
&[],
&StopTokens::default(),
&decode_token,
)
.expect("json mode has a legal token here"),
);
assert_ne!(chosen, 0, "seed {seed} sampled a non-JSON-safe token");
}
}
#[test]
fn a_plain_request_is_not_masked() {
let logits = vec![9.0, 1.0, 0.5];
let mut state = SampleState::new(7);
let chosen = token(
step(
&mut state,
&logits,
¶ms(false, 0.0),
&[],
&StopTokens::default(),
&decode_token,
)
.expect("an unconstrained request cannot fail"),
);
assert_eq!(chosen, 0);
}
#[test]
fn a_grammar_masks_the_argmax_it_forbids() {
let logits = vec![1.0, 9.0, 0.5, 0.1, 0.0];
let mut state = SampleState::new(7);
let chosen = token(
step(
&mut state,
&logits,
&with_grammar(params(false, 0.0), r#"root ::= "ac""#),
&[],
&letter_stops(),
&letter,
)
.expect("\"a\" is legal here"),
);
assert_eq!(chosen, 0, "the grammar's only legal first token is \"a\"");
assert!(state.grammar_started());
}
#[test]
fn a_grammar_advances_between_steps() {
let logits = vec![1.0, 9.0, 0.5, 0.1, 0.0];
let params = with_grammar(params(false, 0.0), r#"root ::= "ac""#);
let mut state = SampleState::new(7);
let first = token(
step(&mut state, &logits, ¶ms, &[], &letter_stops(), &letter)
.expect("\"a\" is legal here"),
);
let second = token(
step(
&mut state,
&logits,
¶ms,
&[first],
&letter_stops(),
&letter,
)
.expect("\"c\" is legal here"),
);
assert_eq!(
second, 2,
"the grammar did not advance past its first token"
);
}
#[test]
fn a_grammar_holds_end_of_generation_back_until_it_is_satisfied() {
let logits = vec![1.0, 0.5, 0.4, 0.1, 9.0];
let params = with_grammar(params(false, 0.0), r#"root ::= "a""#);
let mut state = SampleState::new(7);
let first = token(
step(&mut state, &logits, ¶ms, &[], &letter_stops(), &letter)
.expect("\"a\" is legal here"),
);
assert_eq!(first, 0, "generation was allowed to end unsatisfied");
let second = token(
step(
&mut state,
&logits,
¶ms,
&[first],
&letter_stops(),
&letter,
)
.expect("eog is legal once the parse is complete"),
);
assert_eq!(second, LETTER_EOG);
}
#[test]
fn a_finished_grammar_with_nothing_left_to_say_ends_the_generation() {
let logits = vec![1.0, 9.0, 0.5, 0.1, 0.0];
let params = with_grammar(params(false, 0.0), r#"root ::= "a""#);
let mut state = SampleState::new(7);
let first = token(
step(
&mut state,
&logits,
¶ms,
&[],
&StopTokens::default(),
&letter,
)
.expect("\"a\" is legal here"),
);
assert_eq!(first, 0);
assert_eq!(
step(
&mut state,
&logits,
¶ms,
&[first],
&StopTokens::default(),
&letter,
)
.expect("a completed parse is not an error"),
Step::GrammarComplete
);
}
#[test]
fn a_grammar_no_token_can_satisfy_stops_the_generation() {
let logits = vec![1.0, 9.0, 0.5, 0.1, 0.0];
let mut state = SampleState::new(7);
let err = step(
&mut state,
&logits,
&with_grammar(params(false, 0.0), r#"root ::= "zz""#),
&[],
&letter_stops(),
&letter,
)
.expect_err("no token in this vocabulary starts with z");
assert!(
matches!(err, DecodeError::GrammarConstraint { .. }),
"{err}"
);
}
#[test]
fn a_grammar_and_json_mode_intersect_rather_than_replace_each_other() {
let logits = vec![9.0, 8.0, 0.5];
let mut state = SampleState::new(7);
let chosen = token(
step(
&mut state,
&logits,
&with_grammar(params(true, 0.0), r#"root ::= [0-9]+"#),
&[],
&StopTokens::default(),
&decode_token,
)
.expect("\"0\" is legal under both"),
);
assert_eq!(chosen, 2);
}
#[test]
fn a_constrained_request_may_never_fold_lm_head_into_a_gpu_argmax() {
use crate::generate::greedy_gpu_fold_allowed;
assert!(greedy_gpu_fold_allowed(¶ms(false, 0.0)));
assert!(!greedy_gpu_fold_allowed(¶ms(true, 0.0)));
assert!(!greedy_gpu_fold_allowed(&with_grammar(
params(false, 0.0),
r#"root ::= "a""#
)));
assert!(!greedy_gpu_fold_allowed(¶ms(false, 0.8)));
assert!(!greedy_gpu_fold_allowed(¶ms(true, 0.8)));
assert!(!greedy_gpu_fold_allowed(&with_lazy_grammar(
params(false, 0.0),
r#"root ::= "cd""#,
"c",
)));
}
fn with_lazy_grammar(mut p: GenerationParams, src: &str, trigger: &str) -> GenerationParams {
use ferrox_models::grammar::LazyTriggers;
p.grammar = Some(Arc::new(
Grammar::from_str_with_root(src, "root")
.expect("test grammar parses")
.into_lazy(
LazyTriggers::new()
.with_word(trigger)
.expect("the trigger compiles")
.mandatory(),
)
.expect("there is a trigger"),
));
p
}
#[test]
fn a_mandatory_lazy_grammar_forbids_only_the_ending_until_it_fires() {
let params = with_lazy_grammar(params(false, 0.0), r#"root ::= "cd""#, "c");
let stops = letter_stops();
let mut state = SampleState::new(7);
let logits = vec![5.0, 1.0, 0.5, 0.5, 9.0];
let chosen = token(step(&mut state, &logits, ¶ms, &[], &stops, &letter).unwrap());
assert_eq!(
chosen, 0,
"free text before the trigger, but not the end of the turn"
);
let logits = vec![1.0, 1.0, 9.0, 0.5, 5.0];
let chosen = token(step(&mut state, &logits, ¶ms, &[0], &stops, &letter).unwrap());
assert_eq!(chosen, 2, "the trigger token itself");
let logits = vec![9.0, 9.0, 9.0, 0.0, 9.0];
let chosen = token(step(&mut state, &logits, ¶ms, &[0, 2], &stops, &letter).unwrap());
assert_eq!(
chosen, 3,
"once triggered the grammar constrains: only \"d\" continues \"c\""
);
}
#[test]
fn an_optional_lazy_grammar_lets_the_turn_end_before_the_trigger() {
use ferrox_models::grammar::LazyTriggers;
let mut params = params(false, 0.0);
params.grammar = Some(Arc::new(
Grammar::from_str_with_root(r#"root ::= "cd""#, "root")
.unwrap()
.into_lazy(LazyTriggers::new().with_word("c").unwrap())
.unwrap(),
));
let logits = vec![5.0, 1.0, 0.5, 0.5, 9.0];
let mut state = SampleState::new(7);
let chosen =
token(step(&mut state, &logits, ¶ms, &[], &letter_stops(), &letter).unwrap());
assert_eq!(chosen, LETTER_EOG, "an optional trigger constrains nothing");
}
}