use crate::Tokenizer;
use crate::merges::MergeTable;
use crate::pretokenize;
use crate::specials::{Segment, SpecialTokens};
use crate::vocab::Vocab;
use kopitiam_core::{Error, Result};
#[derive(Debug, Clone)]
pub struct BpeTokenizer {
vocab: Vocab,
byte_ids: [u32; 256],
merges: MergeTable,
add_prefix_space: bool,
specials: SpecialTokens,
}
impl BpeTokenizer {
pub fn from_vocab_and_merges(
vocab_entries: Vec<Vec<u8>>,
merge_pairs: Vec<(Vec<u8>, Vec<u8>)>,
) -> Result<Self> {
Self::from_vocab(Vocab::from_entries(vocab_entries)?, merge_pairs)
}
pub(crate) fn from_vocab(vocab: Vocab, merge_pairs: Vec<(Vec<u8>, Vec<u8>)>) -> Result<Self> {
let merges = MergeTable::build(&merge_pairs, |b| vocab.id_of(b))?;
let mut byte_ids = [0u32; 256];
for b in 0u16..=255 {
let b = b as u8;
byte_ids[b as usize] = vocab.id_of(&[b]).ok_or_else(|| Error::MalformedModel {
format: "bpe-vocab",
reason: format!(
"byte-level vocab has no single-byte token for byte {b:#04x}; \
byte-level BPE requires all 256 bytes as base tokens"
),
})?;
}
Ok(Self {
vocab,
byte_ids,
merges,
add_prefix_space: false,
specials: SpecialTokens::new(),
})
}
#[must_use]
pub fn with_add_prefix_space(mut self, add_prefix_space: bool) -> Self {
self.add_prefix_space = add_prefix_space;
self
}
pub fn add_special_token(&mut self, content: impl Into<String>, id: u32) -> Result<()> {
let content = content.into();
self.vocab.insert(id, content.clone().into_bytes())?;
self.specials.register(content, id);
Ok(())
}
pub fn special_token_id(&self, content: &str) -> Option<u32> {
self.specials.id_of(content)
}
pub fn token_id(&self, token: &[u8]) -> Option<u32> {
self.vocab.id_of(token)
}
fn chunk_to_symbols(&self, chunk: &str) -> Vec<u32> {
chunk.bytes().map(|b| self.byte_ids[b as usize]).collect()
}
fn merge(&self, mut symbols: Vec<u32>) -> Vec<u32> {
loop {
let mut best: Option<(usize, u32, u32)> = None; for i in 0..symbols.len().saturating_sub(1) {
if let Some(rule) = self.merges.get(symbols[i], symbols[i + 1])
&& best.is_none_or(|(_, best_rank, _)| rule.rank < best_rank)
{
best = Some((i, rule.rank, rule.merged_id));
}
}
let Some((i, _, merged_id)) = best else {
return symbols;
};
symbols[i] = merged_id;
symbols.remove(i + 1);
}
}
}
impl Tokenizer for BpeTokenizer {
fn encode(&self, text: &str) -> Result<Vec<u32>> {
if text.is_empty() {
return Ok(Vec::new());
}
let prefixed;
let text = if self.add_prefix_space && !text.starts_with(' ') {
prefixed = format!(" {text}");
prefixed.as_str()
} else {
text
};
let mut ids = Vec::new();
for segment in self.specials.split(text) {
match segment {
Segment::Special(id) => ids.push(id),
Segment::Text(span) => {
for (start, end) in pretokenize::split(span) {
let symbols = self.chunk_to_symbols(&span[start..end]);
ids.extend(self.merge(symbols));
}
}
}
}
Ok(ids)
}
fn decode(&self, ids: &[u32]) -> Result<String> {
let mut bytes = Vec::new();
for &id in ids {
let token_bytes = self.vocab.bytes_of(id).ok_or(Error::IndexOutOfBounds {
dim: 0,
index: id as usize,
len: self.vocab.len(),
})?;
bytes.extend_from_slice(token_bytes);
}
Ok(String::from_utf8_lossy(&bytes).into_owned())
}
fn vocab_size(&self) -> usize {
self.vocab.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
fn tiny_tokenizer() -> BpeTokenizer {
let mut vocab: Vec<Vec<u8>> = (0u16..=255).map(|b| vec![b as u8]).collect();
vocab.push(b"ab".to_vec()); vocab.push(b"abc".to_vec()); let merges = vec![
(b"a".to_vec(), b"b".to_vec()),
(b"ab".to_vec(), b"c".to_vec()),
];
BpeTokenizer::from_vocab_and_merges(vocab, merges).expect("valid tiny tokenizer")
}
fn byte_id(tok: &BpeTokenizer, b: u8) -> u32 {
tok.token_id(&[b]).unwrap()
}
#[test]
fn a_plus_b_then_ab_plus_c_merges_produce_the_single_token_abc() {
let tok = tiny_tokenizer();
let ids = tok.encode("abc").unwrap();
assert_eq!(ids, vec![tok.token_id(b"abc").unwrap()]);
}
#[test]
fn merges_apply_in_rank_order_not_leftmost_scan_order() {
let mut vocab: Vec<Vec<u8>> = (0u16..=255).map(|b| vec![b as u8]).collect();
vocab.push(b"bc".to_vec()); vocab.push(b"ab".to_vec()); let merges = vec![
(b"b".to_vec(), b"c".to_vec()), (b"a".to_vec(), b"b".to_vec()), ];
let tok = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
let ids = tok.encode("abc").unwrap();
let expected = vec![tok.token_id(b"a").unwrap(), tok.token_id(b"bc").unwrap()];
assert_eq!(
ids, expected,
"expected rank-0 (b,c) to win over the leftmost-but-lower-priority (a,b) pair"
);
let wrong_leftmost_greedy_answer =
vec![tok.token_id(b"ab").unwrap(), tok.token_id(b"c").unwrap()];
assert_ne!(ids, wrong_leftmost_greedy_answer);
}
#[test]
fn vocab_size_counts_every_registered_id() {
let tok = tiny_tokenizer();
assert_eq!(tok.vocab_size(), 258);
}
#[test]
fn empty_string_encodes_to_empty() {
let tok = tiny_tokenizer();
assert_eq!(tok.encode("").unwrap(), Vec::<u32>::new());
}
#[test]
fn decode_of_empty_is_empty() {
let tok = tiny_tokenizer();
assert_eq!(tok.decode(&[]).unwrap(), "");
}
#[test]
fn decode_rejects_unknown_id_gracefully() {
let tok = tiny_tokenizer();
let err = tok.decode(&[99_999]).unwrap_err();
assert!(matches!(err, Error::IndexOutOfBounds { .. }));
}
#[test]
fn special_tokens_are_never_split_by_bpe() {
let mut tok = tiny_tokenizer();
tok.add_special_token("<|endoftext|>", 300).unwrap();
let ids = tok.encode("ab<|endoftext|>c").unwrap();
assert_eq!(ids.last().copied(), Some(byte_id(&tok, b'c')));
assert!(ids.contains(&300));
assert_eq!(ids.iter().filter(|&&id| id == 300).count(), 1);
}
#[test]
fn decode_of_encode_is_byte_exact_for_ascii() {
let tok = tiny_tokenizer();
let s = "abcabc hello world";
assert_eq!(tok.decode(&tok.encode(s).unwrap()).unwrap(), s);
}
}