use rustc_hash::FxHashMap;
use super::any_tokenizer::{AnyTokenizer, Backend};
use super::policy::SpecialPolicy;
use super::spm::{SpmPrefixScheme, SpmTokenizer, NEVER_MERGE};
use super::tokenizer::{
Tokenizer, TokenizerError, CL100K_BASE_PATTERN, DEEPSEEK_V3_PATTERNS, GPT2_PATTERN,
LLAMA3_PATTERN, MISTRAL_V3_PATTERN, O200K_BASE_PATTERN,
};
use super::vocab::{load_spm_vocab, place_special_pieces};
use super::whisper::{whisper_special_tokens, WhisperVariant};
pub const CL100K_BASE_VOCAB: &[u8] = include_bytes!("../../vocabs/cl100k_base.tiktoken");
pub const O200K_BASE_VOCAB: &[u8] = include_bytes!("../../vocabs/o200k_base.tiktoken");
pub const LLAMA3_VOCAB: &[u8] = include_bytes!("../../vocabs/llama3.tiktoken");
pub const DEEPSEEK_V3_VOCAB: &[u8] = include_bytes!("../../vocabs/deepseek_v3.tiktoken");
pub const MISTRAL_SPM_VOCAB: &[u8] = include_bytes!("../../vocabs/mistral.spm");
pub const MISTRAL_V2_SPM_VOCAB: &[u8] = include_bytes!("../../vocabs/mistral_v2.spm");
pub const MISTRAL_V3_VOCAB: &[u8] = include_bytes!("../../vocabs/mistral_v3_tekken.tiktoken");
pub const WHISPER_VOCAB: &[u8] = include_bytes!("../../vocabs/whisper.tiktoken");
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum PretrainedVocab {
Cl100kBase,
O200kBase,
Llama3,
DeepseekV3,
MistralV1,
MistralV2,
MistralV3,
WhisperV1,
WhisperV2,
WhisperV3,
}
impl PretrainedVocab {
pub fn from_name(name: &str) -> Option<Self> {
match name {
"cl100k_base" => Some(Self::Cl100kBase),
"o200k_base" => Some(Self::O200kBase),
"llama3" | "llama3.1" | "llama3.2" | "llama3.3" => Some(Self::Llama3),
"deepseek_v3" | "deepseek-v3" => Some(Self::DeepseekV3),
"mistral" | "mistral_v1" => Some(Self::MistralV1),
"mistral_v2" => Some(Self::MistralV2),
"mistral_v3" => Some(Self::MistralV3),
"whisper_v1" | "whisper-v1" | "whisper-multilingual-v1" => Some(Self::WhisperV1),
"whisper" | "whisper_v2" | "whisper-v2" | "whisper-multilingual" => {
Some(Self::WhisperV2)
}
"whisper_v3" | "whisper-v3" | "whisper-large-v3" => Some(Self::WhisperV3),
_ => None,
}
}
pub fn supported_names() -> &'static [&'static str] {
&[
"cl100k_base",
"o200k_base",
"llama3",
"llama3.1",
"llama3.2",
"llama3.3",
"deepseek_v3",
"deepseek-v3",
"mistral",
"mistral_v1",
"mistral_v2",
"mistral_v3",
"whisper",
"whisper_v1",
"whisper_v2",
"whisper_v3",
]
}
}
pub fn from_pretrained(name: &str) -> Result<AnyTokenizer, TokenizerError> {
from_vocab(resolve_vocab(name)?)
}
fn resolve_vocab(name: &str) -> Result<PretrainedVocab, TokenizerError> {
PretrainedVocab::from_name(name).ok_or_else(|| {
TokenizerError::UnknownPretrained(format!(
"{}. Supported: {}",
name,
PretrainedVocab::supported_names().join(", ")
))
})
}
pub fn from_vocab(vocab: PretrainedVocab) -> Result<AnyTokenizer, TokenizerError> {
let special = special_tokens(vocab);
let named = special.clone();
match vocab {
PretrainedVocab::MistralV1 => {
return spm_from_vocab(MISTRAL_SPM_VOCAB, vocab, special, named)
}
PretrainedVocab::MistralV2 => {
return spm_from_vocab(MISTRAL_V2_SPM_VOCAB, vocab, special, named)
}
_ => {}
}
let pats = patterns(vocab).unwrap_or(&[]);
let tokenizer = match vocab {
PretrainedVocab::Cl100kBase => {
Tokenizer::from_bytes_chain(CL100K_BASE_VOCAB, pats, special)
}
PretrainedVocab::O200kBase => Tokenizer::from_bytes_chain(O200K_BASE_VOCAB, pats, special),
PretrainedVocab::Llama3 => Tokenizer::from_bytes_chain(LLAMA3_VOCAB, pats, special),
PretrainedVocab::DeepseekV3 => {
Tokenizer::from_bytes_byte_level_chain(DEEPSEEK_V3_VOCAB, pats, special)
}
PretrainedVocab::MistralV1 | PretrainedVocab::MistralV2 => {
return Err(TokenizerError::UnknownPretrained(
"Mistral V1/V2 take the SPM backend and are routed earlier".to_owned(),
))
}
PretrainedVocab::MistralV3 => {
Tokenizer::from_bytes_byte_level_chain(MISTRAL_V3_VOCAB, pats, special)
}
PretrainedVocab::WhisperV1 | PretrainedVocab::WhisperV2 | PretrainedVocab::WhisperV3 => {
Tokenizer::from_bytes_byte_level_chain(WHISPER_VOCAB, pats, special)
}
}?;
Ok(AnyTokenizer::new(
Backend::Bpe(tokenizer.with_added_token_matching(true)),
SpecialPolicy::boundary(None, None, Some(eos_token_id(vocab)), named),
))
}
fn spm_from_vocab(
data: &[u8],
vocab: PretrainedVocab,
special: FxHashMap<String, u32>,
named: FxHashMap<String, u32>,
) -> Result<AnyTokenizer, TokenizerError> {
let (mut pieces, mut scores) = load_spm_vocab(data)?;
place_special_pieces(&mut pieces, &special)?;
scores.resize(pieces.len(), NEVER_MERGE);
let eos = eos_token_id(vocab);
let tokenizer = SpmTokenizer::new(pieces, scores, bos_token_id(vocab), Some(eos))?
.with_prefix_scheme(spm_prefix_scheme(vocab))
.with_added_tokens(&special)?;
Ok(AnyTokenizer::new(
Backend::Spm(tokenizer),
SpecialPolicy::boundary(None, None, Some(eos), named),
))
}
fn spm_prefix_scheme(vocab: PretrainedVocab) -> SpmPrefixScheme {
match vocab {
PretrainedVocab::MistralV1 => SpmPrefixScheme::AfterEachSpecial,
_ => SpmPrefixScheme::Once,
}
}
pub fn patterns(vocab: PretrainedVocab) -> Option<&'static [&'static str]> {
match vocab {
PretrainedVocab::Cl100kBase => Some(&[CL100K_BASE_PATTERN]),
PretrainedVocab::O200kBase => Some(&[O200K_BASE_PATTERN]),
PretrainedVocab::Llama3 => Some(&[LLAMA3_PATTERN]),
PretrainedVocab::DeepseekV3 => Some(DEEPSEEK_V3_PATTERNS),
PretrainedVocab::MistralV1 | PretrainedVocab::MistralV2 => None,
PretrainedVocab::MistralV3 => Some(&[MISTRAL_V3_PATTERN]),
PretrainedVocab::WhisperV1 | PretrainedVocab::WhisperV2 | PretrainedVocab::WhisperV3 => {
Some(&[GPT2_PATTERN])
}
}
}
pub fn uses_byte_level(vocab: PretrainedVocab) -> bool {
matches!(
vocab,
PretrainedVocab::DeepseekV3
| PretrainedVocab::WhisperV1
| PretrainedVocab::WhisperV2
| PretrainedVocab::WhisperV3
)
}
pub fn eos_token_id(vocab: PretrainedVocab) -> u32 {
match vocab {
PretrainedVocab::Cl100kBase => 100257, PretrainedVocab::O200kBase => 199999, PretrainedVocab::Llama3 => 128001, PretrainedVocab::DeepseekV3 => 1, PretrainedVocab::MistralV1 | PretrainedVocab::MistralV2 | PretrainedVocab::MistralV3 => 2, PretrainedVocab::WhisperV1 => WhisperVariant::V1Multilingual.eos_token_id(),
PretrainedVocab::WhisperV2 => WhisperVariant::V2Multilingual.eos_token_id(),
PretrainedVocab::WhisperV3 => WhisperVariant::V3Multilingual.eos_token_id(),
}
}
pub fn eos_token_id_by_name(name: &str) -> u32 {
PretrainedVocab::from_name(name)
.map(eos_token_id)
.unwrap_or(0)
}
pub fn bos_token_id(vocab: PretrainedVocab) -> Option<u32> {
match vocab {
PretrainedVocab::Cl100kBase => None, PretrainedVocab::O200kBase => None, PretrainedVocab::Llama3 => Some(128000), PretrainedVocab::DeepseekV3 => Some(0), PretrainedVocab::MistralV1 | PretrainedVocab::MistralV2 | PretrainedVocab::MistralV3 => {
Some(1)
} PretrainedVocab::WhisperV1 | PretrainedVocab::WhisperV2 | PretrainedVocab::WhisperV3 => {
None
}
}
}
pub fn bos_token_id_by_name(name: &str) -> Option<u32> {
PretrainedVocab::from_name(name).and_then(bos_token_id)
}
pub fn pad_token_id(vocab: PretrainedVocab) -> Option<u32> {
match vocab {
PretrainedVocab::Cl100kBase => Some(100316), PretrainedVocab::O200kBase => Some(200058), PretrainedVocab::Llama3 => Some(128339), PretrainedVocab::DeepseekV3 => Some(2), PretrainedVocab::MistralV1 => Some(32039), PretrainedVocab::MistralV2 => Some(32807), PretrainedVocab::MistralV3 => Some(131111), PretrainedVocab::WhisperV1 | PretrainedVocab::WhisperV2 | PretrainedVocab::WhisperV3 => {
None
}
}
}
pub fn base_vocab_size(vocab: PretrainedVocab) -> u32 {
match vocab {
PretrainedVocab::Cl100kBase => CL100K_BASE_BASE_VOCAB_SIZE,
PretrainedVocab::O200kBase => O200K_BASE_BASE_VOCAB_SIZE,
PretrainedVocab::Llama3 => LLAMA3_BASE_VOCAB_SIZE,
PretrainedVocab::DeepseekV3 => DEEPSEEK_V3_BASE_VOCAB_SIZE,
PretrainedVocab::MistralV1 => MISTRAL_V1_BASE_VOCAB_SIZE,
PretrainedVocab::MistralV2 => MISTRAL_V2_BASE_VOCAB_SIZE,
PretrainedVocab::MistralV3 => MISTRAL_V3_BASE_VOCAB_SIZE,
PretrainedVocab::WhisperV1 => WhisperVariant::V1Multilingual.vocab_size() as u32,
PretrainedVocab::WhisperV2 => WhisperVariant::V2Multilingual.vocab_size() as u32,
PretrainedVocab::WhisperV3 => WhisperVariant::V3Multilingual.vocab_size() as u32,
}
}
pub fn base_vocab_size_by_name(name: &str) -> Result<u32, TokenizerError> {
resolve_vocab(name).map(base_vocab_size)
}
pub fn special_tokens(vocab: PretrainedVocab) -> FxHashMap<String, u32> {
match vocab {
PretrainedVocab::Cl100kBase => cl100k_base_special_tokens(),
PretrainedVocab::O200kBase => o200k_base_special_tokens(),
PretrainedVocab::Llama3 => llama3_special_tokens(),
PretrainedVocab::DeepseekV3 => deepseek_v3_special_tokens(),
PretrainedVocab::MistralV1 => mistral_v1_special_tokens(),
PretrainedVocab::MistralV2 => mistral_v2_special_tokens(),
PretrainedVocab::MistralV3 => mistral_v3_special_tokens(),
PretrainedVocab::WhisperV1 => whisper_special_tokens(WhisperVariant::V1Multilingual),
PretrainedVocab::WhisperV2 => whisper_special_tokens(WhisperVariant::V2Multilingual),
PretrainedVocab::WhisperV3 => whisper_special_tokens(WhisperVariant::V3Multilingual),
}
}
const CL100K_BASE_BASE_VOCAB_SIZE: u32 = 100277;
const O200K_BASE_BASE_VOCAB_SIZE: u32 = 200019;
pub fn cl100k_base_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<|endoftext|>".to_string(), 100257);
special.insert("<|fim_prefix|>".to_string(), 100258);
special.insert("<|fim_middle|>".to_string(), 100259);
special.insert("<|fim_suffix|>".to_string(), 100260);
special.insert(
"<|endofprompt|>".to_string(),
CL100K_BASE_BASE_VOCAB_SIZE - 1,
);
insert_agent_tokens(&mut special, CL100K_BASE_BASE_VOCAB_SIZE);
special
}
pub fn o200k_base_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<|endoftext|>".to_string(), 199999);
special.insert(
"<|endofprompt|>".to_string(),
O200K_BASE_BASE_VOCAB_SIZE - 1,
);
insert_agent_tokens(&mut special, O200K_BASE_BASE_VOCAB_SIZE);
special
}
const LLAMA3_BASE_VOCAB_SIZE: u32 = 128256;
pub fn llama3_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<|begin_of_text|>".to_string(), 128000);
special.insert("<|end_of_text|>".to_string(), 128001);
special.insert("<|reserved_special_token_0|>".to_string(), 128002);
special.insert("<|reserved_special_token_1|>".to_string(), 128003);
special.insert("<|finetune_right_pad_id|>".to_string(), 128004);
special.insert("<|step_id|>".to_string(), 128005);
special.insert("<|start_header_id|>".to_string(), 128006);
special.insert("<|end_header_id|>".to_string(), 128007);
special.insert("<|eom_id|>".to_string(), 128008);
special.insert("<|eot_id|>".to_string(), 128009);
special.insert("<|python_tag|>".to_string(), 128010);
special.insert("<|image|>".to_string(), LLAMA3_BASE_VOCAB_SIZE);
special.insert("<|/image|>".to_string(), LLAMA3_BASE_VOCAB_SIZE + 1);
special.insert("<|audio|>".to_string(), LLAMA3_BASE_VOCAB_SIZE + 2);
special.insert("<|/audio|>".to_string(), LLAMA3_BASE_VOCAB_SIZE + 3);
special.insert("<|video|>".to_string(), LLAMA3_BASE_VOCAB_SIZE + 4);
special.insert("<|/video|>".to_string(), LLAMA3_BASE_VOCAB_SIZE + 5);
insert_agent_tokens_llama3(&mut special, 128300);
special
}
const DEEPSEEK_V3_BASE_VOCAB_SIZE: u32 = 128815;
pub fn deepseek_v3_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<|begin▁of▁sentence|>".to_string(), 0);
special.insert("<|end▁of▁sentence|>".to_string(), 1);
special.insert("<|▁pad▁|>".to_string(), 2);
special.insert("<think>".to_string(), 128798);
special.insert("</think>".to_string(), 128799);
special.insert("<|fim▁hole|>".to_string(), 128800);
special.insert("<|fim▁begin|>".to_string(), 128801);
special.insert("<|fim▁end|>".to_string(), 128802);
special.insert("<|User|>".to_string(), 128803);
special.insert("<|Assistant|>".to_string(), 128804);
special.insert("<|EOT|>".to_string(), 128805);
special.insert("<|tool▁calls▁begin|>".to_string(), 128806);
special.insert("<|tool▁calls▁end|>".to_string(), 128807);
special.insert("<|tool▁call▁begin|>".to_string(), 128808);
special.insert("<|tool▁call▁end|>".to_string(), 128809);
special.insert("<|tool▁outputs▁begin|>".to_string(), 128810);
special.insert("<|tool▁outputs▁end|>".to_string(), 128811);
special.insert("<|tool▁output▁begin|>".to_string(), 128812);
special.insert("<|tool▁output▁end|>".to_string(), 128813);
special.insert(
"<|tool▁sep|>".to_string(),
DEEPSEEK_V3_BASE_VOCAB_SIZE - 1,
);
insert_agent_tokens(&mut special, 128900);
special
}
const MISTRAL_V1_BASE_VOCAB_SIZE: u32 = 32000;
const MISTRAL_V2_BASE_VOCAB_SIZE: u32 = 32768;
const MISTRAL_V3_BASE_VOCAB_SIZE: u32 = 131072;
pub fn mistral_v1_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<unk>".to_string(), 0);
special.insert("<s>".to_string(), 1);
special.insert("</s>".to_string(), 2);
insert_agent_tokens(&mut special, MISTRAL_V1_BASE_VOCAB_SIZE);
special
}
pub fn mistral_v2_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<unk>".to_string(), 0);
special.insert("<s>".to_string(), 1);
special.insert("</s>".to_string(), 2);
special.insert("[INST]".to_string(), 3);
special.insert("[/INST]".to_string(), 4);
special.insert("[TOOL_CALLS]".to_string(), 5);
special.insert("[AVAILABLE_TOOLS]".to_string(), 6);
special.insert("[/AVAILABLE_TOOLS]".to_string(), 7);
special.insert("[TOOL_RESULTS]".to_string(), 8);
special.insert("[/TOOL_RESULTS]".to_string(), 9);
insert_agent_tokens(&mut special, MISTRAL_V2_BASE_VOCAB_SIZE);
special
}
pub fn mistral_v3_special_tokens() -> FxHashMap<String, u32> {
let mut special = FxHashMap::default();
special.insert("<unk>".to_string(), 0);
special.insert("<s>".to_string(), 1);
special.insert("</s>".to_string(), 2);
special.insert("[INST]".to_string(), 3);
special.insert("[/INST]".to_string(), 4);
special.insert("[AVAILABLE_TOOLS]".to_string(), 5);
special.insert("[/AVAILABLE_TOOLS]".to_string(), 6);
special.insert("[TOOL_RESULTS]".to_string(), 7);
special.insert("[/TOOL_RESULTS]".to_string(), 8);
special.insert("[TOOL_CALLS]".to_string(), 9);
insert_agent_tokens(&mut special, MISTRAL_V3_BASE_VOCAB_SIZE);
special
}
fn insert_agent_tokens(special: &mut FxHashMap<String, u32>, base: u32) {
special.insert("<|system|>".to_string(), base);
special.insert("<|user|>".to_string(), base + 1);
special.insert("<|assistant|>".to_string(), base + 2);
special.insert("<|im_start|>".to_string(), base + 3);
special.insert("<|im_end|>".to_string(), base + 4);
special.insert("<|think|>".to_string(), base + 5);
special.insert("<|/think|>".to_string(), base + 6);
special.insert("<|plan|>".to_string(), base + 7);
special.insert("<|/plan|>".to_string(), base + 8);
special.insert("<|step|>".to_string(), base + 9);
special.insert("<|/step|>".to_string(), base + 10);
special.insert("<|act|>".to_string(), base + 11);
special.insert("<|/act|>".to_string(), base + 12);
special.insert("<|observe|>".to_string(), base + 13);
special.insert("<|/observe|>".to_string(), base + 14);
special.insert("<|function|>".to_string(), base + 15);
special.insert("<|/function|>".to_string(), base + 16);
special.insert("<|result|>".to_string(), base + 17);
special.insert("<|/result|>".to_string(), base + 18);
special.insert("<|error|>".to_string(), base + 19);
special.insert("<|/error|>".to_string(), base + 20);
special.insert("<|code|>".to_string(), base + 21);
special.insert("<|/code|>".to_string(), base + 22);
special.insert("<|output|>".to_string(), base + 23);
special.insert("<|/output|>".to_string(), base + 24);
special.insert("<|lang|>".to_string(), base + 25);
special.insert("<|/lang|>".to_string(), base + 26);
special.insert("<|context|>".to_string(), base + 27);
special.insert("<|/context|>".to_string(), base + 28);
special.insert("<|quote|>".to_string(), base + 29);
special.insert("<|/quote|>".to_string(), base + 30);
special.insert("<|cite|>".to_string(), base + 31);
special.insert("<|/cite|>".to_string(), base + 32);
special.insert("<|source|>".to_string(), base + 33);
special.insert("<|/source|>".to_string(), base + 34);
special.insert("<|memory|>".to_string(), base + 35);
special.insert("<|/memory|>".to_string(), base + 36);
special.insert("<|recall|>".to_string(), base + 37);
special.insert("<|/recall|>".to_string(), base + 38);
special.insert("<|pad|>".to_string(), base + 39);
special.insert("<|stop|>".to_string(), base + 40);
special.insert("<|sep|>".to_string(), base + 41);
special.insert("<|image|>".to_string(), base + 42);
special.insert("<|/image|>".to_string(), base + 43);
special.insert("<|audio|>".to_string(), base + 44);
special.insert("<|/audio|>".to_string(), base + 45);
special.insert("<|video|>".to_string(), base + 46);
special.insert("<|/video|>".to_string(), base + 47);
special.insert("<|title|>".to_string(), base + 48);
special.insert("<|/title|>".to_string(), base + 49);
special.insert("<|section|>".to_string(), base + 50);
special.insert("<|/section|>".to_string(), base + 51);
special.insert("<|summary|>".to_string(), base + 52);
special.insert("<|/summary|>".to_string(), base + 53);
}
fn insert_agent_tokens_llama3(special: &mut FxHashMap<String, u32>, base: u32) {
special.insert("<|system|>".to_string(), base);
special.insert("<|user|>".to_string(), base + 1);
special.insert("<|assistant|>".to_string(), base + 2);
special.insert("<|im_start|>".to_string(), base + 3);
special.insert("<|im_end|>".to_string(), base + 4);
special.insert("<|think|>".to_string(), base + 5);
special.insert("<|/think|>".to_string(), base + 6);
special.insert("<|plan|>".to_string(), base + 7);
special.insert("<|/plan|>".to_string(), base + 8);
special.insert("<|step|>".to_string(), base + 9);
special.insert("<|/step|>".to_string(), base + 10);
special.insert("<|act|>".to_string(), base + 11);
special.insert("<|/act|>".to_string(), base + 12);
special.insert("<|observe|>".to_string(), base + 13);
special.insert("<|/observe|>".to_string(), base + 14);
special.insert("<|function|>".to_string(), base + 15);
special.insert("<|/function|>".to_string(), base + 16);
special.insert("<|result|>".to_string(), base + 17);
special.insert("<|/result|>".to_string(), base + 18);
special.insert("<|error|>".to_string(), base + 19);
special.insert("<|/error|>".to_string(), base + 20);
special.insert("<|code|>".to_string(), base + 21);
special.insert("<|/code|>".to_string(), base + 22);
special.insert("<|output|>".to_string(), base + 23);
special.insert("<|/output|>".to_string(), base + 24);
special.insert("<|lang|>".to_string(), base + 25);
special.insert("<|/lang|>".to_string(), base + 26);
special.insert("<|context|>".to_string(), base + 27);
special.insert("<|/context|>".to_string(), base + 28);
special.insert("<|quote|>".to_string(), base + 29);
special.insert("<|/quote|>".to_string(), base + 30);
special.insert("<|cite|>".to_string(), base + 31);
special.insert("<|/cite|>".to_string(), base + 32);
special.insert("<|source|>".to_string(), base + 33);
special.insert("<|/source|>".to_string(), base + 34);
special.insert("<|memory|>".to_string(), base + 35);
special.insert("<|/memory|>".to_string(), base + 36);
special.insert("<|recall|>".to_string(), base + 37);
special.insert("<|/recall|>".to_string(), base + 38);
special.insert("<|pad|>".to_string(), base + 39);
special.insert("<|stop|>".to_string(), base + 40);
special.insert("<|sep|>".to_string(), base + 41);
special.insert("<|title|>".to_string(), base + 48);
special.insert("<|/title|>".to_string(), base + 49);
special.insert("<|section|>".to_string(), base + 50);
special.insert("<|/section|>".to_string(), base + 51);
special.insert("<|summary|>".to_string(), base + 52);
special.insert("<|/summary|>".to_string(), base + 53);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::core::Tokenize;
fn bpe(tokenizer: AnyTokenizer) -> Tokenizer {
match tokenizer.into_backend() {
Backend::Bpe(t) => t,
_ => panic!("this vocabulary does not load as byte-pair encoding"),
}
}
fn spm_piece_id(vocab_data: &[u8], piece: &str) -> u32 {
let (pieces, _) = load_spm_vocab(vocab_data).expect("vocabulary loads");
let id = pieces
.iter()
.position(|p| p == piece)
.unwrap_or_else(|| panic!("{piece:?} is not in the vocabulary"));
id as u32
}
#[test]
fn test_from_pretrained_llama3() {
let tokenizer = from_pretrained("llama3").unwrap();
assert!(tokenizer.vocab_size() > 100000);
}
#[test]
fn test_from_pretrained_cl100k() {
let tokenizer = from_pretrained("cl100k_base").unwrap();
assert!(tokenizer.vocab_size() > 90000);
}
#[test]
fn test_from_pretrained_whisper_variants() {
for (name, variant) in [
("whisper_v1", WhisperVariant::V1Multilingual),
("whisper", WhisperVariant::V2Multilingual), ("whisper_v2", WhisperVariant::V2Multilingual),
("whisper-v3", WhisperVariant::V3Multilingual),
] {
let tok = from_pretrained(name).unwrap_or_else(|e| panic!("{name}: {e}"));
assert_eq!(tok.vocab_size(), variant.vocab_size(), "{name} vocab_size");
assert_eq!(bpe(tok).encoder().len(), 50257, "{name} base vocab size");
}
}
#[test]
fn test_policy_is_passthrough_but_knows_its_specials() {
let tok = from_pretrained("llama3").unwrap();
assert_eq!(tok.eos_token_id(), Some(128001));
assert!(tok.is_eos(128001));
assert_eq!(tok.special_token_id("<|eot_id|>"), Some(128009));
let text = "Hello, world!";
assert_eq!(tok.encode(text), tok.encode_raw(text));
}
#[test]
fn test_encode_batch_matches_individual() {
let tok = from_pretrained("llama3").unwrap();
let texts = ["Hello, world!", "", "<|eot_id|>after", "你好世界"];
let batch = tok.encode_batch(&texts);
assert_eq!(batch.len(), texts.len());
for (got, text) in batch.iter().zip(texts) {
assert_eq!(got, &tok.encode(text), "batch mismatch for {text:?}");
}
assert!(batch[2].starts_with(&[128009]));
}
#[test]
fn test_whisper_special_tokens_wired() {
let tok = from_pretrained("whisper_v3").unwrap();
assert_eq!(tok.encode("<|en|>"), vec![50259]);
assert_eq!(
tok.encode("<|transcribe|>"),
vec![WhisperVariant::V3Multilingual.transcribe_token_id()]
);
assert_eq!(tok.encode("<|yue|>"), vec![50259 + 99]);
}
#[test]
fn test_whisper_roundtrip() {
let tok = from_pretrained("whisper").unwrap();
let text = "Hello, world! 123 héllo";
assert_eq!(tok.decode(&tok.encode(text)).unwrap(), text);
}
#[test]
fn test_whisper_name_mapping() {
assert_eq!(
PretrainedVocab::from_name("whisper"),
Some(PretrainedVocab::WhisperV2)
);
assert_eq!(
PretrainedVocab::from_name("whisper-large-v3"),
Some(PretrainedVocab::WhisperV3)
);
assert_eq!(PretrainedVocab::from_name("whisper.en"), None);
}
#[test]
fn test_eos_token_ids() {
assert_eq!(eos_token_id(PretrainedVocab::Cl100kBase), 100257);
assert_eq!(eos_token_id(PretrainedVocab::O200kBase), 199999);
assert_eq!(eos_token_id(PretrainedVocab::Llama3), 128001);
assert_eq!(eos_token_id(PretrainedVocab::DeepseekV3), 1);
assert_eq!(eos_token_id(PretrainedVocab::MistralV1), 2);
}
#[test]
fn test_vocab_from_name() {
assert_eq!(
PretrainedVocab::from_name("llama3"),
Some(PretrainedVocab::Llama3)
);
assert_eq!(
PretrainedVocab::from_name("llama3.1"),
Some(PretrainedVocab::Llama3)
);
assert_eq!(
PretrainedVocab::from_name("deepseek_v3"),
Some(PretrainedVocab::DeepseekV3)
);
assert_eq!(
PretrainedVocab::from_name("mistral"),
Some(PretrainedVocab::MistralV1)
);
assert_eq!(PretrainedVocab::from_name("unknown"), None);
}
#[test]
fn test_from_pretrained_mistral() {
let tokenizer = from_pretrained("mistral").unwrap();
assert!(tokenizer.vocab_size() >= 31000);
}
#[test]
fn test_mistral_encode_decode() {
let tokenizer = from_pretrained("mistral").unwrap();
let text = "Hello, world!";
let tokens = tokenizer.encode(text);
assert!(!tokens.is_empty());
let decoded = tokenizer.decode(&tokens).unwrap();
assert_eq!(decoded, text, "Encoding should be reversible");
}
#[test]
fn test_mistral_never_shatters_the_word_boundary_marker() {
for (name, data) in [
("mistral", MISTRAL_SPM_VOCAB),
("mistral_v2", MISTRAL_V2_SPM_VOCAB),
] {
let shattered = [
spm_piece_id(data, "<0xE2>"),
spm_piece_id(data, "<0x96>"),
spm_piece_id(data, "<0x81>"),
];
let tokenizer = from_pretrained(name).unwrap();
let ids = tokenizer.encode("the sourdough starter rose overnight");
assert!(
!ids.windows(3).any(|w| w == shattered.as_slice()),
"{name}: word boundary shattered into byte tokens {shattered:?} in {ids:?}"
);
}
}
#[test]
fn test_mistral_reaches_whole_word_pieces() {
let the = spm_piece_id(MISTRAL_SPM_VOCAB, "▁the");
let sour = spm_piece_id(MISTRAL_SPM_VOCAB, "▁sour");
let ids = from_pretrained("mistral").unwrap().encode("the sourdough");
assert!(ids.contains(&the), "▁the ({the}) missing from {ids:?}");
assert!(ids.contains(&sour), "▁sour ({sour}) missing from {ids:?}");
}
#[test]
fn test_mistral_round_trips_a_sentence() {
let tokenizer = from_pretrained("mistral").unwrap();
let text = "The quick brown fox jumps over the lazy dog.";
let decoded = tokenizer.decode(&tokenizer.encode(text)).unwrap();
assert_eq!(decoded, text);
}
#[test]
fn test_base_vocab_size_matches_reference() {
assert_eq!(base_vocab_size(PretrainedVocab::Cl100kBase), 100277); assert_eq!(base_vocab_size(PretrainedVocab::O200kBase), 200019); assert_eq!(base_vocab_size(PretrainedVocab::Llama3), 128256); assert_eq!(base_vocab_size(PretrainedVocab::DeepseekV3), 128815); assert_eq!(base_vocab_size(PretrainedVocab::MistralV1), 32000); assert_eq!(base_vocab_size(PretrainedVocab::MistralV2), 32768); assert_eq!(base_vocab_size(PretrainedVocab::MistralV3), 131072); assert_eq!(
base_vocab_size(PretrainedVocab::WhisperV1),
WhisperVariant::V1Multilingual.vocab_size() as u32
);
assert_eq!(
base_vocab_size(PretrainedVocab::WhisperV2),
WhisperVariant::V2Multilingual.vocab_size() as u32
);
assert_eq!(
base_vocab_size(PretrainedVocab::WhisperV3),
WhisperVariant::V3Multilingual.vocab_size() as u32
);
}
#[test]
fn test_base_vocab_size_never_exceeds_extended_vocab_size() {
for (name, vocab) in [
("cl100k_base", PretrainedVocab::Cl100kBase),
("o200k_base", PretrainedVocab::O200kBase),
("llama3", PretrainedVocab::Llama3),
("deepseek_v3", PretrainedVocab::DeepseekV3),
("mistral_v1", PretrainedVocab::MistralV1),
("mistral_v2", PretrainedVocab::MistralV2),
("mistral_v3", PretrainedVocab::MistralV3),
("whisper_v1", PretrainedVocab::WhisperV1),
("whisper_v2", PretrainedVocab::WhisperV2),
("whisper_v3", PretrainedVocab::WhisperV3),
] {
let extended = from_vocab(vocab).unwrap().vocab_size() as u32;
let base = base_vocab_size(vocab);
assert!(
base <= extended,
"{name}: base_vocab_size {base} exceeds extended vocab_size {extended}"
);
}
}
#[test]
fn test_no_agent_token_id_below_base_vocab_size() {
for (name, vocab) in [
("cl100k_base", PretrainedVocab::Cl100kBase),
("o200k_base", PretrainedVocab::O200kBase),
("llama3", PretrainedVocab::Llama3),
("deepseek_v3", PretrainedVocab::DeepseekV3),
("mistral_v1", PretrainedVocab::MistralV1),
("mistral_v2", PretrainedVocab::MistralV2),
("mistral_v3", PretrainedVocab::MistralV3),
] {
let base = base_vocab_size(vocab);
for name_and_id in agent_token_ids_in(vocab) {
let (token, id) = name_and_id;
assert!(
id >= base,
"{name}: agent token {token:?} has id {id}, below base_vocab_size {base}"
);
}
}
}
const AGENT_TOKEN_NAMES: [&str; 54] = [
"<|system|>",
"<|user|>",
"<|assistant|>",
"<|im_start|>",
"<|im_end|>",
"<|think|>",
"<|/think|>",
"<|plan|>",
"<|/plan|>",
"<|step|>",
"<|/step|>",
"<|act|>",
"<|/act|>",
"<|observe|>",
"<|/observe|>",
"<|function|>",
"<|/function|>",
"<|result|>",
"<|/result|>",
"<|error|>",
"<|/error|>",
"<|code|>",
"<|/code|>",
"<|output|>",
"<|/output|>",
"<|lang|>",
"<|/lang|>",
"<|context|>",
"<|/context|>",
"<|quote|>",
"<|/quote|>",
"<|cite|>",
"<|/cite|>",
"<|source|>",
"<|/source|>",
"<|memory|>",
"<|/memory|>",
"<|recall|>",
"<|/recall|>",
"<|pad|>",
"<|stop|>",
"<|sep|>",
"<|image|>",
"<|/image|>",
"<|audio|>",
"<|/audio|>",
"<|video|>",
"<|/video|>",
"<|title|>",
"<|/title|>",
"<|section|>",
"<|/section|>",
"<|summary|>",
"<|/summary|>",
];
fn agent_token_ids_in(vocab: PretrainedVocab) -> Vec<(String, u32)> {
let all = special_tokens(vocab);
AGENT_TOKEN_NAMES
.iter()
.filter_map(|name| all.get(*name).map(|&id| (name.to_string(), id)))
.collect()
}
}