use crate::grammar::pda::{CompiledGrammar, GrammarState, SimResult, simulate_token};
pub const MAX_GRAMMAR_STATES: usize = 256;
pub struct VocabPartition {
masks: Vec<u64>,
mask_stride: usize,
vocab_size: usize,
states: Vec<GrammarState>,
context_dependent: Vec<usize>,
context_dependent_by_state: Vec<Option<Vec<usize>>>,
}
impl VocabPartition {
pub fn build(
grammar: &CompiledGrammar,
grammar_states: Vec<GrammarState>,
vocab_bytes: &[Vec<u8>],
) -> Self {
let vocab_size = vocab_bytes.len();
let mask_stride = vocab_size.div_ceil(64);
let num_states = grammar_states.len();
if num_states > MAX_GRAMMAR_STATES {
tracing::warn!(
"grammar has {} states (max {}); first {} will be precomputed",
num_states,
MAX_GRAMMAR_STATES,
MAX_GRAMMAR_STATES
);
}
let effective_states = num_states.min(MAX_GRAMMAR_STATES);
let mut masks = vec![0u64; effective_states * mask_stride];
let mut ctx_dep_set = std::collections::HashSet::new();
let mut context_dependent_by_state = Vec::with_capacity(effective_states);
let context_entry_budget =
masks.len().saturating_mul(std::mem::size_of::<u64>()) / std::mem::size_of::<usize>();
let mut context_entries_stored = 0usize;
for (state_id, grammar_state) in grammar_states[..effective_states].iter().enumerate() {
let mut state_context_dependent = Vec::new();
for (token_id, token_bytes) in vocab_bytes.iter().enumerate() {
if token_bytes.is_empty() {
continue;
}
let (sim_result, _) = simulate_token(grammar_state, grammar, token_bytes);
match sim_result {
SimResult::Accept => {
let word = token_id / 64;
let bit = token_id % 64;
masks[state_id * mask_stride + word] |= 1u64 << bit;
}
SimResult::ContextDependent => {
ctx_dep_set.insert(token_id);
state_context_dependent.push(token_id);
let word = token_id / 64;
let bit = token_id % 64;
masks[state_id * mask_stride + word] |= 1u64 << bit;
}
SimResult::Reject => {
}
}
}
state_context_dependent.shrink_to_fit();
let stored_capacity = state_context_dependent.capacity();
if context_entries_stored.saturating_add(stored_capacity) <= context_entry_budget {
context_entries_stored += stored_capacity;
context_dependent_by_state.push(Some(state_context_dependent));
} else {
context_dependent_by_state.push(None);
}
}
let mut context_dependent: Vec<usize> = ctx_dep_set.into_iter().collect();
context_dependent.sort_unstable();
Self {
masks,
mask_stride,
vocab_size,
states: grammar_states,
context_dependent,
context_dependent_by_state,
}
}
pub fn apply_mask(&self, state_id: usize, logits: &mut [f32]) {
debug_assert!(
logits.len() >= self.vocab_size,
"logits slice shorter than vocab_size"
);
if state_id >= self.states.len().min(MAX_GRAMMAR_STATES) {
for l in logits[..self.vocab_size].iter_mut() {
*l = f32::NEG_INFINITY;
}
return;
}
let mask_base = state_id * self.mask_stride;
for word_idx in 0..self.mask_stride {
let mask_word = self.masks[mask_base + word_idx];
let base_token = word_idx * 64;
if mask_word == u64::MAX {
continue;
}
if mask_word == 0 {
let end = (base_token + 64).min(self.vocab_size);
for l in logits[base_token..end].iter_mut() {
*l = f32::NEG_INFINITY;
}
continue;
}
for bit in 0..64u32 {
let token_idx = base_token + bit as usize;
if token_idx >= self.vocab_size {
break;
}
if mask_word & (1u64 << bit) == 0 {
logits[token_idx] = f32::NEG_INFINITY;
}
}
}
}
pub fn context_dependent_ids(&self) -> &[usize] {
&self.context_dependent
}
pub(crate) fn context_dependent_ids_for_state(&self, state_id: usize) -> &[usize] {
self.context_dependent_by_state
.get(state_id)
.and_then(Option::as_deref)
.unwrap_or(&self.context_dependent)
}
pub fn num_states(&self) -> usize {
self.states.len().min(MAX_GRAMMAR_STATES)
}
pub fn grammar_state(&self, state_id: usize) -> Option<&GrammarState> {
self.states.get(state_id)
}
pub(crate) fn any_allowed_token(
&self,
state_id: usize,
mut predicate: impl FnMut(usize) -> bool,
) -> bool {
if state_id >= self.states.len().min(MAX_GRAMMAR_STATES) {
return false;
}
let mask_base = state_id * self.mask_stride;
for word_idx in 0..self.mask_stride {
let mut mask_word = self.masks[mask_base + word_idx];
while mask_word != 0 {
let bit = mask_word.trailing_zeros() as usize;
let token_id = word_idx * 64 + bit;
if token_id < self.vocab_size && predicate(token_id) {
return true;
}
mask_word &= mask_word - 1;
}
}
false
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::grammar::pda::{
CompiledGrammar, GrammarBuilder, GrammarState, Rule, StepResult, Symbol, advance_byte,
};
fn or_grammar() -> CompiledGrammar {
let mut b = GrammarBuilder::new();
b.add_rule(
"root",
vec![vec![Symbol::Terminal(b'a')], vec![Symbol::Terminal(b'b')]],
);
b.build()
}
fn ab_vocab() -> Vec<Vec<u8>> {
vec![b"a".to_vec(), b"b".to_vec()]
}
fn abc_vocab() -> Vec<Vec<u8>> {
vec![b"a".to_vec(), b"b".to_vec(), b"c".to_vec()]
}
#[test]
fn build_basic_mask() {
let grammar = or_grammar();
let states = vec![GrammarState::initial()];
let vocab = ab_vocab();
let partition = VocabPartition::build(&grammar, states, &vocab);
assert_eq!(partition.num_states(), 1);
}
#[test]
fn apply_mask_allows_correct_tokens() {
let grammar = or_grammar();
let states = vec![GrammarState::initial()];
let vocab = abc_vocab();
let partition = VocabPartition::build(&grammar, states, &vocab);
let mut logits = vec![1.0f32, 2.0f32, 3.0f32];
partition.apply_mask(0, &mut logits);
assert!(logits[0] > f32::NEG_INFINITY, "token 'a' should be allowed");
assert!(logits[1] > f32::NEG_INFINITY, "token 'b' should be allowed");
assert_eq!(logits[2], f32::NEG_INFINITY, "token 'c' should be blocked");
}
#[test]
fn apply_mask_unknown_state_blocks_all() {
let grammar = or_grammar();
let states = vec![GrammarState::initial()];
let vocab = ab_vocab();
let partition = VocabPartition::build(&grammar, states, &vocab);
let mut logits = vec![1.0f32, 2.0f32];
partition.apply_mask(99, &mut logits);
assert_eq!(logits[0], f32::NEG_INFINITY);
assert_eq!(logits[1], f32::NEG_INFINITY);
}
#[test]
fn mask_all_zeros_fills_neg_inf() {
let grammar = CompiledGrammar {
rules: vec![Rule {
name: "root".to_string(),
alts: vec![],
}],
};
let states = vec![GrammarState::initial()];
let vocab = ab_vocab();
let partition = VocabPartition::build(&grammar, states, &vocab);
let mut logits = vec![1.0f32, 2.0f32];
partition.apply_mask(0, &mut logits);
assert_eq!(logits[0], f32::NEG_INFINITY);
assert_eq!(logits[1], f32::NEG_INFINITY);
}
#[test]
fn mask_all_ones_preserves_logits() {
let mut builder = GrammarBuilder::new();
builder.add_rule("root", vec![vec![Symbol::AnyByte]]);
let grammar = builder.build();
let states = vec![GrammarState::initial()];
let vocab = abc_vocab();
let partition = VocabPartition::build(&grammar, states, &vocab);
let mut logits = vec![1.0f32, 2.0f32, 3.0f32];
partition.apply_mask(0, &mut logits);
for &l in &logits {
assert!(l > f32::NEG_INFINITY);
}
}
#[test]
fn bitmask_and_correctness() {
let grammar = or_grammar();
let mut vocab: Vec<Vec<u8>> = vec![b"a".to_vec(), b"b".to_vec()];
vocab.extend((2..65).map(|_| b"c".to_vec()));
assert_eq!(vocab.len(), 65);
let states = vec![GrammarState::initial()];
let partition = VocabPartition::build(&grammar, states, &vocab);
let mut logits = vec![1.0f32; 65];
partition.apply_mask(0, &mut logits);
assert!(logits[0] > f32::NEG_INFINITY, "token 0 allowed");
assert!(logits[1] > f32::NEG_INFINITY, "token 1 allowed");
for i in 2..65 {
assert_eq!(logits[i], f32::NEG_INFINITY, "token {i} blocked");
}
}
#[test]
fn empty_token_skipped() {
let grammar = or_grammar();
let vocab = vec![b"a".to_vec(), vec![], b"b".to_vec()];
let states = vec![GrammarState::initial()];
let partition = VocabPartition::build(&grammar, states, &vocab);
let mut logits = vec![1.0f32; 3];
partition.apply_mask(0, &mut logits);
assert!(logits[0] > f32::NEG_INFINITY);
assert_eq!(logits[1], f32::NEG_INFINITY); assert!(logits[2] > f32::NEG_INFINITY);
}
#[test]
fn context_dependent_ids_are_partitioned_by_state() {
let mut builder = GrammarBuilder::new();
builder.add_rule(
"root",
vec![b"abcd".iter().copied().map(Symbol::Terminal).collect()],
);
let grammar = builder.build();
let state0 = GrammarState::initial();
let mut state1 = state0.clone();
assert_eq!(
advance_byte(&mut state1, &grammar, b'a'),
StepResult::Accepted
);
let mut state2 = state1.clone();
assert_eq!(
advance_byte(&mut state2, &grammar, b'b'),
StepResult::Accepted
);
let vocab = vec![b"ax".to_vec(), b"bx".to_vec(), b"cx".to_vec()];
let partition = VocabPartition::build(&grammar, vec![state0, state1, state2], &vocab);
assert_eq!(partition.context_dependent_ids(), &[0, 1, 2]);
assert_eq!(partition.context_dependent_ids_for_state(0), &[0]);
assert_eq!(partition.context_dependent_ids_for_state(1), &[1]);
assert_eq!(partition.context_dependent_ids_for_state(2), &[2]);
assert_eq!(
partition.context_dependent_ids_for_state(usize::MAX),
&[0, 1, 2],
"unknown states must use the conservative global union"
);
}
#[test]
fn dense_state_lists_fall_back_within_mask_sized_budget() {
let mut builder = GrammarBuilder::new();
builder.add_rule(
"root",
vec![b"aaaa".iter().copied().map(Symbol::Terminal).collect()],
);
let grammar = builder.build();
let state0 = GrammarState::initial();
let mut state1 = state0.clone();
assert_eq!(
advance_byte(&mut state1, &grammar, b'a'),
StepResult::Accepted
);
let mut state2 = state1.clone();
assert_eq!(
advance_byte(&mut state2, &grammar, b'a'),
StepResult::Accepted
);
let vocab = vec![b"ax".to_vec(); 128];
let partition = VocabPartition::build(&grammar, vec![state0, state1, state2], &vocab);
assert_eq!(partition.context_dependent_ids().len(), 128);
assert!(
partition
.context_dependent_by_state
.iter()
.all(Option::is_none),
"dense local lists must use the global fallback instead of exceeding the mask-sized \
storage budget"
);
for state_id in 0..3 {
assert_eq!(
partition.context_dependent_ids_for_state(state_id),
partition.context_dependent_ids()
);
}
}
}