use std::collections::HashMap;
#[cfg(test)]
mod tests;
pub mod special_tokens {
pub const MULTILINGUAL_VOCAB_THRESHOLD: usize = 51865;
pub const EOT_ENGLISH: u32 = 50256;
pub const SOT_ENGLISH: u32 = 50257;
pub const EOT_MULTILINGUAL: u32 = 50257;
pub const SOT_MULTILINGUAL: u32 = 50258;
pub const LANG_BASE_MULTILINGUAL: u32 = 50259;
pub const TRANSCRIBE_MULTILINGUAL: u32 = 50359;
pub const NO_TIMESTAMPS_MULTILINGUAL: u32 = 50363;
pub const EOT: u32 = EOT_MULTILINGUAL;
pub const SOT: u32 = SOT_MULTILINGUAL;
pub const LANG_BASE: u32 = LANG_BASE_MULTILINGUAL;
pub const TRANSLATE: u32 = 50358;
pub const TRANSCRIBE: u32 = TRANSCRIBE_MULTILINGUAL;
pub const SPEAKER_TURN: u32 = 50360;
pub const PREV: u32 = 50361;
pub const NO_SPEECH: u32 = 50362;
pub const NO_TIMESTAMPS: u32 = NO_TIMESTAMPS_MULTILINGUAL;
pub const TIMESTAMP_BASE: u32 = 50364;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct SpecialTokens {
pub eot: u32,
pub sot: u32,
pub lang_base: u32,
pub transcribe: u32,
pub no_timestamps: u32,
pub timestamp_base: u32,
pub is_multilingual: bool,
}
impl SpecialTokens {
#[must_use]
pub fn for_vocab_size(n_vocab: usize) -> Self {
if n_vocab >= MULTILINGUAL_VOCAB_THRESHOLD {
Self::multilingual()
} else {
Self::english_only()
}
}
#[must_use]
pub const fn multilingual() -> Self {
Self {
eot: EOT_MULTILINGUAL,
sot: SOT_MULTILINGUAL,
lang_base: LANG_BASE_MULTILINGUAL,
transcribe: TRANSCRIBE_MULTILINGUAL,
no_timestamps: NO_TIMESTAMPS_MULTILINGUAL,
timestamp_base: 50364,
is_multilingual: true,
}
}
#[must_use]
pub const fn english_only() -> Self {
Self {
eot: EOT_ENGLISH,
sot: SOT_ENGLISH,
lang_base: 50258, transcribe: 50358,
no_timestamps: 50362,
timestamp_base: 50363,
is_multilingual: false,
}
}
#[must_use]
pub fn initial_tokens(&self) -> [u32; 4] {
[
self.sot,
self.lang_base, self.transcribe,
self.no_timestamps,
]
}
}
impl Default for SpecialTokens {
fn default() -> Self {
Self::multilingual()
}
}
#[must_use]
pub fn language_token(lang_code: &str) -> Option<u32> {
language_offset(lang_code).map(|offset| LANG_BASE + offset)
}
#[must_use]
pub const fn is_timestamp(token_id: u32) -> bool {
token_id >= TIMESTAMP_BASE
}
#[must_use]
pub fn timestamp_to_seconds(token_id: u32) -> Option<f32> {
if token_id >= TIMESTAMP_BASE {
Some((token_id - TIMESTAMP_BASE) as f32 * 0.02)
} else {
None
}
}
#[must_use]
#[allow(clippy::cast_possible_truncation)]
pub fn language_offset(lang_code: &str) -> Option<u32> {
crate::detection::SUPPORTED_LANGUAGES
.iter()
.position(|&c| c == lang_code)
.map(|i| i as u32)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct MergeRule {
pub first: Vec<u8>,
pub second: Vec<u8>,
}
impl MergeRule {
#[must_use]
pub fn new(first: Vec<u8>, second: Vec<u8>) -> Self {
Self { first, second }
}
#[must_use]
pub fn merged(&self) -> Vec<u8> {
let mut result = self.first.clone();
result.extend_from_slice(&self.second);
result
}
}
#[derive(Debug, Clone)]
pub struct Vocabulary {
id_to_bytes: Vec<Vec<u8>>,
bytes_to_id: HashMap<Vec<u8>, u32>,
merge_rules: Vec<MergeRule>,
merge_lookup: HashMap<(Vec<u8>, Vec<u8>), u32>,
}
impl Vocabulary {
#[must_use]
pub fn new() -> Self {
Self {
id_to_bytes: Vec::new(),
bytes_to_id: HashMap::new(),
merge_rules: Vec::new(),
merge_lookup: HashMap::new(),
}
}
#[must_use]
pub fn with_base_tokens() -> Self {
let mut vocab = Self::new();
for byte in 0..=255u8 {
vocab.add_token(vec![byte]);
}
vocab
}
pub fn add_token(&mut self, bytes: Vec<u8>) -> u32 {
let id = self.id_to_bytes.len() as u32;
self.bytes_to_id.insert(bytes.clone(), id);
self.id_to_bytes.push(bytes);
id
}
pub fn add_merge(&mut self, first: Vec<u8>, second: Vec<u8>) -> u32 {
let rule = MergeRule::new(first.clone(), second.clone());
let merged = rule.merged();
let merged_id = if let Some(&id) = self.bytes_to_id.get(&merged) {
id
} else {
self.add_token(merged)
};
self.merge_lookup.insert((first, second), merged_id);
self.merge_rules.push(rule);
merged_id
}
#[must_use]
pub fn len(&self) -> usize {
self.id_to_bytes.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.id_to_bytes.is_empty()
}
#[must_use]
pub fn get_bytes(&self, token_id: u32) -> Option<&[u8]> {
self.id_to_bytes.get(token_id as usize).map(Vec::as_slice)
}
#[must_use]
pub fn get_id(&self, bytes: &[u8]) -> Option<u32> {
self.bytes_to_id.get(bytes).copied()
}
#[must_use]
pub fn get_merge(&self, first: &[u8], second: &[u8]) -> Option<u32> {
self.merge_lookup
.get(&(first.to_vec(), second.to_vec()))
.copied()
}
#[must_use]
pub fn merge_priority(&self, first: &[u8], second: &[u8]) -> Option<usize> {
self.merge_rules
.iter()
.position(|r| r.first.as_slice() == first && r.second.as_slice() == second)
}
#[must_use]
pub fn decode(&self, tokens: &[u32]) -> Option<String> {
if tokens.is_empty() {
return Some(String::new());
}
let mut bytes = Vec::new();
for &token_id in tokens {
if token_id >= special_tokens::EOT {
continue;
}
let token_bytes = self.get_bytes(token_id)?;
bytes.extend_from_slice(token_bytes);
}
Some(String::from_utf8_lossy(&bytes).into_owned())
}
#[must_use]
pub fn num_merges(&self) -> usize {
self.merge_rules.len()
}
#[must_use]
pub fn to_bytes(&self) -> Vec<u8> {
let mut bytes = Vec::new();
bytes.extend_from_slice(&(self.id_to_bytes.len() as u32).to_le_bytes());
bytes.extend_from_slice(&(self.merge_rules.len() as u32).to_le_bytes());
for token_bytes in &self.id_to_bytes {
let len = token_bytes.len() as u16;
bytes.extend_from_slice(&len.to_le_bytes());
bytes.extend_from_slice(token_bytes);
}
for rule in &self.merge_rules {
let first_len = rule.first.len() as u16;
bytes.extend_from_slice(&first_len.to_le_bytes());
bytes.extend_from_slice(&rule.first);
let second_len = rule.second.len() as u16;
bytes.extend_from_slice(&second_len.to_le_bytes());
bytes.extend_from_slice(&rule.second);
}
bytes
}
fn read_length_prefixed(data: &[u8], offset: &mut usize) -> Option<Vec<u8>> {
if *offset + 2 > data.len() {
return None;
}
let len = u16::from_le_bytes([data[*offset], data[*offset + 1]]) as usize;
*offset += 2;
if *offset + len > data.len() {
return None;
}
let bytes = data[*offset..*offset + len].to_vec();
*offset += len;
Some(bytes)
}
#[must_use]
pub fn from_bytes(data: &[u8]) -> Option<Self> {
if data.len() < 8 {
return None;
}
let n_tokens = u32::from_le_bytes([data[0], data[1], data[2], data[3]]) as usize;
let n_merges = u32::from_le_bytes([data[4], data[5], data[6], data[7]]) as usize;
let mut offset = 8;
let mut vocab = Self::new();
for _ in 0..n_tokens {
let token_bytes = Self::read_length_prefixed(data, &mut offset)?;
vocab.add_token(token_bytes);
}
for _ in 0..n_merges {
let first = Self::read_length_prefixed(data, &mut offset)?;
let second = Self::read_length_prefixed(data, &mut offset)?;
vocab.merge_lookup.insert(
(first.clone(), second.clone()),
vocab.id_to_bytes.len() as u32,
);
vocab.merge_rules.push(MergeRule::new(first, second));
}
Some(vocab)
}
#[must_use]
pub fn merge_rules(&self) -> &[MergeRule] {
&self.merge_rules
}
}
impl Default for Vocabulary {
fn default() -> Self {
Self::new()
}
}