1pub mod basetenkenizer;
5pub mod cache;
6pub mod fastokens;
7pub mod hf;
8pub mod tiktoken;
9
10use 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::FastTokenizer;
27pub use hf::HuggingFaceTokenizer;
28pub use tiktoken::TikTokenTokenizer;
29pub use traits::DecodeResult;
30
31pub type TokenIdType = u32;
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub struct EncodeSegment<'a> {
36 pub text: &'a str,
37 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#[derive(Debug)]
63pub enum TokenizerType {
64 HuggingFace(String),
65 TikToken(String),
66}
67
68pub type Offsets = (usize, usize);
70
71#[derive(Debug, Clone)]
73pub enum Encoding {
74 Hf(Box<tokenizers::tokenizer::Encoding>),
76 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 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 #[derive(Debug, Clone, PartialEq, Eq, strum::EnumIs)]
133 pub enum DecodeResult {
134 Complete(String),
138 Partial(String),
141 }
142
143 impl DecodeResult {
144 pub fn as_str(&self) -> &str {
146 match self {
147 DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
148 }
149 }
150
151 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 pub trait Decoder: Send + Sync {
181 fn decode(
182 &self,
183 token_ids: &[TokenIdType],
184 skip_special_tokens: bool,
185 ) -> Result<DecodeResult>;
186 }
187
188 pub trait Tokenizer: Encoder + Decoder {
189 fn validate_prefix_cache(&self) -> Result<()> {
194 Err(Error::msg("tokenizer does not support prefix caching"))
195 }
196
197 fn with_options(self, options: TokenizerOptions) -> Self
202 where
203 Self: Sized,
204 {
205 let _ = options;
206 self
207 }
208 fn vocab_size(&self) -> Option<usize> {
212 None
213 }
214
215 fn token_to_id(&self, _token: &str) -> Result<Option<TokenIdType>> {
225 Err(Error::msg(
226 "tokenizer backend does not support token lookup",
227 ))
228 }
229
230 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
240 Err(Error::msg(
241 "tokenizer backend does not support special token enumeration",
242 ))
243 }
244
245 fn num_special_tokens_added(&self) -> Result<usize> {
256 Err(Error::msg(
257 "tokenizer backend does not support special token accounting",
258 ))
259 }
260 }
262}
263
264pub fn file_json_field<T: serde::de::DeserializeOwned>(
265 json_file_path: &Path,
266 field_name: &str,
267) -> anyhow::Result<T> {
268 let file = File::open(json_file_path)
269 .with_context(|| format!("Failed to open file: {:?}", json_file_path))?;
270 let reader = BufReader::new(file);
271
272 let json_data: serde_json::Value = serde_json::from_reader(reader)
273 .with_context(|| format!("Failed to parse JSON from file: {:?}", json_file_path))?;
274
275 let map = json_data.as_object().ok_or_else(|| {
276 anyhow::anyhow!("JSON root is not an object in file: {:?}", json_file_path)
277 })?;
278
279 let field_value = map.get(field_name).ok_or_else(|| {
280 anyhow::anyhow!(
281 "Field '{}' not found in JSON file: {:?}",
282 field_name,
283 json_file_path
284 )
285 })?;
286
287 serde_json::from_value(field_value.clone()).with_context(|| {
288 format!(
289 "Failed to deserialize field '{}' (value: {:?}) to the expected type from file: {:?}",
290 field_name, field_value, json_file_path
291 )
292 })
293}
294
295pub fn log_json_err(filename: &str, json: &str, err: &serde_json::Error) {
296 const ERROR_PREFIX: &str = ">> ";
297
298 if !(err.is_syntax() || err.is_data()) {
299 return;
300 }
301
302 let line = err.line().saturating_sub(1);
303 let column = err.column().saturating_sub(1);
304
305 let json_lines: Vec<&str> = json.lines().collect();
306 if json_lines.is_empty() {
307 tracing::error!("JSON parsing error in {filename}: File is empty.");
308 return;
309 }
310
311 let start_index = line.saturating_sub(2);
312 let end_index = line.saturating_add(3).min(json_lines.len());
313
314 let mut context_lines: Vec<String> = (start_index..end_index)
315 .map(|i| {
316 if i == line {
317 format!("{ERROR_PREFIX}{}", json_lines[i])
318 } else {
319 format!("{:06} {}", i + 1, json_lines[i])
320 }
321 })
322 .collect();
323
324 let col_indicator = "_".to_string().repeat(column + ERROR_PREFIX.len()) + "^";
325 let error_in_context_idx = line - start_index;
326 if error_in_context_idx < context_lines.len() {
327 context_lines.insert(error_in_context_idx + 1, col_indicator);
328 }
329
330 tracing::error!(
331 "JSON parsing error in {filename}: Line {}, column {}:\n{}",
332 err.line(),
333 err.column(),
334 context_lines.join("\n")
335 );
336}
337
338impl Encoding {
339 pub fn get_hash(&self) -> u64 {
340 let mut hasher = DefaultHasher::new();
341 self.hash(&mut hasher);
342 hasher.finish()
343 }
344}
345
346#[derive(Debug, Clone, Copy, Default)]
350pub struct TokenizerOptions {
351 pub add_special_tokens: bool,
358}
359
360#[derive(Clone)]
362pub struct Tokenizer(Arc<dyn traits::Tokenizer>);
363
364impl Tokenizer {
365 pub fn from_file(file_path: &str) -> Result<Tokenizer> {
366 Ok(Tokenizer(create_tokenizer_from_file(file_path)?))
367 }
368
369 pub fn from_file_with_options(file_path: &str, options: TokenizerOptions) -> Result<Tokenizer> {
370 Ok(Tokenizer(create_tokenizer_from_file_with_options(
371 file_path, options,
372 )?))
373 }
374
375 pub fn decode_stream(
377 &self,
378 prompt_token_ids: &[TokenIdType],
379 skip_special_tokens: bool,
380 ) -> DecodeStream {
381 DecodeStream::new(self.0.clone(), prompt_token_ids, skip_special_tokens)
382 }
383}
384
385impl Deref for Tokenizer {
386 type Target = Arc<dyn traits::Tokenizer>;
387
388 fn deref(&self) -> &Self::Target {
389 &self.0
390 }
391}
392
393impl From<Arc<dyn traits::Tokenizer>> for Tokenizer {
394 fn from(tokenizer: Arc<dyn traits::Tokenizer>) -> Self {
395 Tokenizer(tokenizer)
396 }
397}
398
399impl<T> From<Arc<T>> for Tokenizer
400where
401 T: traits::Tokenizer + 'static, {
403 fn from(tokenizer: Arc<T>) -> Self {
404 Tokenizer(tokenizer)
405 }
406}
407
408pub fn create_tokenizer_from_file(file_path: &str) -> Result<Arc<dyn traits::Tokenizer>> {
415 create_tokenizer_from_file_with_options(file_path, Default::default())
416}
417
418pub fn create_tokenizer_from_file_with_options(
425 file_path: &str,
426 options: TokenizerOptions,
427) -> Result<Arc<dyn traits::Tokenizer>> {
428 use traits::Tokenizer as _;
429
430 let path = Path::new(file_path);
431 let extension = path
432 .extension()
433 .and_then(std::ffi::OsStr::to_str)
434 .ok_or_else(|| Error::msg("Failed to read file extension".to_string()))?;
435
436 match extension {
437 "json" => {
438 let tokenizer = HuggingFaceTokenizer::from_file(file_path)?.with_options(options);
439 Ok(Arc::new(tokenizer))
440 }
441 "model" | "tiktoken" => {
442 let tokenizer = TikTokenTokenizer::from_file_auto(file_path)?.with_options(options);
443 Ok(Arc::new(tokenizer))
444 }
445 _ => Err(Error::msg(format!(
446 "Unsupported tokenizer file type: .{extension}"
447 ))),
448 }
449}
450
451const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
457
458pub struct DecodeStream {
464 tokenizer: Arc<dyn traits::Tokenizer>,
466
467 skip_special_tokens: bool,
468 all_token_ids: Vec<u32>,
480
481 prefix_offset: usize,
482
483 read_offset: usize,
484
485 has_emitted: bool,
487}
488
489impl DecodeStream {
490 pub fn new(
491 tokenizer: Arc<dyn traits::Tokenizer>,
492 prompt_token_ids: &[TokenIdType],
493 skip_special_tokens: bool,
494 ) -> Self {
495 let context_start = prompt_token_ids
498 .len()
499 .saturating_sub(INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET);
500 let prompt_token_ids = prompt_token_ids[context_start..].to_vec();
501 let num_input_tokens = prompt_token_ids.len();
502 Self {
503 tokenizer,
504 skip_special_tokens,
505 all_token_ids: prompt_token_ids,
506 prefix_offset: 0,
507 read_offset: num_input_tokens,
508 has_emitted: false,
509 }
510 }
511
512 pub fn step(&mut self, id: u32) -> Result<Option<String>> {
525 self.all_token_ids.push(id);
526
527 let prefix_text: String = self
528 .tokenizer
529 .decode(
530 &self.all_token_ids[self.prefix_offset..self.read_offset],
531 self.skip_special_tokens,
532 )?
533 .into();
534
535 let new_result = self.tokenizer.decode(
536 &self.all_token_ids[self.prefix_offset..],
537 self.skip_special_tokens,
538 )?;
539
540 let new_text = new_result.as_str();
541 let is_partial = new_result.is_partial();
542
543 if self.has_emitted && !is_partial && !new_text.starts_with(prefix_text.as_str()) {
546 return Err(Error::msg(
547 "incremental decoding rewrote already emitted text",
548 ));
549 }
550
551 if new_text.len() > prefix_text.len() && !is_partial {
552 let requested_split = prefix_text.len();
553 let split = if self.has_emitted {
554 requested_split
557 } else {
558 new_text.floor_char_boundary(requested_split)
561 };
562
563 let emitted = new_text[split..].to_string();
564
565 self.prefix_offset = self.read_offset;
566 self.read_offset = self.all_token_ids.len();
567 self.has_emitted = true;
568
569 Ok(Some(emitted))
570 } else {
571 Ok(None)
572 }
573 }
574}
575
576#[cfg(test)]
577mod decode_stream_unicode_tests {
578 use super::{DecodeResult, DecodeStream, Encoding, Result, TokenIdType};
579 use std::sync::Arc;
580
581 struct RewritingTokenizer {
582 prefix_text: &'static str,
583 rewritten_text: &'static str,
584 prefix_is_partial: bool,
585 }
586
587 impl super::traits::Encoder for RewritingTokenizer {
588 fn encode(&self, _input: &str) -> Result<Encoding> {
589 Ok(Encoding::Sp(vec![]))
590 }
591
592 fn encode_batch(&self, _inputs: &[&str]) -> Result<Vec<Encoding>> {
593 Ok(vec![])
594 }
595 }
596
597 impl super::traits::Decoder for RewritingTokenizer {
598 fn decode(
599 &self,
600 token_ids: &[TokenIdType],
601 _skip_special_tokens: bool,
602 ) -> Result<DecodeResult> {
603 let result = match token_ids.len() {
604 1 if self.prefix_is_partial => DecodeResult::Partial(self.prefix_text.to_string()),
605 1 => DecodeResult::Complete(self.prefix_text.to_string()),
606 2 => DecodeResult::Complete(self.rewritten_text.to_string()),
607 _ => DecodeResult::Complete(String::new()),
608 };
609 Ok(result)
610 }
611 }
612
613 impl super::traits::Tokenizer for RewritingTokenizer {}
614
615 #[test]
616 fn prompt_suffix_preserves_decode_inputs_and_output() {
617 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(
618 super::HuggingFaceTokenizer::from_file(concat!(
619 env!("CARGO_MANIFEST_DIR"),
620 "/tests/data/minimal-bpe/tokenizer.json"
621 ))
622 .unwrap(),
623 );
624 for prompt_len in [0usize, 1, 4, 5, 6, 775_168] {
625 let prompt: Vec<_> = (0..prompt_len)
626 .map(|index| 1 + (index % 22) as u32)
627 .collect();
628 for skip_special_tokens in [false, true] {
629 let mut stream = DecodeStream::new(tokenizer.clone(), &prompt, skip_special_tokens);
630 let mut full = DecodeStream {
632 tokenizer: tokenizer.clone(),
633 skip_special_tokens,
634 all_token_ids: prompt.clone(),
635 prefix_offset: prompt_len
636 .saturating_sub(super::INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET),
637 read_offset: prompt_len,
638 has_emitted: false,
639 };
640 assert!(stream.all_token_ids.capacity() <= 5);
641 for token in [5, 9, 12, 12, 13, 0, 1, 17, 13, 14, 12, 8, 2] {
642 assert_eq!(
643 &stream.all_token_ids[stream.prefix_offset..stream.read_offset],
644 &full.all_token_ids[full.prefix_offset..full.read_offset],
645 );
646 assert_eq!(
647 &stream.all_token_ids[stream.prefix_offset..],
648 &full.all_token_ids[full.prefix_offset..],
649 );
650 let actual = stream.step(token).map_err(|error| error.to_string());
651 let expected = full.step(token).map_err(|error| error.to_string());
652 assert_eq!(actual, expected, "prompt length {prompt_len}");
653 if actual.is_err() {
654 break;
655 }
656 }
657 }
658 }
659 }
660
661 #[test]
662 fn allows_boundary_recovery_before_generated_text_is_emitted() {
663 for (prefix_text, rewritten_text, expected) in [
664 ("㺄馉凓鄗\u{FFFD}", "㺄馉凓鄗𫷲", "𫷲"),
665 (
666 "JUnitworkflow Completion intuition\u{FFFD}",
667 "JUnitworkflow Completion intuition𝟙",
668 "𝟙",
669 ),
670 ] {
671 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
672 prefix_text,
673 rewritten_text,
674 prefix_is_partial: true,
675 });
676 let mut stream = DecodeStream::new(tokenizer, &[1], false);
677
678 assert_eq!(stream.step(2).unwrap(), Some(expected.to_string()));
679 }
680 }
681
682 #[test]
683 fn errors_when_incremental_decode_rewrites_emitted_text() {
684 for (prefix_text, rewritten_text) in [
685 ("abcde", "abcdXY"),
686 ("abcde", "abcdX"),
687 ("abcde", "abcd"),
688 ("㺄馉凓鄗abc", "㺄馉凓鄗𫷲"),
689 (
690 "JUnitworkflow Completion intuitionabc",
691 "JUnitworkflow Completion intuition𝟙",
692 ),
693 ] {
694 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
695 prefix_text,
696 rewritten_text,
697 prefix_is_partial: false,
698 });
699 let mut stream = DecodeStream::new(tokenizer, &[], false);
700
701 assert_eq!(stream.step(1).unwrap(), Some(prefix_text.to_string()));
702
703 let error = stream.step(2).unwrap_err();
704 assert!(error.to_string().contains("already emitted text"));
705 }
706 }
707
708 #[test]
709 fn emits_suffix_when_incremental_decode_preserves_emitted_text() {
710 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
711 prefix_text: "abcde",
712 rewritten_text: "abcdeXY",
713 prefix_is_partial: false,
714 });
715 let mut stream = DecodeStream::new(tokenizer, &[], false);
716
717 assert_eq!(stream.step(1).unwrap(), Some("abcde".to_string()));
718 assert_eq!(stream.step(2).unwrap(), Some("XY".to_string()));
719 }
720}
721
722pub struct Sequence {
724 tokenizer: Tokenizer,
726
727 token_ids: Vec<TokenIdType>,
729
730 prefix_offset: usize,
732
733 read_offset: usize,
735}
736
737impl std::fmt::Debug for Sequence {
738 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
739 f.debug_struct("Sequence")
740 .field("tokenizer", &"Arc<dyn Tokenizer>")
741 .field(
742 "token_ids",
743 &format_args!("{}", {
744 let token_ids = self.token_ids();
745 if token_ids.len() <= 20 {
746 format!("{:?}", token_ids)
747 } else {
748 let first_ten = &token_ids[..10];
749 let last_ten = &token_ids[token_ids.len() - 10..];
750 format!("{:?} ... {:?}", first_ten, last_ten)
751 }
752 }),
753 )
754 .field("prefix_offset", &self.prefix_offset)
755 .field("read_offset", &self.read_offset)
756 .field("token count", &self.token_ids.len())
757 .finish()
758 }
759}
760
761impl Sequence {
762 pub fn new(tokenizer: Tokenizer) -> Self {
763 Self {
764 tokenizer,
765 token_ids: Vec::new(),
766 prefix_offset: 0,
767 read_offset: 0,
768 }
769 }
770
771 pub fn is_empty(&self) -> bool {
772 self.token_ids.is_empty()
773 }
774
775 pub fn len(&self) -> usize {
776 self.token_ids.len()
777 }
778
779 pub fn clear(&mut self) {
780 self.token_ids.clear();
781 self.prefix_offset = 0;
782 self.read_offset = 0;
783 }
784
785 pub fn append_text(&mut self, input: &str) -> Result<()> {
786 let encoding = self.tokenizer.encode(input)?;
791 self.token_ids.extend(encoding.token_ids());
792 Ok(())
793 }
794
795 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
799 self.token_ids.push(token_id);
800 let prefix_text: String = self
803 .tokenizer
804 .decode(&self.token_ids[self.prefix_offset..self.read_offset], false)?
805 .into();
806
807 let new_result = self
808 .tokenizer
809 .decode(&self.token_ids[self.prefix_offset..], false)?;
810
811 let new_text = new_result.as_str();
812
813 let mut prefix_text_len = prefix_text.len();
817 while !new_text.is_char_boundary(prefix_text_len) && prefix_text_len > 0 {
818 prefix_text_len -= 1;
819 }
820 let prefix_text_len = prefix_text_len;
821
822 if new_text.len() > prefix_text.len() {
823 if new_result.is_partial() {
824 return Ok("".to_string());
825 } else {
826 let new_text = new_text[prefix_text_len..]
828 .to_string()
829 .replace('\u{FFFD}', "");
830 self.prefix_offset = self.read_offset;
831 self.read_offset = self.token_ids.len();
832 return Ok(new_text);
833 }
834 }
835
836 Ok("".to_string())
837 }
838
839 pub fn tokenizer(&self) -> Tokenizer {
840 self.tokenizer.clone()
841 }
842
843 pub fn token_ids(&self) -> &[TokenIdType] {
844 &self.token_ids
845 }
846
847 pub fn text(&self) -> Result<String> {
848 Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
852 }
853}
854
855pub enum SequenceDecoderOutput {
858 Text(String),
860
861 Held,
864
865 Stopped,
868
869 StoppedWithText(String),
873}
874
875#[derive(Debug)]
881pub struct StopSequenceDecoder {
882 sequence: Sequence,
884
885 stop_token_ids_visible: Vec<TokenIdType>,
888
889 stop_token_ids_hidden: Vec<TokenIdType>,
892
893 #[allow(dead_code)]
896 stop_sequences_visible: Vec<String>,
897
898 stop_sequences_hidden: Vec<String>,
901
902 stopped: bool,
905
906 state: String,
909}
910
911impl StopSequenceDecoder {
912 pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
914 StopSequenceDecoderBuilder::new(tokenizer)
915 }
916
917 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<SequenceDecoderOutput> {
919 if self.stopped {
920 return Err(Error::msg("Decoder is stopped"));
921 }
922
923 let text = self.sequence.append_token_id(token_id)?;
925
926 self.state.push_str(text.as_str());
928
929 let mut stop: bool = false;
930 let mut visible: bool = false;
931
932 if self.stop_token_ids_visible.contains(&token_id) {
933 stop = true;
934 visible = true;
935 }
936
937 if self.stop_token_ids_hidden.contains(&token_id) {
938 stop = true;
939 visible = false;
940 }
941
942 if stop {
943 self.stopped = true;
944 let state = std::mem::take(&mut self.state);
945 if visible {
946 return Ok(SequenceDecoderOutput::StoppedWithText(state));
947 }
948 return Ok(SequenceDecoderOutput::Stopped);
949 }
950
951 for stop_sequence in self.stop_sequences_hidden.iter() {
953 if stop_sequence.starts_with(&self.state) {
954 if stop_sequence == &self.state {
955 self.stopped = true;
957 return Ok(SequenceDecoderOutput::Stopped);
958 } else {
959 return Ok(SequenceDecoderOutput::Held);
960 }
961 }
962 }
963
964 let state = std::mem::take(&mut self.state);
965 Ok(SequenceDecoderOutput::Text(state))
966 }
967
968 pub fn is_empty(&self) -> bool {
969 self.sequence.token_ids.is_empty()
970 }
971
972 pub fn len(&self) -> usize {
973 self.sequence.token_ids.len()
974 }
975
976 pub fn is_complete(&self) -> bool {
977 self.stopped
978 }
979
980 pub fn close(&mut self) {
981 self.stopped = true;
982 }
983}
984
985pub struct StopSequenceDecoderBuilder {
986 tokenizer: Tokenizer,
987 stop_token_ids_visible: Vec<TokenIdType>,
988 stop_token_ids_hidden: Vec<TokenIdType>,
989 stop_sequences_visible: Vec<String>,
990 stop_sequences_hidden: Vec<String>,
991}
992
993impl StopSequenceDecoderBuilder {
994 pub fn new(tokenizer: Tokenizer) -> Self {
995 Self {
996 tokenizer,
997 stop_token_ids_visible: Vec::new(),
998 stop_token_ids_hidden: Vec::new(),
999 stop_sequences_visible: Vec::new(),
1000 stop_sequences_hidden: Vec::new(),
1001 }
1002 }
1003
1004 pub fn add_stop_token_id_visible(mut self, token_id: TokenIdType) -> Self {
1006 self.stop_token_ids_visible.push(token_id);
1007 self
1008 }
1009
1010 pub fn add_stop_token_ids_visible(mut self, token_ids: &[TokenIdType]) -> Self {
1013 self.stop_token_ids_visible.extend(token_ids);
1014 self
1015 }
1016
1017 pub fn add_stop_token_id_hidden(mut self, token_id: TokenIdType) -> Self {
1019 self.stop_token_ids_hidden.push(token_id);
1020 self
1021 }
1022
1023 pub fn add_stop_token_ids_hidden(mut self, token_ids: &[TokenIdType]) -> Self {
1026 self.stop_token_ids_hidden.extend(token_ids);
1027 self
1028 }
1029
1030 pub fn add_stop_sequence_visible(mut self, text: &str) -> Self {
1031 self.stop_sequences_visible.push(text.to_string());
1032 self
1033 }
1034
1035 pub fn add_stop_sequences_visible(mut self, strings: &[&str]) -> Self {
1036 self.stop_sequences_visible
1037 .extend(strings.iter().map(|text| text.to_string()));
1038 self
1039 }
1040
1041 pub fn add_stop_sequence_hidden(mut self, text: &str) -> Self {
1042 self.stop_sequences_hidden.push(text.to_string());
1043 self
1044 }
1045
1046 pub fn add_stop_sequences_hidden(mut self, strings: &[&str]) -> Self {
1047 self.stop_sequences_hidden
1048 .extend(strings.iter().map(|text| text.to_string()));
1049 self
1050 }
1051
1052 pub fn build(self) -> Result<StopSequenceDecoder> {
1053 Ok(StopSequenceDecoder {
1054 sequence: Sequence::new(self.tokenizer.clone()),
1055 stop_token_ids_visible: self.stop_token_ids_visible,
1056 stop_token_ids_hidden: self.stop_token_ids_hidden,
1057 stop_sequences_visible: self.stop_sequences_visible,
1058 stop_sequences_hidden: self.stop_sequences_hidden,
1059 stopped: false,
1060 state: String::new(),
1061 })
1062 }
1063}