use alloc::collections::{BTreeMap, BinaryHeap};
use alloc::{
format,
string::{String, ToString},
vec,
vec::Vec,
};
use core::cmp::Reverse;
use crate::error::TokenizerError;
use crate::pretokenize::{decode_byte_level, encode_byte_level, gpt2_split};
use crate::tokenizer::{split_words, TokenId, Tokenizer};
const UNK: &str = "<unk>";
#[derive(Debug, Clone)]
pub struct BpeTokenizer {
vocab: BTreeMap<String, TokenId>,
id_to_token: BTreeMap<TokenId, String>,
merge_ranks: BTreeMap<(String, String), usize>,
byte_level: bool,
special_tokens: BTreeMap<String, TokenId>,
special_id_to_token: BTreeMap<TokenId, String>,
#[cfg(feature = "normalization")]
normalize: Option<crate::normalize::NormalizationForm>,
}
impl BpeTokenizer {
#[must_use]
pub fn from_vocab_merges(
vocab: BTreeMap<String, TokenId>,
merges: Vec<(String, String)>,
) -> Self {
let merge_ranks = merges
.into_iter()
.enumerate()
.map(|(rank, pair)| (pair, rank))
.collect();
let id_to_token = vocab
.iter()
.map(|(token, &id)| (id, token.clone()))
.collect();
Self {
vocab,
id_to_token,
merge_ranks,
byte_level: false,
special_tokens: BTreeMap::new(),
special_id_to_token: BTreeMap::new(),
#[cfg(feature = "normalization")]
normalize: None,
}
}
#[must_use]
pub fn with_byte_level(mut self) -> Self {
self.byte_level = true;
self
}
#[cfg(feature = "normalization")]
#[must_use]
pub fn with_normalization(mut self, form: crate::normalize::NormalizationForm) -> Self {
self.normalize = Some(form);
self
}
#[must_use]
pub fn with_special_tokens(mut self, specials: BTreeMap<String, TokenId>) -> Self {
self.special_id_to_token = specials.iter().map(|(t, &id)| (id, t.clone())).collect();
self.special_tokens = specials;
self
}
#[must_use]
pub fn vocab_size(&self) -> usize {
self.vocab.len()
}
fn tokenize_word(&self, word: &str) -> Vec<String> {
let mut symbols: Vec<String> = word.chars().map(|c| c.to_string()).collect();
if symbols.len() < 2 {
return symbols;
}
let n = symbols.len();
let mut prev: Vec<Option<usize>> = (0..n).map(|i| i.checked_sub(1)).collect();
let mut next: Vec<Option<usize>> = (0..n).map(|i| (i + 1 < n).then_some(i + 1)).collect();
let mut alive = vec![true; n];
let mut heap: BinaryHeap<Reverse<(usize, usize)>> = BinaryHeap::new();
for i in 0..n - 1 {
if let Some(&rank) = self
.merge_ranks
.get(&(symbols[i].clone(), symbols[i + 1].clone()))
{
heap.push(Reverse((rank, i)));
}
}
while let Some(Reverse((rank, i))) = heap.pop() {
if !alive[i] {
continue;
}
let Some(j) = next[i] else { continue };
if !alive[j] {
continue;
}
match self
.merge_ranks
.get(&(symbols[i].clone(), symbols[j].clone()))
{
Some(&cur) if cur == rank => {}
_ => continue, }
symbols[i] = format!("{}{}", symbols[i], symbols[j]);
alive[j] = false;
next[i] = next[j];
if let Some(k) = next[j] {
prev[k] = Some(i);
}
if let Some(p) = prev[i] {
if let Some(&r) = self
.merge_ranks
.get(&(symbols[p].clone(), symbols[i].clone()))
{
heap.push(Reverse((r, p)));
}
}
if let Some(k) = next[i] {
if let Some(&r) = self
.merge_ranks
.get(&(symbols[i].clone(), symbols[k].clone()))
{
heap.push(Reverse((r, i)));
}
}
}
let mut out = Vec::new();
let mut cur = Some(0usize);
while let Some(i) = cur {
if alive[i] {
out.push(core::mem::take(&mut symbols[i]));
}
cur = next[i];
}
out
}
fn push_symbol(&self, sub: String, ids: &mut Vec<TokenId>) -> Result<(), TokenizerError> {
match self.vocab.get(&sub) {
Some(&id) => ids.push(id),
None => match self.vocab.get(UNK) {
Some(&unk) => ids.push(unk),
None => return Err(TokenizerError::UnknownToken(sub)),
},
}
Ok(())
}
fn encode_text(&self, text: &str, ids: &mut Vec<TokenId>) -> Result<(), TokenizerError> {
if self.byte_level {
for piece in gpt2_split(text) {
let encoded = encode_byte_level(&piece);
for sub in self.tokenize_word(&encoded) {
self.push_symbol(sub, ids)?;
}
}
} else {
for word in split_words(text) {
for sub in self.tokenize_word(word) {
self.push_symbol(sub, ids)?;
}
}
}
Ok(())
}
fn split_specials<'a>(&self, text: &'a str) -> Vec<Segment<'a>> {
if self.special_tokens.is_empty() {
return vec![Segment::Text(text)];
}
let mut specials: Vec<(&String, &TokenId)> = self.special_tokens.iter().collect();
specials.sort_by_key(|s| core::cmp::Reverse(s.0.len()));
let mut segments = Vec::new();
let mut cursor = 0usize;
let mut run_start = 0usize;
while cursor < text.len() {
let mut matched = None;
for (tok, &id) in &specials {
if text[cursor..].starts_with(tok.as_str()) {
matched = Some((tok.len(), id));
break;
}
}
if let Some((len, id)) = matched {
if run_start < cursor {
segments.push(Segment::Text(&text[run_start..cursor]));
}
segments.push(Segment::Special(id));
cursor += len;
run_start = cursor;
} else {
cursor += text[cursor..].chars().next().map_or(1, char::len_utf8);
}
}
if run_start < text.len() {
segments.push(Segment::Text(&text[run_start..]));
}
segments
}
#[cfg(feature = "std")]
pub fn from_files(vocab_path: &str, merges_path: &str) -> Result<Self, TokenizerError> {
let vocab_text = std::fs::read_to_string(vocab_path)?;
let merges_text = std::fs::read_to_string(merges_path)?;
let vocab = crate::tokenizer::parse_vocab_lines(&vocab_text);
let mut merges = Vec::new();
for line in merges_text.lines() {
if line.starts_with("#version") || line.trim().is_empty() {
continue;
}
let parts: Vec<&str> = line.split_whitespace().collect();
if parts.len() != 2 {
continue; }
merges.push((parts[0].to_string(), parts[1].to_string()));
}
Ok(Self::from_vocab_merges(vocab, merges))
}
}
enum Segment<'a> {
Text(&'a str),
Special(TokenId),
}
impl Tokenizer for BpeTokenizer {
fn encode(&self, text: &str) -> Result<Vec<TokenId>, TokenizerError> {
#[cfg(feature = "normalization")]
let normalized;
#[cfg(feature = "normalization")]
let text = if let Some(form) = self.normalize {
normalized = crate::normalize::normalize(text, form);
normalized.as_str()
} else {
text
};
let mut ids = Vec::new();
for segment in self.split_specials(text) {
match segment {
Segment::Special(id) => ids.push(id),
Segment::Text(run) => self.encode_text(run, &mut ids)?,
}
}
Ok(ids)
}
fn decode(&self, ids: &[TokenId]) -> Result<String, TokenizerError> {
let mut out = String::new();
let mut buf = String::new();
let flush = |buf: &mut String, out: &mut String| -> Result<(), TokenizerError> {
if buf.is_empty() {
return Ok(());
}
if self.byte_level {
let decoded = decode_byte_level(buf).ok_or_else(|| {
TokenizerError::MalformedFile("invalid byte-level sequence".to_string())
})?;
out.push_str(&decoded);
} else {
out.push_str(buf);
}
buf.clear();
Ok(())
};
for &id in ids {
if let Some(special) = self.special_id_to_token.get(&id) {
flush(&mut buf, &mut out)?;
out.push_str(special);
continue;
}
match self.id_to_token.get(&id) {
Some(token) => buf.push_str(token),
None => return Err(TokenizerError::UnknownToken(id.to_string())),
}
}
flush(&mut buf, &mut out)?;
Ok(out)
}
}