mod algorithm;
use crate::{
Method, utok,
vocab::{CollectedVocab, CompressedVocab},
};
use std::{
collections::{HashMap, HashSet},
iter::zip,
ops::Deref,
pin::Pin,
ptr::NonNull,
};
pub struct Bpe {
_vocabs: Pin<Box<[u8]>>,
tokens: Box<[TokenMeta]>,
sorted_pieces: Box<[utok]>,
bytes: Box<[utok; 256]>,
unk: utok,
}
struct TokenMeta {
ptr: NonNull<u8>,
len: u32,
rank: u32,
}
unsafe impl Send for TokenMeta {}
unsafe impl Sync for TokenMeta {}
impl Deref for TokenMeta {
type Target = [u8];
#[inline]
fn deref(&self) -> &Self::Target {
unsafe { std::slice::from_raw_parts(self.ptr.as_ptr(), self.len as _) }
}
}
impl Bpe {
pub fn from_tokenizer_model(model: &[u8]) -> Self {
let offsets = (0..)
.scan(0usize, |offset, _| match &model[*offset..] {
[10, total_len, 10, content @ ..] => {
let total_len = *total_len as usize;
*offset += total_len + 2;
Some(&content[..total_len - 2])
}
[..] => None,
})
.collect::<Vec<_>>();
let vocabs = offsets.iter().map(|slice| {
let &&[len, ref content @ ..] = slice else {
unreachable!()
};
std::str::from_utf8(&content[..len as usize]).unwrap()
});
let scores = offsets.iter().map(|slice| {
let len = slice[0] as usize;
let ptr = slice[len + 2..].as_ptr().cast::<f32>();
unsafe { ptr.read_unaligned() }
});
Self::from_collected_vocab(
CollectedVocab::collect(vocabs.into_iter().map(|s| s.as_bytes()), 0),
scores,
0,
)
}
pub fn new<'a>(
vocabs: impl IntoIterator<Item = &'a str>,
scores: impl IntoIterator<Item = f32>,
is_byte: impl IntoIterator<Item = bool>,
unk: utok,
) -> Self {
Self::from_collected_vocab(
CollectedVocab::collect_with_hint(
vocabs.into_iter().map(|s| s.as_bytes()),
is_byte,
unk,
),
scores,
unk,
)
}
fn from_collected_vocab(
vocab: CollectedVocab,
scores: impl IntoIterator<Item = f32>,
unk: utok,
) -> Self {
let CollectedVocab {
vocabs,
total_len,
bytes,
} = vocab;
let CompressedVocab { vocabs, slices } = CompressedVocab::new(&vocabs, total_len);
let scores = scores.into_iter().collect::<Vec<_>>();
assert_eq!(
slices.len(),
scores.len(),
"scores size mismatch with vocab size"
);
let tokens = zip(slices, rank(&scores))
.map(|((off, len), rank)| TokenMeta {
ptr: unsafe { NonNull::new_unchecked(vocabs[off..].as_ptr().cast_mut()) },
len: len as _,
rank,
})
.collect::<Box<_>>();
let bytes_set = bytes.iter().chain(&[unk]).cloned().collect::<HashSet<_>>();
let mut sorted_pieces = (0..tokens.len() as utok)
.filter(|i| !bytes_set.contains(i))
.collect::<Box<_>>();
sorted_pieces.sort_unstable_by_key(|&i| &*tokens[i as usize]);
Self {
_vocabs: vocabs,
tokens,
sorted_pieces,
bytes,
unk,
}
}
pub fn inaccessible(&self) -> HashMap<&str, utok> {
self.sorted_pieces
.iter()
.filter_map(|&t| {
let s = unsafe { std::str::from_utf8_unchecked(self.token(t)) };
if self.encode(s).into_iter().nth(1).is_some() {
Some((s, t))
} else {
None
}
})
.collect()
}
#[inline]
fn find_piece(&self, piece: &[u8]) -> Option<utok> {
match self
.sorted_pieces
.binary_search_by_key(&piece, |&i| self.token(i))
{
Ok(i) => Some(self.sorted_pieces[i]),
Err(_) => match *piece {
[b] => Some(self.bytes[b as usize]),
[..] => None,
},
}
}
#[inline(always)]
fn token(&self, token: utok) -> &TokenMeta {
&self.tokens[token as usize]
}
}
impl Method for Bpe {
#[inline]
fn unk_token(&self) -> utok {
self.unk
}
#[inline]
fn vocab_size(&self) -> usize {
self.tokens.len()
}
#[inline]
fn internal_special(&self) -> impl IntoIterator<Item = (&str, utok)> {
self.inaccessible()
}
#[inline]
fn encode(&self, text: &str) -> impl IntoIterator<Item = utok> + '_ {
let mut tokenizer = self.begin_merge(text);
while tokenizer.merge() {}
tokenizer.into_iter()
}
#[inline]
fn decode(&self, token: utok) -> &[u8] {
self.token(token)
}
}
fn rank(scores: &[f32]) -> impl IntoIterator<Item = u32> + '_ {
use std::{
cmp::Ordering,
collections::{BTreeMap, BTreeSet},
};
#[derive(Debug)]
struct FloatOrd(f32);
impl Eq for FloatOrd {}
impl PartialEq for FloatOrd {
#[inline]
fn eq(&self, other: &Self) -> bool {
self.0.total_cmp(&other.0).is_eq()
}
}
impl PartialOrd for FloatOrd {
#[inline]
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
impl Ord for FloatOrd {
#[inline]
fn cmp(&self, other: &Self) -> Ordering {
self.0.total_cmp(&other.0)
}
}
let map = scores
.iter()
.copied()
.map(FloatOrd)
.collect::<BTreeSet<_>>()
.into_iter()
.rev()
.enumerate()
.map(|(i, f)| (f, i as u32))
.collect::<BTreeMap<_, _>>();
scores.iter().map(move |&f| map[&FloatOrd(f)])
}
#[cfg(test)]
mod bpe_tests {
use super::*;
#[test]
fn test() {
if let Ok(buf) = std::fs::read("tokenizer.model") {
let bpe = Bpe::from_tokenizer_model(&buf);
let inaccessible = bpe.inaccessible();
println!(
"bpe: detected {} tokens, compressed to {} bytes",
bpe.vocab_size(),
bpe._vocabs.len(),
);
println!("inaccessible: {inaccessible:#?}")
}
}
fn test_bpe() -> Bpe {
Bpe::new(
[
"<unk>", "a", "b", "c", "d", "ab", "ac", "ad", "bd", "bcd",
],
[
0., 1., 1., 1., 1., 1.1, 1.2, 1.3, 1.4, 10.,
],
[false; 10],
0,
)
}
#[test]
fn test_bpe_new() {
let bpe = test_bpe();
assert_eq!(bpe.vocab_size(), 10);
}
#[test]
fn test_bpe_unk_token() {
let bpe = test_bpe();
assert_eq!(bpe.unk_token(), 0);
}
#[test]
fn test_bpe_encode() {
let bpe = test_bpe();
let encoded: Vec<_> = bpe.encode("abd").into_iter().collect();
assert_eq!(encoded, [1, 8]); }
#[test]
fn test_bpe_decode() {
let bpe = test_bpe();
assert_eq!(bpe.decode(3), b"c");
assert_eq!(bpe.decode(6), b"ac");
assert_eq!(bpe.decode(9), b"bcd");
assert_eq!(bpe.decode(0), b"<unk>");
}
#[test]
fn test_bpe_encode_decode() {
let bpe = test_bpe();
let text = "abcdx";
let encoded: Vec<_> = bpe.encode(text).into_iter().collect();
assert_eq!(encoded, [5, 3, 4, 0]);
let decoded: Vec<_> = encoded
.iter()
.flat_map(|&t| bpe.decode(t).iter().copied())
.collect();
assert_eq!(std::str::from_utf8(&decoded), Ok("abcd<unk>"))
}
#[test]
fn test_bpe_inaccessible() {
let bpe = test_bpe();
let inaccessible = bpe.inaccessible();
println!("Inaccessible tokens: {:?}", inaccessible);
assert!(
!inaccessible.contains_key("d"),
"Token 'd' should be accessible"
);
assert_eq!(
inaccessible.get("bcd"),
Some(&9),
"Token 'bcd' should be inaccessible"
);
assert!(
!inaccessible.contains_key("ab"),
"Token 'ab' should be accessible"
);
}
#[test]
fn test_bpe_with_byte_tokens() {
let vocabs = ["a", "b", "<0x41>", "<0x42>"];
let scores = [1.0, 1.0, 1.0, 1.0];
let is_byte = [false, false, true, true];
let bpe = Bpe::new(vocabs, scores, is_byte, 0);
let encoded: Vec<_> = bpe.encode("aAB").into_iter().collect();
assert_eq!(encoded, [0, 2, 3], "Expected 3 tokens for input 'aAB'")
}
}