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::{CacheTokenUsage, CacheTokenUsageFn, CachedTokenizer, L1CacheStats};
23pub use fastokens::FastTokenizer;
24pub use hf::HuggingFaceTokenizer;
25pub use tiktoken::TikTokenTokenizer;
26pub use traits::DecodeResult;
27
28pub type TokenIdType = u32;
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub struct EncodeSegment<'a> {
33 pub text: &'a str,
34 pub allow_special: bool,
39}
40
41impl<'a> EncodeSegment<'a> {
42 pub const fn new(text: &'a str, allow_special: bool) -> Self {
43 Self {
44 text,
45 allow_special,
46 }
47 }
48
49 pub const fn ordinary(text: &'a str) -> Self {
50 Self::new(text, false)
51 }
52
53 pub const fn control(text: &'a str) -> Self {
54 Self::new(text, true)
55 }
56}
57
58#[derive(Debug)]
60pub enum TokenizerType {
61 HuggingFace(String),
62 TikToken(String),
63}
64
65pub type Offsets = (usize, usize);
67
68#[derive(Debug, Clone)]
70pub enum Encoding {
71 Hf(Box<tokenizers::tokenizer::Encoding>),
73 Sp(Vec<TokenIdType>),
75}
76
77impl Encoding {
78 pub fn token_ids(&self) -> &[u32] {
79 match self {
80 Encoding::Hf(inner) => inner.get_ids(),
81 Encoding::Sp(inner) => inner,
82 }
83 }
84}
85
86impl Hash for Encoding {
87 fn hash<H: Hasher>(&self, state: &mut H) {
88 self.token_ids().hash(state);
89 }
90}
91
92pub mod traits {
93 use super::*;
94
95 pub trait Encoder: Send + Sync {
96 fn encode(&self, input: &str) -> Result<Encoding>;
97 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>>;
98
99 fn encode_segments(&self, _segments: &[EncodeSegment<'_>]) -> Result<Encoding> {
117 Err(Error::msg(
118 "tokenizer backend does not support segmented encoding",
119 ))
120 }
121 }
122
123 #[derive(Debug, Clone, PartialEq, Eq, strum::EnumIs)]
130 pub enum DecodeResult {
131 Complete(String),
135 Partial(String),
138 }
139
140 impl DecodeResult {
141 pub fn as_str(&self) -> &str {
143 match self {
144 DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
145 }
146 }
147
148 pub fn from_decoded(text: String) -> Self {
150 if text.ends_with('\u{FFFD}') {
151 DecodeResult::Partial(text)
152 } else {
153 DecodeResult::Complete(text)
154 }
155 }
156 }
157
158 impl From<String> for DecodeResult {
159 fn from(text: String) -> Self {
160 DecodeResult::from_decoded(text)
161 }
162 }
163
164 impl From<DecodeResult> for String {
165 fn from(result: DecodeResult) -> Self {
166 match result {
167 DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
168 }
169 }
170 }
171
172 pub trait Decoder: Send + Sync {
178 fn decode(
179 &self,
180 token_ids: &[TokenIdType],
181 skip_special_tokens: bool,
182 ) -> Result<DecodeResult>;
183 }
184
185 pub trait Tokenizer: Encoder + Decoder {
186 fn validate_prefix_cache(&self) -> Result<()> {
191 Err(Error::msg("tokenizer does not support prefix caching"))
192 }
193
194 fn with_options(self, options: TokenizerOptions) -> Self
199 where
200 Self: Sized,
201 {
202 let _ = options;
203 self
204 }
205 fn vocab_size(&self) -> Option<usize> {
209 None
210 }
211
212 fn token_to_id(&self, _token: &str) -> Result<Option<TokenIdType>> {
222 Err(Error::msg(
223 "tokenizer backend does not support token lookup",
224 ))
225 }
226
227 fn special_token_ids(&self) -> Result<Vec<TokenIdType>> {
237 Err(Error::msg(
238 "tokenizer backend does not support special token enumeration",
239 ))
240 }
241
242 fn num_special_tokens_added(&self) -> Result<usize> {
253 Err(Error::msg(
254 "tokenizer backend does not support special token accounting",
255 ))
256 }
257 }
259}
260
261pub fn file_json_field<T: serde::de::DeserializeOwned>(
262 json_file_path: &Path,
263 field_name: &str,
264) -> anyhow::Result<T> {
265 let file = File::open(json_file_path)
266 .with_context(|| format!("Failed to open file: {:?}", json_file_path))?;
267 let reader = BufReader::new(file);
268
269 let json_data: serde_json::Value = serde_json::from_reader(reader)
270 .with_context(|| format!("Failed to parse JSON from file: {:?}", json_file_path))?;
271
272 let map = json_data.as_object().ok_or_else(|| {
273 anyhow::anyhow!("JSON root is not an object in file: {:?}", json_file_path)
274 })?;
275
276 let field_value = map.get(field_name).ok_or_else(|| {
277 anyhow::anyhow!(
278 "Field '{}' not found in JSON file: {:?}",
279 field_name,
280 json_file_path
281 )
282 })?;
283
284 serde_json::from_value(field_value.clone()).with_context(|| {
285 format!(
286 "Failed to deserialize field '{}' (value: {:?}) to the expected type from file: {:?}",
287 field_name, field_value, json_file_path
288 )
289 })
290}
291
292pub fn log_json_err(filename: &str, json: &str, err: &serde_json::Error) {
293 const ERROR_PREFIX: &str = ">> ";
294
295 if !(err.is_syntax() || err.is_data()) {
296 return;
297 }
298
299 let line = err.line().saturating_sub(1);
300 let column = err.column().saturating_sub(1);
301
302 let json_lines: Vec<&str> = json.lines().collect();
303 if json_lines.is_empty() {
304 tracing::error!("JSON parsing error in {filename}: File is empty.");
305 return;
306 }
307
308 let start_index = line.saturating_sub(2);
309 let end_index = line.saturating_add(3).min(json_lines.len());
310
311 let mut context_lines: Vec<String> = (start_index..end_index)
312 .map(|i| {
313 if i == line {
314 format!("{ERROR_PREFIX}{}", json_lines[i])
315 } else {
316 format!("{:06} {}", i + 1, json_lines[i])
317 }
318 })
319 .collect();
320
321 let col_indicator = "_".to_string().repeat(column + ERROR_PREFIX.len()) + "^";
322 let error_in_context_idx = line - start_index;
323 if error_in_context_idx < context_lines.len() {
324 context_lines.insert(error_in_context_idx + 1, col_indicator);
325 }
326
327 tracing::error!(
328 "JSON parsing error in {filename}: Line {}, column {}:\n{}",
329 err.line(),
330 err.column(),
331 context_lines.join("\n")
332 );
333}
334
335impl Encoding {
336 pub fn get_hash(&self) -> u64 {
337 let mut hasher = DefaultHasher::new();
338 self.hash(&mut hasher);
339 hasher.finish()
340 }
341}
342
343#[derive(Debug, Clone, Copy, Default)]
347pub struct TokenizerOptions {
348 pub add_special_tokens: bool,
355}
356
357#[derive(Clone)]
359pub struct Tokenizer(Arc<dyn traits::Tokenizer>);
360
361impl Tokenizer {
362 pub fn from_file(file_path: &str) -> Result<Tokenizer> {
363 Ok(Tokenizer(create_tokenizer_from_file(file_path)?))
364 }
365
366 pub fn from_file_with_options(file_path: &str, options: TokenizerOptions) -> Result<Tokenizer> {
367 Ok(Tokenizer(create_tokenizer_from_file_with_options(
368 file_path, options,
369 )?))
370 }
371
372 pub fn decode_stream(
374 &self,
375 prompt_token_ids: &[TokenIdType],
376 skip_special_tokens: bool,
377 ) -> DecodeStream {
378 DecodeStream::new(self.0.clone(), prompt_token_ids, skip_special_tokens)
379 }
380}
381
382impl Deref for Tokenizer {
383 type Target = Arc<dyn traits::Tokenizer>;
384
385 fn deref(&self) -> &Self::Target {
386 &self.0
387 }
388}
389
390impl From<Arc<dyn traits::Tokenizer>> for Tokenizer {
391 fn from(tokenizer: Arc<dyn traits::Tokenizer>) -> Self {
392 Tokenizer(tokenizer)
393 }
394}
395
396impl<T> From<Arc<T>> for Tokenizer
397where
398 T: traits::Tokenizer + 'static, {
400 fn from(tokenizer: Arc<T>) -> Self {
401 Tokenizer(tokenizer)
402 }
403}
404
405pub fn create_tokenizer_from_file(file_path: &str) -> Result<Arc<dyn traits::Tokenizer>> {
412 create_tokenizer_from_file_with_options(file_path, Default::default())
413}
414
415pub fn create_tokenizer_from_file_with_options(
422 file_path: &str,
423 options: TokenizerOptions,
424) -> Result<Arc<dyn traits::Tokenizer>> {
425 use traits::Tokenizer as _;
426
427 let path = Path::new(file_path);
428 let extension = path
429 .extension()
430 .and_then(std::ffi::OsStr::to_str)
431 .ok_or_else(|| Error::msg("Failed to read file extension".to_string()))?;
432
433 match extension {
434 "json" => {
435 let tokenizer = HuggingFaceTokenizer::from_file(file_path)?.with_options(options);
436 Ok(Arc::new(tokenizer))
437 }
438 "model" | "tiktoken" => {
439 let tokenizer = TikTokenTokenizer::from_file_auto(file_path)?.with_options(options);
440 Ok(Arc::new(tokenizer))
441 }
442 _ => Err(Error::msg(format!(
443 "Unsupported tokenizer file type: .{extension}"
444 ))),
445 }
446}
447
448const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
454
455pub struct DecodeStream {
461 tokenizer: Arc<dyn traits::Tokenizer>,
463
464 skip_special_tokens: bool,
465 all_token_ids: Vec<u32>,
477
478 prefix_offset: usize,
479
480 read_offset: usize,
481
482 has_emitted: bool,
484}
485
486impl DecodeStream {
487 pub fn new(
488 tokenizer: Arc<dyn traits::Tokenizer>,
489 prompt_token_ids: &[TokenIdType],
490 skip_special_tokens: bool,
491 ) -> Self {
492 let context_start = prompt_token_ids
495 .len()
496 .saturating_sub(INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET);
497 let prompt_token_ids = prompt_token_ids[context_start..].to_vec();
498 let num_input_tokens = prompt_token_ids.len();
499 Self {
500 tokenizer,
501 skip_special_tokens,
502 all_token_ids: prompt_token_ids,
503 prefix_offset: 0,
504 read_offset: num_input_tokens,
505 has_emitted: false,
506 }
507 }
508
509 pub fn step(&mut self, id: u32) -> Result<Option<String>> {
522 self.all_token_ids.push(id);
523
524 let prefix_text: String = self
525 .tokenizer
526 .decode(
527 &self.all_token_ids[self.prefix_offset..self.read_offset],
528 self.skip_special_tokens,
529 )?
530 .into();
531
532 let new_result = self.tokenizer.decode(
533 &self.all_token_ids[self.prefix_offset..],
534 self.skip_special_tokens,
535 )?;
536
537 let new_text = new_result.as_str();
538 let is_partial = new_result.is_partial();
539
540 if self.has_emitted && !is_partial && !new_text.starts_with(prefix_text.as_str()) {
543 return Err(Error::msg(
544 "incremental decoding rewrote already emitted text",
545 ));
546 }
547
548 if new_text.len() > prefix_text.len() && !is_partial {
549 let requested_split = prefix_text.len();
550 let split = if self.has_emitted {
551 requested_split
554 } else {
555 new_text.floor_char_boundary(requested_split)
558 };
559
560 let emitted = new_text[split..].to_string();
561
562 self.prefix_offset = self.read_offset;
563 self.read_offset = self.all_token_ids.len();
564 self.has_emitted = true;
565
566 Ok(Some(emitted))
567 } else {
568 Ok(None)
569 }
570 }
571}
572
573#[cfg(test)]
574mod decode_stream_unicode_tests {
575 use super::{DecodeResult, DecodeStream, Encoding, Result, TokenIdType};
576 use std::sync::Arc;
577
578 struct RewritingTokenizer {
579 prefix_text: &'static str,
580 rewritten_text: &'static str,
581 prefix_is_partial: bool,
582 }
583
584 impl super::traits::Encoder for RewritingTokenizer {
585 fn encode(&self, _input: &str) -> Result<Encoding> {
586 Ok(Encoding::Sp(vec![]))
587 }
588
589 fn encode_batch(&self, _inputs: &[&str]) -> Result<Vec<Encoding>> {
590 Ok(vec![])
591 }
592 }
593
594 impl super::traits::Decoder for RewritingTokenizer {
595 fn decode(
596 &self,
597 token_ids: &[TokenIdType],
598 _skip_special_tokens: bool,
599 ) -> Result<DecodeResult> {
600 let result = match token_ids.len() {
601 1 if self.prefix_is_partial => DecodeResult::Partial(self.prefix_text.to_string()),
602 1 => DecodeResult::Complete(self.prefix_text.to_string()),
603 2 => DecodeResult::Complete(self.rewritten_text.to_string()),
604 _ => DecodeResult::Complete(String::new()),
605 };
606 Ok(result)
607 }
608 }
609
610 impl super::traits::Tokenizer for RewritingTokenizer {}
611
612 #[test]
613 fn prompt_suffix_preserves_decode_inputs_and_output() {
614 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(
615 super::HuggingFaceTokenizer::from_file(concat!(
616 env!("CARGO_MANIFEST_DIR"),
617 "/tests/data/minimal-bpe/tokenizer.json"
618 ))
619 .unwrap(),
620 );
621 for prompt_len in [0usize, 1, 4, 5, 6, 775_168] {
622 let prompt: Vec<_> = (0..prompt_len)
623 .map(|index| 1 + (index % 22) as u32)
624 .collect();
625 for skip_special_tokens in [false, true] {
626 let mut stream = DecodeStream::new(tokenizer.clone(), &prompt, skip_special_tokens);
627 let mut full = DecodeStream {
629 tokenizer: tokenizer.clone(),
630 skip_special_tokens,
631 all_token_ids: prompt.clone(),
632 prefix_offset: prompt_len
633 .saturating_sub(super::INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET),
634 read_offset: prompt_len,
635 has_emitted: false,
636 };
637 assert!(stream.all_token_ids.capacity() <= 5);
638 for token in [5, 9, 12, 12, 13, 0, 1, 17, 13, 14, 12, 8, 2] {
639 assert_eq!(
640 &stream.all_token_ids[stream.prefix_offset..stream.read_offset],
641 &full.all_token_ids[full.prefix_offset..full.read_offset],
642 );
643 assert_eq!(
644 &stream.all_token_ids[stream.prefix_offset..],
645 &full.all_token_ids[full.prefix_offset..],
646 );
647 let actual = stream.step(token).map_err(|error| error.to_string());
648 let expected = full.step(token).map_err(|error| error.to_string());
649 assert_eq!(actual, expected, "prompt length {prompt_len}");
650 if actual.is_err() {
651 break;
652 }
653 }
654 }
655 }
656 }
657
658 #[test]
659 fn allows_boundary_recovery_before_generated_text_is_emitted() {
660 for (prefix_text, rewritten_text, expected) in [
661 ("㺄馉凓鄗\u{FFFD}", "㺄馉凓鄗𫷲", "𫷲"),
662 (
663 "JUnitworkflow Completion intuition\u{FFFD}",
664 "JUnitworkflow Completion intuition𝟙",
665 "𝟙",
666 ),
667 ] {
668 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
669 prefix_text,
670 rewritten_text,
671 prefix_is_partial: true,
672 });
673 let mut stream = DecodeStream::new(tokenizer, &[1], false);
674
675 assert_eq!(stream.step(2).unwrap(), Some(expected.to_string()));
676 }
677 }
678
679 #[test]
680 fn errors_when_incremental_decode_rewrites_emitted_text() {
681 for (prefix_text, rewritten_text) in [
682 ("abcde", "abcdXY"),
683 ("abcde", "abcdX"),
684 ("abcde", "abcd"),
685 ("㺄馉凓鄗abc", "㺄馉凓鄗𫷲"),
686 (
687 "JUnitworkflow Completion intuitionabc",
688 "JUnitworkflow Completion intuition𝟙",
689 ),
690 ] {
691 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
692 prefix_text,
693 rewritten_text,
694 prefix_is_partial: false,
695 });
696 let mut stream = DecodeStream::new(tokenizer, &[], false);
697
698 assert_eq!(stream.step(1).unwrap(), Some(prefix_text.to_string()));
699
700 let error = stream.step(2).unwrap_err();
701 assert!(error.to_string().contains("already emitted text"));
702 }
703 }
704
705 #[test]
706 fn emits_suffix_when_incremental_decode_preserves_emitted_text() {
707 let tokenizer: Arc<dyn super::traits::Tokenizer> = Arc::new(RewritingTokenizer {
708 prefix_text: "abcde",
709 rewritten_text: "abcdeXY",
710 prefix_is_partial: false,
711 });
712 let mut stream = DecodeStream::new(tokenizer, &[], false);
713
714 assert_eq!(stream.step(1).unwrap(), Some("abcde".to_string()));
715 assert_eq!(stream.step(2).unwrap(), Some("XY".to_string()));
716 }
717}
718
719pub struct Sequence {
721 tokenizer: Tokenizer,
723
724 token_ids: Vec<TokenIdType>,
726
727 prefix_offset: usize,
729
730 read_offset: usize,
732}
733
734impl std::fmt::Debug for Sequence {
735 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
736 f.debug_struct("Sequence")
737 .field("tokenizer", &"Arc<dyn Tokenizer>")
738 .field(
739 "token_ids",
740 &format_args!("{}", {
741 let token_ids = self.token_ids();
742 if token_ids.len() <= 20 {
743 format!("{:?}", token_ids)
744 } else {
745 let first_ten = &token_ids[..10];
746 let last_ten = &token_ids[token_ids.len() - 10..];
747 format!("{:?} ... {:?}", first_ten, last_ten)
748 }
749 }),
750 )
751 .field("prefix_offset", &self.prefix_offset)
752 .field("read_offset", &self.read_offset)
753 .field("token count", &self.token_ids.len())
754 .finish()
755 }
756}
757
758impl Sequence {
759 pub fn new(tokenizer: Tokenizer) -> Self {
760 Self {
761 tokenizer,
762 token_ids: Vec::new(),
763 prefix_offset: 0,
764 read_offset: 0,
765 }
766 }
767
768 pub fn is_empty(&self) -> bool {
769 self.token_ids.is_empty()
770 }
771
772 pub fn len(&self) -> usize {
773 self.token_ids.len()
774 }
775
776 pub fn clear(&mut self) {
777 self.token_ids.clear();
778 self.prefix_offset = 0;
779 self.read_offset = 0;
780 }
781
782 pub fn append_text(&mut self, input: &str) -> Result<()> {
783 let encoding = self.tokenizer.encode(input)?;
788 self.token_ids.extend(encoding.token_ids());
789 Ok(())
790 }
791
792 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
796 self.token_ids.push(token_id);
797 let prefix_text: String = self
800 .tokenizer
801 .decode(&self.token_ids[self.prefix_offset..self.read_offset], false)?
802 .into();
803
804 let new_result = self
805 .tokenizer
806 .decode(&self.token_ids[self.prefix_offset..], false)?;
807
808 let new_text = new_result.as_str();
809
810 let mut prefix_text_len = prefix_text.len();
814 while !new_text.is_char_boundary(prefix_text_len) && prefix_text_len > 0 {
815 prefix_text_len -= 1;
816 }
817 let prefix_text_len = prefix_text_len;
818
819 if new_text.len() > prefix_text.len() {
820 if new_result.is_partial() {
821 return Ok("".to_string());
822 } else {
823 let new_text = new_text[prefix_text_len..]
825 .to_string()
826 .replace('\u{FFFD}', "");
827 self.prefix_offset = self.read_offset;
828 self.read_offset = self.token_ids.len();
829 return Ok(new_text);
830 }
831 }
832
833 Ok("".to_string())
834 }
835
836 pub fn tokenizer(&self) -> Tokenizer {
837 self.tokenizer.clone()
838 }
839
840 pub fn token_ids(&self) -> &[TokenIdType] {
841 &self.token_ids
842 }
843
844 pub fn text(&self) -> Result<String> {
845 Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
849 }
850}
851
852pub enum SequenceDecoderOutput {
855 Text(String),
857
858 Held,
861
862 Stopped,
865
866 StoppedWithText(String),
870}
871
872#[derive(Debug)]
878pub struct StopSequenceDecoder {
879 sequence: Sequence,
881
882 stop_token_ids_visible: Vec<TokenIdType>,
885
886 stop_token_ids_hidden: Vec<TokenIdType>,
889
890 #[allow(dead_code)]
893 stop_sequences_visible: Vec<String>,
894
895 stop_sequences_hidden: Vec<String>,
898
899 stopped: bool,
902
903 state: String,
906}
907
908impl StopSequenceDecoder {
909 pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
911 StopSequenceDecoderBuilder::new(tokenizer)
912 }
913
914 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<SequenceDecoderOutput> {
916 if self.stopped {
917 return Err(Error::msg("Decoder is stopped"));
918 }
919
920 let text = self.sequence.append_token_id(token_id)?;
922
923 self.state.push_str(text.as_str());
925
926 let mut stop: bool = false;
927 let mut visible: bool = false;
928
929 if self.stop_token_ids_visible.contains(&token_id) {
930 stop = true;
931 visible = true;
932 }
933
934 if self.stop_token_ids_hidden.contains(&token_id) {
935 stop = true;
936 visible = false;
937 }
938
939 if stop {
940 self.stopped = true;
941 let state = std::mem::take(&mut self.state);
942 if visible {
943 return Ok(SequenceDecoderOutput::StoppedWithText(state));
944 }
945 return Ok(SequenceDecoderOutput::Stopped);
946 }
947
948 for stop_sequence in self.stop_sequences_hidden.iter() {
950 if stop_sequence.starts_with(&self.state) {
951 if stop_sequence == &self.state {
952 self.stopped = true;
954 return Ok(SequenceDecoderOutput::Stopped);
955 } else {
956 return Ok(SequenceDecoderOutput::Held);
957 }
958 }
959 }
960
961 let state = std::mem::take(&mut self.state);
962 Ok(SequenceDecoderOutput::Text(state))
963 }
964
965 pub fn is_empty(&self) -> bool {
966 self.sequence.token_ids.is_empty()
967 }
968
969 pub fn len(&self) -> usize {
970 self.sequence.token_ids.len()
971 }
972
973 pub fn is_complete(&self) -> bool {
974 self.stopped
975 }
976
977 pub fn close(&mut self) {
978 self.stopped = true;
979 }
980}
981
982pub struct StopSequenceDecoderBuilder {
983 tokenizer: Tokenizer,
984 stop_token_ids_visible: Vec<TokenIdType>,
985 stop_token_ids_hidden: Vec<TokenIdType>,
986 stop_sequences_visible: Vec<String>,
987 stop_sequences_hidden: Vec<String>,
988}
989
990impl StopSequenceDecoderBuilder {
991 pub fn new(tokenizer: Tokenizer) -> Self {
992 Self {
993 tokenizer,
994 stop_token_ids_visible: Vec::new(),
995 stop_token_ids_hidden: Vec::new(),
996 stop_sequences_visible: Vec::new(),
997 stop_sequences_hidden: Vec::new(),
998 }
999 }
1000
1001 pub fn add_stop_token_id_visible(mut self, token_id: TokenIdType) -> Self {
1003 self.stop_token_ids_visible.push(token_id);
1004 self
1005 }
1006
1007 pub fn add_stop_token_ids_visible(mut self, token_ids: &[TokenIdType]) -> Self {
1010 self.stop_token_ids_visible.extend(token_ids);
1011 self
1012 }
1013
1014 pub fn add_stop_token_id_hidden(mut self, token_id: TokenIdType) -> Self {
1016 self.stop_token_ids_hidden.push(token_id);
1017 self
1018 }
1019
1020 pub fn add_stop_token_ids_hidden(mut self, token_ids: &[TokenIdType]) -> Self {
1023 self.stop_token_ids_hidden.extend(token_ids);
1024 self
1025 }
1026
1027 pub fn add_stop_sequence_visible(mut self, text: &str) -> Self {
1028 self.stop_sequences_visible.push(text.to_string());
1029 self
1030 }
1031
1032 pub fn add_stop_sequences_visible(mut self, strings: &[&str]) -> Self {
1033 self.stop_sequences_visible
1034 .extend(strings.iter().map(|text| text.to_string()));
1035 self
1036 }
1037
1038 pub fn add_stop_sequence_hidden(mut self, text: &str) -> Self {
1039 self.stop_sequences_hidden.push(text.to_string());
1040 self
1041 }
1042
1043 pub fn add_stop_sequences_hidden(mut self, strings: &[&str]) -> Self {
1044 self.stop_sequences_hidden
1045 .extend(strings.iter().map(|text| text.to_string()));
1046 self
1047 }
1048
1049 pub fn build(self) -> Result<StopSequenceDecoder> {
1050 Ok(StopSequenceDecoder {
1051 sequence: Sequence::new(self.tokenizer.clone()),
1052 stop_token_ids_visible: self.stop_token_ids_visible,
1053 stop_token_ids_hidden: self.stop_token_ids_hidden,
1054 stop_sequences_visible: self.stop_sequences_visible,
1055 stop_sequences_hidden: self.stop_sequences_hidden,
1056 stopped: false,
1057 state: String::new(),
1058 })
1059 }
1060}