use daggrs::{DoubleArrayAhoCorasick, MatchKind, Trie};
use foldhash::HashMap as FoldHashMap;
use chunk::chunk;
use smallvec::SmallVec;
use std::collections::VecDeque;
use std::thread;
use crate::types::{Split, TokenId};
const PARALLEL_THRESHOLD: usize = 10_000;
const MAX_CACHED_TOKEN_LEN: usize = 16;
const ENCODE_ITER_BUFFER_SIZE: usize = 8;
#[inline(always)]
fn pack_pair(left: TokenId, right: TokenId) -> u64 {
((left as u64) << 32) | (right as u64)
}
const RANK_MERGE_MAX_LEN: usize = 32;
const DENSE_PAIR_BOUND: u32 = 512;
const PAIR_EMPTY_KEY: u64 = u64::MAX;
#[derive(Clone)]
struct RankPairTable {
entries: Box<[PairEntry]>,
mask: usize,
dense: Box<[u32]>,
byte_to_base: [TokenId; 256],
}
#[derive(Clone, Copy)]
#[repr(C)]
struct PairEntry {
key: u64,
merged: TokenId,
_pad: u32,
}
impl RankPairTable {
fn build(pair_lookup: &FoldHashMap<u64, TokenId>, token_bytes: &[Vec<u8>]) -> Option<Self> {
if pair_lookup.is_empty() {
return None;
}
let mut byte_to_base = [u32::MAX; 256];
for (id, bytes) in token_bytes.iter().enumerate() {
if bytes.len() == 1 && byte_to_base[bytes[0] as usize] == u32::MAX {
byte_to_base[bytes[0] as usize] = id as TokenId;
}
}
if byte_to_base.iter().any(|&id| id == u32::MAX) {
return None; }
for (&key, &merged) in pair_lookup.iter() {
let left = (key >> 32) as u32;
let right = key as u32;
if merged <= left || merged <= right {
return None;
}
}
let slots = (pair_lookup.len() * 2).next_power_of_two();
let mut entries = vec![
PairEntry { key: PAIR_EMPTY_KEY, merged: 0, _pad: 0 };
slots
]
.into_boxed_slice();
let mask = slots - 1;
let dense_len = (DENSE_PAIR_BOUND * DENSE_PAIR_BOUND) as usize;
let mut dense = vec![u32::MAX; dense_len].into_boxed_slice();
for (&key, &merged) in pair_lookup.iter() {
let left = (key >> 32) as u32;
let right = key as u32;
if left < DENSE_PAIR_BOUND && right < DENSE_PAIR_BOUND {
let idx = ((left << 9) | right) as usize;
if merged < dense[idx] {
dense[idx] = merged;
}
}
let mut i = Self::home_slot(key, mask);
loop {
let e = &mut entries[i];
if e.key == PAIR_EMPTY_KEY {
e.key = key;
e.merged = merged;
break;
}
if e.key == key {
if merged < e.merged {
e.merged = merged;
}
break;
}
i = (i + 1) & mask;
}
}
Some(Self { entries, mask, dense, byte_to_base })
}
#[inline(always)]
fn home_slot(key: u64, mask: usize) -> usize {
let h = key.wrapping_mul(0x9E37_79B9_7F4A_7C15);
((h >> 32) as usize) & mask
}
#[inline(always)]
fn merged_id(&self, left: TokenId, right: TokenId) -> u32 {
if left < DENSE_PAIR_BOUND && right < DENSE_PAIR_BOUND {
return self.dense[((left << 9) | right) as usize];
}
self.merged_id_flat(left, right)
}
#[inline(always)]
fn merged_id_flat(&self, left: TokenId, right: TokenId) -> u32 {
let key = pack_pair(left, right);
let mut i = Self::home_slot(key, self.mask);
loop {
let e = self.entries[i];
if e.key == key {
return e.merged;
}
if e.key == PAIR_EMPTY_KEY {
return u32::MAX;
}
i = (i + 1) & self.mask;
}
}
}
#[inline]
fn split_at_boundaries(text: &[u8]) -> Vec<&[u8]> {
let num_cpus = thread::available_parallelism()
.map(|p| p.get())
.unwrap_or(1);
let target_size = text.len() / num_cpus;
chunk(text)
.size(target_size)
.delimiters(b" \n")
.prefix()
.collect()
}
pub struct EncodeIter<'a> {
encoder: &'a BacktrackingBytePairEncoder,
text: &'a [u8],
pos: usize,
buffer: VecDeque<TokenId>,
bitfield: Bitfield,
next_token: Option<TokenId>,
done: bool,
}
impl<'a> EncodeIter<'a> {
pub(crate) fn new(encoder: &'a BacktrackingBytePairEncoder, text: &'a [u8]) -> Self {
let n = text.len();
let next_token = if text.is_empty() {
None
} else {
encoder.next_match(text)
};
Self {
encoder,
text,
pos: 0,
buffer: VecDeque::with_capacity(ENCODE_ITER_BUFFER_SIZE + 1),
bitfield: Bitfield::new(n + 1),
next_token,
done: text.is_empty(),
}
}
fn encode_one_token(&mut self) -> bool {
let Some(mut token) = self.next_token else {
return false;
};
let last = self.buffer.back().copied();
loop {
let token_len = self.encoder.token_len(token);
let end_pos = self.pos + token_len;
let is_reachable = self.bitfield.is_set(end_pos);
let is_compatible = last
.map(|last_token| self.encoder.is_valid_pair(last_token, token))
.unwrap_or(true);
if is_reachable && is_compatible {
self.buffer.push_back(token);
self.pos = end_pos;
self.next_token = self.encoder.next_match(&self.text[self.pos..]);
return true;
} else if let Some(shorter) = self.encoder.next_prefix(token) {
token = shorter;
} else {
self.bitfield.clear(self.pos);
if let Some(last_token) = self.buffer.pop_back() {
self.pos -= self.encoder.token_len(last_token);
self.next_token = Some(last_token);
return false;
} else {
self.next_token = None;
return false;
}
}
}
}
}
impl Iterator for EncodeIter<'_> {
type Item = TokenId;
fn next(&mut self) -> Option<TokenId> {
if self.done {
return self.buffer.pop_front();
}
while self.buffer.len() < ENCODE_ITER_BUFFER_SIZE {
if !self.encode_one_token() {
if self.next_token.is_none() {
self.done = true;
break;
}
}
}
self.buffer.pop_front()
}
}
impl std::iter::FusedIterator for EncodeIter<'_> {}
#[derive(Clone)]
pub struct BacktrackingBytePairEncoder {
split_table: Vec<Split>,
pair_lookup: FoldHashMap<u64, TokenId>,
token_lengths: Vec<u8>,
num_base_tokens: usize,
matcher: DoubleArrayAhoCorasick,
next_prefix_match: Vec<TokenId>,
token_cache: FoldHashMap<Vec<u8>, TokenId>,
rank_table: Option<RankPairTable>,
}
impl BacktrackingBytePairEncoder {
pub fn from_merges(
merges: &[(TokenId, TokenId)],
base_tokens: &[Vec<u8>],
) -> (Self, Vec<Vec<u8>>) {
Self::from_merges_with_added(merges, base_tokens, &[])
}
pub fn from_vocab_and_merges(
vocab: &[(u32, Vec<u8>)],
merges: &[(TokenId, TokenId)],
num_base_tokens: usize,
) -> (Self, Vec<Vec<u8>>) {
let token_bytes: Vec<Vec<u8>> = vocab.iter().map(|(_, bytes)| bytes.clone()).collect();
let bytes_to_id: FoldHashMap<Vec<u8>, TokenId> = vocab
.iter()
.map(|(id, bytes)| (bytes.clone(), *id))
.collect();
let mut pair_lookup = FoldHashMap::default();
let mut merge_creates: FoldHashMap<TokenId, (TokenId, TokenId)> = FoldHashMap::default();
for &(left, right) in merges.iter() {
let mut merged_bytes = token_bytes[left as usize].clone();
merged_bytes.extend_from_slice(&token_bytes[right as usize]);
if let Some(&merged_id) = bytes_to_id.get(&merged_bytes) {
pair_lookup.insert(pack_pair(left, right), merged_id);
merge_creates.entry(merged_id).or_insert((left, right));
}
}
let mut split_table: Vec<Split> = Vec::with_capacity(vocab.len());
for (id, _) in vocab.iter() {
let id = *id as TokenId;
if let Some(&(left, right)) = merge_creates.get(&id) {
split_table.push(Split::merge(left, right));
} else {
split_table.push(Split::base(id));
}
}
let (matcher, next_prefix_match) = Self::build_matcher_and_prefixes(&token_bytes);
let token_lengths = Self::build_token_lengths(&token_bytes);
let mut token_cache = FoldHashMap::default();
for (token_id, bytes) in token_bytes.iter().enumerate() {
if bytes.len() <= MAX_CACHED_TOKEN_LEN {
token_cache.insert(bytes.clone(), token_id as TokenId);
}
}
let rank_table = RankPairTable::build(&pair_lookup, &token_bytes);
let encoder = Self {
split_table,
pair_lookup,
token_lengths,
num_base_tokens,
matcher,
next_prefix_match,
token_cache,
rank_table,
};
(encoder, token_bytes)
}
pub fn from_merges_with_added(
merges: &[(TokenId, TokenId)],
base_tokens: &[Vec<u8>],
added_tokens: &[(u32, Vec<u8>)],
) -> (Self, Vec<Vec<u8>>) {
let num_base_tokens = base_tokens.len();
let mut split_table: Vec<Split> = (0..num_base_tokens as TokenId)
.map(Split::base)
.collect();
let mut token_bytes: Vec<Vec<u8>> = base_tokens.to_vec();
let mut pair_lookup = FoldHashMap::default();
let mut added_sorted: Vec<_> = added_tokens.to_vec();
added_sorted.sort_by_key(|(id, _)| *id);
let mut added_iter = added_sorted.into_iter().peekable();
for &(left, right) in merges.iter() {
let next_id = split_table.len() as TokenId;
while let Some(&(added_id, _)) = added_iter.peek() {
if added_id <= next_id {
let (_, bytes) = added_iter.next().unwrap();
split_table.push(Split::base(split_table.len() as TokenId));
token_bytes.push(bytes);
} else {
break;
}
}
let new_id = split_table.len() as TokenId;
split_table.push(Split::merge(left, right));
pair_lookup.insert(pack_pair(left, right), new_id);
let mut bytes = token_bytes[left as usize].clone();
bytes.extend_from_slice(&token_bytes[right as usize]);
token_bytes.push(bytes);
}
for (_, bytes) in added_iter {
split_table.push(Split::base(split_table.len() as TokenId));
token_bytes.push(bytes);
}
let (matcher, next_prefix_match) = Self::build_matcher_and_prefixes(&token_bytes);
let token_lengths = Self::build_token_lengths(&token_bytes);
let mut token_cache = FoldHashMap::default();
for (token_id, bytes) in token_bytes.iter().enumerate() {
if bytes.len() <= MAX_CACHED_TOKEN_LEN {
token_cache.insert(bytes.clone(), token_id as TokenId);
}
}
let rank_table = RankPairTable::build(&pair_lookup, &token_bytes);
let encoder = Self {
split_table,
pair_lookup,
token_lengths,
num_base_tokens,
matcher,
next_prefix_match,
token_cache,
rank_table,
};
(encoder, token_bytes)
}
pub fn from_parts(
split_table: Vec<Split>,
pair_lookup: FoldHashMap<u64, TokenId>,
token_lengths: Vec<u8>,
num_base_tokens: usize,
matcher: DoubleArrayAhoCorasick,
next_prefix_match: Vec<TokenId>,
token_bytes: &[Vec<u8>],
) -> Self {
let mut token_cache = FoldHashMap::default();
for (token_id, bytes) in token_bytes.iter().enumerate() {
if bytes.len() <= MAX_CACHED_TOKEN_LEN {
token_cache.insert(bytes.clone(), token_id as TokenId);
}
}
let rank_table = RankPairTable::build(&pair_lookup, token_bytes);
Self {
split_table,
pair_lookup,
token_lengths,
num_base_tokens,
matcher,
next_prefix_match,
token_cache,
rank_table,
}
}
fn build_matcher_and_prefixes(token_bytes: &[Vec<u8>]) -> (DoubleArrayAhoCorasick, Vec<TokenId>) {
let mut trie = Trie::new();
for (id, bytes) in token_bytes.iter().enumerate() {
trie.add(bytes, id as TokenId);
}
trie.build(MatchKind::LeftmostLongest);
let matcher = trie.compile();
let next_prefix_match: Vec<TokenId> = token_bytes
.iter()
.map(|token| {
if token.len() <= 1 {
u32::MAX
} else {
let prefix = &token[..token.len() - 1];
matcher
.find_iter(prefix)
.next()
.map(|m| m.pattern_id)
.unwrap_or(u32::MAX)
}
})
.collect();
(matcher, next_prefix_match)
}
fn build_token_lengths(token_bytes: &[Vec<u8>]) -> Vec<u8> {
token_bytes
.iter()
.map(|t| t.len().min(255) as u8)
.collect()
}
pub fn split_table(&self) -> &[Split] {
&self.split_table
}
pub fn matcher(&self) -> &DoubleArrayAhoCorasick {
&self.matcher
}
pub fn next_prefix_match_table(&self) -> &[TokenId] {
&self.next_prefix_match
}
#[inline]
pub fn is_valid_pair(&self, mut token1: TokenId, mut token2: TokenId) -> bool {
let mut limit = u32::MAX;
loop {
if let Some(&combined) = self.pair_lookup.get(&pack_pair(token1, token2)) {
if combined < limit {
return false;
}
}
if token1 > token2 {
limit = token1;
let right = self.split_table[token1 as usize].right;
if right == token1 {
limit = token2 + 1;
let left = self.split_table[token2 as usize].left;
if left + 1 == limit {
return true;
}
token2 = left;
} else {
token1 = right;
}
} else {
limit = token2 + 1;
let left = self.split_table[token2 as usize].left;
if left + 1 == limit {
limit = token1;
let right = self.split_table[token1 as usize].right;
if right == limit {
return true;
}
token1 = right;
} else {
token2 = left;
}
}
}
}
#[inline]
pub fn token_len(&self, token: TokenId) -> usize {
self.token_lengths[token as usize] as usize
}
pub fn vocab_size(&self) -> usize {
self.token_lengths.len()
}
pub fn num_base_tokens(&self) -> usize {
self.num_base_tokens
}
#[inline]
pub fn encode_into(&self, text: &[u8], cache: Option<&mut PretokenCache>, out: &mut Vec<TokenId>) {
self.encode_piece_into(text, text, cache, out)
}
#[inline]
pub fn encode_piece_into(&self, doc: &[u8], piece: &[u8], cache: Option<&mut PretokenCache>, out: &mut Vec<TokenId>) {
if piece.is_empty() {
return;
}
if let Some(c) = cache {
if piece.len() <= CACHE_KEY_MAX {
let (lo, hi) = key_words_within(doc, piece);
if c.get_with_key(lo, hi, out) {
return;
}
return self.encode_cache_miss(piece, lo, hi, c, out);
}
}
self.encode_uncached(piece, out)
}
#[inline(never)]
fn encode_cache_miss(&self, text: &[u8], lo: u64, hi: u64, cache: &mut PretokenCache, out: &mut Vec<TokenId>) {
debug_assert!(text.len() <= CACHE_KEY_MAX);
if let Some(&token_id) = self.token_cache.get(text) {
cache.insert_with_key(lo, hi, &[token_id]);
out.push(token_id);
return;
}
let start = out.len();
if self.rank_table.is_some() {
self.encode_rank_merge(text, out);
} else {
self.encode_sequential_into(text, out);
}
let toks = &out[start..];
if !toks.is_empty() && toks.len() <= CACHE_MAX_TOKENS {
cache.insert_with_key(lo, hi, toks);
}
}
#[inline(never)]
fn encode_uncached(&self, text: &[u8], out: &mut Vec<TokenId>) {
if text.len() <= MAX_CACHED_TOKEN_LEN {
if let Some(&token_id) = self.token_cache.get(text) {
out.push(token_id);
return;
}
}
if text.len() >= PARALLEL_THRESHOLD {
out.extend(self.encode(text));
return;
}
if text.len() <= RANK_MERGE_MAX_LEN && self.rank_table.is_some() {
return self.encode_rank_merge(text, out);
}
self.encode_sequential_into(text, out);
}
pub fn has_rank_merge(&self) -> bool {
self.rank_table.is_some()
}
#[doc(hidden)]
pub fn encode_rank_merge(&self, text: &[u8], out: &mut Vec<TokenId>) {
self.encode_rank_merge_impl::<true>(text, out)
}
#[doc(hidden)]
pub fn encode_rank_merge_flat(&self, text: &[u8], out: &mut Vec<TokenId>) {
self.encode_rank_merge_impl::<false>(text, out)
}
#[inline(always)]
fn encode_rank_merge_impl<const DENSE: bool>(&self, text: &[u8], out: &mut Vec<TokenId>) {
let table = self.rank_table.as_ref().expect("rank table unavailable");
let mut toks: SmallVec<[TokenId; RANK_MERGE_MAX_LEN]> = text
.iter()
.map(|&b| table.byte_to_base[b as usize])
.collect();
while toks.len() > 1 {
let mut best_rank = u32::MAX;
let mut best_i = usize::MAX;
for i in 0..toks.len() - 1 {
let m = if DENSE {
table.merged_id(toks[i], toks[i + 1])
} else {
table.merged_id_flat(toks[i], toks[i + 1])
};
if m < best_rank {
best_rank = m;
best_i = i;
}
}
if best_i == usize::MAX {
break;
}
let left = toks[best_i];
let right = toks[best_i + 1];
let mut w = best_i;
let mut i = best_i;
let n = toks.len();
while i < n {
if i + 1 < n && toks[i] == left && toks[i + 1] == right {
toks[w] = best_rank;
i += 2;
} else {
toks[w] = toks[i];
i += 1;
}
w += 1;
}
toks.truncate(w);
}
out.extend_from_slice(&toks);
}
pub fn encode(&self, text: &[u8]) -> Vec<TokenId> {
if text.is_empty() {
return Vec::new();
}
if text.len() <= MAX_CACHED_TOKEN_LEN {
if let Some(&token_id) = self.token_cache.get(text) {
return vec![token_id];
}
}
if text.len() < PARALLEL_THRESHOLD {
return self.encode_sequential(text);
}
let chunks = split_at_boundaries(text);
if chunks.len() == 1 {
return self.encode_sequential(chunks[0]);
}
let results: Vec<Vec<TokenId>> = thread::scope(|s| {
let handles: Vec<_> = chunks
.iter()
.map(|chunk| s.spawn(|| self.encode_sequential(chunk)))
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
let total: usize = results.iter().map(|v| v.len()).sum();
let mut output = Vec::with_capacity(total);
for chunk in results {
output.extend(chunk);
}
output
}
pub fn encode_iter<'a>(&'a self, text: &'a [u8]) -> EncodeIter<'a> {
EncodeIter::new(self, text)
}
pub fn encode_batch(&self, texts: &[&[u8]]) -> Vec<Vec<TokenId>> {
if texts.is_empty() {
return Vec::new();
}
let num_cpus = thread::available_parallelism()
.map(|p| p.get())
.unwrap_or(1);
if texts.len() <= num_cpus || num_cpus == 1 {
if num_cpus == 1 {
return texts.iter().map(|t| self.encode_sequential(t)).collect();
}
return thread::scope(|s| {
let handles: Vec<_> = texts
.iter()
.map(|text| s.spawn(|| self.encode_sequential(text)))
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
}
let chunk_size = (texts.len() + num_cpus - 1) / num_cpus;
thread::scope(|s| {
let handles: Vec<_> = texts
.chunks(chunk_size)
.map(|chunk| {
s.spawn(|| {
chunk
.iter()
.map(|t| self.encode_sequential(t))
.collect::<Vec<_>>()
})
})
.collect();
handles
.into_iter()
.flat_map(|h| h.join().unwrap())
.collect()
})
}
fn encode_sequential(&self, text: &[u8]) -> Vec<TokenId> {
if text.is_empty() {
return Vec::new();
}
if text.len() <= MAX_CACHED_TOKEN_LEN {
if let Some(&token_id) = self.token_cache.get(text) {
return vec![token_id];
}
}
let mut out = Vec::new();
self.encode_sequential_into(text, &mut out);
out
}
#[doc(hidden)]
pub fn encode_sequential_into(&self, text: &[u8], out: &mut Vec<TokenId>) {
let n = text.len();
let mut tokens: SmallVec<[TokenId; 16]> = SmallVec::new();
let mut bitfield = Bitfield::new(n + 1);
let mut pos = 0;
let mut next_token = self.next_match(&text[pos..]);
while let Some(mut token) = next_token {
let last = tokens.last().copied();
loop {
let token_len = self.token_len(token);
let end_pos = pos + token_len;
let is_reachable = bitfield.is_set(end_pos);
let is_compatible = last
.map(|last_token| self.is_valid_pair(last_token, token))
.unwrap_or(true);
if is_reachable && is_compatible {
tokens.push(token);
pos = end_pos;
next_token = self.next_match(&text[pos..]);
break;
} else if let Some(shorter) = self.next_prefix(token) {
token = shorter;
} else {
bitfield.clear(pos);
if let Some(last_token) = tokens.pop() {
pos -= self.token_len(last_token);
}
next_token = last;
break;
}
}
}
out.extend_from_slice(&tokens);
}
#[doc(hidden)]
#[inline]
pub fn token_cache_get(&self, text: &[u8]) -> Option<TokenId> {
self.token_cache.get(text).copied()
}
#[doc(hidden)]
#[inline]
pub fn encode_backtrack_into(&self, text: &[u8], out: &mut Vec<TokenId>) {
self.encode_sequential_into(text, out);
}
#[inline]
fn next_match(&self, text: &[u8]) -> Option<TokenId> {
self.matcher.find_iter(text).next().map(|m| m.pattern_id)
}
#[inline]
fn next_prefix(&self, token: TokenId) -> Option<TokenId> {
let prefix = self.next_prefix_match[token as usize];
if prefix == u32::MAX {
None
} else {
Some(prefix)
}
}
}
pub struct PretokenCache {
entries: Box<[CacheEntry]>,
mask: usize,
}
const CACHE_KEY_MAX: usize = 15;
const CACHE_BITS_DEFAULT: usize = 16; const CACHE_PROBES: usize = 4;
const CACHE_MAX_TOKENS: usize = 3;
fn cache_bits() -> usize {
static BITS: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*BITS.get_or_init(|| {
std::env::var("TOKIE_CACHE_BITS").ok()
.and_then(|v| v.parse().ok())
.filter(|&b| (10..=24).contains(&b))
.unwrap_or(CACHE_BITS_DEFAULT)
})
}
#[derive(Clone, Copy)]
#[repr(C)]
struct CacheEntry {
key_lo: u64,
key_hi: u64,
toks: [TokenId; CACHE_MAX_TOKENS],
ntok: u32,
}
#[inline(always)]
fn load_le_partial(p: &[u8]) -> u64 {
let len = p.len();
debug_assert!((1..=7).contains(&len));
if len >= 4 {
let a = u32::from_le_bytes(p[..4].try_into().unwrap()) as u64;
let b = u32::from_le_bytes(p[len - 4..].try_into().unwrap()) as u64;
a | (b << ((len - 4) * 8))
} else {
let a = p[0] as u64;
let b = (p[len / 2] as u64) << ((len / 2) * 8);
let c = (p[len - 1] as u64) << ((len - 1) * 8);
a | b | c
}
}
#[inline(always)]
fn key_words_within(doc: &[u8], piece: &[u8]) -> (u64, u64) {
let len = piece.len();
debug_assert!((1..=CACHE_KEY_MAX).contains(&len));
let start = (piece.as_ptr() as usize).wrapping_sub(doc.as_ptr() as usize);
if start <= doc.len() && doc.len() - start >= 16 {
let raw = u128::from_le_bytes(doc[start..start + 16].try_into().unwrap());
let masked = raw & (u128::MAX >> (128 - 8 * len));
(masked as u64, (masked >> 64) as u64 | ((len as u64) << 56))
} else {
key_words(piece)
}
}
#[inline(always)]
fn key_words(bytes: &[u8]) -> (u64, u64) {
let len = bytes.len();
debug_assert!((1..=CACHE_KEY_MAX).contains(&len));
let (lo, mut hi) = if len >= 8 {
let lo = u64::from_le_bytes(bytes[..8].try_into().unwrap());
let hi = if len > 8 { load_le_partial(&bytes[8..]) } else { 0 };
(lo, hi)
} else {
(load_le_partial(bytes), 0)
};
hi |= (len as u64) << 56;
(lo, hi)
}
impl PretokenCache {
pub const KEY_MAX: usize = CACHE_KEY_MAX;
pub const MAX_TOKENS: usize = CACHE_MAX_TOKENS;
pub fn new() -> Self {
let empty = CacheEntry { key_lo: 0, key_hi: 0, toks: [0; CACHE_MAX_TOKENS], ntok: 0 };
let n = 1usize << cache_bits();
let entries = vec![empty; n].into_boxed_slice();
#[cfg(target_os = "linux")]
{
const MADV_HUGEPAGE: i32 = 14;
unsafe extern "C" {
fn madvise(addr: *mut core::ffi::c_void, length: usize, advice: i32) -> i32;
}
unsafe {
madvise(
entries.as_ptr() as *mut core::ffi::c_void,
n * std::mem::size_of::<CacheEntry>(),
MADV_HUGEPAGE,
);
}
}
Self { entries, mask: n - 1 }
}
pub fn clear(&mut self) {
let empty = CacheEntry { key_lo: 0, key_hi: 0, toks: [0; CACHE_MAX_TOKENS], ntok: 0 };
self.entries.fill(empty);
}
#[inline(always)]
fn slot(&self, lo: u64, hi: u64) -> usize {
let h = (lo ^ 0x9E37_79B9_7F4A_7C15)
.wrapping_mul(0xA076_1D64_78BD_642F)
^ hi.wrapping_mul(0xE703_7ED1_A0B4_28DB);
((h ^ (h >> 32)) as usize) & self.mask
}
#[inline(always)]
pub fn get(&self, bytes: &[u8], out: &mut Vec<TokenId>) -> bool {
let (lo, hi) = key_words(bytes);
self.get_with_key(lo, hi, out)
}
#[doc(hidden)]
#[inline(always)]
pub fn key_of(bytes: &[u8]) -> (u64, u64) {
key_words(bytes)
}
#[doc(hidden)]
#[inline(always)]
pub fn get_with_key(&self, lo: u64, hi: u64, out: &mut Vec<TokenId>) -> bool {
let mut i = self.slot(lo, hi);
for _ in 0..CACHE_PROBES {
let e = &self.entries[i];
if e.key_lo == lo && e.key_hi == hi {
out.reserve(CACHE_MAX_TOKENS);
unsafe {
let len = out.len();
std::ptr::copy_nonoverlapping(e.toks.as_ptr(), out.as_mut_ptr().add(len), CACHE_MAX_TOKENS);
out.set_len(len + e.ntok as usize);
}
return true;
}
if e.key_hi == 0 {
return false;
}
i = (i + 1) & self.mask;
}
false
}
#[inline]
pub fn insert(&mut self, bytes: &[u8], toks: &[TokenId]) {
if bytes.is_empty() || bytes.len() > CACHE_KEY_MAX || toks.is_empty() || toks.len() > CACHE_MAX_TOKENS {
return;
}
let (lo, hi) = key_words(bytes);
self.insert_with_key(lo, hi, toks);
}
#[inline]
pub fn insert_with_key(&mut self, lo: u64, hi: u64, toks: &[TokenId]) {
debug_assert!(!toks.is_empty() && toks.len() <= CACHE_MAX_TOKENS);
let home = self.slot(lo, hi);
let mut i = home;
let mut target = home;
for _ in 0..CACHE_PROBES {
let e = &self.entries[i];
if e.key_hi == 0 || (e.key_lo == lo && e.key_hi == hi) {
target = i;
break;
}
i = (i + 1) & self.mask;
}
let e = &mut self.entries[target];
e.key_lo = lo;
e.key_hi = hi;
e.ntok = toks.len() as u32;
e.toks[..toks.len()].copy_from_slice(toks);
}
}
impl Default for PretokenCache {
fn default() -> Self {
Self::new()
}
}
struct Bitfield {
bits: SmallVec<[u64; 4]>,
}
impl Bitfield {
fn new(size: usize) -> Self {
let num_words = (size + 63) / 64;
let mut bits = SmallVec::new();
bits.resize(num_words, u64::MAX);
Self { bits }
}
#[inline]
fn clear(&mut self, pos: usize) {
let word = pos / 64;
let bit = pos % 64;
self.bits[word] &= !(1 << bit);
}
#[inline]
fn is_set(&self, pos: usize) -> bool {
let word = pos / 64;
let bit = pos % 64;
(self.bits[word] >> bit) & 1 != 0
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::decoder::VocabDecoder;
#[test]
fn test_encode_into_with_cache_matches_encode() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1), (3, 2)]; let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
let pieces: Vec<&[u8]> = vec![
b"abc", b"ab", b"ba", b"cab", b"abcabcabc", b"a", b"",
b"cccab", b"abcabcabcabcabca", ];
let mut cache = PretokenCache::new();
for pass in 0..2 {
for &p in &pieces {
let expect = encoder.encode(p);
let mut got = Vec::new();
encoder.encode_into(p, Some(&mut cache), &mut got);
assert_eq!(got, expect, "pass {pass}, piece {:?}", p);
}
}
for &p in &pieces {
let mut got = Vec::new();
encoder.encode_into(p, None, &mut got);
assert_eq!(got, encoder.encode(p), "no-cache piece {:?}", p);
}
}
#[test]
fn test_key_words_within_matches_standalone() {
let doc: Vec<u8> = (0..64u8).map(|i| i.wrapping_mul(37).wrapping_add(11)).collect();
for start in 0..doc.len() {
for len in 1..=CACHE_KEY_MAX {
if start + len > doc.len() {
break;
}
let piece = &doc[start..start + len];
assert_eq!(
key_words_within(&doc, piece),
key_words(piece),
"start {start} len {len}"
);
}
}
let outside = vec![0xABu8; 7];
assert_eq!(key_words_within(&doc, &outside), key_words(&outside));
let doc2 = [0xFFu8; 32];
for len in 1..=CACHE_KEY_MAX {
assert_eq!(key_words_within(&doc2, &doc2[3..3 + len]), key_words(&doc2[3..3 + len]));
}
}
#[test]
fn test_encode_piece_into_doc_context_matches_encode() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1), (3, 2)]; let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
let doc: Vec<u8> = b"abcabbacababcabcabcabcabcababcacab".to_vec();
let ranges: Vec<(usize, usize)> = vec![
(0, 3), (3, 5), (5, 8), (8, 11), (11, 14), (14, 17), (11, 14), (5, 25), (doc.len() - 2, doc.len()), (doc.len() - 1, doc.len()), (7, 7), ];
let mut cache = PretokenCache::new();
for pass in 0..2 {
for &(s, e) in &ranges {
let piece = &doc[s..e];
let expect = encoder.encode(piece);
let mut got = Vec::new();
encoder.encode_piece_into(&doc, piece, Some(&mut cache), &mut got);
assert_eq!(got, expect, "pass {pass}, range {s}..{e}");
let mut got2 = Vec::new();
encoder.encode_piece_into(&doc, piece, None, &mut got2);
assert_eq!(got2, expect, "no-cache, range {s}..{e}");
}
}
}
#[test]
fn test_from_merges() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1), (3, 2)];
let (encoder, token_bytes) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
let decoder = VocabDecoder::new(token_bytes);
assert_eq!(encoder.vocab_size(), 5);
assert_eq!(encoder.num_base_tokens(), 3);
assert_eq!(decoder.token_to_bytes(0), b"a");
assert_eq!(decoder.token_to_bytes(3), b"ab");
assert_eq!(decoder.token_to_bytes(4), b"abc");
}
#[test]
fn test_is_valid_pair() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1)];
let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
assert!(!encoder.is_valid_pair(0, 1));
assert!(encoder.is_valid_pair(3, 2));
assert!(encoder.is_valid_pair(1, 2));
}
#[test]
fn test_encode_merged_token() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1)];
let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
assert_eq!(encoder.encode(b"ab"), vec![3]);
assert_eq!(encoder.encode(b"abc"), vec![3, 2]);
}
#[test]
fn test_early_exit() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1), (3, 2)];
let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
assert_eq!(encoder.encode(b"a"), vec![0]);
assert_eq!(encoder.encode(b"ab"), vec![3]);
assert_eq!(encoder.encode(b"abc"), vec![4]);
}
#[test]
fn test_encode_decode_roundtrip() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c'], vec![b'd']];
let merges = vec![(0, 1), (2, 3), (4, 5)];
let (encoder, token_bytes) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
let decoder = VocabDecoder::new(token_bytes);
for text in [b"abcd".as_slice(), b"ab", b"cd", b"abcdabcd", b"a", b""] {
let encoded = encoder.encode(text);
let decoded = decoder.decode(&encoded);
assert_eq!(decoded, text);
}
}
fn byte_complete_encoder() -> BacktrackingBytePairEncoder {
let base_tokens: Vec<Vec<u8>> = (0u16..256).map(|b| vec![b as u8]).collect();
let a = b'a' as TokenId;
let merges = vec![
(a, a + 1), (256, a + 2), (b'l' as u32, b'l' as u32), (b'h' as u32, b'e' as u32), (258, b'o' as u32), (259, 260), (b' ' as u32, b't' as u32), (262, 259), ];
let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
encoder
}
#[test]
fn test_rank_merge_matches_backtracking() {
let encoder = byte_complete_encoder();
assert!(encoder.has_rank_merge());
let cases: Vec<&[u8]> = vec![
b"abc", b"ab", b"hello", b"hhello", b"llllll", b"aaabbb",
b" the", b" the the", b"abcabcabc", b"xyz", b"\x00\xff\xfe",
];
for text in cases {
let mut daac = Vec::new();
encoder.encode_sequential_into(text, &mut daac);
let mut rank = Vec::new();
encoder.encode_rank_merge(text, &mut rank);
assert_eq!(rank, daac, "piece {:?}", text);
let mut flat = Vec::new();
encoder.encode_rank_merge_flat(text, &mut flat);
assert_eq!(flat, daac, "flat probe, piece {:?}", text);
}
}
#[test]
fn test_rank_merge_fuzz_matches_backtracking() {
let encoder = byte_complete_encoder();
let mut state = 0x853C_49E6_748F_EA9Bu64;
let mut next = move || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
state
};
for _ in 0..5000 {
let len = 1 + (next() as usize) % 40;
let bytes: Vec<u8> = (0..len).map(|_| next() as u8).collect();
let mut daac = Vec::new();
encoder.encode_sequential_into(&bytes, &mut daac);
let mut rank = Vec::new();
encoder.encode_rank_merge(&bytes, &mut rank);
assert_eq!(rank, daac, "fuzz bytes {:?}", bytes);
}
}
#[test]
fn test_rank_merge_disabled_for_non_byte_vocab() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c']];
let merges = vec![(0, 1), (3, 2)];
let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
assert!(!encoder.has_rank_merge());
let mut out = Vec::new();
encoder.encode_into(b"abcab", None, &mut out);
assert_eq!(out, encoder.encode(b"abcab"));
}
#[test]
fn test_encode_iter_matches_encode() {
let base_tokens = vec![vec![b'a'], vec![b'b'], vec![b'c'], vec![b'd']];
let merges = vec![(0, 1), (2, 3), (4, 5)];
let (encoder, _) = BacktrackingBytePairEncoder::from_merges(&merges, &base_tokens);
for text in [b"".as_slice(), b"a", b"ab", b"abcd", b"abcdabcdabcdabcdabcd"] {
let encoded = encoder.encode(text);
let iter_encoded: Vec<_> = encoder.encode_iter(text).collect();
assert_eq!(encoded, iter_encoded);
}
}
}