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>,
}
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();
for (state_id, grammar_state) in grammar_states[..effective_states].iter().enumerate() {
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);
let word = token_id / 64;
let bit = token_id % 64;
masks[state_id * mask_stride + word] |= 1u64 << bit;
}
SimResult::Reject => {
}
}
}
}
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,
}
}
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 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)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::grammar::pda::{Alt, CompiledGrammar, GrammarBuilder, GrammarState, Rule, Symbol};
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);
}
}