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 with_options(self, options: TokenizerOptions) -> Self
138 where
139 Self: Sized,
140 {
141 let _ = options;
142 self
143 }
144 }
147}
148
149pub fn file_json_field<T: serde::de::DeserializeOwned>(
150 json_file_path: &Path,
151 field_name: &str,
152) -> anyhow::Result<T> {
153 let file = File::open(json_file_path)
154 .with_context(|| format!("Failed to open file: {:?}", json_file_path))?;
155 let reader = BufReader::new(file);
156
157 let json_data: serde_json::Value = serde_json::from_reader(reader)
158 .with_context(|| format!("Failed to parse JSON from file: {:?}", json_file_path))?;
159
160 let map = json_data.as_object().ok_or_else(|| {
161 anyhow::anyhow!("JSON root is not an object in file: {:?}", json_file_path)
162 })?;
163
164 let field_value = map.get(field_name).ok_or_else(|| {
165 anyhow::anyhow!(
166 "Field '{}' not found in JSON file: {:?}",
167 field_name,
168 json_file_path
169 )
170 })?;
171
172 serde_json::from_value(field_value.clone()).with_context(|| {
173 format!(
174 "Failed to deserialize field '{}' (value: {:?}) to the expected type from file: {:?}",
175 field_name, field_value, json_file_path
176 )
177 })
178}
179
180pub fn log_json_err(filename: &str, json: &str, err: &serde_json::Error) {
181 const ERROR_PREFIX: &str = ">> ";
182
183 if !(err.is_syntax() || err.is_data()) {
184 return;
185 }
186
187 let line = err.line().saturating_sub(1);
188 let column = err.column().saturating_sub(1);
189
190 let json_lines: Vec<&str> = json.lines().collect();
191 if json_lines.is_empty() {
192 tracing::error!("JSON parsing error in {filename}: File is empty.");
193 return;
194 }
195
196 let start_index = line.saturating_sub(2);
197 let end_index = line.saturating_add(3).min(json_lines.len());
198
199 let mut context_lines: Vec<String> = (start_index..end_index)
200 .map(|i| {
201 if i == line {
202 format!("{ERROR_PREFIX}{}", json_lines[i])
203 } else {
204 format!("{:06} {}", i + 1, json_lines[i])
205 }
206 })
207 .collect();
208
209 let col_indicator = "_".to_string().repeat(column + ERROR_PREFIX.len()) + "^";
210 let error_in_context_idx = line - start_index;
211 if error_in_context_idx < context_lines.len() {
212 context_lines.insert(error_in_context_idx + 1, col_indicator);
213 }
214
215 tracing::error!(
216 "JSON parsing error in {filename}: Line {}, column {}:\n{}",
217 err.line(),
218 err.column(),
219 context_lines.join("\n")
220 );
221}
222
223impl Encoding {
224 pub fn get_hash(&self) -> u64 {
225 let mut hasher = DefaultHasher::new();
226 self.hash(&mut hasher);
227 hasher.finish()
228 }
229}
230
231#[derive(Debug, Clone, Copy, Default)]
235pub struct TokenizerOptions {
236 pub add_special_tokens: bool,
242}
243
244#[derive(Clone)]
246pub struct Tokenizer(Arc<dyn traits::Tokenizer>);
247
248impl Tokenizer {
249 pub fn from_file(file_path: &str) -> Result<Tokenizer> {
250 Ok(Tokenizer(create_tokenizer_from_file(file_path)?))
251 }
252
253 pub fn from_file_with_options(file_path: &str, options: TokenizerOptions) -> Result<Tokenizer> {
254 Ok(Tokenizer(create_tokenizer_from_file_with_options(
255 file_path, options,
256 )?))
257 }
258
259 pub fn decode_stream(
261 &self,
262 prompt_token_ids: &[TokenIdType],
263 skip_special_tokens: bool,
264 ) -> DecodeStream {
265 DecodeStream::new(self.0.clone(), prompt_token_ids, skip_special_tokens)
266 }
267}
268
269impl Deref for Tokenizer {
270 type Target = Arc<dyn traits::Tokenizer>;
271
272 fn deref(&self) -> &Self::Target {
273 &self.0
274 }
275}
276
277impl From<Arc<dyn traits::Tokenizer>> for Tokenizer {
278 fn from(tokenizer: Arc<dyn traits::Tokenizer>) -> Self {
279 Tokenizer(tokenizer)
280 }
281}
282
283impl<T> From<Arc<T>> for Tokenizer
284where
285 T: traits::Tokenizer + 'static, {
287 fn from(tokenizer: Arc<T>) -> Self {
288 Tokenizer(tokenizer)
289 }
290}
291
292pub fn create_tokenizer_from_file(file_path: &str) -> Result<Arc<dyn traits::Tokenizer>> {
299 create_tokenizer_from_file_with_options(file_path, Default::default())
300}
301
302pub fn create_tokenizer_from_file_with_options(
309 file_path: &str,
310 options: TokenizerOptions,
311) -> Result<Arc<dyn traits::Tokenizer>> {
312 use traits::Tokenizer as _;
313
314 let path = Path::new(file_path);
315 let extension = path
316 .extension()
317 .and_then(std::ffi::OsStr::to_str)
318 .ok_or_else(|| Error::msg("Failed to read file extension".to_string()))?;
319
320 match extension {
321 "json" => {
322 let tokenizer = HuggingFaceTokenizer::from_file(file_path)?.with_options(options);
323 Ok(Arc::new(tokenizer))
324 }
325 "model" | "tiktoken" => {
326 let tokenizer = TikTokenTokenizer::from_file_auto(file_path)?.with_options(options);
327 Ok(Arc::new(tokenizer))
328 }
329 _ => Err(Error::msg(format!(
330 "Unsupported tokenizer file type: .{extension}"
331 ))),
332 }
333}
334
335const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
341
342pub struct DecodeStream {
348 tokenizer: Arc<dyn traits::Tokenizer>,
350
351 skip_special_tokens: bool,
352 all_token_ids: Vec<u32>,
364
365 prefix_offset: usize,
366
367 read_offset: usize,
368}
369
370impl DecodeStream {
371 pub fn new(
372 tokenizer: Arc<dyn traits::Tokenizer>,
373 prompt_token_ids: &[TokenIdType],
374 skip_special_tokens: bool,
375 ) -> Self {
376 let num_input_tokens = prompt_token_ids.len();
377 let prompt_token_ids = prompt_token_ids.to_vec();
378 Self {
379 tokenizer,
380 skip_special_tokens,
381 all_token_ids: prompt_token_ids,
382 prefix_offset: num_input_tokens
383 .saturating_sub(INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET),
384 read_offset: num_input_tokens,
385 }
386 }
387
388 pub fn step(&mut self, id: u32) -> Result<Option<String>> {
398 self.all_token_ids.push(id);
399
400 let prefix_text: String = self
401 .tokenizer
402 .decode(
403 &self.all_token_ids[self.prefix_offset..self.read_offset],
404 self.skip_special_tokens,
405 )?
406 .into();
407
408 let new_result = self.tokenizer.decode(
409 &self.all_token_ids[self.prefix_offset..],
410 self.skip_special_tokens,
411 )?;
412
413 let new_text = new_result.as_str();
414 if new_text.len() > prefix_text.len() && !new_result.is_partial() {
415 let emitted = new_text[prefix_text.len()..].to_string();
416
417 self.prefix_offset = self.read_offset;
418 self.read_offset = self.all_token_ids.len();
419
420 Ok(Some(emitted))
421 } else {
422 Ok(None)
423 }
424 }
425}
426
427pub struct Sequence {
429 tokenizer: Tokenizer,
431
432 token_ids: Vec<TokenIdType>,
434
435 prefix_offset: usize,
437
438 read_offset: usize,
440}
441
442impl std::fmt::Debug for Sequence {
443 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
444 f.debug_struct("Sequence")
445 .field("tokenizer", &"Arc<dyn Tokenizer>")
446 .field(
447 "token_ids",
448 &format_args!("{}", {
449 let token_ids = self.token_ids();
450 if token_ids.len() <= 20 {
451 format!("{:?}", token_ids)
452 } else {
453 let first_ten = &token_ids[..10];
454 let last_ten = &token_ids[token_ids.len() - 10..];
455 format!("{:?} ... {:?}", first_ten, last_ten)
456 }
457 }),
458 )
459 .field("prefix_offset", &self.prefix_offset)
460 .field("read_offset", &self.read_offset)
461 .field("token count", &self.token_ids.len())
462 .finish()
463 }
464}
465
466impl Sequence {
467 pub fn new(tokenizer: Tokenizer) -> Self {
468 Self {
469 tokenizer,
470 token_ids: Vec::new(),
471 prefix_offset: 0,
472 read_offset: 0,
473 }
474 }
475
476 pub fn is_empty(&self) -> bool {
477 self.token_ids.is_empty()
478 }
479
480 pub fn len(&self) -> usize {
481 self.token_ids.len()
482 }
483
484 pub fn clear(&mut self) {
485 self.token_ids.clear();
486 self.prefix_offset = 0;
487 self.read_offset = 0;
488 }
489
490 pub fn append_text(&mut self, input: &str) -> Result<()> {
491 let encoding = self.tokenizer.encode(input)?;
496 self.token_ids.extend(encoding.token_ids());
497 Ok(())
498 }
499
500 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
504 self.token_ids.push(token_id);
505 let prefix_text: String = self
508 .tokenizer
509 .decode(&self.token_ids[self.prefix_offset..self.read_offset], false)?
510 .into();
511
512 let new_result = self
513 .tokenizer
514 .decode(&self.token_ids[self.prefix_offset..], false)?;
515
516 let new_text = new_result.as_str();
517
518 let mut prefix_text_len = prefix_text.len();
522 while !new_text.is_char_boundary(prefix_text_len) && prefix_text_len > 0 {
523 prefix_text_len -= 1;
524 }
525 let prefix_text_len = prefix_text_len;
526
527 if new_text.len() > prefix_text.len() {
528 if new_result.is_partial() {
529 return Ok("".to_string());
530 } else {
531 let new_text = new_text[prefix_text_len..]
533 .to_string()
534 .replace('\u{FFFD}', "");
535 self.prefix_offset = self.read_offset;
536 self.read_offset = self.token_ids.len();
537 return Ok(new_text);
538 }
539 }
540
541 Ok("".to_string())
542 }
543
544 pub fn tokenizer(&self) -> Tokenizer {
545 self.tokenizer.clone()
546 }
547
548 pub fn token_ids(&self) -> &[TokenIdType] {
549 &self.token_ids
550 }
551
552 pub fn text(&self) -> Result<String> {
553 Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
557 }
558}
559
560pub enum SequenceDecoderOutput {
563 Text(String),
565
566 Held,
569
570 Stopped,
573
574 StoppedWithText(String),
578}
579
580#[derive(Debug)]
586pub struct StopSequenceDecoder {
587 sequence: Sequence,
589
590 stop_token_ids_visible: Vec<TokenIdType>,
593
594 stop_token_ids_hidden: Vec<TokenIdType>,
597
598 #[allow(dead_code)]
601 stop_sequences_visible: Vec<String>,
602
603 stop_sequences_hidden: Vec<String>,
606
607 stopped: bool,
610
611 state: String,
614}
615
616impl StopSequenceDecoder {
617 pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
619 StopSequenceDecoderBuilder::new(tokenizer)
620 }
621
622 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<SequenceDecoderOutput> {
624 if self.stopped {
625 return Err(Error::msg("Decoder is stopped"));
626 }
627
628 let text = self.sequence.append_token_id(token_id)?;
630
631 self.state.push_str(text.as_str());
633
634 let mut stop: bool = false;
635 let mut visible: bool = false;
636
637 if self.stop_token_ids_visible.contains(&token_id) {
638 stop = true;
639 visible = true;
640 }
641
642 if self.stop_token_ids_hidden.contains(&token_id) {
643 stop = true;
644 visible = false;
645 }
646
647 if stop {
648 self.stopped = true;
649 let state = std::mem::take(&mut self.state);
650 if visible {
651 return Ok(SequenceDecoderOutput::StoppedWithText(state));
652 }
653 return Ok(SequenceDecoderOutput::Stopped);
654 }
655
656 for stop_sequence in self.stop_sequences_hidden.iter() {
658 if stop_sequence.starts_with(&self.state) {
659 if stop_sequence == &self.state {
660 self.stopped = true;
662 return Ok(SequenceDecoderOutput::Stopped);
663 } else {
664 return Ok(SequenceDecoderOutput::Held);
665 }
666 }
667 }
668
669 let state = std::mem::take(&mut self.state);
670 Ok(SequenceDecoderOutput::Text(state))
671 }
672
673 pub fn is_empty(&self) -> bool {
674 self.sequence.token_ids.is_empty()
675 }
676
677 pub fn len(&self) -> usize {
678 self.sequence.token_ids.len()
679 }
680
681 pub fn is_complete(&self) -> bool {
682 self.stopped
683 }
684
685 pub fn close(&mut self) {
686 self.stopped = true;
687 }
688}
689
690pub struct StopSequenceDecoderBuilder {
691 tokenizer: Tokenizer,
692 stop_token_ids_visible: Vec<TokenIdType>,
693 stop_token_ids_hidden: Vec<TokenIdType>,
694 stop_sequences_visible: Vec<String>,
695 stop_sequences_hidden: Vec<String>,
696}
697
698impl StopSequenceDecoderBuilder {
699 pub fn new(tokenizer: Tokenizer) -> Self {
700 Self {
701 tokenizer,
702 stop_token_ids_visible: Vec::new(),
703 stop_token_ids_hidden: Vec::new(),
704 stop_sequences_visible: Vec::new(),
705 stop_sequences_hidden: Vec::new(),
706 }
707 }
708
709 pub fn add_stop_token_id_visible(mut self, token_id: TokenIdType) -> Self {
711 self.stop_token_ids_visible.push(token_id);
712 self
713 }
714
715 pub fn add_stop_token_ids_visible(mut self, token_ids: &[TokenIdType]) -> Self {
718 self.stop_token_ids_visible.extend(token_ids);
719 self
720 }
721
722 pub fn add_stop_token_id_hidden(mut self, token_id: TokenIdType) -> Self {
724 self.stop_token_ids_hidden.push(token_id);
725 self
726 }
727
728 pub fn add_stop_token_ids_hidden(mut self, token_ids: &[TokenIdType]) -> Self {
731 self.stop_token_ids_hidden.extend(token_ids);
732 self
733 }
734
735 pub fn add_stop_sequence_visible(mut self, text: &str) -> Self {
736 self.stop_sequences_visible.push(text.to_string());
737 self
738 }
739
740 pub fn add_stop_sequences_visible(mut self, strings: &[&str]) -> Self {
741 self.stop_sequences_visible
742 .extend(strings.iter().map(|text| text.to_string()));
743 self
744 }
745
746 pub fn add_stop_sequence_hidden(mut self, text: &str) -> Self {
747 self.stop_sequences_hidden.push(text.to_string());
748 self
749 }
750
751 pub fn add_stop_sequences_hidden(mut self, strings: &[&str]) -> Self {
752 self.stop_sequences_hidden
753 .extend(strings.iter().map(|text| text.to_string()));
754 self
755 }
756
757 pub fn build(self) -> Result<StopSequenceDecoder> {
758 Ok(StopSequenceDecoder {
759 sequence: Sequence::new(self.tokenizer.clone()),
760 stop_token_ids_visible: self.stop_token_ids_visible,
761 stop_token_ids_hidden: self.stop_token_ids_hidden,
762 stop_sequences_visible: self.stop_sequences_visible,
763 stop_sequences_hidden: self.stop_sequences_hidden,
764 stopped: false,
765 state: String::new(),
766 })
767 }
768}