use std::collections::HashMap;
use super::scored_vocab::ScoredVocab;
use super::special::{SpecialTokenTable, TextOrSpecial};
use super::{SpecialTokens, TokenizerLoadError};
const TABLE_PIECE_LENGTH: usize = 0;
const TABLE_TOKEN_ID: usize = 1;
const TABLE_SCORE: usize = 2;
const TABLE_PIECE_ID: usize = 3;
const INVALID_SCORE: i32 = -20_000_000;
const UNKNOWN_SCORE: i32 = -10_000_000;
const GGML_TOKEN_TYPE_BYTE: i64 = 6;
pub struct GgufPlamo2Tokenizer {
vocab: ScoredVocab,
bytes: [u32; 256],
is_byte: Vec<bool>,
to_suffix_id: HashMap<u64, i32>,
table: Vec<[i32; 4]>,
special_tokens: SpecialTokenTable,
}
impl GgufPlamo2Tokenizer {
pub fn from_gguf(file: &impl frink_gguf::TensorSource) -> Result<Self, TokenizerLoadError> {
let vocab = ScoredVocab::from_gguf(file)?;
let n = vocab.len();
let mut is_byte = vec![false; n];
if let Some(frink_gguf::GgufValue::Array(items)) =
file.metadata("tokenizer.ggml.token_type")
{
for (flag, v) in is_byte.iter_mut().zip(items) {
let ty = match v {
frink_gguf::GgufValue::I32(t) => *t as i64,
frink_gguf::GgufValue::U32(t) => *t as i64,
_ => continue,
};
*flag = ty == GGML_TOKEN_TYPE_BYTE;
}
}
let mut bytes = [u32::MAX; 256];
let mut suffix_to_score: HashMap<String, f32> = HashMap::new();
for (id, text) in vocab.tokens().iter().enumerate() {
if is_byte[id] {
if let Some(b) = byte_token_value(text) {
bytes[b as usize] = id as u32;
}
continue;
}
let (_, score) = vocab
.lookup(text)
.expect("token text is in its own vocabulary");
suffix_to_score.insert(text.clone(), score);
let cpts: Vec<char> = text.chars().collect();
for i in 1..cpts.len() {
let suffix: String = cpts[i..].iter().collect();
suffix_to_score.entry(suffix).or_insert(f32::NAN);
}
}
if let Some(missing) = bytes.iter().position(|&id| id == u32::MAX) {
return Err(TokenizerLoadError::Plamo2ByteTokenMissing {
byte: missing as u8,
});
}
let mut suffixes: Vec<&str> = suffix_to_score.keys().map(String::as_str).collect();
suffixes.push("");
suffixes.sort_by(|a, b| {
let ra = a.bytes().rev();
let rb = b.bytes().rev();
ra.cmp(rb)
});
let mut suffix_to_id: HashMap<&str, i32> = HashMap::with_capacity(suffixes.len());
let mut to_suffix_id: HashMap<u64, i32> = HashMap::new();
let mut num_pieces: i32 = 0;
for &suffix in &suffixes {
suffix_to_id.insert(suffix, num_pieces);
if suffix.is_empty() {
num_pieces += 1;
continue;
}
let mut chars = suffix.char_indices();
let (_, first) = chars.next().expect("non-empty");
let rest = match chars.next() {
Some((at, _)) => &suffix[at..],
None => "",
};
let rest_id = *suffix_to_id
.get(rest)
.expect("a suffix's rest sorts before it in reversed order");
to_suffix_id.insert(piece_code(first, rest_id), num_pieces);
let mut rows = 1;
for (end, _) in suffix.char_indices().skip(1) {
if suffix_to_score.contains_key(&suffix[..end]) {
rows += 1;
}
}
rows += 1; num_pieces += rows;
}
let mut table: Vec<[i32; 4]> = Vec::with_capacity(num_pieces as usize);
for &suffix in &suffixes {
let mut ends: Vec<usize> = suffix.char_indices().skip(1).map(|(i, _)| i).collect();
ends.push(suffix.len());
for &end in ends.iter().rev() {
if end == 0 {
continue;
}
let piece = &suffix[..end];
let Some(&score) = suffix_to_score.get(piece) else {
continue;
};
let mut row = [0i32; 4];
row[TABLE_PIECE_LENGTH] = piece.chars().count() as i32;
row[TABLE_TOKEN_ID] = vocab.id_of(piece).map_or(-1, |id| id as i32);
row[TABLE_SCORE] = if score.is_finite() {
(score * 1e4).round() as i32
} else {
INVALID_SCORE
};
row[TABLE_PIECE_ID] = suffix_to_id[piece];
table.push(row);
}
let mut sentinel = [0i32; 4];
sentinel[TABLE_PIECE_LENGTH] = 1;
sentinel[TABLE_TOKEN_ID] = -1;
sentinel[TABLE_SCORE] = UNKNOWN_SCORE;
table.push(sentinel);
}
debug_assert_eq!(table.len(), num_pieces as usize);
let special_tokens = SpecialTokenTable::from_gguf(file, vocab.tokens());
Ok(Self {
vocab,
bytes,
is_byte,
to_suffix_id,
table,
special_tokens,
})
}
pub fn vocab_size(&self) -> usize {
self.vocab.len()
}
pub fn encode(&self, text: &str, specials: SpecialTokens) -> Vec<u32> {
self.special_tokens
.split(text, specials)
.into_iter()
.flat_map(|seg| match seg {
TextOrSpecial::Special(id) => vec![id],
TextOrSpecial::Text(t) => self.encode_run(t),
})
.collect()
}
fn encode_run(&self, text: &str) -> Vec<u32> {
let mut cpts: Vec<char> = text.chars().collect();
if cpts.first() == Some(&'\u{FEFF}') {
cpts.remove(0);
}
if cpts.is_empty() {
return Vec::new();
}
let n = cpts.len();
let mut scores = vec![1i64 << 60; n + 1];
scores[n] = 0;
let mut path = vec![(0i32, -1i32, 0i32); n + 1];
let mut suffix_id: i32 = 0;
for i in (0..n).rev() {
let c = cpts[i];
for p in (suffix_id as usize)..self.table.len() {
let code = piece_code(c, self.table[p][TABLE_PIECE_ID]);
suffix_id = self.to_suffix_id.get(&code).copied().unwrap_or(0);
if suffix_id > 0 || self.table[p][TABLE_SCORE] == UNKNOWN_SCORE {
break;
}
}
for p in (suffix_id as usize)..self.table.len() {
let score = self.table[p][TABLE_SCORE];
if score > INVALID_SCORE {
let len = self.table[p][TABLE_PIECE_LENGTH];
let s = scores[i + len as usize] - score as i64;
if s < scores[i] {
scores[i] = s;
let mut count = path[i + len as usize].2 + 1;
if score == UNKNOWN_SCORE {
count += utf8_extra_bytes(c);
}
path[i] = (len, self.table[p][TABLE_TOKEN_ID], count);
}
}
if score == UNKNOWN_SCORE {
break;
}
}
}
let mut out = Vec::with_capacity(path[0].2.max(0) as usize);
let mut pos = 0;
while pos < n {
let (len, id, _) = path[pos];
if id >= 0 {
out.push(id as u32);
} else {
let mut buf = [0u8; 4];
for &b in cpts[pos].encode_utf8(&mut buf).as_bytes() {
out.push(self.bytes[b as usize]);
}
}
debug_assert!(len > 0, "every position advances");
pos += len.max(1) as usize;
}
out
}
pub fn decode(&self, ids: &[u32]) -> String {
String::from_utf8_lossy(&self.decode_bytes(ids)).into_owned()
}
pub fn decode_bytes(&self, ids: &[u32]) -> Vec<u8> {
let mut out = Vec::new();
for &id in ids {
let Some(text) = self.vocab.token(id) else {
continue;
};
if self.is_byte.get(id as usize).copied().unwrap_or(false) {
if let Some(b) = byte_token_value(text) {
out.push(b);
continue;
}
}
out.extend_from_slice(text.as_bytes());
}
out
}
}
fn byte_token_value(text: &str) -> Option<u8> {
let hex = text.strip_prefix("<0x")?.strip_suffix('>')?;
if hex.len() != 2 {
return None;
}
u8::from_str_radix(hex, 16).ok()
}
fn piece_code(c: char, rest_id: i32) -> u64 {
((c as u64) << 32) | (rest_id as u32 as u64)
}
fn utf8_extra_bytes(c: char) -> i32 {
let c = c as u32;
(c >= 0x80) as i32 + (c >= 0x800) as i32 + (c >= 0x10000) as i32
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tokenizer::scored_vocab::MetadataOnlyGguf as MetaOnly;
fn fixture() -> GgufPlamo2Tokenizer {
let mut tokens: Vec<String> = vec![
"<|plamo:unk|>".into(),
"<|plamo:bos|>".into(),
"<|plamo:eos|>".into(),
"<|plamo:pad|>".into(),
];
let mut types: Vec<i32> = vec![2, 3, 3, 3];
for b in 0..256u32 {
tokens.push(format!("<0x{b:02X}>"));
types.push(6);
}
let pieces: &[(&str, f32)] = &[
("a", -3.0),
("b", -3.0),
("ab", -4.0),
("abc", -2.0),
("bc", -3.5),
("c", -3.0),
(" the", -1.5),
("the", -2.5),
(" ", -4.0),
("日本", -2.0),
("本", -5.0),
];
let mut scores = vec![0.0f32; tokens.len()];
for (t, s) in pieces {
tokens.push((*t).to_string());
types.push(1);
scores.push(*s);
}
let toks: Vec<&str> = tokens.iter().map(String::as_str).collect();
let meta = MetaOnly::new()
.with_tokens(&toks)
.with_scores(&scores)
.with(
"tokenizer.ggml.token_type",
frink_gguf::GgufValue::Array(
types.into_iter().map(frink_gguf::GgufValue::I32).collect(),
),
);
GgufPlamo2Tokenizer::from_gguf(&meta).unwrap()
}
fn id(t: &GgufPlamo2Tokenizer, s: &str) -> u32 {
t.vocab.id_of(s).unwrap()
}
#[test]
fn picks_the_highest_scoring_segmentation_not_the_greedy_one() {
let t = fixture();
assert_eq!(t.encode("abc", SpecialTokens::AsText), vec![id(&t, "abc")]);
assert_eq!(
t.encode("the the", SpecialTokens::AsText),
vec![id(&t, "the"), id(&t, " the")]
);
assert_eq!(
t.encode("日本", SpecialTokens::AsText),
vec![id(&t, "日本")]
);
}
#[test]
fn a_character_no_piece_covers_falls_back_to_its_utf8_bytes() {
let t = fixture();
assert_eq!(
t.encode("z", SpecialTokens::AsText),
vec![t.bytes[b'z' as usize]]
);
let got = t.encode("語", SpecialTokens::AsText);
let want: Vec<u32> = "語".bytes().map(|b| t.bytes[b as usize]).collect();
assert_eq!(got, want);
assert_eq!(t.decode(&got), "語");
assert_eq!(
t.decode(&t.encode("abc z 日本", SpecialTokens::AsText)),
"abc z 日本"
);
}
#[test]
fn a_missing_byte_token_is_refused_at_load() {
let toks = ["<|plamo:unk|>", "<0x00>", "a"];
let meta = MetaOnly::new()
.with_tokens(&toks)
.with_scores(&[0.0, 0.0, -1.0])
.with(
"tokenizer.ggml.token_type",
frink_gguf::GgufValue::Array(
[2, 6, 1]
.into_iter()
.map(frink_gguf::GgufValue::I32)
.collect(),
),
);
assert!(matches!(
GgufPlamo2Tokenizer::from_gguf(&meta),
Err(TokenizerLoadError::Plamo2ByteTokenMissing { byte: 1 })
));
}
}