use std::sync::Arc;
use tiktoken_rs::CoreBPE;
use tiktoken_rs::cl100k_base;
pub trait Tokenizer: Send + Sync {
fn count(&self, text: &str) -> u32;
fn name(&self) -> &str;
}
pub struct TiktokenTokenizer {
bpe: Arc<CoreBPE>,
name: &'static str,
}
impl TiktokenTokenizer {
pub fn cl100k() -> Result<Self, crate::ClassifiedError> {
let bpe = cl100k_base()
.map_err(|e| crate::ClassifiedError::Runtime(format!("tiktoken cl100k_base: {e}")))?;
Ok(Self {
bpe: Arc::new(bpe),
name: "tiktoken.cl100k",
})
}
}
impl Tokenizer for TiktokenTokenizer {
fn count(&self, text: &str) -> u32 {
self.bpe
.encode_with_special_tokens(text)
.len()
.try_into()
.unwrap_or(u32::MAX / 2)
}
fn name(&self) -> &str {
self.name
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct ApproxTokenizer;
impl Tokenizer for ApproxTokenizer {
fn count(&self, text: &str) -> u32 {
let chars = text.chars().count();
(chars as u32).div_ceil(4)
}
fn name(&self) -> &str {
"approx.chars_over_4"
}
}
#[cfg(test)]
#[allow(clippy::expect_used, clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn approx_counts_nonzero_for_ascii() {
let t = ApproxTokenizer;
assert_eq!(t.count(""), 0);
assert_eq!(t.count("abcd"), 1);
assert_eq!(t.count("abcde"), 2);
}
#[test]
fn approx_counts_multibyte_by_char_not_byte() {
let t = ApproxTokenizer;
assert_eq!(t.count("中"), 1);
assert_eq!(t.count("中国加油"), 1);
}
#[test]
fn tiktoken_counts_standard_phrase() {
let t = TiktokenTokenizer::cl100k().unwrap();
let n = t.count("Hello world");
assert!(n > 0 && n < 10, "expected a small token count, got {n}");
}
#[test]
fn tiktoken_empty_is_zero() {
let t = TiktokenTokenizer::cl100k().unwrap();
assert_eq!(t.count(""), 0);
}
}