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                },
325                byte_strings,
326            }
327        }
328    }
329
330    impl Tokenizer for ByteTokenizer {
331        fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
332            Ok(text.bytes().map(|b| TokenId::new(b as u32)).collect())
333        }
334        fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
335            let mut out = String::new();
336            for t in tokens {
337                let idx = t.get() as usize;
338                if idx < 256 {
339                    out.push(idx as u8 as char);
340                }
341            }
342            Ok(out)
343        }
344        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
345            self.decode(&[next], false)
346        }
347        fn vocab_size(&self) -> usize {
348            257
349        }
350        fn special_tokens(&self) -> &SpecialTokens {
351            &self.special
352        }
353        fn token_id(&self, text: &str) -> Option<TokenId> {
354            if text.len() == 1 {
355                Some(TokenId::new(text.bytes().next().unwrap() as u32))
356            } else {
357                None
358            }
359        }
360        fn token_text(&self, token_id: TokenId) -> Option<&str> {
361            self.byte_strings
362                .get(token_id.get() as usize)
363                .map(|s| s.as_str())
364        }
365        fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
366            Ok(String::new())
367        }
368        fn info(&self) -> TokenizerInfo {
369            TokenizerInfo {
370                tokenizer_type: TokenizerType::Custom,
371                vocab_size: 257,
372                special_tokens: self.special.clone(),
373                supports_incremental: true,
374                supports_chat_template: false,
375                max_token_length: Some(1),
376                model_name: Some("byte-tokenizer-test".into()),
377            }
378        }
379    }
380
381    fn processor(pattern: &str) -> RegexGuidedProcessor {
382        let tok: Arc<dyn Tokenizer> = Arc::new(ByteTokenizer::new());
383        RegexGuidedProcessor::new(pattern, tok, Some(TokenId::new(256))).unwrap()
384    }
385
386    struct TinyTokenizer {
387        special: SpecialTokens,
388        strings: Vec<String>,
389    }
390
391    impl TinyTokenizer {
392        fn new(strings: Vec<&str>, eos: u32) -> Self {
393            Self {
394                special: SpecialTokens {
395                    bos_token: None,
396                    eos_token: Some(TokenId::new(eos)),
397                    unk_token: None,
398                    pad_token: None,
399                    sep_token: None,
400                    cls_token: None,
401                    mask_token: None,
402                },
403                strings: strings.into_iter().map(str::to_string).collect(),
404            }
405        }
406    }
407
408    impl Tokenizer for TinyTokenizer {
409        fn encode(&self, text: &str, _add_special: bool) -> Result<Vec<TokenId>> {
410            Ok(text
411                .chars()
412                .filter_map(|ch| self.token_id(&ch.to_string()))
413                .collect())
414        }
415
416        fn decode(&self, tokens: &[TokenId], _skip_special: bool) -> Result<String> {
417            Ok(tokens
418                .iter()
419                .filter_map(|token| self.token_text(*token))
420                .collect::<Vec<_>>()
421                .join(""))
422        }
423
424        fn decode_incremental(&self, _prev: &[TokenId], next: TokenId) -> Result<String> {
425            self.decode(&[next], false)
426        }
427
428        fn vocab_size(&self) -> usize {
429            self.strings.len()
430        }
431
432        fn special_tokens(&self) -> &SpecialTokens {
433            &self.special
434        }
435
436        fn token_id(&self, text: &str) -> Option<TokenId> {
437            self.strings
438                .iter()
439                .position(|value| value == text)
440                .map(|idx| TokenId::new(idx as u32))
441        }
442
443        fn token_text(&self, token_id: TokenId) -> Option<&str> {
444            self.strings
445                .get(token_id.get() as usize)
446                .map(|value| value.as_str())
447        }
448
449        fn apply_chat_template(&self, _messages: &[ChatMessage]) -> Result<String> {
450            Ok(String::new())
451        }
452
453        fn info(&self) -> TokenizerInfo {
454            TokenizerInfo {
455                tokenizer_type: TokenizerType::Custom,
456                vocab_size: self.strings.len(),
457                special_tokens: self.special.clone(),
458                supports_incremental: true,
459                supports_chat_template: false,
460                max_token_length: Some(1),
461                model_name: Some("tiny-tokenizer-test".into()),
462            }
463        }
464    }
465
466    #[test]
467    fn digits_only_allows_digits_at_start() {
468        let p = processor(r"[0-9]+");
469        let mut logits = vec![0.0f32; 257];
470        p.mask_logits(&mut logits);
471        for b in 0u8..=255 {
472            let expected_allowed = b.is_ascii_digit();
473            let got = logits[b as usize].is_finite();
474            assert_eq!(
475                got, expected_allowed,
476                "byte {b:?} ({}): expected allowed={expected_allowed}, got={got}",
477                b as char
478            );
479        }
480        // EOS is NOT yet allowed (pattern requires >=1 digit).
481        assert!(logits[256].is_infinite() && logits[256].is_sign_negative());
482    }
483
484    #[test]
485    fn digits_only_allows_eos_after_a_digit() {
486        let p = processor(r"[0-9]+");
487        p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
488        let mut logits = vec![0.0f32; 257];
489        p.mask_logits(&mut logits);
490        assert!(
491            logits[256].is_finite(),
492            "EOS should be allowed after a digit"
493        );
494        assert!(
495            logits[b'7' as usize].is_finite(),
496            "another digit still allowed"
497        );
498        assert!(logits[b'a' as usize].is_infinite(), "alpha still forbidden");
499    }
500
501    #[test]
502    fn dead_state_disables_masking_instead_of_forcing_eos() {
503        let p = processor(r"[0-9]+");
504        // Feed an invalid token ("a") — DFA dies.
505        p.advance_with_tokens(&[TokenId::new(b'a' as u32)]);
506        let mut logits = vec![0.0f32; 257];
507        logits[256] = 7.0;
508        p.mask_logits(&mut logits);
509        assert_eq!(logits[256], 7.0, "EOS should not be forced");
510        assert_eq!(logits[b'a' as usize], 0.0, "mask should be disabled");
511    }
512
513    #[test]
514    fn no_extension_before_accept_uses_best_non_eos_token() {
515        let tok: Arc<dyn Tokenizer> = Arc::new(TinyTokenizer::new(vec!["a", "b", "</s>"], 2));
516        let p = RegexGuidedProcessor::new("z", tok, Some(TokenId::new(2))).unwrap();
517        let mut logits = vec![0.1, 3.0, 99.0];
518        p.mask_logits(&mut logits);
519
520        assert!(
521            logits[0].is_infinite(),
522            "lower non-EOS token should be masked"
523        );
524        assert!(logits[1].is_finite(), "best non-EOS token should remain");
525        assert!(
526            logits[2].is_infinite(),
527            "EOS must stay masked before the regex accepts"
528        );
529    }
530
531    #[test]
532    fn reset_restores_initial_state() {
533        let p = processor(r"[0-9]+");
534        p.advance_with_tokens(&[TokenId::new(b'3' as u32)]);
535        assert!(p.can_accept());
536        p.reset().unwrap();
537        assert!(!p.can_accept(), "fresh state should not accept empty input");
538    }
539
540    #[test]
541    fn hex_prefix_pattern() {
542        let p = processor(r"0x[0-9a-f]+");
543        let mut logits = vec![0.0f32; 257];
544        p.mask_logits(&mut logits);
545        assert!(logits[b'0' as usize].is_finite(), "'0' starts the pattern");
546        assert!(logits[b'1' as usize].is_infinite(), "'1' can't start");
547        p.advance_with_tokens(&[TokenId::new(b'0' as u32), TokenId::new(b'x' as u32)]);
548        let mut logits = vec![0.0f32; 257];
549        p.mask_logits(&mut logits);
550        assert!(logits[b'a' as usize].is_finite());
551        assert!(logits[b'f' as usize].is_finite());
552        assert!(logits[b'g' as usize].is_infinite());
553        assert!(logits[b'9' as usize].is_finite());
554    }
555}