Skip to main content

dynamo_tokenizers/
lib.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4pub mod basetenkenizer;
5pub mod cache;
6pub mod fastokens;
7pub mod hf;
8pub mod tiktoken;
9
10// TODO: Add tokenizer benchmarks
11// TODO: Enable README.md as a module doc
12// #[doc = include_str!("../README.md")]
13
14use std::hash::{DefaultHasher, Hash, Hasher};
15use std::sync::Arc;
16use std::{fs::File, io::BufReader, ops::Deref, path::Path};
17
18use anyhow::Context as _;
19pub use anyhow::{Error, Result};
20
21pub use basetenkenizer::BasetenTokenizer;
22pub use cache::{
23    CacheTokenUsage, CacheTokenUsageFn, CachedTokenizer, L1CacheStats, SharedTokenizerCache,
24    SharedTokenizerCacheStats,
25};
26pub use fastokens::{FastTikTokenTokenizer, FastTokenizer};
27pub use hf::HuggingFaceTokenizer;
28pub use tiktoken::TikTokenTokenizer;
29pub use traits::DecodeResult;
30
31pub type TokenIdType = u32;
32
33/// A rendered prompt segment with an explicit trust boundary for special tokens.
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub struct EncodeSegment<'a> {
36    pub text: &'a str,
37    /// Recognize added/control tokens in this trusted renderer output.
38    ///
39    /// Set this to `false` for user, tool, and attribute content so text that
40    /// resembles a control token is encoded as ordinary text.
41    pub allow_special: bool,
42}
43
44impl<'a> EncodeSegment<'a> {
45    pub const fn new(text: &'a str, allow_special: bool) -> Self {
46        Self {
47            text,
48            allow_special,
49        }
50    }
51
52    pub const fn ordinary(text: &'a str) -> Self {
53        Self::new(text, false)
54    }
55
56    pub const fn control(text: &'a str) -> Self {
57        Self::new(text, true)
58    }
59}
60
61/// Represents the type of tokenizer being used
62#[derive(Debug)]
63pub enum TokenizerType {
64    HuggingFace(String),
65    TikToken(String),
66}
67
68/// character offsets in the original text
69pub type Offsets = (usize, usize);
70
71/// Contains the results of tokenizing text: token IDs, string tokens, and their spans
72#[derive(Debug, Clone)]
73pub enum Encoding {
74    /// Hugging Face
75    Hf(Box<tokenizers::tokenizer::Encoding>),
76    /// Sentence Piece
77    Sp(Vec<TokenIdType>),
78}
79
80impl Encoding {
81    pub fn token_ids(&self) -> &[u32] {
82        match self {
83            Encoding::Hf(inner) => inner.get_ids(),
84            Encoding::Sp(inner) => inner,
85        }
86    }
87}
88
89impl Hash for Encoding {
90    fn hash<H: Hasher>(&self, state: &mut H) {
91        self.token_ids().hash(state);
92    }
93}
94
95pub mod traits {
96    use super::*;
97
98    pub trait Encoder: Send + Sync {
99        fn encode(&self, input: &str) -> Result<Encoding>;
100        fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>>;
101
102        /// Encode Kimi K3-style renderer segments while preserving trusted
103        /// control-token and untrusted content boundaries.
104        ///
105        /// Each segment controls whether added/control tokens are recognized
106        /// through [`EncodeSegment::allow_special`]. This prevents text in user,
107        /// tool, or attribute content from becoming structural when it happens
108        /// to resemble a control token.
109        ///
110        /// The Baseten backend preserves legacy tiktoken behavior by splitting
111        /// each segment into chunks of at most 400,000 characters and splitting
112        /// whitespace/non-whitespace runs at 25,000 characters. Independent
113        /// chunks are encoded through the Rayon thread pool, then concatenated
114        /// in input order. Tokenizer post-processing is applied once after all
115        /// segment IDs have been joined.
116        ///
117        /// Backends must not implement this by flattening the segments first,
118        /// because that discards the special-token trust boundary.
119        fn encode_segments(&self, _segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
120            Err(Error::msg(
121                "tokenizer backend does not support segmented encoding",
122            ))
123        }
124    }
125
126    /// Result of decoding token IDs to text.
127    ///
128    /// Distinguishes between fully valid UTF-8 output and output that contains
129    /// trailing incomplete multi-byte sequences (represented as U+FFFD).
130    /// This lets callers like `DecodeStream::step()` decide whether to emit or
131    /// buffer without resorting to hardcoded replacement-character string checks.
132    #[derive(Debug, Clone, PartialEq, Eq, strum::EnumIs)]
133    pub enum DecodeResult {
134        /// No trailing incomplete multi-byte sequences (text does not end with U+FFFD).
135        /// Note: the string may still contain *interior* U+FFFD characters from
136        /// mid-stream invalid byte sequences; only trailing status is tracked here.
137        Complete(String),
138        /// The decoded string ends with U+FFFD, indicating incomplete trailing
139        /// multi-byte bytes that may be completed by subsequent tokens.
140        Partial(String),
141    }
142
143    impl DecodeResult {
144        /// Returns a reference to the inner string.
145        pub fn as_str(&self) -> &str {
146            match self {
147                DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
148            }
149        }
150
151        /// Construct from a decoded string: `Partial` if it ends with U+FFFD, else `Complete`.
152        pub fn from_decoded(text: String) -> Self {
153            if text.ends_with('\u{FFFD}') {
154                DecodeResult::Partial(text)
155            } else {
156                DecodeResult::Complete(text)
157            }
158        }
159    }
160
161    impl From<String> for DecodeResult {
162        fn from(text: String) -> Self {
163            DecodeResult::from_decoded(text)
164        }
165    }
166
167    impl From<DecodeResult> for String {
168        fn from(result: DecodeResult) -> Self {
169            match result {
170                DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
171            }
172        }
173    }
174
175    /// Implementations must ensure that partial multi-byte sequences produce U+FFFD
176    /// (`\u{FFFD}`) in the output rather than returning `Err`. This is commonly achieved
177    /// via `String::from_utf8_lossy` (tiktoken) or library-internal byte-fallback handling
178    /// (HuggingFace). `DecodeStream::step()` relies on `DecodeResult::Partial` to detect
179    /// incomplete sequences and buffer tokens until the full character arrives.
180    pub trait Decoder: Send + Sync {
181        /// Whether appending tokens can still reinterpret the decoded suffix.
182        /// Callers must buffer an unstable suffix until a boundary or end of input.
183        fn has_unstable_suffix(
184            &self,
185            _token_ids: &[TokenIdType],
186            _skip_special_tokens: bool,
187        ) -> bool {
188            false
189        }
190
191        fn decode(
192            &self,
193            token_ids: &[TokenIdType],
194            skip_special_tokens: bool,
195        ) -> Result<DecodeResult>;
196    }
197
198    pub trait Tokenizer: Encoder + Decoder {
199        /// Validate that this tokenizer can be safely wrapped in the prefix cache.
200        ///
201        /// Implementations must explicitly opt in by returning `Ok(())`, or
202        /// return an error explaining why the tokenizer is incompatible.
203        fn validate_prefix_cache(&self) -> Result<()> {
204            Err(Error::msg("tokenizer does not support prefix caching"))
205        }
206
207        /// Apply construction-time [`TokenizerOptions`].
208        ///
209        /// The default implementation ignores the options — correct for
210        /// tokenizers with no applicable option.
211        fn with_options(self, options: TokenizerOptions) -> Self
212        where
213            Self: Sized,
214        {
215            let _ = options;
216            self
217        }
218        /// Vocabulary cardinality including added tokens, when the backend
219        /// can expose one. `None` for backends without a bounded id space or
220        /// vocabulary introspection.
221        fn vocab_size(&self) -> Option<usize> {
222            None
223        }
224
225        /// Resolve a token string to its vocabulary id, when the backend
226        /// supports lookup. `Ok(None)` when the token is not in the
227        /// vocabulary. `Err` when this backend cannot do id lookup at all —
228        /// kept distinct from a genuine vocabulary miss so callers can tell
229        /// "unsupported" from "looked up, not found".
230        ///
231        /// Defaults to unsupported, matching `encode_segments` and
232        /// `validate_prefix_cache`: only backends that can perform lookup
233        /// need to override it.
234        fn token_to_id(&self, _token: &str) -> Result<Option<TokenIdType>> {
235            Err(Error::msg(
236                "tokenizer backend does not support token lookup",
237            ))
238        }
239
240        /// Ids of added tokens marked special (e.g. BOS/EOS/PAD/control
241        /// tokens), as distinct from ordinary vocabulary tokens. `Ok(vec![])`
242        /// for backends that genuinely have none, or that do not distinguish
243        /// special from ordinary vocabulary ids. `Err` when the backend
244        /// cannot enumerate added tokens at all — an empty `Vec` alone
245        /// cannot carry that distinction.
246        ///
247        /// Defaults to unsupported, matching `token_to_id`: only backends
248        /// that can enumerate special ids need to override it.
249        fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
250            Err(Error::msg(
251                "tokenizer backend does not support special token enumeration",
252            ))
253        }
254
255        /// Count of special tokens `encode`'s `add_special_tokens: true`
256        /// path would add to a bare encoding (e.g. BOS/EOS), available
257        /// without performing an encode. `Ok(0)` for backends that
258        /// genuinely add none. `Err` when the backend cannot determine this
259        /// at all — a plain `0` alone cannot carry that distinction, and
260        /// this value feeds token-budget accounting where a silent `0`
261        /// would under-count rather than fail loudly.
262        ///
263        /// Defaults to unsupported, matching `token_to_id`: only backends
264        /// that can determine this count need to override it.
265        fn num_special_tokens_added(&self) -> Result<usize> {
266            Err(Error::msg(
267                "tokenizer backend does not support special token accounting",
268            ))
269        }
270        // fn make_unique_clone(&self) -> Box<dyn Tokenizer>;
271    }
272}
273
274pub fn file_json_field<T: serde::de::DeserializeOwned>(
275    json_file_path: &Path,
276    field_name: &str,
277) -> anyhow::Result<T> {
278    let file = File::open(json_file_path)
279        .with_context(|| format!("Failed to open file: {:?}", json_file_path))?;
280    let reader = BufReader::new(file);
281
282    let json_data: serde_json::Value = serde_json::from_reader(reader)
283        .with_context(|| format!("Failed to parse JSON from file: {:?}", json_file_path))?;
284
285    let map = json_data.as_object().ok_or_else(|| {
286        anyhow::anyhow!("JSON root is not an object in file: {:?}", json_file_path)
287    })?;
288
289    let field_value = map.get(field_name).ok_or_else(|| {
290        anyhow::anyhow!(
291            "Field '{}' not found in JSON file: {:?}",
292            field_name,
293            json_file_path
294        )
295    })?;
296
297    serde_json::from_value(field_value.clone()).with_context(|| {
298        format!(
299            "Failed to deserialize field '{}' (value: {:?}) to the expected type from file: {:?}",
300            field_name, field_value, json_file_path
301        )
302    })
303}
304
305pub fn log_json_err(filename: &str, json: &str, err: &serde_json::Error) {
306    const ERROR_PREFIX: &str = ">>     ";
307
308    if !(err.is_syntax() || err.is_data()) {
309        return;
310    }
311
312    let line = err.line().saturating_sub(1);
313    let column = err.column().saturating_sub(1);
314
315    let json_lines: Vec<&str> = json.lines().collect();
316    if json_lines.is_empty() {
317        tracing::error!("JSON parsing error in {filename}: File is empty.");
318        return;
319    }
320
321    let start_index = line.saturating_sub(2);
322    let end_index = line.saturating_add(3).min(json_lines.len());
323
324    let mut context_lines: Vec<String> = (start_index..end_index)
325        .map(|i| {
326            if i == line {
327                format!("{ERROR_PREFIX}{}", json_lines[i])
328            } else {
329                format!("{:06} {}", i + 1, json_lines[i])
330            }
331        })
332        .collect();
333
334    let col_indicator = "_".to_string().repeat(column + ERROR_PREFIX.len()) + "^";
335    let error_in_context_idx = line - start_index;
336    if error_in_context_idx < context_lines.len() {
337        context_lines.insert(error_in_context_idx + 1, col_indicator);
338    }
339
340    tracing::error!(
341        "JSON parsing error in {filename}: Line {}, column {}:\n{}",
342        err.line(),
343        err.column(),
344        context_lines.join("\n")
345    );
346}
347
348impl Encoding {
349    pub fn get_hash(&self) -> u64 {
350        let mut hasher = DefaultHasher::new();
351        self.hash(&mut hasher);
352        hasher.finish()
353    }
354}
355
356/// Construction options for [`Tokenizer::from_file_with_options`] /
357/// [`create_tokenizer_from_file`], applied to concrete tokenizers via
358/// [`traits::Tokenizer::with_options`].
359#[derive(Debug, Clone, Copy, Default)]
360pub struct TokenizerOptions {
361    /// Ask the tokenizer to add its declared special tokens (e.g. BOS/EOS via
362    /// its post-processor) during `encode`, `encode_batch`, and supported
363    /// `encode_segments` calls.
364    /// Defaults to `false`, the historical behavior.
365    ///
366    /// Applicable to Hugging Face and Baseten tokenizers.
367    pub add_special_tokens: bool,
368}
369
370/// Main tokenizer wrapper that provides a unified interface for different tokenizer implementations
371#[derive(Clone)]
372pub struct Tokenizer(Arc<dyn traits::Tokenizer>);
373
374impl Tokenizer {
375    pub fn from_file(file_path: &str) -> Result<Tokenizer> {
376        Ok(Tokenizer(create_tokenizer_from_file(file_path)?))
377    }
378
379    pub fn from_file_with_options(file_path: &str, options: TokenizerOptions) -> Result<Tokenizer> {
380        Ok(Tokenizer(create_tokenizer_from_file_with_options(
381            file_path, options,
382        )?))
383    }
384
385    /// Create a stateful sequence object for decoding token_ids into text.
386    /// Append the result of [`DecodeStream::finish`] when input ends; `step` may
387    /// retain a trailing byte-fallback run even when it currently forms valid UTF-8.
388    pub fn decode_stream(
389        &self,
390        prompt_token_ids: &[TokenIdType],
391        skip_special_tokens: bool,
392    ) -> DecodeStream {
393        DecodeStream::new(self.0.clone(), prompt_token_ids, skip_special_tokens)
394    }
395}
396
397impl Deref for Tokenizer {
398    type Target = Arc<dyn traits::Tokenizer>;
399
400    fn deref(&self) -> &Self::Target {
401        &self.0
402    }
403}
404
405impl From<Arc<dyn traits::Tokenizer>> for Tokenizer {
406    fn from(tokenizer: Arc<dyn traits::Tokenizer>) -> Self {
407        Tokenizer(tokenizer)
408    }
409}
410
411impl<T> From<Arc<T>> for Tokenizer
412where
413    T: traits::Tokenizer + 'static, // 'static is required to ensure T can be safely put into an Arc
414{
415    fn from(tokenizer: Arc<T>) -> Self {
416        Tokenizer(tokenizer)
417    }
418}
419
420/// Create a tokenizer from a file path to a tokenizer file.
421/// The file extension is used to determine the tokenizer type.
422/// Supported file types are:
423/// - json: HuggingFace tokenizer
424/// - model, tiktoken: tiktoken BPE tokenizer (requires `config.json` with a supported
425///   `model_type` in the same directory; currently: kimi, kimi_k2, kimi_k25, kimi_k3)
426pub fn create_tokenizer_from_file(file_path: &str) -> Result<Arc<dyn traits::Tokenizer>> {
427    create_tokenizer_from_file_with_options(file_path, Default::default())
428}
429
430/// Create a tokenizer from a file path to a tokenizer file with additional tokenizer option.
431/// The file extension is used to determine the tokenizer type.
432/// Supported file types are:
433/// - json: HuggingFace tokenizer
434/// - model, tiktoken: tiktoken BPE tokenizer (requires `config.json` with a supported
435///   `model_type` in the same directory; currently: kimi, kimi_k2, kimi_k25, kimi_k3)
436pub fn create_tokenizer_from_file_with_options(
437    file_path: &str,
438    options: TokenizerOptions,
439) -> Result<Arc<dyn traits::Tokenizer>> {
440    use traits::Tokenizer as _;
441
442    let path = Path::new(file_path);
443    let extension = path
444        .extension()
445        .and_then(std::ffi::OsStr::to_str)
446        .ok_or_else(|| Error::msg("Failed to read file extension".to_string()))?;
447
448    match extension {
449        "json" => {
450            let tokenizer = HuggingFaceTokenizer::from_file(file_path)?.with_options(options);
451            Ok(Arc::new(tokenizer))
452        }
453        "model" | "tiktoken" => {
454            let tokenizer = TikTokenTokenizer::from_file_auto(file_path)?.with_options(options);
455            Ok(Arc::new(tokenizer))
456        }
457        _ => Err(Error::msg(format!(
458            "Unsupported tokenizer file type: .{extension}"
459        ))),
460    }
461}
462
463// With incremental detokenization, we need to consider the final context tokens when handling the initial decode tokens.
464// This is the initial offset from the end of the context that we start decoding from.
465// Both Huggingface TGI and vLLM use this same value.
466// See: https://github.com/huggingface/text-generation-inference/blob/24c2bff65924801ddf90fa24fcc72752d4f45538/server/text_generation_server/models/mamba.py#L169
467// and https://github.com/vllm-project/vllm/blob/da2705198fa19030a25d0bea437f7be6547d47d4/vllm/transformers_utils/detokenizer_utils.py#L51
468const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
469
470/// DecodeStream will keep the state necessary to produce individual chunks of
471/// strings given an input stream of token_ids.
472///
473/// This is necessary because decoding in general cannot achieve that since strings
474/// depend on surrounding ids to provide a valid string. Typically stripping extra spaces.
475pub struct DecodeStream {
476    /// The tokenizer used to decode token_ids
477    tokenizer: Arc<dyn traits::Tokenizer>,
478
479    skip_special_tokens: bool,
480    /// A temporary buffer of the necessary token_ids needed
481    /// to produce valid string chunks.
482    /// This typically contains 3 parts:
483    ///  - read
484    ///  - prefix
485    ///  - rest
486    ///
487    /// Read is the bit necessary to surround the prefix
488    /// so decoding the whole ids produces a valid prefix.
489    /// Prefix is the previously produced string, kept around to trim off of
490    /// the next valid chunk
491    all_token_ids: Vec<u32>,
492
493    prefix_offset: usize,
494
495    read_offset: usize,
496
497    /// Whether any generated text has already been returned to the caller.
498    has_emitted: bool,
499}
500
501impl DecodeStream {
502    pub fn new(
503        tokenizer: Arc<dyn traits::Tokenizer>,
504        prompt_token_ids: &[TokenIdType],
505        skip_special_tokens: bool,
506    ) -> Self {
507        // Earlier prompt tokens are never read by incremental decoding. Keep the
508        // same context suffix and rebase its offsets before copying it.
509        let context_start = prompt_token_ids
510            .len()
511            .saturating_sub(INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET);
512        let prompt_token_ids = prompt_token_ids[context_start..].to_vec();
513        let num_input_tokens = prompt_token_ids.len();
514        Self {
515            tokenizer,
516            skip_special_tokens,
517            all_token_ids: prompt_token_ids,
518            prefix_offset: 0,
519            read_offset: num_input_tokens,
520            has_emitted: false,
521        }
522    }
523
524    /// Step appends a token_id to the internal state and tries to produce a text chunk.
525    ///
526    /// Implementation directly copied from Huggingface's TGI:
527    /// https://github.com/huggingface/text-generation-inference/blob/24c2bff65924801ddf90fa24fcc72752d4f45538/server/text_generation_server/models/model.py#L144
528    ///
529    /// Returning `None` means the given id is not enough to produce a chunk.
530    /// This typically happens with `byte_fallback` options where some tokens do not
531    /// represent valid UTF-8, and only follow-up token_ids will help produce
532    /// a valid chunk.
533    ///
534    /// An error is terminal for this stream because the token may already have
535    /// been appended to its internal state.
536    pub fn step(&mut self, id: u32) -> Result<Option<String>> {
537        self.all_token_ids.push(id);
538
539        if self.tokenizer.has_unstable_suffix(
540            &self.all_token_ids[self.read_offset..],
541            self.skip_special_tokens,
542        ) {
543            return Ok(None);
544        }
545        self.decode_pending(false)
546    }
547
548    /// Decode the remaining suffix once no further tokens can reinterpret it.
549    /// Call this at end-of-input and append its output after all `step` chunks.
550    /// Dropping the stream without finishing discards any buffered text.
551    pub fn finish(&mut self) -> Result<Option<String>> {
552        self.decode_pending(true)
553    }
554
555    fn decode_pending(&mut self, finishing: bool) -> Result<Option<String>> {
556        if self.read_offset == self.all_token_ids.len() {
557            return Ok(None);
558        }
559
560        let prefix_text: String = self
561            .tokenizer
562            .decode(
563                &self.all_token_ids[self.prefix_offset..self.read_offset],
564                self.skip_special_tokens,
565            )?
566            .into();
567
568        let new_result = self.tokenizer.decode(
569            &self.all_token_ids[self.prefix_offset..],
570            self.skip_special_tokens,
571        )?;
572
573        let new_text = new_result.as_str();
574        let is_partial = new_result.is_partial() && !finishing;
575
576        // Once generated text has been returned, decoding must remain append-only.
577        // A complete rewrite cannot be repaired after the caller has seen the old text.
578        if self.has_emitted && !is_partial && !new_text.starts_with(prefix_text.as_str()) {
579            return Err(Error::msg(
580                "incremental decoding rewrote already emitted text",
581            ));
582        }
583
584        if new_text.len() > prefix_text.len() && !is_partial {
585            let requested_split = prefix_text.len();
586            let split = if self.has_emitted {
587                // starts_with() guarantees this is a character boundary and preserves
588                // everything already returned to the caller.
589                requested_split
590            } else {
591                // Rewinding into prompt context is safe because prompt text was not
592                // returned by this stream.
593                new_text.floor_char_boundary(requested_split)
594            };
595
596            let emitted = new_text[split..].to_string();
597
598            self.prefix_offset = self.read_offset;
599            self.read_offset = self.all_token_ids.len();
600            self.has_emitted = true;
601
602            Ok(Some(emitted))
603        } else {
604            Ok(None)
605        }
606    }
607}
608
609#[cfg(test)]
610mod decode_stream_unicode_tests {
611    use super::{DecodeResult, DecodeStream, Encoding, Result, TokenIdType};
612    use std::sync::Arc;
613
614    struct RewritingTokenizer {
615        prefix_text: &'static str,
616        rewritten_text: &'static str,
617        prefix_is_partial: bool,
618    }
619
620    impl super::traits::Encoder for RewritingTokenizer {
621        fn encode(&self, _input: &str) -> Result<Encoding> {
622            Ok(Encoding::Sp(vec![]))
623        }
624
625        fn encode_batch(&self, _inputs: &[&str]) -> Result<Vec<Encoding>> {
626            Ok(vec![])
627        }
628    }
629
630    impl super::traits::Decoder for RewritingTokenizer {
631        fn decode(
632            &self,
633            token_ids: &[TokenIdType],
634            _skip_special_tokens: bool,
635        ) -> Result<DecodeResult> {
636            let result = match token_ids.len() {
637                1 if self.prefix_is_partial => DecodeResult::Partial(self.prefix_text.to_string()),
638                1 => DecodeResult::Complete(self.prefix_text.to_string()),
639                2 => DecodeResult::Complete(self.rewritten_text.to_string()),
640                _ => DecodeResult::Complete(String::new()),
641            };
642            Ok(result)
643        }
644    }
645
646    impl super::traits::Tokenizer for RewritingTokenizer {}
647
648    #[test]
649    fn prompt_suffix_preserves_decode_inputs_and_output() {
650        let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(
651            super::HuggingFaceTokenizer::from_file(concat!(
652                env!("CARGO_MANIFEST_DIR"),
653                "/tests/data/minimal-bpe/tokenizer.json"
654            ))
655            .unwrap(),
656        );
657        for prompt_len in [0usize, 1, 4, 5, 6, 775_168] {
658            let prompt: Vec<_> = (0..prompt_len)
659                .map(|index| 1 + (index % 22) as u32)
660                .collect();
661            for skip_special_tokens in [false, true] {
662                let mut stream = DecodeStream::new(tokenizer.clone(), &prompt, skip_special_tokens);
663                // Reconstruct the previous full-prompt state as a parity oracle.
664                let mut full = DecodeStream {
665                    tokenizer: tokenizer.clone(),
666                    skip_special_tokens,
667                    all_token_ids: prompt.clone(),
668                    prefix_offset: prompt_len
669                        .saturating_sub(super::INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET),
670                    read_offset: prompt_len,
671                    has_emitted: false,
672                };
673                assert!(stream.all_token_ids.capacity() <= 5);
674                for token in [5, 9, 12, 12, 13, 0, 1, 17, 13, 14, 12, 8, 2] {
675                    assert_eq!(
676                        &stream.all_token_ids[stream.prefix_offset..stream.read_offset],
677                        &full.all_token_ids[full.prefix_offset..full.read_offset],
678                    );
679                    assert_eq!(
680                        &stream.all_token_ids[stream.prefix_offset..],
681                        &full.all_token_ids[full.prefix_offset..],
682                    );
683                    let actual = stream.step(token).map_err(|error| error.to_string());
684                    let expected = full.step(token).map_err(|error| error.to_string());
685                    assert_eq!(actual, expected, "prompt length {prompt_len}");
686                    if actual.is_err() {
687                        break;
688                    }
689                }
690            }
691        }
692    }
693
694    #[test]
695    fn allows_boundary_recovery_before_generated_text_is_emitted() {
696        for (prefix_text, rewritten_text, expected) in [
697            ("㺄馉凓鄗\u{FFFD}", "㺄馉凓鄗𫷲", "𫷲"),
698            (
699                "JUnitworkflow Completion intuition\u{FFFD}",
700                "JUnitworkflow Completion intuition𝟙",
701                "𝟙",
702            ),
703        ] {
704            let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
705                prefix_text,
706                rewritten_text,
707                prefix_is_partial: true,
708            });
709            let mut stream = DecodeStream::new(tokenizer, &[1], false);
710
711            assert_eq!(stream.step(2).unwrap(), Some(expected.to_string()));
712        }
713    }
714
715    #[test]
716    fn errors_when_incremental_decode_rewrites_emitted_text() {
717        for (prefix_text, rewritten_text) in [
718            ("abcde", "abcdXY"),
719            ("abcde", "abcdX"),
720            ("abcde", "abcd"),
721            ("㺄馉凓鄗abc", "㺄馉凓鄗𫷲"),
722            (
723                "JUnitworkflow Completion intuitionabc",
724                "JUnitworkflow Completion intuition𝟙",
725            ),
726        ] {
727            let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
728                prefix_text,
729                rewritten_text,
730                prefix_is_partial: false,
731            });
732            let mut stream = DecodeStream::new(tokenizer, &[], false);
733
734            assert_eq!(stream.step(1).unwrap(), Some(prefix_text.to_string()));
735
736            let error = stream.step(2).unwrap_err();
737            assert!(error.to_string().contains("already emitted text"));
738        }
739    }
740
741    #[test]
742    fn emits_suffix_when_incremental_decode_preserves_emitted_text() {
743        let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
744            prefix_text: "abcde",
745            rewritten_text: "abcdeXY",
746            prefix_is_partial: false,
747        });
748        let mut stream = DecodeStream::new(tokenizer, &[], false);
749
750        assert_eq!(stream.step(1).unwrap(), Some("abcde".to_string()));
751        assert_eq!(stream.step(2).unwrap(), Some("XY".to_string()));
752    }
753}
754
755/// Maintains state for an ongoing sequence of tokens and their decoded text
756pub struct Sequence {
757    /// Encodes text -> token_ids
758    tokenizer: Tokenizer,
759
760    /// The current sequence of token ids
761    token_ids: Vec<TokenIdType>,
762
763    /// The position in the current sequence the last decoded token completed
764    prefix_offset: usize,
765
766    /// Current position in the sequence
767    read_offset: usize,
768}
769
770impl std::fmt::Debug for Sequence {
771    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
772        f.debug_struct("Sequence")
773            .field("tokenizer", &"Arc<dyn Tokenizer>")
774            .field(
775                "token_ids",
776                &format_args!("{}", {
777                    let token_ids = self.token_ids();
778                    if token_ids.len() <= 20 {
779                        format!("{:?}", token_ids)
780                    } else {
781                        let first_ten = &token_ids[..10];
782                        let last_ten = &token_ids[token_ids.len() - 10..];
783                        format!("{:?} ... {:?}", first_ten, last_ten)
784                    }
785                }),
786            )
787            .field("prefix_offset", &self.prefix_offset)
788            .field("read_offset", &self.read_offset)
789            .field("token count", &self.token_ids.len())
790            .finish()
791    }
792}
793
794impl Sequence {
795    pub fn new(tokenizer: Tokenizer) -> Self {
796        Self {
797            tokenizer,
798            token_ids: Vec::new(),
799            prefix_offset: 0,
800            read_offset: 0,
801        }
802    }
803
804    pub fn is_empty(&self) -> bool {
805        self.token_ids.is_empty()
806    }
807
808    pub fn len(&self) -> usize {
809        self.token_ids.len()
810    }
811
812    pub fn clear(&mut self) {
813        self.token_ids.clear();
814        self.prefix_offset = 0;
815        self.read_offset = 0;
816    }
817
818    pub fn append_text(&mut self, input: &str) -> Result<()> {
819        // let tokenizer = self.tokenizer.read().map_err(|err| {
820        //     Error::msg(format!("Failed to acquire read lock on tokenizer: {}", err))
821        // })?;
822
823        let encoding = self.tokenizer.encode(input)?;
824        self.token_ids.extend(encoding.token_ids());
825        Ok(())
826    }
827
828    // Based on
829    // https://github.com/huggingface/text-generation-inference/blob/v0.9.4/server/text_generation_server/models/model.py#L62C9-L62C15
830    // under Apache 2.0 license
831    pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
832        self.token_ids.push(token_id);
833        // log::trace!("pushed token_id: {}", token_id);
834
835        let prefix_text: String = self
836            .tokenizer
837            .decode(&self.token_ids[self.prefix_offset..self.read_offset], false)?
838            .into();
839
840        let new_result = self
841            .tokenizer
842            .decode(&self.token_ids[self.prefix_offset..], false)?;
843
844        let new_text = new_result.as_str();
845
846        // if the end character of the previous returned sequence is a multi-byte character
847        // then we can not split the text on that byte offset, so we roll back to the byte offset
848        // of the start of that character
849        let mut prefix_text_len = prefix_text.len();
850        while !new_text.is_char_boundary(prefix_text_len) && prefix_text_len > 0 {
851            prefix_text_len -= 1;
852        }
853        let prefix_text_len = prefix_text_len;
854
855        if new_text.len() > prefix_text.len() {
856            if new_result.is_partial() {
857                return Ok("".to_string());
858            } else {
859                // shift and update the state
860                let new_text = new_text[prefix_text_len..]
861                    .to_string()
862                    .replace('\u{FFFD}', "");
863                self.prefix_offset = self.read_offset;
864                self.read_offset = self.token_ids.len();
865                return Ok(new_text);
866            }
867        }
868
869        Ok("".to_string())
870    }
871
872    pub fn tokenizer(&self) -> Tokenizer {
873        self.tokenizer.clone()
874    }
875
876    pub fn token_ids(&self) -> &[TokenIdType] {
877        &self.token_ids
878    }
879
880    pub fn text(&self) -> Result<String> {
881        // let tokenizer = self.tokenizer.read().map_err(|err| {
882        //     Error::msg(format!("Failed to acquire read lock on tokenizer: {}", err))
883        // })?;
884        Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
885    }
886}
887
888/// The output conditions/values of a SequenceDecoder::add_token_id operation.
889/// Result of decoding a token, indicating whether text was produced or a stop condition was met
890pub enum SequenceDecoderOutput {
891    /// The text for the appended token_id
892    Text(String),
893
894    /// A sequence of token_ids has been partially matched a stop sequence, so the text is held
895    /// until either a match or a divergence
896    Held,
897
898    /// Indicates that a stop sequence has been matched and the decoder is stopped.
899    /// Subsequent calls to append_token_id will return an error
900    Stopped,
901
902    /// Indicates that a stop token_id has been matched and the decoder is stopped.
903    /// Subsequent calls to append_token_id will return an error
904    /// The text for the stop token_id is returned
905    StoppedWithText(String),
906}
907
908/// A Sequence for decoding a stream of token ids into text and detecting stop sequences.
909/// A stop sequence is either a matching token_id or a sequence of texts/strings which match.
910/// Matches happen first at the token-level, then at the sequence-level. Hidden takes precedence
911/// over visible. For example, if you put the same token_id in both `stop_token_ids_visible` and
912/// `stop_token_ids_hidden`, the token_id will be treated as hidden.
913#[derive(Debug)]
914pub struct StopSequenceDecoder {
915    // The current sequence of token ids
916    sequence: Sequence,
917
918    // Stop Tokens - the presence of any one of these should trigger a stop
919    // If found, the text for the matched token will be returned
920    stop_token_ids_visible: Vec<TokenIdType>,
921
922    // Stop Tokens - the presence of any one of these should trigger a stop
923    // If found, the text for the matched token will NOT be returned
924    stop_token_ids_hidden: Vec<TokenIdType>,
925
926    // Stop Words - the presence of any one of these should trigger a stop
927    // If found, the text for the matched token will be returned
928    #[allow(dead_code)]
929    stop_sequences_visible: Vec<String>,
930
931    // Stop Words - the presence of any one of these should trigger a stop
932    // If found, the text for the matched token will NOT be returned
933    stop_sequences_hidden: Vec<String>,
934
935    // If the decoder has observed and returned a stop SequenceDecoderOutput,
936    // futhur calls to append_token_id will return an error
937    stopped: bool,
938
939    // text jail - if a partial stop sequence is being observed, we hold/jail the text
940    // until either the stop sequence is matched or the sequence is reset by a divergence
941    state: String,
942}
943
944impl StopSequenceDecoder {
945    /// Builder object for configurating a StopSequenceDecoder
946    pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
947        StopSequenceDecoderBuilder::new(tokenizer)
948    }
949
950    /// Add a token_id to the sequence and return the SequenceDecoderOutput
951    pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<SequenceDecoderOutput> {
952        if self.stopped {
953            return Err(Error::msg("Decoder is stopped"));
954        }
955
956        // update the sequence
957        let text = self.sequence.append_token_id(token_id)?;
958
959        // append the text to the state
960        self.state.push_str(text.as_str());
961
962        let mut stop: bool = false;
963        let mut visible: bool = false;
964
965        if self.stop_token_ids_visible.contains(&token_id) {
966            stop = true;
967            visible = true;
968        }
969
970        if self.stop_token_ids_hidden.contains(&token_id) {
971            stop = true;
972            visible = false;
973        }
974
975        if stop {
976            self.stopped = true;
977            let state = std::mem::take(&mut self.state);
978            if visible {
979                return Ok(SequenceDecoderOutput::StoppedWithText(state));
980            }
981            return Ok(SequenceDecoderOutput::Stopped);
982        }
983
984        // determine if state matches any of the stop sequences
985        for stop_sequence in self.stop_sequences_hidden.iter() {
986            if stop_sequence.starts_with(&self.state) {
987                if stop_sequence == &self.state {
988                    // on matched stop sequence, we do NOT return the jailed stop sequence
989                    self.stopped = true;
990                    return Ok(SequenceDecoderOutput::Stopped);
991                } else {
992                    return Ok(SequenceDecoderOutput::Held);
993                }
994            }
995        }
996
997        let state = std::mem::take(&mut self.state);
998        Ok(SequenceDecoderOutput::Text(state))
999    }
1000
1001    pub fn is_empty(&self) -> bool {
1002        self.sequence.token_ids.is_empty()
1003    }
1004
1005    pub fn len(&self) -> usize {
1006        self.sequence.token_ids.len()
1007    }
1008
1009    pub fn is_complete(&self) -> bool {
1010        self.stopped
1011    }
1012
1013    pub fn close(&mut self) {
1014        self.stopped = true;
1015    }
1016}
1017
1018pub struct StopSequenceDecoderBuilder {
1019    tokenizer: Tokenizer,
1020    stop_token_ids_visible: Vec<TokenIdType>,
1021    stop_token_ids_hidden: Vec<TokenIdType>,
1022    stop_sequences_visible: Vec<String>,
1023    stop_sequences_hidden: Vec<String>,
1024}
1025
1026impl StopSequenceDecoderBuilder {
1027    pub fn new(tokenizer: Tokenizer) -> Self {
1028        Self {
1029            tokenizer,
1030            stop_token_ids_visible: Vec::new(),
1031            stop_token_ids_hidden: Vec::new(),
1032            stop_sequences_visible: Vec::new(),
1033            stop_sequences_hidden: Vec::new(),
1034        }
1035    }
1036
1037    /// Adds a visible stop token id to the StopSequenceDecoder
1038    pub fn add_stop_token_id_visible(mut self, token_id: TokenIdType) -> Self {
1039        self.stop_token_ids_visible.push(token_id);
1040        self
1041    }
1042
1043    /// Adds a list of visible stop token ids to the StopSequenceDecoder
1044    /// Each token_id is added as for an individual match
1045    pub fn add_stop_token_ids_visible(mut self, token_ids: &[TokenIdType]) -> Self {
1046        self.stop_token_ids_visible.extend(token_ids);
1047        self
1048    }
1049
1050    /// Adds a hidden stop token id to the StopSequenceDecoder
1051    pub fn add_stop_token_id_hidden(mut self, token_id: TokenIdType) -> Self {
1052        self.stop_token_ids_hidden.push(token_id);
1053        self
1054    }
1055
1056    /// Adds a list of hidden stop token ids to the StopSequenceDecoder
1057    /// Each token_id is added as for an individual match
1058    pub fn add_stop_token_ids_hidden(mut self, token_ids: &[TokenIdType]) -> Self {
1059        self.stop_token_ids_hidden.extend(token_ids);
1060        self
1061    }
1062
1063    pub fn add_stop_sequence_visible(mut self, text: &str) -> Self {
1064        self.stop_sequences_visible.push(text.to_string());
1065        self
1066    }
1067
1068    pub fn add_stop_sequences_visible(mut self, strings: &[&str]) -> Self {
1069        self.stop_sequences_visible
1070            .extend(strings.iter().map(|text| text.to_string()));
1071        self
1072    }
1073
1074    pub fn add_stop_sequence_hidden(mut self, text: &str) -> Self {
1075        self.stop_sequences_hidden.push(text.to_string());
1076        self
1077    }
1078
1079    pub fn add_stop_sequences_hidden(mut self, strings: &[&str]) -> Self {
1080        self.stop_sequences_hidden
1081            .extend(strings.iter().map(|text| text.to_string()));
1082        self
1083    }
1084
1085    pub fn build(self) -> Result<StopSequenceDecoder> {
1086        Ok(StopSequenceDecoder {
1087            sequence: Sequence::new(self.tokenizer.clone()),
1088            stop_token_ids_visible: self.stop_token_ids_visible,
1089            stop_token_ids_hidden: self.stop_token_ids_hidden,
1090            stop_sequences_visible: self.stop_sequences_visible,
1091            stop_sequences_hidden: self.stop_sequences_hidden,
1092            stopped: false,
1093            state: String::new(),
1094        })
1095    }
1096}