use std::cell::RefCell;
use std::sync::atomic::{AtomicU64, Ordering};
use daggrs::{DoubleArrayAhoCorasick, MatchKind, Trie};
use foldhash::HashMap as FoldHashMap;
use smallvec::SmallVec;
use crate::types::TokenId;
static NEXT_CACHE_ID: AtomicU64 = AtomicU64::new(1);
#[inline]
fn utf8_char_len(b: u8) -> usize {
match b {
0..=0x7F => 1,
0xC0..=0xDF => 2,
0xE0..=0xEF => 3,
0xF0..=0xFF => 4,
_ => 1,
}
}
const METASPACE: [u8; 3] = [0xE2, 0x96, 0x81];
const MAX_CACHED_TOKEN_LEN: usize = 16;
const FRONT_BITS: u32 = 18;
const SHORT_KEY_MAX: usize = 15;
#[inline]
fn compute_unit_split_safe(token_bytes: &[Vec<u8>]) -> bool {
!token_bytes
.iter()
.any(|bytes| memchr::memmem::find_iter(bytes, &METASPACE).any(|pos| pos > 0))
}
thread_local! {
static THREAD_UNIT_CACHE: RefCell<Option<(usize, UnigramPieceCache)>> =
const { RefCell::new(None) };
}
pub struct UnigramPieceCache {
arena: Vec<TokenId>,
front_keys: Box<[u128]>,
front_vals: Box<[(u32, u32)]>,
short: FoldHashMap<u128, (u32, u32)>,
long: FoldHashMap<Box<[u8]>, (u32, u32)>,
}
impl UnigramPieceCache {
pub fn new() -> Self {
let n = 1usize << FRONT_BITS;
Self {
arena: Vec::new(),
front_keys: vec![0u128; n].into_boxed_slice(),
front_vals: vec![(0u32, 0u32); n].into_boxed_slice(),
short: FoldHashMap::default(),
long: FoldHashMap::default(),
}
}
pub fn clear(&mut self) {
self.arena.clear();
self.front_keys.fill(0);
self.front_vals.fill((0, 0));
self.short.clear();
self.long.clear();
}
#[inline]
fn front_index(key: u128) -> usize {
let folded = (key as u64) ^ ((key >> 64) as u64);
let h = folded.wrapping_mul(0x9E37_79B9_7F4A_7C15);
(h >> (64 - FRONT_BITS)) as usize
}
#[inline]
fn pack_key(bytes: &[u8]) -> Option<u128> {
let n = bytes.len();
if n == 0 || n > SHORT_KEY_MAX {
return None;
}
let mut lanes = [0u8; 16];
lanes[..n].copy_from_slice(bytes);
Some(u128::from_le_bytes(lanes) | ((n as u128) << 120))
}
#[inline]
fn lookup(&mut self, unit: &[u8], out: &mut Vec<TokenId>) -> bool {
if let Some(key) = Self::pack_key(unit) {
let idx = Self::front_index(key);
if self.front_keys[idx] == key {
let (offset, len) = self.front_vals[idx];
let start = offset as usize;
out.extend_from_slice(&self.arena[start..start + len as usize]);
return true;
}
if let Some(&(offset, len)) = self.short.get(&key) {
self.front_keys[idx] = key;
self.front_vals[idx] = (offset, len);
let start = offset as usize;
out.extend_from_slice(&self.arena[start..start + len as usize]);
return true;
}
false
} else if let Some(&(offset, len)) = self.long.get(unit) {
let start = offset as usize;
out.extend_from_slice(&self.arena[start..start + len as usize]);
true
} else {
false
}
}
#[inline]
fn insert(&mut self, unit: &[u8], toks: &[TokenId]) {
if unit.is_empty() || toks.is_empty() {
return;
}
let offset = self.arena.len() as u32;
let len = toks.len() as u32;
self.arena.extend_from_slice(toks);
if let Some(key) = Self::pack_key(unit) {
self.short.insert(key, (offset, len));
let idx = Self::front_index(key);
self.front_keys[idx] = key;
self.front_vals[idx] = (offset, len);
} else {
self.long.insert(unit.to_vec().into_boxed_slice(), (offset, len));
}
}
}
impl Default for UnigramPieceCache {
fn default() -> Self {
Self::new()
}
}
#[derive(Clone)]
pub struct UnigramEncoder {
matcher: DoubleArrayAhoCorasick,
scores: Vec<f64>,
unk_token: TokenId,
byte_tokens: [TokenId; 256],
token_lengths: Vec<u16>,
vocab_size: usize,
token_cache: FoldHashMap<Vec<u8>, TokenId>,
has_byte_fallback: bool,
cache_id: u64,
unit_split_safe: bool,
}
impl std::fmt::Debug for UnigramEncoder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("UnigramEncoder")
.field("vocab_size", &self.vocab_size)
.field("unk_token", &self.unk_token)
.field("unit_split_safe", &self.unit_split_safe)
.finish()
}
}
impl UnigramEncoder {
pub fn from_vocab_with_scores(
vocab: &[(u32, Vec<u8>, f64)],
unk_token: TokenId,
) -> (Self, Vec<Vec<u8>>) {
let token_bytes: Vec<Vec<u8>> = vocab.iter().map(|(_, bytes, _)| bytes.clone()).collect();
let scores: Vec<f64> = vocab.iter().map(|(_, _, score)| *score).collect();
let mut byte_tokens = [u32::MAX; 256];
for (id, bytes, _) in vocab {
if bytes.len() == 6 && bytes.starts_with(b"<0x") && bytes.ends_with(b">") {
if let Ok(byte_val) = u8::from_str_radix(
std::str::from_utf8(&bytes[3..5]).unwrap_or(""),
16,
) {
if byte_tokens[byte_val as usize] == u32::MAX {
byte_tokens[byte_val as usize] = *id;
}
}
}
}
let mut trie = Trie::new();
for (id, bytes, _) in vocab {
if !bytes.is_empty() {
trie.add(bytes, *id);
}
}
trie.build(MatchKind::Overlapping);
let matcher = trie.compile();
let token_lengths: Vec<u16> = token_bytes.iter().map(|b| b.len() as u16).collect();
let mut token_cache = FoldHashMap::default();
for (id, bytes, _) in vocab {
if bytes.len() <= MAX_CACHED_TOKEN_LEN {
token_cache.insert(bytes.clone(), *id);
}
}
let has_byte_fallback = byte_tokens.iter().any(|&t| t != u32::MAX);
let encoder = Self {
matcher,
scores,
unk_token,
byte_tokens,
token_lengths,
vocab_size: vocab.len(),
token_cache,
has_byte_fallback,
cache_id: NEXT_CACHE_ID.fetch_add(1, Ordering::Relaxed),
unit_split_safe: compute_unit_split_safe(&token_bytes),
};
(encoder, token_bytes)
}
pub fn from_parts(
matcher: DoubleArrayAhoCorasick,
scores: Vec<f64>,
unk_token: TokenId,
byte_tokens: [TokenId; 256],
token_lengths: Vec<u16>,
token_bytes: &[Vec<u8>],
) -> Self {
let vocab_size = scores.len();
let mut token_cache = FoldHashMap::default();
for (id, bytes) in token_bytes.iter().enumerate() {
if bytes.len() <= MAX_CACHED_TOKEN_LEN {
token_cache.insert(bytes.clone(), id as TokenId);
}
}
let has_byte_fallback = byte_tokens.iter().any(|&t| t != u32::MAX);
Self {
matcher,
scores,
unk_token,
byte_tokens,
token_lengths,
vocab_size,
token_cache,
has_byte_fallback,
cache_id: NEXT_CACHE_ID.fetch_add(1, Ordering::Relaxed),
unit_split_safe: compute_unit_split_safe(token_bytes),
}
}
pub fn vocab_size(&self) -> usize {
self.vocab_size
}
pub fn num_base_tokens(&self) -> usize {
self.vocab_size
}
pub fn unk_token(&self) -> TokenId {
self.unk_token
}
pub fn scores(&self) -> &[f64] {
&self.scores
}
pub fn byte_tokens(&self) -> &[TokenId; 256] {
&self.byte_tokens
}
pub fn token_lengths(&self) -> &[u16] {
&self.token_lengths
}
pub fn matcher(&self) -> &DoubleArrayAhoCorasick {
&self.matcher
}
#[inline]
pub fn unit_split_safe(&self) -> bool {
self.unit_split_safe
}
#[inline]
pub fn token_len(&self, token: TokenId) -> usize {
self.token_lengths[token as usize] as usize
}
#[inline]
pub fn is_valid_pair(&self, _token1: TokenId, _token2: TokenId) -> bool {
true
}
pub fn encode(&self, text: &[u8]) -> Vec<TokenId> {
let mut out = Vec::with_capacity(text.len() / 3);
self.encode_into(text, None, &mut out);
out
}
pub fn encode_into(
&self,
text: &[u8],
cache: Option<&mut UnigramPieceCache>,
out: &mut Vec<TokenId>,
) {
if text.is_empty() {
return;
}
if !self.unit_split_safe {
let toks = self.viterbi_unit(text);
self.append_collapsing_unk(out, &toks);
return;
}
match cache {
Some(cache) => self.encode_units_into(text, cache, out),
None => {
let key = self.cache_id as usize;
THREAD_UNIT_CACHE.with(|slot| {
let mut slot = slot.borrow_mut();
let needs_new = match slot.as_ref() {
Some((k, _)) => *k != key,
None => true,
};
if needs_new {
*slot = Some((key, UnigramPieceCache::new()));
}
self.encode_units_into(text, &mut slot.as_mut().unwrap().1, out);
});
}
}
}
fn encode_units_into(
&self,
text: &[u8],
cache: &mut UnigramPieceCache,
out: &mut Vec<TokenId>,
) {
let mut unit_start = 0usize;
for pos in memchr::memmem::find_iter(text, &METASPACE) {
if pos > unit_start {
self.encode_unit_cached(&text[unit_start..pos], cache, out);
unit_start = pos;
}
}
if unit_start < text.len() {
self.encode_unit_cached(&text[unit_start..], cache, out);
}
}
#[inline]
fn encode_unit_cached(
&self,
unit: &[u8],
cache: &mut UnigramPieceCache,
out: &mut Vec<TokenId>,
) {
if unit.is_empty() {
return;
}
let start_len = out.len();
if cache.lookup(unit, out) {
self.collapse_unk_join(out, start_len);
return;
}
let tokens = self.viterbi_unit(unit);
cache.insert(unit, &tokens);
self.append_collapsing_unk(out, &tokens);
}
#[inline]
fn append_collapsing_unk(&self, out: &mut Vec<TokenId>, tokens: &[TokenId]) {
if tokens.is_empty() {
return;
}
let mut start = 0;
if out.last() == Some(&self.unk_token) {
while start < tokens.len() && tokens[start] == self.unk_token {
start += 1;
}
}
out.extend_from_slice(&tokens[start..]);
}
#[inline]
fn collapse_unk_join(&self, out: &mut Vec<TokenId>, appended_at: usize) {
if appended_at == 0 || appended_at >= out.len() {
return;
}
if out[appended_at - 1] == self.unk_token && out[appended_at] == self.unk_token {
let mut end = appended_at;
while end < out.len() && out[end] == self.unk_token {
end += 1;
}
out.drain(appended_at..end);
}
}
fn viterbi_unit(&self, text: &[u8]) -> Vec<TokenId> {
let n = text.len();
if n == 0 {
return Vec::new();
}
let mut best_score = vec![f64::NEG_INFINITY; n + 1];
let mut backptr: Vec<(TokenId, usize)> = vec![(0, 0); n + 1];
best_score[0] = 0.0;
let unk_penalty = if self.has_byte_fallback {
self.scores[self.unk_token as usize]
} else {
-100.0
};
type MatchList = SmallVec<[(usize, TokenId); 8]>;
let mut matches_at: Vec<MatchList> = vec![SmallVec::new(); n];
for m in self.matcher.find_iter(text) {
matches_at[m.start].push((m.end, m.pattern_id));
}
for pos in 0..n {
if best_score[pos] == f64::NEG_INFINITY {
continue;
}
let current_score = best_score[pos];
let has_match = !matches_at[pos].is_empty();
for &(end, token_id) in &matches_at[pos] {
let token_score = self.scores[token_id as usize];
let new_score = current_score + token_score;
if new_score > best_score[end] {
best_score[end] = new_score;
backptr[end] = (token_id, pos);
}
}
let byte_val = text[pos];
let byte_token = self.byte_tokens[byte_val as usize];
if byte_token != u32::MAX {
let token_score = self.scores[byte_token as usize];
let new_score = current_score + token_score;
if new_score > best_score[pos + 1] {
best_score[pos + 1] = new_score;
backptr[pos + 1] = (byte_token, pos);
}
} else if !has_match {
let char_len = utf8_char_len(text[pos]);
let end = (pos + char_len).min(n);
let new_score = current_score + unk_penalty;
if new_score > best_score[end] {
best_score[end] = new_score;
backptr[end] = (self.unk_token, pos);
}
}
}
if best_score[n] == f64::NEG_INFINITY {
return self.encode_with_unk_bridging(text);
}
self.collect_tokens_from_backptr(&backptr, n)
}
#[inline]
fn collect_tokens_from_backptr(&self, backptr: &[(TokenId, usize)], end: usize) -> Vec<TokenId> {
let mut tokens = Vec::new();
let mut pos = end;
while pos > 0 {
let (token_id, start_pos) = backptr[pos];
tokens.push(token_id);
pos = start_pos;
}
tokens.reverse();
if tokens.contains(&self.unk_token) {
tokens.dedup_by(|a, b| *a == self.unk_token && *b == self.unk_token);
}
tokens
}
fn encode_with_unk_bridging(&self, text: &[u8]) -> Vec<TokenId> {
let n = text.len();
let mut tokens = Vec::new();
let mut pos = 0;
while pos < n {
let max_len = (n - pos).min(MAX_CACHED_TOKEN_LEN);
let mut best_match: Option<(usize, TokenId)> = None;
for len in (1..=max_len).rev() {
let substr = &text[pos..pos + len];
if let Some(&token_id) = self.token_cache.get(substr) {
best_match = Some((len, token_id));
break;
}
}
let remaining = &text[pos..];
if let Some(m) = self.matcher.find_iter(remaining).next() {
if m.start == 0 && (best_match.is_none() || m.end > best_match.unwrap().0) {
best_match = Some((m.end, m.pattern_id));
}
}
if let Some((len, token_id)) = best_match {
tokens.push(token_id);
pos += len;
} else {
let byte_val = text[pos];
let byte_token = self.byte_tokens[byte_val as usize];
if byte_token != u32::MAX {
tokens.push(byte_token);
} else {
tokens.push(self.unk_token);
}
pos += 1;
}
}
tokens
}
pub fn encode_chunked(&self, text: &[u8], _chunk_size: usize) -> Vec<TokenId> {
self.encode(text)
}
#[inline]
pub fn encode_chunked_default(&self, text: &[u8]) -> Vec<TokenId> {
self.encode(text)
}
pub fn encode_single(&self, text: &[u8]) -> Vec<TokenId> {
let mut out = Vec::with_capacity(text.len() / 3);
let mut unit_start = 0usize;
for pos in memchr::memmem::find_iter(text, &METASPACE) {
if pos > unit_start {
let toks = self.viterbi_unit(&text[unit_start..pos]);
self.append_collapsing_unk(&mut out, &toks);
unit_start = pos;
}
}
if unit_start < text.len() {
let toks = self.viterbi_unit(&text[unit_start..]);
self.append_collapsing_unk(&mut out, &toks);
}
out
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_unigram() {
let vocab = vec![
(0, b"h".to_vec(), -1.0),
(1, b"e".to_vec(), -1.0),
(2, b"l".to_vec(), -1.0),
(3, b"o".to_vec(), -1.0),
(4, b"hell".to_vec(), -2.0),
(5, b"hello".to_vec(), -3.0), ];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
assert_eq!(encoder.encode(b"hello"), vec![5]);
assert_eq!(encoder.encode(b"h"), vec![0]);
}
#[test]
fn test_viterbi_chooses_best_path() {
let vocab = vec![
(0, b"a".to_vec(), -0.1), (1, b"b".to_vec(), -0.1), (2, b"ab".to_vec(), -10.0), ];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
assert_eq!(encoder.encode(b"ab"), vec![0, 1]);
}
#[test]
fn test_byte_fallback() {
let vocab = vec![
(0, b"<0x00>".to_vec(), -5.0),
(1, b"<0x01>".to_vec(), -5.0),
(2, b"<0xFF>".to_vec(), -5.0),
(3, b"hello".to_vec(), -1.0),
];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
assert_eq!(encoder.byte_tokens[0x00], 0);
assert_eq!(encoder.byte_tokens[0x01], 1);
assert_eq!(encoder.byte_tokens[0xFF], 2);
}
#[test]
fn test_empty_input() {
let vocab = vec![(0, b"a".to_vec(), -1.0)];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
let empty: Vec<TokenId> = vec![];
assert_eq!(encoder.encode(b""), empty);
}
#[test]
fn test_vocab_size() {
let vocab = vec![
(0, b"a".to_vec(), -1.0),
(1, b"b".to_vec(), -1.0),
(2, b"c".to_vec(), -1.0),
];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
assert_eq!(encoder.vocab_size(), 3);
assert_eq!(encoder.num_base_tokens(), 3);
}
#[test]
fn test_unit_cache_matches_cold() {
let mark = "▁".as_bytes();
let vocab = vec![
(0, mark.to_vec(), -1.0),
(1, [mark, b"a"].concat(), -0.5),
(2, b"a".to_vec(), -1.0),
(3, b"b".to_vec(), -1.0),
(4, [mark, b"b"].concat(), -0.5),
];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
let text = [mark, b"a", mark, b"b", mark, b"a"].concat();
let cold = encoder.encode_single(&text);
let mut cache = UnigramPieceCache::new();
let mut warm = Vec::new();
encoder.encode_into(&text, Some(&mut cache), &mut warm);
assert_eq!(cold, warm);
let mut warm2 = Vec::new();
encoder.encode_into(&text, Some(&mut cache), &mut warm2);
assert_eq!(cold, warm2);
}
fn mark() -> &'static [u8] {
"▁".as_bytes()
}
fn interior_metaspace_vocab() -> (UnigramEncoder, Vec<TokenId>) {
let vocab = vec![
(0u32, [mark(), b"a", mark(), b"b"].concat(), -0.5),
(1u32, [mark(), b"a"].concat(), -3.0),
(2u32, [mark(), b"b"].concat(), -3.0),
(3u32, b"a".to_vec(), -5.0),
(4u32, b"b".to_vec(), -5.0),
(5u32, mark().to_vec(), -5.0),
];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 5);
(encoder, vec![0])
}
#[test]
fn test_interior_metaspace_forces_whole_string_viterbi() {
let (encoder, whole_string_expected) = interior_metaspace_vocab();
let text = [mark(), b"a", mark(), b"b"].concat();
assert!(
!encoder.unit_split_safe(),
"interior-`▁` token must make the vocab unit-split-unsafe"
);
assert_eq!(
encoder.encode(&text),
whole_string_expected,
"guarded encode must match whole-string Viterbi"
);
let naive_split = encoder.encode_single(&text);
assert_eq!(
naive_split,
vec![1, 2],
"naive per-unit split picks [▁a, ▁b] — the wrong segmentation"
);
assert_ne!(
encoder.encode(&text),
naive_split,
"guard must produce a different (correct) result than the naive split"
);
let mut cache = UnigramPieceCache::new();
let mut warm = Vec::new();
encoder.encode_into(&text, Some(&mut cache), &mut warm);
assert_eq!(warm, whole_string_expected, "cache path must honor the guard");
}
#[test]
fn test_normal_vocab_is_unit_split_safe() {
let vocab = vec![
(0u32, mark().to_vec(), -1.0),
(1u32, [mark(), b"the"].concat(), -0.5),
(2u32, [mark(), b"a"].concat(), -0.5),
(3u32, b"t".to_vec(), -2.0),
(4u32, b"h".to_vec(), -2.0),
(5u32, b"e".to_vec(), -2.0),
(6u32, b"a".to_vec(), -2.0),
];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
assert!(
encoder.unit_split_safe(),
"vocab without interior `▁` must be unit-split-safe"
);
let text = [mark(), b"the", mark(), b"a"].concat();
let cold = encoder.encode_single(&text);
assert_eq!(encoder.encode(&text), cold);
assert_eq!(encoder.encode(&text), vec![1, 2]);
}
#[test]
fn test_leading_metaspace_is_safe() {
let vocab = vec![
(0u32, mark().to_vec(), -1.0),
(1u32, [mark(), b"word"].concat(), -0.5),
(2u32, b"w".to_vec(), -2.0),
];
let (encoder, _) = UnigramEncoder::from_vocab_with_scores(&vocab, 0);
assert!(encoder.unit_split_safe());
}
}