Skip to main content

ferrum_sampler/
guided.rs

1//! Regex-guided decoding — hard token masking via DFA.
2//!
3//! Given a regex pattern and a tokenizer vocab, build a DFA and at each
4//! sampling step compute which tokens can extend the currently accepted
5//! prefix without leaving the language. Invalid tokens get `-INFINITY` so
6//! the downstream sampler cannot pick them, regardless of temperature or
7//! top-k/top-p.
8//!
9//! This is the "outlines"-style approach: convert the constraint to a
10//! finite automaton, walk it byte-by-byte per token to decide validity.
11//! No schema → regex transformation here — that belongs a layer up (see
12//! `ResponseFormat::JsonSchema` handling).
13//!
14//! # Design notes
15//!
16//! * The DFA is built once at request admission — regex compilation is the
17//!   expensive step (~1-5 ms for short patterns, scales with ambiguity).
18//! * Per-step cost is O(vocab_size · avg_token_bytes). For a 150k vocab
19//!   with ~5 byte tokens that's ~750k state transitions per sampling step.
20//!   Fine for single requests; we'll add a cached (state, token) → (valid,
21//!   next_state) transition table if this becomes a bottleneck.
22//! * End-of-string: once the DFA can accept, EOS becomes a valid choice.
23//!   If the pattern is "open" (e.g. `.*`) EOS is always allowed.
24
25use std::sync::Arc;
26
27use ferrum_interfaces::sampler::{LogitsProcessor, ProcessorPriority, SamplingContext};
28use ferrum_interfaces::tokenizer::Tokenizer;
29use ferrum_types::{FerrumError, Result, TokenId};
30use parking_lot::Mutex;
31use regex_automata::{
32    dfa::{dense::DFA, Automaton, StartKind},
33    util::{primitives::StateID, start::Config as StartConfig},
34    Anchored,
35};
36
37/// Hard-mask regex constraint processor.
38///
39/// Build with `RegexGuidedProcessor::new(pattern, tokenizer, eos_token)`;
40/// use by adding as a high-priority logits processor before temperature /
41/// top-k / top-p.
42pub struct RegexGuidedProcessor {
43    dfa: DFA<Vec<u32>>,
44    /// Current DFA state — advanced lazily per `process()` call from the
45    /// generated tokens accumulated so far.
46    state: Mutex<DfaPosition>,
47    /// Precomputed per-token byte sequences. Indexed by token id.
48    token_bytes: Vec<Vec<u8>>,
49    /// Optional EOS id — always allowed once the DFA can accept.
50    eos_token: Option<TokenId>,
51    /// Number of tokens already consumed into `state`.
52    consumed: Mutex<usize>,
53}
54
55impl std::fmt::Debug for RegexGuidedProcessor {
56    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57        f.debug_struct("RegexGuidedProcessor")
58            .field("vocab_size", &self.token_bytes.len())
59            .field("consumed", &*self.consumed.lock())
60            .finish()
61    }
62}
63
64#[derive(Copy, Clone, Debug)]
65struct DfaPosition {
66    state: StateID,
67    /// Set once the DFA enters a dead (non-accepting, non-escapable) state.
68    /// At that point no token is valid and we must rely on the sampler's
69    /// fallback (we emit EOS to terminate the sequence gracefully).
70    dead: bool,
71}
72
73impl RegexGuidedProcessor {
74    /// Build a guided-decoding processor for `pattern`. `tokenizer` is used
75    /// to map each vocab entry to its byte representation.
76    pub fn new(
77        pattern: &str,
78        tokenizer: Arc<dyn Tokenizer + Send + Sync>,
79        eos_token: Option<TokenId>,
80    ) -> Result<Self> {
81        // `Anchored::Yes` only pins the match start; without `\z` the DFA
82        // happily keeps accepting after the match completes, and tokens that
83        // would produce "valid match + garbage" wouldn't get masked. Wrap
84        // the user pattern so the full generated output must match.
85        let wrapped = format!(r"(?:{pattern})\z");
86        let dfa = DFA::builder()
87            .configure(DFA::config().start_kind(StartKind::Anchored))
88            .build(&wrapped)
89            .map_err(|e| FerrumError::invalid_request(format!("regex compile: {e}")))?;
90
91        let start = dfa
92            .start_state(&StartConfig::new().anchored(Anchored::Yes))
93            .map_err(|e| FerrumError::invalid_request(format!("regex start state: {e}")))?;
94
95        // Decode every token in the vocab once. The regex constrains the
96        // generated text, so this must use the tokenizer's decoded surface
97        // form rather than the raw vocab entry (for example byte-level/BPE
98        // vocab entries may contain internal markers that are not emitted).
99        let vocab_size = tokenizer.vocab_size();
100        let mut token_bytes = Vec::with_capacity(vocab_size);
101        for i in 0..vocab_size {
102            let id = TokenId::new(i as u32);
103            let bytes = tokenizer
104                .decode(&[id], false)
105                .ok()
106                .or_else(|| tokenizer.token_text(id).map(str::to_string))
107                .map(String::into_bytes)
108                .unwrap_or_default();
109            token_bytes.push(bytes);
110        }
111
112        Ok(Self {
113            dfa,
114            state: Mutex::new(DfaPosition {
115                state: start,
116                dead: false,
117            }),
118            token_bytes,
119            eos_token,
120            consumed: Mutex::new(0),
121        })
122    }
123
124    /// Reset for a new generation.
125    pub fn reset(&self) -> Result<()> {
126        let start = self
127            .dfa
128            .start_state(&StartConfig::new().anchored(Anchored::Yes))
129            .map_err(|e| FerrumError::internal(format!("regex start state: {e}")))?;
130        *self.state.lock() = DfaPosition {
131            state: start,
132            dead: false,
133        };
134        *self.consumed.lock() = 0;
135        Ok(())
136    }
137
138    /// Check if the pattern can currently accept (i.e. EOS is valid here).
139    pub fn can_accept(&self) -> bool {
140        let pos = *self.state.lock();
141        !pos.dead && self.dfa.is_match_state(self.dfa.next_eoi_state(pos.state))
142    }
143
144    /// Walk the DFA over `bytes` starting from `state`. Returns the new
145    /// state, or `None` if a dead state is reached partway through — i.e.
146    /// this byte sequence cannot extend the current match.
147    fn advance(&self, mut state: StateID, bytes: &[u8]) -> Option<StateID> {
148        for &b in bytes {
149            state = self.dfa.next_state(state, b);
150            if self.dfa.is_dead_state(state) {
151                return None;
152            }
153        }
154        Some(state)
155    }
156
157    /// Apply the hard mask to `logits` given the current DFA state.
158    pub fn mask_logits(&self, logits: &mut [f32]) {
159        let pos = *self.state.lock();
160        if pos.dead {
161            // The generated prefix already left the DFA language. At this
162            // point a hard mask cannot repair the output; forcing EOS would
163            // actively truncate structured responses. Let generation proceed
164            // and leave final API/schema validation to reject bad output.
165            return;
166        }
167
168        let pattern_done = self.dfa.is_match_state(self.dfa.next_eoi_state(pos.state));
169
170        // Mask tokens individually. `&self.token_bytes` is O(vocab) once; a
171        // cached per-state transition table would amortise further, but this
172        // keeps the hot path allocation-free for now.
173        let vocab = logits.len().min(self.token_bytes.len());
174        let mut any_allowed = false;
175        let mut fallback_token: Option<(usize, f32)> = None;
176        for idx in 0..vocab {
177            let is_eos = self.eos_token.map_or(false, |e| e.get() as usize == idx);
178            let bytes = &self.token_bytes[idx];
179            let original_logit = logits[idx];
180            if !is_eos
181                && original_logit.is_finite()
182                && fallback_token
183                    .map(|(_, best)| original_logit > best)
184                    .unwrap_or(true)
185            {
186                fallback_token = Some((idx, original_logit));
187            }
188
189            let allowed = if is_eos {
190                pattern_done
191            } else if bytes.is_empty() {
192                // Unknown / special token outside the regex alphabet — only
193                // let it through if pattern can already accept (so it can't
194                // block a valid termination).
195                pattern_done
196            } else {
197                self.advance(pos.state, bytes).is_some()
198            };
199
200            if allowed {
201                any_allowed = true;
202            } else {
203                logits[idx] = f32::NEG_INFINITY;
204            }
205        }
206
207        // No token in the vocab can extend the current match. If the regex is
208        // already complete, EOS is the correct clean terminator. Otherwise do
209        // not force an invalid premature EOS; choose the model's best non-EOS
210        // token and let later API/schema validation reject the response if the
211        // fallback cannot recover. This can happen when single-token decode
212        // surfaces differ from full-sequence decode for byte-level tokenizers.
213        if !any_allowed {
214            if pattern_done {
215                self.force_eos(logits);
216            } else if let Some((token, _)) = fallback_token {
217                self.force_only_token(logits, token);
218            }
219        }
220    }
221
222    fn force_eos(&self, logits: &mut [f32]) {
223        if let Some(eos) = self.eos_token {
224            let eos_idx = eos.get() as usize;
225            for (i, l) in logits.iter_mut().enumerate() {
226                *l = if i == eos_idx { 0.0 } else { f32::NEG_INFINITY };
227            }
228        }
229    }
230
231    fn force_only_token(&self, logits: &mut [f32], token: usize) {
232        for (i, l) in logits.iter_mut().enumerate() {
233            *l = if i == token { 0.0 } else { f32::NEG_INFINITY };
234        }
235    }
236
237    /// Public wrapper around `advance_with_tokens` for direct callers
238    /// (the engine applies the mask inline rather than via
239    /// `LogitsProcessor::process`).
240    pub fn advance_with_tokens_public(&self, tokens: &[TokenId]) {
241        self.advance_with_tokens(tokens);
242    }
243
244    /// Advance the stored state by consuming tokens that were decided after
245    /// the last `process()` call. The engine calls this in `process()` with
246    /// the full generated-tokens list; we skip the prefix we've already seen.
247    fn advance_with_tokens(&self, tokens: &[TokenId]) {
248        let mut consumed = self.consumed.lock();
249        if *consumed >= tokens.len() {
250            return;
251        }
252        let mut pos = self.state.lock();
253        for &tok in &tokens[*consumed..] {
254            if pos.dead {
255                break;
256            }
257            let idx = tok.get() as usize;
258            if idx >= self.token_bytes.len() {
259                continue;
260            }
261            let bytes = &self.token_bytes[idx];
262            // EOS terminates cleanly — leave state where it is.
263            if self.eos_token.map_or(false, |e| e == tok) {
264                continue;
265            }
266            if let Some(next) = self.advance(pos.state, bytes) {
267                pos.state = next;
268            } else {
269                pos.dead = true;
270            }
271        }
272        *consumed = tokens.len();
273    }
274}
275
276impl LogitsProcessor for RegexGuidedProcessor {
277    fn process(&self, ctx: &mut SamplingContext) -> Result<()> {
278        self.advance_with_tokens(ctx.previous_tokens);
279        self.mask_logits(ctx.logits);
280        Ok(())
281    }
282
283    fn name(&self) -> &str {
284        "regex_guided"
285    }
286
287    fn priority(&self) -> ProcessorPriority {
288        // Apply before temperature / top-k / top-p — those should only see
289        // logits for *valid* tokens.
290        ProcessorPriority::High
291    }
292}
293
294#[cfg(test)]
295mod tests {
296    use super::*;
297    use ferrum_interfaces::tokenizer::{ChatMessage, TokenizerInfo, TokenizerType};
298    use ferrum_types::{SpecialTokens, TokenId};
299
300    /// Tiny tokenizer: each ASCII character 0..=255 is a single token, plus
301    /// an EOS token at 256. Matches how byte-level BPEs decompose in the
302    /// worst case, so the test is a lower bound on the real-world case.
303    struct ByteTokenizer {
304        special: SpecialTokens,
305        byte_strings: Vec<String>,
306    }
307
308    impl ByteTokenizer {
309        fn new() -> Self {
310            let mut byte_strings = Vec::with_capacity(257);
311            for b in 0u8..=255 {
312                byte_strings.push(String::from_utf8(vec![b]).unwrap_or_default());
313            }
314            byte_strings.push("</s>".to_string());
315            Self {
316                special: SpecialTokens {
317                    bos_token: None,
318                    eos_token: Some(TokenId::new(256)),
319                    unk_token: None,
320                    pad_token: None,
321                    sep_token: None,
322                    cls_token: None,
323                    mask_token: None,
324                    extra_eos_tokens: Vec::new(),
325                },
326                byte_strings,
327            }
328        }
329    }
330
331    impl Tokenizer for ByteTokenizer {
332        fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
333            Ok(text.bytes().map(|b| TokenId::new(b as u32)).collect())
334        }
335        fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
336            let mut out = String::new();
337            for t in tokens {
338                let idx = t.get() as usize;
339                if idx < 256 {
340                    out.push(idx as u8 as char);
341                }
342            }
343            Ok(out)
344        }
345        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
346            self.decode(&[next], false)
347        }
348        fn vocab_size(&self) -> usize {
349            257
350        }
351        fn special_tokens(&self) -> &SpecialTokens {
352            &self.special
353        }
354        fn token_id(&self, text: &str) -> Option<TokenId> {
355            if text.len() == 1 {
356                Some(TokenId::new(text.bytes().next().unwrap() as u32))
357            } else {
358                None
359            }
360        }
361        fn token_text(&self, token_id: TokenId) -> Option<&str> {
362            self.byte_strings
363                .get(token_id.get() as usize)
364                .map(|s| s.as_str())
365        }
366        fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
367            Ok(String::new())
368        }
369        fn info(&self) -> TokenizerInfo {
370            TokenizerInfo {
371                tokenizer_type: TokenizerType::Custom,
372                vocab_size: 257,
373                special_tokens: self.special.clone(),
374                supports_incremental: true,
375                supports_chat_template: false,
376                max_token_length: Some(1),
377                model_name: Some("byte-tokenizer-test".into()),
378            }
379        }
380    }
381
382    fn processor(pattern: &str) -> RegexGuidedProcessor {
383        let tok: Arc<dyn Tokenizer> = Arc::new(ByteTokenizer::new());
384        RegexGuidedProcessor::new(pattern, tok, Some(TokenId::new(256))).unwrap()
385    }
386
387    struct TinyTokenizer {
388        special: SpecialTokens,
389        strings: Vec<String>,
390    }
391
392    impl TinyTokenizer {
393        fn new(strings: Vec<&str>, eos: u32) -> Self {
394            Self {
395                special: SpecialTokens {
396                    bos_token: None,
397                    eos_token: Some(TokenId::new(eos)),
398                    unk_token: None,
399                    pad_token: None,
400                    sep_token: None,
401                    cls_token: None,
402                    mask_token: None,
403                    extra_eos_tokens: Vec::new(),
404                },
405                strings: strings.into_iter().map(str::to_string).collect(),
406            }
407        }
408    }
409
410    impl Tokenizer for TinyTokenizer {
411        fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
412            Ok(text
413                .chars()
414                .filter_map(|ch| self.token_id(&ch.to_string()))
415                .collect())
416        }
417
418        fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
419            Ok(tokens
420                .iter()
421                .filter_map(|token| self.token_text(*token))
422                .collect::<Vec<_>>()
423                .join(""))
424        }
425
426        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
427            self.decode(&[next], false)
428        }
429
430        fn vocab_size(&self) -> usize {
431            self.strings.len()
432        }
433
434        fn special_tokens(&self) -> &SpecialTokens {
435            &self.special
436        }
437
438        fn token_id(&self, text: &str) -> Option<TokenId> {
439            self.strings
440                .iter()
441                .position(|value| value == text)
442                .map(|idx| TokenId::new(idx as u32))
443        }
444
445        fn token_text(&self, token_id: TokenId) -> Option<&str> {
446            self.strings
447                .get(token_id.get() as usize)
448                .map(|value| value.as_str())
449        }
450
451        fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
452            Ok(String::new())
453        }
454
455        fn info(&self) -> TokenizerInfo {
456            TokenizerInfo {
457                tokenizer_type: TokenizerType::Custom,
458                vocab_size: self.strings.len(),
459                special_tokens: self.special.clone(),
460                supports_incremental: true,
461                supports_chat_template: false,
462                max_token_length: Some(1),
463                model_name: Some("tiny-tokenizer-test".into()),
464            }
465        }
466    }
467
468    #[test]
469    fn digits_only_allows_digits_at_start() {
470        let p = processor(r"[0-9]+");
471        let mut logits = vec![0.0f32; 257];
472        p.mask_logits(&mut logits);
473        for b in 0u8..=255 {
474            let expected_allowed = b.is_ascii_digit();
475            let got = logits[b as usize].is_finite();
476            assert_eq!(
477                got, expected_allowed,
478                "byte {b:?} ({}): expected allowed={expected_allowed}, got={got}",
479                b as char
480            );
481        }
482        // EOS is NOT yet allowed (pattern requires >=1 digit).
483        assert!(logits[256].is_infinite() && logits[256].is_sign_negative());
484    }
485
486    #[test]
487    fn digits_only_allows_eos_after_a_digit() {
488        let p = processor(r"[0-9]+");
489        p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
490        let mut logits = vec![0.0f32; 257];
491        p.mask_logits(&mut logits);
492        assert!(
493            logits[256].is_finite(),
494            "EOS should be allowed after a digit"
495        );
496        assert!(
497            logits[b'7' as usize].is_finite(),
498            "another digit still allowed"
499        );
500        assert!(logits[b'a' as usize].is_infinite(), "alpha still forbidden");
501    }
502
503    #[test]
504    fn dead_state_disables_masking_instead_of_forcing_eos() {
505        let p = processor(r"[0-9]+");
506        // Feed an invalid token ("a") — DFA dies.
507        p.advance_with_tokens(&[TokenId::new(b'a' as u32)]);
508        let mut logits = vec![0.0f32; 257];
509        logits[256] = 7.0;
510        p.mask_logits(&mut logits);
511        assert_eq!(logits[256], 7.0, "EOS should not be forced");
512        assert_eq!(logits[b'a' as usize], 0.0, "mask should be disabled");
513    }
514
515    #[test]
516    fn no_extension_before_accept_uses_best_non_eos_token() {
517        let tok: Arc<dyn Tokenizer> = Arc::new(TinyTokenizer::new(vec!["a", "b", "</s>"], 2));
518        let p = RegexGuidedProcessor::new("z", tok, Some(TokenId::new(2))).unwrap();
519        let mut logits = vec![0.1, 3.0, 99.0];
520        p.mask_logits(&mut logits);
521
522        assert!(
523            logits[0].is_infinite(),
524            "lower non-EOS token should be masked"
525        );
526        assert!(logits[1].is_finite(), "best non-EOS token should remain");
527        assert!(
528            logits[2].is_infinite(),
529            "EOS must stay masked before the regex accepts"
530        );
531    }
532
533    #[test]
534    fn reset_restores_initial_state() {
535        let p = processor(r"[0-9]+");
536        p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
537        assert!(p.can_accept());
538        p.reset().unwrap();
539        assert!(!p.can_accept(), "fresh state should not accept empty input");
540    }
541
542    #[test]
543    fn hex_prefix_pattern() {
544        let p = processor(r"0x[0-9a-f]+");
545        let mut logits = vec![0.0f32; 257];
546        p.mask_logits(&mut logits);
547        assert!(logits[b'0' as usize].is_finite(), "'0' starts the pattern");
548        assert!(logits[b'1' as usize].is_infinite(), "'1' can't start");
549        p.advance_with_tokens(&[TokenId::new(b'0' as u32), TokenId::new(b'x' as u32)]);
550        let mut logits = vec![0.0f32; 257];
551        p.mask_logits(&mut logits);
552        assert!(logits[b'a' as usize].is_finite());
553        assert!(logits[b'f' as usize].is_finite());
554        assert!(logits[b'g' as usize].is_infinite());
555        assert!(logits[b'9' as usize].is_finite());
556    }
557}