1pub mod cache;
5pub mod fastokens;
6pub mod hf;
7pub mod tiktoken;
8
9use std::hash::{DefaultHasher, Hash, Hasher};
14use std::sync::Arc;
15use std::{fs::File, io::BufReader, ops::Deref, path::Path};
16
17use anyhow::Context as _;
18pub use anyhow::{Error, Result};
19
20pub use cache::{CacheTokenUsage, CacheTokenUsageFn, CachedTokenizer, L1CacheStats};
21pub use fastokens::FastTokenizer;
22pub use hf::HuggingFaceTokenizer;
23pub use tiktoken::TikTokenTokenizer;
24pub use traits::DecodeResult;
25
26pub type TokenIdType = u32;
27
28#[derive(Debug)]
30pub enum TokenizerType {
31 HuggingFace(String),
32 TikToken(String),
33}
34
35pub type Offsets = (usize, usize);
37
38#[derive(Debug, Clone)]
40pub enum Encoding {
41 Hf(Box<tokenizers::tokenizer::Encoding>),
43 Sp(Vec<TokenIdType>),
45}
46
47impl Encoding {
48 pub fn token_ids(&self) -> &[u32] {
49 match self {
50 Encoding::Hf(inner) => inner.get_ids(),
51 Encoding::Sp(inner) => inner,
52 }
53 }
54}
55
56impl Hash for Encoding {
57 fn hash<H: Hasher>(&self, state: &mut H) {
58 self.token_ids().hash(state);
59 }
60}
61
62pub mod traits {
63 use super::*;
64
65 pub trait Encoder: Send + Sync {
66 fn encode(&self, input: &str) -> Result<Encoding>;
67 fn encode_batch(&self, inputs: &[&str]) -> Result<Vec<Encoding>>;
68 }
69
70 #[derive(Debug, Clone, PartialEq, Eq, strum::EnumIs)]
77 pub enum DecodeResult {
78 Complete(String),
82 Partial(String),
85 }
86
87 impl DecodeResult {
88 pub fn as_str(&self) -> &str {
90 match self {
91 DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
92 }
93 }
94
95 pub fn from_decoded(text: String) -> Self {
97 if text.ends_with('\u{FFFD}') {
98 DecodeResult::Partial(text)
99 } else {
100 DecodeResult::Complete(text)
101 }
102 }
103 }
104
105 impl From<String> for DecodeResult {
106 fn from(text: String) -> Self {
107 DecodeResult::from_decoded(text)
108 }
109 }
110
111 impl From<DecodeResult> for String {
112 fn from(result: DecodeResult) -> Self {
113 match result {
114 DecodeResult::Complete(s) | DecodeResult::Partial(s) => s,
115 }
116 }
117 }
118
119 pub trait Decoder: Send + Sync {
125 fn decode(
126 &self,
127 token_ids: &[TokenIdType],
128 skip_special_tokens: bool,
129 ) -> Result<DecodeResult>;
130 }
131
132 pub trait Tokenizer: Encoder + Decoder {
133 fn validate_prefix_cache(&self) -> Result<()> {
138 Err(Error::msg("tokenizer does not support prefix caching"))
139 }
140
141 fn with_options(self, options: TokenizerOptions) -> Self
146 where
147 Self: Sized,
148 {
149 let _ = options;
150 self
151 }
152 }
155}
156
157pub fn file_json_field<T: serde::de::DeserializeOwned>(
158 json_file_path: &Path,
159 field_name: &str,
160) -> anyhow::Result<T> {
161 let file = File::open(json_file_path)
162 .with_context(|| format!("Failed to open file: {:?}", json_file_path))?;
163 let reader = BufReader::new(file);
164
165 let json_data: serde_json::Value = serde_json::from_reader(reader)
166 .with_context(|| format!("Failed to parse JSON from file: {:?}", json_file_path))?;
167
168 let map = json_data.as_object().ok_or_else(|| {
169 anyhow::anyhow!("JSON root is not an object in file: {:?}", json_file_path)
170 })?;
171
172 let field_value = map.get(field_name).ok_or_else(|| {
173 anyhow::anyhow!(
174 "Field '{}' not found in JSON file: {:?}",
175 field_name,
176 json_file_path
177 )
178 })?;
179
180 serde_json::from_value(field_value.clone()).with_context(|| {
181 format!(
182 "Failed to deserialize field '{}' (value: {:?}) to the expected type from file: {:?}",
183 field_name, field_value, json_file_path
184 )
185 })
186}
187
188pub fn log_json_err(filename: &str, json: &str, err: &serde_json::Error) {
189 const ERROR_PREFIX: &str = ">> ";
190
191 if !(err.is_syntax() || err.is_data()) {
192 return;
193 }
194
195 let line = err.line().saturating_sub(1);
196 let column = err.column().saturating_sub(1);
197
198 let json_lines: Vec<&str> = json.lines().collect();
199 if json_lines.is_empty() {
200 tracing::error!("JSON parsing error in {filename}: File is empty.");
201 return;
202 }
203
204 let start_index = line.saturating_sub(2);
205 let end_index = line.saturating_add(3).min(json_lines.len());
206
207 let mut context_lines: Vec<String> = (start_index..end_index)
208 .map(|i| {
209 if i == line {
210 format!("{ERROR_PREFIX}{}", json_lines[i])
211 } else {
212 format!("{:06} {}", i + 1, json_lines[i])
213 }
214 })
215 .collect();
216
217 let col_indicator = "_".to_string().repeat(column + ERROR_PREFIX.len()) + "^";
218 let error_in_context_idx = line - start_index;
219 if error_in_context_idx < context_lines.len() {
220 context_lines.insert(error_in_context_idx + 1, col_indicator);
221 }
222
223 tracing::error!(
224 "JSON parsing error in {filename}: Line {}, column {}:\n{}",
225 err.line(),
226 err.column(),
227 context_lines.join("\n")
228 );
229}
230
231impl Encoding {
232 pub fn get_hash(&self) -> u64 {
233 let mut hasher = DefaultHasher::new();
234 self.hash(&mut hasher);
235 hasher.finish()
236 }
237}
238
239#[derive(Debug, Clone, Copy, Default)]
243pub struct TokenizerOptions {
244 pub add_special_tokens: bool,
250}
251
252#[derive(Clone)]
254pub struct Tokenizer(Arc<dyn traits::Tokenizer>);
255
256impl Tokenizer {
257 pub fn from_file(file_path: &str) -> Result<Tokenizer> {
258 Ok(Tokenizer(create_tokenizer_from_file(file_path)?))
259 }
260
261 pub fn from_file_with_options(file_path: &str, options: TokenizerOptions) -> Result<Tokenizer> {
262 Ok(Tokenizer(create_tokenizer_from_file_with_options(
263 file_path, options,
264 )?))
265 }
266
267 pub fn decode_stream(
269 &self,
270 prompt_token_ids: &[TokenIdType],
271 skip_special_tokens: bool,
272 ) -> DecodeStream {
273 DecodeStream::new(self.0.clone(), prompt_token_ids, skip_special_tokens)
274 }
275}
276
277impl Deref for Tokenizer {
278 type Target = Arc<dyn traits::Tokenizer>;
279
280 fn deref(&self) -> &Self::Target {
281 &self.0
282 }
283}
284
285impl From<Arc<dyn traits::Tokenizer>> for Tokenizer {
286 fn from(tokenizer: Arc<dyn traits::Tokenizer>) -> Self {
287 Tokenizer(tokenizer)
288 }
289}
290
291impl<T> From<Arc<T>> for Tokenizer
292where
293 T: traits::Tokenizer + 'static, {
295 fn from(tokenizer: Arc<T>) -> Self {
296 Tokenizer(tokenizer)
297 }
298}
299
300pub fn create_tokenizer_from_file(file_path: &str) -> Result<Arc<dyn traits::Tokenizer>> {
307 create_tokenizer_from_file_with_options(file_path, Default::default())
308}
309
310pub fn create_tokenizer_from_file_with_options(
317 file_path: &str,
318 options: TokenizerOptions,
319) -> Result<Arc<dyn traits::Tokenizer>> {
320 use traits::Tokenizer as _;
321
322 let path = Path::new(file_path);
323 let extension = path
324 .extension()
325 .and_then(std::ffi::OsStr::to_str)
326 .ok_or_else(|| Error::msg("Failed to read file extension".to_string()))?;
327
328 match extension {
329 "json" => {
330 let tokenizer = HuggingFaceTokenizer::from_file(file_path)?.with_options(options);
331 Ok(Arc::new(tokenizer))
332 }
333 "model" | "tiktoken" => {
334 let tokenizer = TikTokenTokenizer::from_file_auto(file_path)?.with_options(options);
335 Ok(Arc::new(tokenizer))
336 }
337 _ => Err(Error::msg(format!(
338 "Unsupported tokenizer file type: .{extension}"
339 ))),
340 }
341}
342
343const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
349
350pub struct DecodeStream {
356 tokenizer: Arc<dyn traits::Tokenizer>,
358
359 skip_special_tokens: bool,
360 all_token_ids: Vec<u32>,
372
373 prefix_offset: usize,
374
375 read_offset: usize,
376}
377
378impl DecodeStream {
379 pub fn new(
380 tokenizer: Arc<dyn traits::Tokenizer>,
381 prompt_token_ids: &[TokenIdType],
382 skip_special_tokens: bool,
383 ) -> Self {
384 let num_input_tokens = prompt_token_ids.len();
385 let prompt_token_ids = prompt_token_ids.to_vec();
386 Self {
387 tokenizer,
388 skip_special_tokens,
389 all_token_ids: prompt_token_ids,
390 prefix_offset: num_input_tokens
391 .saturating_sub(INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET),
392 read_offset: num_input_tokens,
393 }
394 }
395
396 pub fn step(&mut self, id: u32) -> Result<Option<String>> {
406 self.all_token_ids.push(id);
407
408 let prefix_text: String = self
409 .tokenizer
410 .decode(
411 &self.all_token_ids[self.prefix_offset..self.read_offset],
412 self.skip_special_tokens,
413 )?
414 .into();
415
416 let new_result = self.tokenizer.decode(
417 &self.all_token_ids[self.prefix_offset..],
418 self.skip_special_tokens,
419 )?;
420
421 let new_text = new_result.as_str();
422 if new_text.len() > prefix_text.len() && !new_result.is_partial() {
423 let emitted = new_text[prefix_text.len()..].to_string();
424
425 self.prefix_offset = self.read_offset;
426 self.read_offset = self.all_token_ids.len();
427
428 Ok(Some(emitted))
429 } else {
430 Ok(None)
431 }
432 }
433}
434
435pub struct Sequence {
437 tokenizer: Tokenizer,
439
440 token_ids: Vec<TokenIdType>,
442
443 prefix_offset: usize,
445
446 read_offset: usize,
448}
449
450impl std::fmt::Debug for Sequence {
451 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
452 f.debug_struct("Sequence")
453 .field("tokenizer", &"Arc<dyn Tokenizer>")
454 .field(
455 "token_ids",
456 &format_args!("{}", {
457 let token_ids = self.token_ids();
458 if token_ids.len() <= 20 {
459 format!("{:?}", token_ids)
460 } else {
461 let first_ten = &token_ids[..10];
462 let last_ten = &token_ids[token_ids.len() - 10..];
463 format!("{:?} ... {:?}", first_ten, last_ten)
464 }
465 }),
466 )
467 .field("prefix_offset", &self.prefix_offset)
468 .field("read_offset", &self.read_offset)
469 .field("token count", &self.token_ids.len())
470 .finish()
471 }
472}
473
474impl Sequence {
475 pub fn new(tokenizer: Tokenizer) -> Self {
476 Self {
477 tokenizer,
478 token_ids: Vec::new(),
479 prefix_offset: 0,
480 read_offset: 0,
481 }
482 }
483
484 pub fn is_empty(&self) -> bool {
485 self.token_ids.is_empty()
486 }
487
488 pub fn len(&self) -> usize {
489 self.token_ids.len()
490 }
491
492 pub fn clear(&mut self) {
493 self.token_ids.clear();
494 self.prefix_offset = 0;
495 self.read_offset = 0;
496 }
497
498 pub fn append_text(&mut self, input: &str) -> Result<()> {
499 let encoding = self.tokenizer.encode(input)?;
504 self.token_ids.extend(encoding.token_ids());
505 Ok(())
506 }
507
508 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
512 self.token_ids.push(token_id);
513 let prefix_text: String = self
516 .tokenizer
517 .decode(&self.token_ids[self.prefix_offset..self.read_offset], false)?
518 .into();
519
520 let new_result = self
521 .tokenizer
522 .decode(&self.token_ids[self.prefix_offset..], false)?;
523
524 let new_text = new_result.as_str();
525
526 let mut prefix_text_len = prefix_text.len();
530 while !new_text.is_char_boundary(prefix_text_len) && prefix_text_len > 0 {
531 prefix_text_len -= 1;
532 }
533 let prefix_text_len = prefix_text_len;
534
535 if new_text.len() > prefix_text.len() {
536 if new_result.is_partial() {
537 return Ok("".to_string());
538 } else {
539 let new_text = new_text[prefix_text_len..]
541 .to_string()
542 .replace('\u{FFFD}', "");
543 self.prefix_offset = self.read_offset;
544 self.read_offset = self.token_ids.len();
545 return Ok(new_text);
546 }
547 }
548
549 Ok("".to_string())
550 }
551
552 pub fn tokenizer(&self) -> Tokenizer {
553 self.tokenizer.clone()
554 }
555
556 pub fn token_ids(&self) -> &[TokenIdType] {
557 &self.token_ids
558 }
559
560 pub fn text(&self) -> Result<String> {
561 Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
565 }
566}
567
568pub enum SequenceDecoderOutput {
571 Text(String),
573
574 Held,
577
578 Stopped,
581
582 StoppedWithText(String),
586}
587
588#[derive(Debug)]
594pub struct StopSequenceDecoder {
595 sequence: Sequence,
597
598 stop_token_ids_visible: Vec<TokenIdType>,
601
602 stop_token_ids_hidden: Vec<TokenIdType>,
605
606 #[allow(dead_code)]
609 stop_sequences_visible: Vec<String>,
610
611 stop_sequences_hidden: Vec<String>,
614
615 stopped: bool,
618
619 state: String,
622}
623
624impl StopSequenceDecoder {
625 pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
627 StopSequenceDecoderBuilder::new(tokenizer)
628 }
629
630 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<SequenceDecoderOutput> {
632 if self.stopped {
633 return Err(Error::msg("Decoder is stopped"));
634 }
635
636 let text = self.sequence.append_token_id(token_id)?;
638
639 self.state.push_str(text.as_str());
641
642 let mut stop: bool = false;
643 let mut visible: bool = false;
644
645 if self.stop_token_ids_visible.contains(&token_id) {
646 stop = true;
647 visible = true;
648 }
649
650 if self.stop_token_ids_hidden.contains(&token_id) {
651 stop = true;
652 visible = false;
653 }
654
655 if stop {
656 self.stopped = true;
657 let state = std::mem::take(&mut self.state);
658 if visible {
659 return Ok(SequenceDecoderOutput::StoppedWithText(state));
660 }
661 return Ok(SequenceDecoderOutput::Stopped);
662 }
663
664 for stop_sequence in self.stop_sequences_hidden.iter() {
666 if stop_sequence.starts_with(&self.state) {
667 if stop_sequence == &self.state {
668 self.stopped = true;
670 return Ok(SequenceDecoderOutput::Stopped);
671 } else {
672 return Ok(SequenceDecoderOutput::Held);
673 }
674 }
675 }
676
677 let state = std::mem::take(&mut self.state);
678 Ok(SequenceDecoderOutput::Text(state))
679 }
680
681 pub fn is_empty(&self) -> bool {
682 self.sequence.token_ids.is_empty()
683 }
684
685 pub fn len(&self) -> usize {
686 self.sequence.token_ids.len()
687 }
688
689 pub fn is_complete(&self) -> bool {
690 self.stopped
691 }
692
693 pub fn close(&mut self) {
694 self.stopped = true;
695 }
696}
697
698pub struct StopSequenceDecoderBuilder {
699 tokenizer: Tokenizer,
700 stop_token_ids_visible: Vec<TokenIdType>,
701 stop_token_ids_hidden: Vec<TokenIdType>,
702 stop_sequences_visible: Vec<String>,
703 stop_sequences_hidden: Vec<String>,
704}
705
706impl StopSequenceDecoderBuilder {
707 pub fn new(tokenizer: Tokenizer) -> Self {
708 Self {
709 tokenizer,
710 stop_token_ids_visible: Vec::new(),
711 stop_token_ids_hidden: Vec::new(),
712 stop_sequences_visible: Vec::new(),
713 stop_sequences_hidden: Vec::new(),
714 }
715 }
716
717 pub fn add_stop_token_id_visible(mut self, token_id: TokenIdType) -> Self {
719 self.stop_token_ids_visible.push(token_id);
720 self
721 }
722
723 pub fn add_stop_token_ids_visible(mut self, token_ids: &[TokenIdType]) -> Self {
726 self.stop_token_ids_visible.extend(token_ids);
727 self
728 }
729
730 pub fn add_stop_token_id_hidden(mut self, token_id: TokenIdType) -> Self {
732 self.stop_token_ids_hidden.push(token_id);
733 self
734 }
735
736 pub fn add_stop_token_ids_hidden(mut self, token_ids: &[TokenIdType]) -> Self {
739 self.stop_token_ids_hidden.extend(token_ids);
740 self
741 }
742
743 pub fn add_stop_sequence_visible(mut self, text: &str) -> Self {
744 self.stop_sequences_visible.push(text.to_string());
745 self
746 }
747
748 pub fn add_stop_sequences_visible(mut self, strings: &[&str]) -> Self {
749 self.stop_sequences_visible
750 .extend(strings.iter().map(|text| text.to_string()));
751 self
752 }
753
754 pub fn add_stop_sequence_hidden(mut self, text: &str) -> Self {
755 self.stop_sequences_hidden.push(text.to_string());
756 self
757 }
758
759 pub fn add_stop_sequences_hidden(mut self, strings: &[&str]) -> Self {
760 self.stop_sequences_hidden
761 .extend(strings.iter().map(|text| text.to_string()));
762 self
763 }
764
765 pub fn build(self) -> Result<StopSequenceDecoder> {
766 Ok(StopSequenceDecoder {
767 sequence: Sequence::new(self.tokenizer.clone()),
768 stop_token_ids_visible: self.stop_token_ids_visible,
769 stop_token_ids_hidden: self.stop_token_ids_hidden,
770 stop_sequences_visible: self.stop_sequences_visible,
771 stop_sequences_hidden: self.stop_sequences_hidden,
772 stopped: false,
773 state: String::new(),
774 })
775 }
776}