use crate::error::FocrResult;
use super::Tokenizer;
use super::tiktoken::Tiktoken;
pub trait TokenizerOps {
fn encode(&self, text: &str) -> FocrResult<Vec<u32>>;
fn decode(&self, ids: &[u32]) -> FocrResult<String>;
fn decode_skip_special(&self, ids: &[u32]) -> FocrResult<String>;
#[must_use]
fn bos_id(&self) -> u32;
#[must_use]
fn eos_id(&self) -> u32;
#[must_use]
fn token_to_id(&self, content: &str) -> Option<u32>;
}
impl TokenizerOps for Tokenizer {
fn encode(&self, text: &str) -> FocrResult<Vec<u32>> {
Tokenizer::encode(self, text)
}
fn decode(&self, ids: &[u32]) -> FocrResult<String> {
Tokenizer::decode(self, ids)
}
fn decode_skip_special(&self, ids: &[u32]) -> FocrResult<String> {
Tokenizer::decode_skip_special(self, ids)
}
fn bos_id(&self) -> u32 {
Tokenizer::bos_id(self)
}
fn eos_id(&self) -> u32 {
Tokenizer::eos_id(self)
}
fn token_to_id(&self, content: &str) -> Option<u32> {
Tokenizer::token_to_id(self, content)
}
}
impl TokenizerOps for Tiktoken {
fn encode(&self, text: &str) -> FocrResult<Vec<u32>> {
Tiktoken::encode(self, text)
}
fn decode(&self, ids: &[u32]) -> FocrResult<String> {
Tiktoken::decode(self, ids)
}
fn decode_skip_special(&self, ids: &[u32]) -> FocrResult<String> {
Tiktoken::decode_skip_special(self, ids)
}
fn bos_id(&self) -> u32 {
Tiktoken::bos_id(self)
}
fn eos_id(&self) -> u32 {
Tiktoken::eos_id(self)
}
fn token_to_id(&self, content: &str) -> Option<u32> {
Tiktoken::token_to_id(self, content)
}
}
#[cfg(test)]
mod tests {
use super::super::{special, tests as bpe_tests};
use super::*;
use crate::tokenizer::tiktoken;
const _OBJECT_SAFE: fn(&dyn TokenizerOps) = |_| {};
fn bpe() -> Tokenizer {
Tokenizer::from_json_bytes(bpe_tests::tiny_json().as_bytes()).expect("tiny tokenizer loads")
}
fn b64_encode(bytes: &[u8]) -> String {
const ALPHABET: &[u8; 64] =
b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
let mut out = String::new();
for chunk in bytes.chunks(3) {
let b = [
chunk[0],
*chunk.get(1).unwrap_or(&0),
*chunk.get(2).unwrap_or(&0),
];
let n = (u32::from(b[0]) << 16) | (u32::from(b[1]) << 8) | u32::from(b[2]);
out.push(ALPHABET[(n >> 18) as usize & 63] as char);
out.push(ALPHABET[(n >> 12) as usize & 63] as char);
out.push(if chunk.len() > 1 {
ALPHABET[(n >> 6) as usize & 63] as char
} else {
'='
});
out.push(if chunk.len() > 2 {
ALPHABET[n as usize & 63] as char
} else {
'='
});
}
out
}
fn synthetic_qwen_tiktoken() -> Vec<u8> {
let mut file = String::new();
for b in 0u8..=255 {
file.push_str(&b64_encode(&[b]));
file.push(' ');
file.push_str(&b.to_string());
file.push('\n');
}
file.push_str(&format!("{} 256\n", b64_encode(b"ab")));
file.push_str(&format!("{} 257\n", b64_encode(b"abc")));
for r in 258u32..151_643 {
let filler = [0xC0, 0xC1, (r >> 16) as u8, (r >> 8) as u8, r as u8];
file.push_str(&format!("{} {r}\n", b64_encode(&filler)));
}
file.into_bytes()
}
fn tik() -> Tiktoken {
Tiktoken::from_qwen_tiktoken(&synthetic_qwen_tiktoken())
.expect("synthetic qwen.tiktoken loads")
}
fn assert_round_trip(tk: &dyn TokenizerOps, text: &str) {
let ids = tk.encode(text).expect("encode");
assert_eq!(
tk.decode(&ids).expect("decode"),
text,
"round trip {text:?}"
);
}
#[test]
fn both_impls_round_trip_through_dyn() {
let bpe = bpe();
let tik = tik();
for text in ["abc", "ab", " a"] {
assert_round_trip(&bpe, text);
}
for text in ["abc", "abd", "Hello, world! 123", "café"] {
assert_round_trip(&tik, text);
}
}
#[test]
fn dyn_dispatch_matches_inherent_exactly() {
let bpe = bpe();
let dyn_bpe: &dyn TokenizerOps = &bpe;
assert_eq!(
dyn_bpe.encode("ab<image>c").unwrap(),
Tokenizer::encode(&bpe, "ab<image>c").unwrap()
);
let tik = tik();
let dyn_tik: &dyn TokenizerOps = &tik;
assert_eq!(
dyn_tik.encode("ab<|endoftext|>").unwrap(),
Tiktoken::encode(&tik, "ab<|endoftext|>").unwrap()
);
}
#[test]
fn bpe_encode_and_merges_through_trait() {
let t = bpe();
let t: &dyn TokenizerOps = &t;
assert_eq!(t.encode("abc").unwrap(), vec![8]);
assert_eq!(t.encode("ab").unwrap(), vec![7]);
assert_eq!(t.encode("ba").unwrap(), vec![1, 0]);
assert_eq!(t.encode("ab<image>c").unwrap(), vec![7, 128815, 2]);
}
#[test]
fn bpe_special_handling_through_trait() {
let t = bpe();
let t: &dyn TokenizerOps = &t;
let ids = t.encode("ab<image>c").unwrap();
assert_eq!(t.decode(&ids).unwrap(), "ab<image>c");
assert_eq!(t.decode_skip_special(&ids).unwrap(), "abc");
let ids2 = t.encode("a<|x|>b").unwrap();
assert_eq!(t.decode_skip_special(&ids2).unwrap(), "a<|x|>b");
}
#[test]
fn bpe_id_lookups_through_trait() {
let t = bpe();
let t: &dyn TokenizerOps = &t;
assert_eq!(t.bos_id(), special::BOS);
assert_eq!(t.eos_id(), special::EOS);
assert_eq!(t.token_to_id("<image>"), Some(special::IMAGE));
assert_eq!(t.token_to_id("abc"), Some(8));
assert_eq!(t.token_to_id("no-such-token"), None);
}
#[test]
fn tiktoken_encode_and_merges_through_trait() {
let t = tik();
let t: &dyn TokenizerOps = &t;
assert_eq!(t.encode("abc").unwrap(), vec![257]);
assert_eq!(t.encode("abd").unwrap(), vec![256, u32::from(b'd')]);
assert_eq!(
t.encode("ba").unwrap(),
vec![u32::from(b'b'), u32::from(b'a')]
);
assert_eq!(
t.encode("12").unwrap(),
vec![u32::from(b'1'), u32::from(b'2')]
);
}
#[test]
fn tiktoken_special_handling_through_trait() {
let t = tik();
let t: &dyn TokenizerOps = &t;
assert_eq!(
t.encode("ab<|endoftext|>").unwrap(),
vec![256, tiktoken::ENDOFTEXT]
);
let ids = t.encode("ab<img></img>").unwrap();
assert_eq!(ids, vec![256, tiktoken::IMG_START, tiktoken::IMG_END]);
assert_eq!(t.decode(&ids).unwrap(), "ab<img></img>");
assert_eq!(t.decode_skip_special(&ids).unwrap(), "ab");
}
#[test]
fn tiktoken_id_lookups_through_trait() {
let t = tik();
let t: &dyn TokenizerOps = &t;
assert_eq!(t.bos_id(), tiktoken::ENDOFTEXT);
assert_eq!(t.eos_id(), tiktoken::ENDOFTEXT);
assert_eq!(t.token_to_id("<imgpad>"), Some(tiktoken::IMG_PAD));
assert_eq!(t.token_to_id("<|im_end|>"), Some(tiktoken::IM_END));
assert_eq!(t.token_to_id("ab"), Some(256));
assert_eq!(t.token_to_id("no-such-token"), None);
}
#[test]
fn real_baidu_tokenizer_through_trait() {
let Some(t) = bpe_tests::load_real() else {
return;
};
let t: &dyn TokenizerOps = &t;
assert_eq!(t.encode("<image>").unwrap(), vec![special::IMAGE]);
assert_eq!(t.bos_id(), special::BOS);
assert_eq!(t.eos_id(), special::EOS);
assert_eq!(t.token_to_id("<|grounding|>"), Some(special::GROUNDING));
assert_round_trip(t, "The quick brown fox jumps over the lazy dog.");
}
fn load_real_tiktoken() -> Option<Tiktoken> {
let p = std::env::var("FOCR_GOT_TIKTOKEN").ok()?;
let bytes = std::fs::read(p).ok()?;
Some(Tiktoken::from_qwen_tiktoken(&bytes).expect("real qwen.tiktoken must parse"))
}
#[test]
fn real_got_tiktoken_through_trait() {
let Some(t) = load_real_tiktoken() else {
return;
};
let t: &dyn TokenizerOps = &t;
assert_eq!(
t.encode("1234567890").unwrap(),
vec![16, 17, 18, 19, 20, 21, 22, 23, 24, 15]
);
assert_eq!(t.bos_id(), tiktoken::ENDOFTEXT);
assert_eq!(t.eos_id(), tiktoken::ENDOFTEXT);
assert_eq!(t.token_to_id("<imgpad>"), Some(tiktoken::IMG_PAD));
assert_round_trip(t, "Hello, world!");
}
}