use base64::{Engine as _, engine::general_purpose};
use rustc_hash::FxHashMap;
use std::collections::HashMap;
use std::path::Path;
use tiktoken_rs::CoreBPE;
use crate::audio::{Audio, AudioConfig, AudioEncoder, AudioEncoding};
use crate::config::{ModelData, TekkenConfig, TokenInfo, TokenizerVersion};
use crate::errors::{Result, TokenizerError};
use crate::special_tokens::{SpecialTokenInfo, SpecialTokenPolicy, SpecialTokens};
pub struct Tekkenizer {
tekkenizer: CoreBPE,
vocab_size: usize,
num_special_tokens: usize,
version: TokenizerVersion,
pattern: String,
special_tokens: Vec<SpecialTokenInfo>,
special_tokens_map: HashMap<String, usize>,
vocab: Vec<String>,
vocab_tokens: Vec<TokenInfo>,
audio_config: Option<AudioConfig>,
audio_encoder: Option<AudioEncoder>,
}
impl Tekkenizer {
#[allow(clippy::cast_possible_truncation)]
pub fn new(
vocab: Vec<TokenInfo>,
special_tokens: &Vec<SpecialTokenInfo>,
pattern: &str,
vocab_size: usize,
num_special_tokens: usize,
version: TokenizerVersion,
audio_config: Option<AudioConfig>,
) -> Result<Self> {
if vocab_size > vocab.len() + num_special_tokens {
return Err(TokenizerError::InvalidConfig(format!(
"vocab_size ({}) must be <= vocab.len() ({}) + num_special_tokens ({})",
vocab_size,
vocab.len(),
num_special_tokens
)));
}
let mut token_strings = std::collections::HashSet::new();
for token in special_tokens {
if !token_strings.insert(&token.token_str) {
return Err(TokenizerError::InvalidConfig(format!(
"Duplicate special token: {}",
token.token_str
)));
}
}
if special_tokens.len() > num_special_tokens {
return Err(TokenizerError::InvalidConfig(format!(
"special_tokens.len() ({}) must be <= num_special_tokens ({})",
special_tokens.len(),
num_special_tokens
)));
}
let mut all_special_tokens = special_tokens.clone();
for i in special_tokens.len()..num_special_tokens {
all_special_tokens.push(SpecialTokenInfo {
rank: i,
token_str: format!("<SPECIAL_{i}>"),
is_control: true,
});
}
let inner_vocab_size = vocab_size - num_special_tokens;
let vocab_tokens_copy = vocab.clone();
let mergeable_ranks = reload_mergeable_ranks(vocab, inner_vocab_size)?;
let special_tokens: FxHashMap<String, u32> = FxHashMap::default();
let tekkenizer = CoreBPE::new(mergeable_ranks.clone(), special_tokens, pattern)
.map_err(|e| TokenizerError::InvalidConfig(format!("Failed to create CoreBPE: {e}")))?;
let special_tokens_map: HashMap<String, usize> = all_special_tokens
.iter()
.map(|token| (token.token_str.clone(), token.rank))
.collect();
let rank_to_bytes: FxHashMap<u32, &Vec<u8>> = mergeable_ranks
.iter()
.map(|(bytes, &rank)| (rank, bytes))
.collect();
let vocab_strings: Vec<String> = (0..vocab_size)
.map(|i| {
if i < num_special_tokens {
all_special_tokens[i].token_str.clone()
} else {
#[allow(clippy::cast_possible_truncation)]
let token_id = (i - num_special_tokens) as u32;
rank_to_bytes.get(&token_id).map_or_else(
|| "<?>".to_string(),
|bytes| String::from_utf8_lossy(bytes).to_string(),
)
}
})
.collect();
let audio_encoder = if let Some(ref config) = audio_config {
let audio_token_id = special_tokens_map
.get(SpecialTokens::Audio.as_str())
.ok_or_else(|| {
TokenizerError::TokenNotFound("Audio token not found".to_string())
})?;
let begin_audio_token_id = special_tokens_map
.get(SpecialTokens::BeginAudio.as_str())
.ok_or_else(|| {
TokenizerError::TokenNotFound("BeginAudio token not found".to_string())
})?;
#[allow(clippy::cast_possible_truncation)]
Some(AudioEncoder::new(
config.clone(),
*audio_token_id as u32,
*begin_audio_token_id as u32,
))
} else {
None
};
Ok(Self {
tekkenizer,
vocab_size,
num_special_tokens,
version,
pattern: pattern.to_string(),
special_tokens: all_special_tokens,
special_tokens_map,
vocab: vocab_strings,
vocab_tokens: vocab_tokens_copy,
audio_config,
audio_encoder,
})
}
pub fn from_file<P: AsRef<Path>>(path: P) -> Result<Self> {
let content = std::fs::read_to_string(path)?;
let model_data: ModelData = serde_json::from_str(&content)?;
let version =
TokenizerVersion::from_string(&model_data.config.version).ok_or_else(|| {
TokenizerError::InvalidConfig(format!(
"Unknown version: {}",
model_data.config.version
))
})?;
let special_tokens = model_data.special_tokens.unwrap_or_else(|| {
get_deprecated_special_tokens()
});
Self::new(
model_data.vocab,
&special_tokens,
&model_data.config.pattern,
model_data.config.default_vocab_size,
model_data.config.default_num_special_tokens,
version,
model_data.audio,
)
}
#[must_use]
pub const fn vocab_size(&self) -> usize {
self.vocab_size
}
#[must_use]
pub const fn num_special_tokens(&self) -> usize {
self.num_special_tokens
}
#[must_use]
pub const fn version(&self) -> &TokenizerVersion {
&self.version
}
pub fn bos_id(&self) -> Result<u32> {
self.get_control_token(SpecialTokens::Bos.as_str())
}
pub fn eos_id(&self) -> Result<u32> {
self.get_control_token(SpecialTokens::Eos.as_str())
}
pub fn pad_id(&self) -> Result<u32> {
self.get_control_token(SpecialTokens::Pad.as_str())
}
pub fn unk_id(&self) -> Result<u32> {
self.get_control_token(SpecialTokens::Unk.as_str())
}
#[allow(clippy::cast_possible_truncation)]
pub fn get_control_token(&self, token_str: &str) -> Result<u32> {
self.special_tokens_map
.get(token_str)
.map(|&id| id as u32)
.ok_or_else(|| {
let available_tokens: Vec<&String> = self.special_tokens_map.keys().collect();
TokenizerError::TokenNotFound(format!(
"Unknown control token: '{token_str}'. Available special tokens: {available_tokens:?}",
))
})
}
#[must_use]
pub fn vocab(&self) -> &[String] {
&self.vocab
}
#[allow(clippy::cast_possible_truncation)]
pub fn encode(
&self,
text: &str,
add_beginning_of_sequence: bool,
add_end_of_sequence: bool,
) -> Result<Vec<u32>> {
let (tokens, _) = self
.tekkenizer
.encode(text, &std::collections::HashSet::new());
let mut tokens: Vec<u32> = tokens;
for token in &mut tokens {
*token += self.num_special_tokens as u32;
}
if add_beginning_of_sequence {
let bos_id = self.bos_id()?;
tokens.insert(0, bos_id);
}
if add_end_of_sequence {
let eos_id = self.eos_id()?;
tokens.push(eos_id);
}
Ok(tokens)
}
pub fn decode(
&self,
tokens: &[u32],
special_token_policy: SpecialTokenPolicy,
) -> Result<String> {
let decoded_parts = self.decode_all(tokens, special_token_policy)?;
Ok(decoded_parts.join(""))
}
#[allow(clippy::cast_possible_truncation)]
pub fn decode_all(
&self,
tokens: &[u32],
special_token_policy: SpecialTokenPolicy,
) -> Result<Vec<String>> {
let mut decoded = Vec::new();
let mut current_group = Vec::new();
let mut current_is_special = None;
for &token_id in tokens {
#[allow(clippy::cast_possible_truncation)]
let is_special = token_id < self.num_special_tokens as u32;
if current_is_special.is_none() {
current_is_special = Some(is_special);
}
if current_is_special == Some(is_special) {
current_group.push(token_id);
} else {
if let Some(was_special) = current_is_special {
self.decode_group(
¤t_group,
was_special,
&mut decoded,
special_token_policy,
)?;
}
current_group.clear();
current_group.push(token_id);
current_is_special = Some(is_special);
}
}
if let Some(was_special) = current_is_special {
self.decode_group(
¤t_group,
was_special,
&mut decoded,
special_token_policy,
)?;
}
Ok(decoded)
}
#[allow(clippy::cast_possible_truncation)]
fn decode_group(
&self,
group: &[u32],
is_special: bool,
decoded: &mut Vec<String>,
special_token_policy: SpecialTokenPolicy,
) -> Result<()> {
if is_special {
match special_token_policy {
SpecialTokenPolicy::Raise => {
return Err(TokenizerError::SpecialTokenPolicy(format!(
"Decoding tokens that contain special tokens ({group:?}) is not allowed",
)));
}
SpecialTokenPolicy::Keep => {
for &token_id in group {
decoded.push(self.special_tokens[token_id as usize].token_str.clone());
}
}
SpecialTokenPolicy::Ignore => {
}
}
} else {
#[allow(clippy::cast_possible_truncation)]
let shifted_tokens: Vec<u32> = group
.iter()
.map(|&t| t - self.num_special_tokens as u32)
.collect();
let decoded_text = self
.tekkenizer
.decode(shifted_tokens)
.map_err(|e| TokenizerError::Tokenizers(format!("{e:?}")))?;
decoded.push(decoded_text);
}
Ok(())
}
#[must_use]
pub const fn is_special_token(&self, token_id: u32) -> bool {
(token_id as usize) < self.num_special_tokens
}
#[must_use]
#[allow(clippy::cast_possible_truncation)]
pub const fn is_byte(&self, token_id: u32) -> bool {
#[allow(clippy::cast_possible_truncation)]
if token_id < self.num_special_tokens as u32 {
false
} else {
#[allow(clippy::cast_possible_truncation)]
let shifted_id = token_id - self.num_special_tokens as u32;
shifted_id < 256
}
}
pub fn id_to_piece(&self, token_id: u32) -> Result<String> {
if token_id as usize >= self.vocab_size {
return Err(TokenizerError::InvalidConfig(format!(
"Token ID {} is out of vocabulary range (0-{})",
token_id,
self.vocab_size - 1
)));
}
self.decode(&[token_id], SpecialTokenPolicy::Keep)
}
#[allow(clippy::cast_possible_truncation)]
pub fn id_to_byte_piece(
&self,
token_id: u32,
special_token_policy: SpecialTokenPolicy,
) -> Result<Vec<u8>> {
if token_id as usize >= self.vocab_size {
return Err(TokenizerError::InvalidConfig(format!(
"Token ID {} is out of vocabulary range (0-{})",
token_id,
self.vocab_size - 1
)));
}
#[allow(clippy::cast_possible_truncation)]
if token_id < self.num_special_tokens as u32 {
match special_token_policy {
SpecialTokenPolicy::Keep => Ok(self.special_tokens[token_id as usize]
.token_str
.as_bytes()
.to_vec()),
SpecialTokenPolicy::Raise => Err(TokenizerError::SpecialTokenPolicy(format!(
"Token ID {} is a special token ({}), cannot convert to byte piece with Raise policy",
token_id, self.special_tokens[token_id as usize].token_str
))),
SpecialTokenPolicy::Ignore => Ok(vec![]),
}
} else {
#[allow(clippy::cast_possible_truncation)]
let shifted_id = token_id - self.num_special_tokens as u32;
match self.tekkenizer.decode(vec![shifted_id]) {
Ok(decoded) => Ok(decoded.as_bytes().to_vec()),
Err(e) => {
if let Some(vocab_entry) = self.vocab.get(token_id as usize) {
Ok(vocab_entry.as_bytes().to_vec())
} else {
Err(TokenizerError::Tokenizers(format!(
"Failed to decode token ID {token_id} to bytes: {e:?}. Token may represent invalid UTF-8 sequence.",
)))
}
}
}
}
}
pub fn encode_audio(&self, audio: Audio) -> Result<AudioEncoding> {
self.audio_encoder.as_ref().map_or_else(
|| {
Err(TokenizerError::Audio(
"Audio encoder not configured".to_string(),
))
},
|encoder| encoder.encode(audio),
)
}
#[must_use]
pub const fn has_audio_support(&self) -> bool {
self.audio_encoder.is_some()
}
#[must_use]
pub const fn audio_config(&self) -> Option<&AudioConfig> {
self.audio_config.as_ref()
}
pub fn to_file<P: AsRef<Path>>(&self, path: P) -> Result<()> {
let model_data = ModelData {
vocab: self.vocab_tokens.clone(),
special_tokens: Some(self.special_tokens.clone()),
config: TekkenConfig {
pattern: self.pattern.clone(),
num_vocab_tokens: self.vocab_size - self.num_special_tokens,
default_vocab_size: self.vocab_size,
default_num_special_tokens: self.num_special_tokens,
version: self.version.as_str().to_string(),
},
audio: self.audio_config.clone(),
};
let json_content = serde_json::to_string_pretty(&model_data)?;
std::fs::write(path, json_content)?;
Ok(())
}
}
#[allow(clippy::cast_possible_truncation)]
fn reload_mergeable_ranks(
vocab: Vec<TokenInfo>,
max_vocab: usize,
) -> Result<FxHashMap<Vec<u8>, u32>> {
let vocab = if vocab.len() > max_vocab {
vocab.into_iter().take(max_vocab).collect()
} else {
vocab
};
let mut ranks = FxHashMap::default();
for token in vocab {
let token_bytes = general_purpose::STANDARD.decode(&token.token_bytes)?;
#[allow(clippy::cast_possible_truncation)]
if token.rank < 256 && token_bytes != vec![token.rank as u8] {
return Err(TokenizerError::InvalidConfig(format!(
"Expected byte token at rank {} to be [{}], got {:?}",
token.rank, token.rank, token_bytes
)));
}
#[allow(clippy::cast_possible_truncation)]
ranks.insert(token_bytes, token.rank as u32);
}
#[allow(clippy::cast_possible_truncation)]
let expected_ranks: std::collections::HashSet<_> = (0..ranks.len() as u32).collect();
let actual_ranks: std::collections::HashSet<_> = ranks.values().copied().collect();
if expected_ranks != actual_ranks {
return Err(TokenizerError::InvalidConfig(
"Vocabulary ranks are not contiguous".to_string(),
));
}
Ok(ranks)
}
#[allow(clippy::too_many_lines)]
fn get_deprecated_special_tokens() -> Vec<SpecialTokenInfo> {
vec![
SpecialTokenInfo {
rank: 0,
token_str: SpecialTokens::Unk.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 1,
token_str: SpecialTokens::Bos.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 2,
token_str: SpecialTokens::Eos.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 3,
token_str: SpecialTokens::BeginInst.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 4,
token_str: SpecialTokens::EndInst.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 5,
token_str: SpecialTokens::BeginTools.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 6,
token_str: SpecialTokens::EndTools.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 7,
token_str: SpecialTokens::BeginToolResults.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 8,
token_str: SpecialTokens::EndToolResults.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 9,
token_str: SpecialTokens::ToolCalls.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 10,
token_str: SpecialTokens::Img.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 11,
token_str: SpecialTokens::Pad.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 12,
token_str: SpecialTokens::ImgBreak.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 13,
token_str: SpecialTokens::ImgEnd.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 14,
token_str: SpecialTokens::Prefix.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 15,
token_str: SpecialTokens::Middle.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 16,
token_str: SpecialTokens::Suffix.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 17,
token_str: SpecialTokens::BeginSystem.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 18,
token_str: SpecialTokens::EndSystem.as_str().to_string(),
is_control: true,
},
SpecialTokenInfo {
rank: 19,
token_str: SpecialTokens::BeginToolContent.as_str().to_string(),
is_control: true,
},
]
}