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 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 fn validate_prefix_cache(&self) -> Result<()> {
204 Err(Error::msg("tokenizer does not support prefix caching"))
205 }
206
207 fn with_options(self, options: TokenizerOptions) -> Self
212 where
213 Self: Sized,
214 {
215 let _ = options;
216 self
217 }
218 fn vocab_size(&self) -> Option<usize> {
222 None
223 }
224
225 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 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 fn num_special_tokens_added(&self) -> Result<usize> {
266 Err(Error::msg(
267 "tokenizer backend does not support special token accounting",
268 ))
269 }
270 }
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#[derive(Debug, Clone, Copy, Default)]
360pub struct TokenizerOptions {
361 pub add_special_tokens: bool,
368}
369
370#[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 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, {
415 fn from(tokenizer: Arc<T>) -> Self {
416 Tokenizer(tokenizer)
417 }
418}
419
420pub 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
430pub 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
463const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
469
470pub struct DecodeStream {
476 tokenizer: Arc<dyn traits::Tokenizer>,
478
479 skip_special_tokens: bool,
480 all_token_ids: Vec<u32>,
492
493 prefix_offset: usize,
494
495 read_offset: usize,
496
497 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 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 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 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 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 requested_split
590 } else {
591 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 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
755pub struct Sequence {
757 tokenizer: Tokenizer,
759
760 token_ids: Vec<TokenIdType>,
762
763 prefix_offset: usize,
765
766 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 encoding = self.tokenizer.encode(input)?;
824 self.token_ids.extend(encoding.token_ids());
825 Ok(())
826 }
827
828 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
832 self.token_ids.push(token_id);
833 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 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 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 Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
885 }
886}
887
888pub enum SequenceDecoderOutput {
891 Text(String),
893
894 Held,
897
898 Stopped,
901
902 StoppedWithText(String),
906}
907
908#[derive(Debug)]
914pub struct StopSequenceDecoder {
915 sequence: Sequence,
917
918 stop_token_ids_visible: Vec<TokenIdType>,
921
922 stop_token_ids_hidden: Vec<TokenIdType>,
925
926 #[allow(dead_code)]
929 stop_sequences_visible: Vec<String>,
930
931 stop_sequences_hidden: Vec<String>,
934
935 stopped: bool,
938
939 state: String,
942}
943
944impl StopSequenceDecoder {
945 pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
947 StopSequenceDecoderBuilder::new(tokenizer)
948 }
949
950 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 let text = self.sequence.append_token_id(token_id)?;
958
959 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 for stop_sequence in self.stop_sequences_hidden.iter() {
986 if stop_sequence.starts_with(&self.state) {
987 if stop_sequence == &self.state {
988 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 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 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 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 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}