use std::collections::HashMap;
use std::str;
use anyhow::{Context as _, Result, anyhow, bail, ensure};
use super::gguf::Gguf;
const NORMAL: u32 = 1;
const UNKNOWN: u32 = 2;
const CONTROL: u32 = 3;
const USER_DEFINED: u32 = 4;
const UNUSED: u32 = 5;
const BYTE: u32 = 6;
#[derive(Clone, Copy)]
struct Merge {
rank: usize,
token: u32,
}
pub(super) struct Tokenizer {
pieces: Vec<Vec<u8>>,
byte_tokens: [Option<u32>; 256],
merges: HashMap<(u32, u32), Merge>,
special: Vec<(String, u32)>,
eos: u32,
splitter: PreTokenizer,
}
#[derive(Clone, Copy)]
enum PreTokenizer {
Gpt4o,
}
impl Tokenizer {
pub(super) fn from_gpt_oss(gguf: &Gguf) -> Result<Self> {
ensure!(
gguf.string("tokenizer.ggml.pre")? == "gpt-4o",
"expected GPT-4o pre-tokenizer"
);
Self::from_gguf(gguf, 201_088, PreTokenizer::Gpt4o)
}
fn from_gguf(gguf: &Gguf, vocabulary: usize, pre_tokenizer: PreTokenizer) -> Result<Self> {
let texts = gguf.strings("tokenizer.ggml.tokens")?.to_vec();
let merge_texts = gguf.strings("tokenizer.ggml.merges")?;
let types = gguf
.i32s("tokenizer.ggml.token_type")?
.iter()
.copied()
.enumerate()
.map(|(index, kind)| {
u32::try_from(kind).with_context(|| format!("token {index} has negative type"))
})
.collect::<Result<Vec<_>>>()?;
ensure!(
texts.len() == vocabulary,
"tokenizer vocabulary has {} entries, expected {vocabulary}",
texts.len(),
);
ensure!(
types.len() == texts.len(),
"token type count differs from vocabulary"
);
let byte_characters = byte_characters();
let byte_values = byte_characters
.iter()
.copied()
.enumerate()
.map(|(byte, character)| {
(
character,
u8::try_from(byte).expect("byte table has 256 entries"),
)
})
.collect::<HashMap<_, _>>();
let mut ids = HashMap::with_capacity(texts.len());
for (index, text) in texts.iter().enumerate() {
let token = u32::try_from(index).context("token index exceeds u32")?;
if ids.insert(text.clone(), token).is_some() {
bail!("duplicate tokenizer token {text:?}");
}
}
let mut pieces = Vec::with_capacity(texts.len());
let mut special = Vec::new();
for (index, (text, kind)) in texts.iter().zip(&types).enumerate() {
let token = u32::try_from(index).context("token index exceeds u32")?;
ensure!(
matches!(
*kind,
NORMAL | UNKNOWN | CONTROL | USER_DEFINED | UNUSED | BYTE
),
"token {index} has unknown type {kind}"
);
if matches!(*kind, CONTROL | USER_DEFINED) {
pieces.push(text.as_bytes().to_vec());
special.push((text.clone(), token));
} else {
pieces.push(decode_token_text(text, *kind, &byte_values));
}
}
special.sort_unstable_by(|left, right| {
right.0.len().cmp(&left.0.len()).then(left.1.cmp(&right.1))
});
let mut byte_tokens = [None; 256];
for (byte, character) in byte_characters.iter().enumerate() {
byte_tokens[byte] = ids.get(&character.to_string()).copied();
ensure!(
byte_tokens[byte].is_some(),
"missing byte token 0x{byte:02x}"
);
}
let merges = Self::build_merges(merge_texts, &ids, &texts)?;
let eos = gguf.u32("tokenizer.ggml.eos_token_id")?;
ensure!(
usize::try_from(eos)? < texts.len(),
"EOS token is outside vocabulary"
);
for token in [
"<|start|>",
"<|channel|>",
"<|message|>",
"<|end|>",
"<|return|>",
] {
ensure!(ids.contains_key(token), "missing Harmony token {token}");
}
Ok(Self {
pieces,
byte_tokens,
merges,
special,
eos,
splitter: pre_tokenizer,
})
}
fn build_merges(
merge_texts: &[String],
ids: &HashMap<String, u32>,
texts: &[String],
) -> Result<HashMap<(u32, u32), Merge>> {
let mut merges = HashMap::with_capacity(merge_texts.len());
for (rank, text) in merge_texts.iter().enumerate() {
let (left, right) = text
.split_once(' ')
.ok_or_else(|| anyhow!("invalid tokenizer merge {rank}: {text:?}"))?;
ensure!(
!left.is_empty() && !right.is_empty() && !right.contains(' '),
"invalid tokenizer merge {rank}"
);
let left = *ids
.get(left)
.ok_or_else(|| anyhow!("merge {rank} has missing left token"))?;
let right = *ids
.get(right)
.ok_or_else(|| anyhow!("merge {rank} has missing right token"))?;
let merged_text = format!(
"{}{}",
texts[usize::try_from(left)?],
texts[usize::try_from(right)?]
);
let token = *ids
.get(&merged_text)
.ok_or_else(|| anyhow!("merge {rank} produces missing token"))?;
merges.entry((left, right)).or_insert(Merge { rank, token });
}
Ok(merges)
}
pub(super) fn encode(&self, text: &str, parse_special: bool) -> Result<Vec<u32>> {
let mut output = Vec::new();
if !parse_special {
self.encode_ordinary(text, &mut output)?;
return Ok(output);
}
let mut start = 0;
while start < text.len() {
let Some((offset, length, token)) = self.next_special(&text[start..]) else {
self.encode_ordinary(&text[start..], &mut output)?;
break;
};
self.encode_ordinary(&text[start..start + offset], &mut output)?;
output.push(token);
start += offset + length;
}
Ok(output)
}
pub(super) fn piece(&self, token: u32) -> Result<&[u8]> {
self.pieces
.get(usize::try_from(token)?)
.map(Vec::as_slice)
.ok_or_else(|| anyhow!("token {token} is outside vocabulary"))
}
pub(super) fn is_eos(&self, token: u32) -> bool {
token == self.eos
}
fn encode_ordinary(&self, text: &str, output: &mut Vec<u32>) -> Result<()> {
for word in self.splitter.split(text) {
let mut symbols = word
.bytes()
.map(|byte| {
self.byte_tokens[usize::from(byte)]
.ok_or_else(|| anyhow!("missing byte token 0x{byte:02x}"))
})
.collect::<Result<Vec<_>>>()?;
self.apply_bpe(&mut symbols);
output.extend(symbols);
}
Ok(())
}
fn apply_bpe(&self, symbols: &mut Vec<u32>) {
while symbols.len() > 1 {
let best = symbols
.windows(2)
.filter_map(|pair| {
self.merges
.get(&(pair[0], pair[1]))
.map(|merge| (merge.rank, pair[0], pair[1], merge.token))
})
.min_by_key(|merge| merge.0);
let Some((_, left, right, merged)) = best else {
break;
};
let mut read = 0;
let mut write = 0;
while read < symbols.len() {
if read + 1 < symbols.len() && symbols[read] == left && symbols[read + 1] == right {
symbols[write] = merged;
read += 2;
} else {
symbols[write] = symbols[read];
read += 1;
}
write += 1;
}
symbols.truncate(write);
}
}
fn next_special(&self, text: &str) -> Option<(usize, usize, u32)> {
self.special
.iter()
.filter_map(|(special, token)| {
text.find(special)
.map(|offset| (offset, special.len(), *token))
})
.min_by(|left, right| left.0.cmp(&right.0).then(right.1.cmp(&left.1)))
}
}
#[derive(Default)]
pub(super) struct Utf8Decoder {
pending: Vec<u8>,
}
impl Utf8Decoder {
pub(super) fn push(&mut self, piece: &[u8]) -> Result<Option<String>> {
self.pending.extend_from_slice(piece);
match str::from_utf8(&self.pending) {
Ok(text) => {
let text = text.to_owned();
self.pending.clear();
Ok((!text.is_empty()).then_some(text))
}
Err(error) if error.error_len().is_none() => {
let complete = error.valid_up_to();
if complete == 0 {
return Ok(None);
}
let text = str::from_utf8(&self.pending[..complete])?.to_owned();
self.pending.drain(..complete);
Ok(Some(text))
}
Err(error) => bail!(
"token pieces contain invalid UTF-8 at byte {}",
error.valid_up_to()
),
}
}
pub(super) fn finish(self) -> Result<String> {
Ok(str::from_utf8(&self.pending)
.context("token pieces end with incomplete UTF-8")?
.to_owned())
}
}
fn decode_token_text(text: &str, kind: u32, byte_values: &HashMap<char, u8>) -> Vec<u8> {
if kind == BYTE
&& let Some(byte) = text
.strip_prefix("<0x")
.and_then(|text| text.strip_suffix('>'))
.and_then(|hex| u8::from_str_radix(hex, 16).ok())
{
return vec![byte];
}
let mut output = Vec::with_capacity(text.len());
for character in text.chars() {
if let Some(byte) = byte_values.get(&character) {
output.push(*byte);
} else {
let mut encoded = [0; 4];
output.extend_from_slice(character.encode_utf8(&mut encoded).as_bytes());
}
}
output
}
fn byte_characters() -> [char; 256] {
let mut characters = ['\0'; 256];
let mut extra = 0_u32;
for byte in u8::MIN..=u8::MAX {
let codepoint = if (b'!'..=b'~').contains(&byte)
|| (0xa1..=0xac).contains(&byte)
|| (0xae..=0xff).contains(&byte)
{
u32::from(byte)
} else {
let codepoint = 256 + extra;
extra += 1;
codepoint
};
characters[usize::from(byte)] = char::from_u32(codepoint).expect("valid GPT-2 byte map");
}
characters
}
impl PreTokenizer {
fn split(self, text: &str) -> Vec<&str> {
match self {
Self::Gpt4o => split_words(text, true),
}
}
}
fn split_words(text: &str, attach_contractions: bool) -> Vec<&str> {
let mut words = Vec::new();
let mut start = 0;
while start < text.len() {
let length = word_len(&text[start..], attach_contractions);
words.push(&text[start..start + length]);
start += length;
}
words
}
fn word_len(text: &str, attach_contractions: bool) -> usize {
const CONTRACTIONS: [&str; 7] = ["'re", "'ve", "'ll", "'s", "'t", "'m", "'d"];
if !attach_contractions
&& let Some(length) = CONTRACTIONS.iter().find_map(|suffix| {
text.get(..suffix.len())
.filter(|prefix| prefix.eq_ignore_ascii_case(suffix))
.map(str::len)
})
{
return length;
}
let first = text.chars().next().expect("nonempty tokenizer input");
if first.is_alphabetic() {
let length = take_while(text, char::is_alphabetic);
return length + contraction_len(&text[length..], attach_contractions);
}
if first != '\r' && first != '\n' && !first.is_alphabetic() && !first.is_numeric() {
let end = first.len_utf8();
if text[end..].chars().next().is_some_and(char::is_alphabetic) {
let letters = take_while(&text[end..], char::is_alphabetic);
return end + letters + contraction_len(&text[end + letters..], attach_contractions);
}
}
if first.is_numeric() {
return text
.chars()
.take_while(|character| character.is_numeric())
.take(3)
.map(char::len_utf8)
.sum();
}
let prefix = usize::from(first == ' ');
let punctuation = take_while(&text[prefix..], |character| {
!character.is_whitespace() && !character.is_alphabetic() && !character.is_numeric()
});
if punctuation > 0 {
let mut length = prefix + punctuation;
length += take_while(&text[length..], |character| {
character == '\r' || character == '\n'
});
return length;
}
if first.is_whitespace() {
let length = take_while(text, char::is_whitespace);
let whitespace = &text[..length];
if let Some((offset, character)) = whitespace
.char_indices()
.rev()
.find(|(_, character)| *character == '\r' || *character == '\n')
{
return offset + character.len_utf8();
}
if text[length..]
.chars()
.next()
.is_some_and(|character| !character.is_whitespace())
&& whitespace.chars().nth_back(1).is_some()
{
return length - whitespace.chars().next_back().unwrap().len_utf8();
}
return length;
}
first.len_utf8()
}
fn contraction_len(text: &str, enabled: bool) -> usize {
const CONTRACTIONS: [&str; 7] = ["'re", "'ve", "'ll", "'s", "'t", "'m", "'d"];
enabled
.then(|| {
CONTRACTIONS.iter().find_map(|suffix| {
text.get(..suffix.len())
.filter(|prefix| prefix.eq_ignore_ascii_case(suffix))
.map(str::len)
})
})
.flatten()
.unwrap_or(0)
}
fn take_while(text: &str, predicate: impl Fn(char) -> bool) -> usize {
text.chars()
.take_while(|character| predicate(*character))
.map(char::len_utf8)
.sum()
}