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 }
136}
137
138pub fn file_json_field<T: serde::de::DeserializeOwned>(
139 json_file_path: &Path,
140 field_name: &str,
141) -> anyhow::Result<T> {
142 let file = File::open(json_file_path)
143 .with_context(|| format!("Failed to open file: {:?}", json_file_path))?;
144 let reader = BufReader::new(file);
145
146 let json_data: serde_json::Value = serde_json::from_reader(reader)
147 .with_context(|| format!("Failed to parse JSON from file: {:?}", json_file_path))?;
148
149 let map = json_data.as_object().ok_or_else(|| {
150 anyhow::anyhow!("JSON root is not an object in file: {:?}", json_file_path)
151 })?;
152
153 let field_value = map.get(field_name).ok_or_else(|| {
154 anyhow::anyhow!(
155 "Field '{}' not found in JSON file: {:?}",
156 field_name,
157 json_file_path
158 )
159 })?;
160
161 serde_json::from_value(field_value.clone()).with_context(|| {
162 format!(
163 "Failed to deserialize field '{}' (value: {:?}) to the expected type from file: {:?}",
164 field_name, field_value, json_file_path
165 )
166 })
167}
168
169pub fn log_json_err(filename: &str, json: &str, err: &serde_json::Error) {
170 const ERROR_PREFIX: &str = ">> ";
171
172 if !(err.is_syntax() || err.is_data()) {
173 return;
174 }
175
176 let line = err.line().saturating_sub(1);
177 let column = err.column().saturating_sub(1);
178
179 let json_lines: Vec<&str> = json.lines().collect();
180 if json_lines.is_empty() {
181 tracing::error!("JSON parsing error in {filename}: File is empty.");
182 return;
183 }
184
185 let start_index = line.saturating_sub(2);
186 let end_index = line.saturating_add(3).min(json_lines.len());
187
188 let mut context_lines: Vec<String> = (start_index..end_index)
189 .map(|i| {
190 if i == line {
191 format!("{ERROR_PREFIX}{}", json_lines[i])
192 } else {
193 format!("{:06} {}", i + 1, json_lines[i])
194 }
195 })
196 .collect();
197
198 let col_indicator = "_".to_string().repeat(column + ERROR_PREFIX.len()) + "^";
199 let error_in_context_idx = line - start_index;
200 if error_in_context_idx < context_lines.len() {
201 context_lines.insert(error_in_context_idx + 1, col_indicator);
202 }
203
204 tracing::error!(
205 "JSON parsing error in {filename}: Line {}, column {}:\n{}",
206 err.line(),
207 err.column(),
208 context_lines.join("\n")
209 );
210}
211
212impl Encoding {
213 pub fn get_hash(&self) -> u64 {
214 let mut hasher = DefaultHasher::new();
215 self.hash(&mut hasher);
216 hasher.finish()
217 }
218}
219
220#[derive(Clone)]
222pub struct Tokenizer(Arc<dyn traits::Tokenizer>);
223
224impl Tokenizer {
225 pub fn from_file(file_path: &str) -> Result<Tokenizer> {
226 Ok(Tokenizer(create_tokenizer_from_file(file_path)?))
227 }
228
229 pub fn decode_stream(
231 &self,
232 prompt_token_ids: &[TokenIdType],
233 skip_special_tokens: bool,
234 ) -> DecodeStream {
235 DecodeStream::new(self.0.clone(), prompt_token_ids, skip_special_tokens)
236 }
237}
238
239impl Deref for Tokenizer {
240 type Target = Arc<dyn traits::Tokenizer>;
241
242 fn deref(&self) -> &Self::Target {
243 &self.0
244 }
245}
246
247impl From<Arc<dyn traits::Tokenizer>> for Tokenizer {
248 fn from(tokenizer: Arc<dyn traits::Tokenizer>) -> Self {
249 Tokenizer(tokenizer)
250 }
251}
252
253impl<T> From<Arc<T>> for Tokenizer
254where
255 T: traits::Tokenizer + 'static, {
257 fn from(tokenizer: Arc<T>) -> Self {
258 Tokenizer(tokenizer)
259 }
260}
261
262pub fn create_tokenizer_from_file(file_path: &str) -> Result<Arc<dyn traits::Tokenizer>> {
269 let path = Path::new(file_path);
270 let extension = path
271 .extension()
272 .and_then(std::ffi::OsStr::to_str)
273 .ok_or_else(|| Error::msg("Failed to read file extension".to_string()))?;
274
275 match extension {
276 "json" => {
277 let tokenizer = HuggingFaceTokenizer::from_file(file_path)?;
278 Ok(Arc::new(tokenizer))
279 }
280 "model" | "tiktoken" => {
281 let tokenizer = TikTokenTokenizer::from_file_auto(file_path)?;
282 Ok(Arc::new(tokenizer))
283 }
284 _ => Err(Error::msg(format!(
285 "Unsupported tokenizer file type: .{extension}"
286 ))),
287 }
288}
289
290const INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET: usize = 5;
296
297pub struct DecodeStream {
303 tokenizer: Arc<dyn traits::Tokenizer>,
305
306 skip_special_tokens: bool,
307 all_token_ids: Vec<u32>,
319
320 prefix_offset: usize,
321
322 read_offset: usize,
323}
324
325impl DecodeStream {
326 pub fn new(
327 tokenizer: Arc<dyn traits::Tokenizer>,
328 prompt_token_ids: &[TokenIdType],
329 skip_special_tokens: bool,
330 ) -> Self {
331 let num_input_tokens = prompt_token_ids.len();
332 let prompt_token_ids = prompt_token_ids.to_vec();
333 Self {
334 tokenizer,
335 skip_special_tokens,
336 all_token_ids: prompt_token_ids,
337 prefix_offset: num_input_tokens
338 .saturating_sub(INITIAL_INCREMENTAL_DETOKENIZATION_OFFSET),
339 read_offset: num_input_tokens,
340 }
341 }
342
343 pub fn step(&mut self, id: u32) -> Result<Option<String>> {
353 self.all_token_ids.push(id);
354
355 let prefix_text: String = self
356 .tokenizer
357 .decode(
358 &self.all_token_ids[self.prefix_offset..self.read_offset],
359 self.skip_special_tokens,
360 )?
361 .into();
362
363 let new_result = self.tokenizer.decode(
364 &self.all_token_ids[self.prefix_offset..],
365 self.skip_special_tokens,
366 )?;
367
368 let new_text = new_result.as_str();
369 if new_text.len() > prefix_text.len() && !new_result.is_partial() {
370 let emitted = new_text[prefix_text.len()..].to_string();
371
372 self.prefix_offset = self.read_offset;
373 self.read_offset = self.all_token_ids.len();
374
375 Ok(Some(emitted))
376 } else {
377 Ok(None)
378 }
379 }
380}
381
382pub struct Sequence {
384 tokenizer: Tokenizer,
386
387 token_ids: Vec<TokenIdType>,
389
390 prefix_offset: usize,
392
393 read_offset: usize,
395}
396
397impl std::fmt::Debug for Sequence {
398 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
399 f.debug_struct("Sequence")
400 .field("tokenizer", &"Arc<dyn Tokenizer>")
401 .field(
402 "token_ids",
403 &format_args!("{}", {
404 let token_ids = self.token_ids();
405 if token_ids.len() <= 20 {
406 format!("{:?}", token_ids)
407 } else {
408 let first_ten = &token_ids[..10];
409 let last_ten = &token_ids[token_ids.len() - 10..];
410 format!("{:?} ... {:?}", first_ten, last_ten)
411 }
412 }),
413 )
414 .field("prefix_offset", &self.prefix_offset)
415 .field("read_offset", &self.read_offset)
416 .field("token count", &self.token_ids.len())
417 .finish()
418 }
419}
420
421impl Sequence {
422 pub fn new(tokenizer: Tokenizer) -> Self {
423 Self {
424 tokenizer,
425 token_ids: Vec::new(),
426 prefix_offset: 0,
427 read_offset: 0,
428 }
429 }
430
431 pub fn is_empty(&self) -> bool {
432 self.token_ids.is_empty()
433 }
434
435 pub fn len(&self) -> usize {
436 self.token_ids.len()
437 }
438
439 pub fn clear(&mut self) {
440 self.token_ids.clear();
441 self.prefix_offset = 0;
442 self.read_offset = 0;
443 }
444
445 pub fn append_text(&mut self, input: &str) -> Result<()> {
446 let encoding = self.tokenizer.encode(input)?;
451 self.token_ids.extend(encoding.token_ids());
452 Ok(())
453 }
454
455 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<String> {
459 self.token_ids.push(token_id);
460 let prefix_text: String = self
463 .tokenizer
464 .decode(&self.token_ids[self.prefix_offset..self.read_offset], false)?
465 .into();
466
467 let new_result = self
468 .tokenizer
469 .decode(&self.token_ids[self.prefix_offset..], false)?;
470
471 let new_text = new_result.as_str();
472
473 let mut prefix_text_len = prefix_text.len();
477 while !new_text.is_char_boundary(prefix_text_len) && prefix_text_len > 0 {
478 prefix_text_len -= 1;
479 }
480 let prefix_text_len = prefix_text_len;
481
482 if new_text.len() > prefix_text.len() {
483 if new_result.is_partial() {
484 return Ok("".to_string());
485 } else {
486 let new_text = new_text[prefix_text_len..]
488 .to_string()
489 .replace('\u{FFFD}', "");
490 self.prefix_offset = self.read_offset;
491 self.read_offset = self.token_ids.len();
492 return Ok(new_text);
493 }
494 }
495
496 Ok("".to_string())
497 }
498
499 pub fn tokenizer(&self) -> Tokenizer {
500 self.tokenizer.clone()
501 }
502
503 pub fn token_ids(&self) -> &[TokenIdType] {
504 &self.token_ids
505 }
506
507 pub fn text(&self) -> Result<String> {
508 Ok(self.tokenizer.decode(&self.token_ids, false)?.into())
512 }
513}
514
515pub enum SequenceDecoderOutput {
518 Text(String),
520
521 Held,
524
525 Stopped,
528
529 StoppedWithText(String),
533}
534
535#[derive(Debug)]
541pub struct StopSequenceDecoder {
542 sequence: Sequence,
544
545 stop_token_ids_visible: Vec<TokenIdType>,
548
549 stop_token_ids_hidden: Vec<TokenIdType>,
552
553 #[allow(dead_code)]
556 stop_sequences_visible: Vec<String>,
557
558 stop_sequences_hidden: Vec<String>,
561
562 stopped: bool,
565
566 state: String,
569}
570
571impl StopSequenceDecoder {
572 pub fn builder(tokenizer: Tokenizer) -> StopSequenceDecoderBuilder {
574 StopSequenceDecoderBuilder::new(tokenizer)
575 }
576
577 pub fn append_token_id(&mut self, token_id: TokenIdType) -> Result<SequenceDecoderOutput> {
579 if self.stopped {
580 return Err(Error::msg("Decoder is stopped"));
581 }
582
583 let text = self.sequence.append_token_id(token_id)?;
585
586 self.state.push_str(text.as_str());
588
589 let mut stop: bool = false;
590 let mut visible: bool = false;
591
592 if self.stop_token_ids_visible.contains(&token_id) {
593 stop = true;
594 visible = true;
595 }
596
597 if self.stop_token_ids_hidden.contains(&token_id) {
598 stop = true;
599 visible = false;
600 }
601
602 if stop {
603 self.stopped = true;
604 let state = std::mem::take(&mut self.state);
605 if visible {
606 return Ok(SequenceDecoderOutput::StoppedWithText(state));
607 }
608 return Ok(SequenceDecoderOutput::Stopped);
609 }
610
611 for stop_sequence in self.stop_sequences_hidden.iter() {
613 if stop_sequence.starts_with(&self.state) {
614 if stop_sequence == &self.state {
615 self.stopped = true;
617 return Ok(SequenceDecoderOutput::Stopped);
618 } else {
619 return Ok(SequenceDecoderOutput::Held);
620 }
621 }
622 }
623
624 let state = std::mem::take(&mut self.state);
625 Ok(SequenceDecoderOutput::Text(state))
626 }
627
628 pub fn is_empty(&self) -> bool {
629 self.sequence.token_ids.is_empty()
630 }
631
632 pub fn len(&self) -> usize {
633 self.sequence.token_ids.len()
634 }
635
636 pub fn is_complete(&self) -> bool {
637 self.stopped
638 }
639
640 pub fn close(&mut self) {
641 self.stopped = true;
642 }
643}
644
645pub struct StopSequenceDecoderBuilder {
646 tokenizer: Tokenizer,
647 stop_token_ids_visible: Vec<TokenIdType>,
648 stop_token_ids_hidden: Vec<TokenIdType>,
649 stop_sequences_visible: Vec<String>,
650 stop_sequences_hidden: Vec<String>,
651}
652
653impl StopSequenceDecoderBuilder {
654 pub fn new(tokenizer: Tokenizer) -> Self {
655 Self {
656 tokenizer,
657 stop_token_ids_visible: Vec::new(),
658 stop_token_ids_hidden: Vec::new(),
659 stop_sequences_visible: Vec::new(),
660 stop_sequences_hidden: Vec::new(),
661 }
662 }
663
664 pub fn add_stop_token_id_visible(mut self, token_id: TokenIdType) -> Self {
666 self.stop_token_ids_visible.push(token_id);
667 self
668 }
669
670 pub fn add_stop_token_ids_visible(mut self, token_ids: &[TokenIdType]) -> Self {
673 self.stop_token_ids_visible.extend(token_ids);
674 self
675 }
676
677 pub fn add_stop_token_id_hidden(mut self, token_id: TokenIdType) -> Self {
679 self.stop_token_ids_hidden.push(token_id);
680 self
681 }
682
683 pub fn add_stop_token_ids_hidden(mut self, token_ids: &[TokenIdType]) -> Self {
686 self.stop_token_ids_hidden.extend(token_ids);
687 self
688 }
689
690 pub fn add_stop_sequence_visible(mut self, text: &str) -> Self {
691 self.stop_sequences_visible.push(text.to_string());
692 self
693 }
694
695 pub fn add_stop_sequences_visible(mut self, strings: &[&str]) -> Self {
696 self.stop_sequences_visible
697 .extend(strings.iter().map(|text| text.to_string()));
698 self
699 }
700
701 pub fn add_stop_sequence_hidden(mut self, text: &str) -> Self {
702 self.stop_sequences_hidden.push(text.to_string());
703 self
704 }
705
706 pub fn add_stop_sequences_hidden(mut self, strings: &[&str]) -> Self {
707 self.stop_sequences_hidden
708 .extend(strings.iter().map(|text| text.to_string()));
709 self
710 }
711
712 pub fn build(self) -> Result<StopSequenceDecoder> {
713 Ok(StopSequenceDecoder {
714 sequence: Sequence::new(self.tokenizer.clone()),
715 stop_token_ids_visible: self.stop_token_ids_visible,
716 stop_token_ids_hidden: self.stop_token_ids_hidden,
717 stop_sequences_visible: self.stop_sequences_visible,
718 stop_sequences_hidden: self.stop_sequences_hidden,
719 stopped: false,
720 state: String::new(),
721 })
722 }
723}