mod algorithm;
use crate::{
functions::{collect_vocabs_with_hint, CompressedVocab},
utok, Method,
};
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 = slice[0] as usize;
std::str::from_utf8(&slice[1..][..len]).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() }
});
let mut i = 0;
let is_byte = std::iter::from_fn(|| {
if i < 3 {
i += 1;
Some(false)
} else if i < 3 + 256 {
i += 1;
Some(true)
} else {
Some(false)
}
});
Self::new(vocabs, scores, is_byte, 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 {
let (vocabs, bytes, total_len) =
collect_vocabs_with_hint(vocabs.into_iter().map(|s| s.as_bytes()), is_byte, unk);
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.build_tokenizer(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(PartialEq, Debug)]
struct FloatOrd(f32);
impl Eq for FloatOrd {}
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)])
}
#[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:#?}");
}
}