use std::{cell::RefCell, collections::HashMap, sync::OnceLock};
use tokenizers::Tokenizer;
pub(crate) use self::bpe_mirror::MergeTable;
#[cfg(test)]
pub(crate) use self::suffix_session::build_meter;
use self::{
bpe_mirror::TailClass,
suffix_session::{INITIAL_CAP, Prefix, Session},
};
use crate::embeddings::granite::error::{Error, Result};
mod bpe_mirror;
mod suffix_session;
#[cfg(test)]
pub(crate) mod encode_meter {
use std::cell::Cell;
thread_local! {
static BYTES: Cell<usize> = const { Cell::new(0) };
static CALLS: Cell<usize> = const { Cell::new(0) };
static SIZES: std::cell::RefCell<Vec<usize>> = const { std::cell::RefCell::new(Vec::new()) };
}
pub(crate) fn reset() {
BYTES.with(|b| b.set(0));
CALLS.with(|c| c.set(0));
SIZES.with(|s| s.borrow_mut().clear());
}
pub(crate) fn sizes() -> Vec<usize> {
SIZES.with(|s| s.borrow().clone())
}
pub(crate) fn get() -> usize {
BYTES.with(Cell::get)
}
pub(crate) fn calls() -> usize {
CALLS.with(Cell::get)
}
pub(crate) fn add(n: usize) {
BYTES.with(|b| b.set(b.get().saturating_add(n)));
CALLS.with(|c| c.set(c.get().saturating_add(1)));
SIZES.with(|s| s.borrow_mut().push(n));
}
}
fn encode_content_len(tok: &Tokenizer, s: &str) -> Result<usize> {
#[cfg(test)]
encode_meter::add(s.len());
tok
.encode(s, false)
.map(|e| e.get_ids().len())
.map_err(Error::Tokenize)
}
pub(crate) struct TokenIndex {
pretoken_ends: Vec<u32>,
count_prefix: Vec<u32>,
digit: Vec<bool>,
direct_only: bool,
}
struct Boundary {
z: usize,
count: usize,
}
impl Boundary {
#[inline(always)]
const fn new(z: usize, count: usize) -> Self {
Self { z, count }
}
#[inline(always)]
const fn z(&self) -> usize {
self.z
}
#[inline(always)]
const fn count(&self) -> usize {
self.count
}
}
enum Resync {
Boundary(Boundary),
WholeQuery(usize),
Direct,
}
impl TokenIndex {
pub(crate) fn build(measure_tok: &Tokenizer, text: &str) -> Result<Self> {
#[cfg(test)]
encode_meter::add(text.len());
let enc = measure_tok.encode(text, false).map_err(Error::Tokenize)?;
let offsets = enc.get_offsets();
let word_ids = enc.get_word_ids();
if text.len() > u32::MAX as usize {
return Ok(Self::direct_only());
}
let byte_level_chars: usize = enc.get_tokens().iter().map(|t| t.chars().count()).sum();
if byte_level_chars != text.len() {
return Ok(Self::direct_only());
}
let added = measure_tok.get_added_vocabulary().get_vocab();
let mut literal_lead = [false; 256];
for lit in added.keys() {
if let Some(&b) = lit.as_bytes().first() {
literal_lead[b as usize] = true;
}
}
if text.bytes().any(|b| literal_lead[b as usize])
&& added.keys().any(|lit| text.contains(lit.as_str()))
{
return Ok(Self::direct_only());
}
let mut pretoken_ends: Vec<u32> = Vec::new();
let mut count_prefix: Vec<u32> = vec![0];
let mut acc: u32 = 0;
let mut expected: u32 = 0;
let mut i = 0usize;
let n_tokens = offsets.len();
while i < n_tokens {
let Some(wid) = word_ids[i] else {
return Ok(Self::direct_only());
};
if wid != expected {
return Ok(Self::direct_only());
}
let mut j = i;
let mut end: usize = 0;
while j < n_tokens && word_ids[j] == Some(wid) {
end = end.max(offsets[j].1);
j += 1;
}
pretoken_ends.push(end as u32);
acc = acc.saturating_add((j - i) as u32);
count_prefix.push(acc);
expected += 1;
i = j;
}
let mut prev: u32 = 0;
for (k, &e) in pretoken_ends.iter().enumerate() {
if k == 0 {
if e == 0 {
return Ok(Self::direct_only());
}
} else if e <= prev {
return Ok(Self::direct_only());
}
prev = e;
}
match pretoken_ends.last() {
Some(&last) if last as usize == text.len() => {}
None if text.is_empty() => {}
_ => return Ok(Self::direct_only()),
}
let mut digit: Vec<bool> = Vec::with_capacity(pretoken_ends.len());
let mut start = 0usize;
for &e in &pretoken_ends {
let word = &text[start..e as usize];
digit.push(!word.is_empty() && word.chars().all(char::is_numeric));
start = e as usize;
}
Ok(Self {
pretoken_ends,
count_prefix,
digit,
direct_only: false,
})
}
pub(crate) const fn is_direct_only(&self) -> bool {
self.direct_only
}
fn direct_only() -> Self {
Self {
pretoken_ends: Vec::new(),
count_prefix: vec![0],
digit: Vec::new(),
direct_only: true,
}
}
pub(crate) fn measure_range(
&self,
tok: &Tokenizer,
text: &str,
a: usize,
b: usize,
) -> Result<usize> {
if a >= b {
return Ok(2);
}
if self.direct_only {
return Ok(encode_content_len(tok, &text[a..b])? + 2);
}
let ends = &self.pretoken_ends;
let n = ends.len();
let a32 = a as u32;
let b32 = b as u32;
let z: usize;
let i: usize;
let mut left_count: usize = 0;
if a == 0 {
(z, i) = (0, 0);
} else {
let p = ends.partition_point(|&e| e <= a32);
if p > 0 && ends[p - 1] == a32 {
(z, i) = (a, p);
} else {
let mut pp = p;
while pp + 1 < n && self.digit[pp + 1] {
pp += 1;
}
match self.left_resync(tok, text, a, b, pp)? {
Resync::Boundary(boundary) => {
let zz = boundary.z();
z = zz;
i = ends.partition_point(|&e| e <= zz as u32);
left_count = boundary.count();
}
Resync::WholeQuery(count) => {
return Ok(count + 2);
}
Resync::Direct => {
return Ok(encode_content_len(tok, &text[a..b])? + 2);
}
}
}
}
let (y, j, right_fragment): (usize, usize, Option<(usize, usize)>);
if b == text.len() {
(y, j, right_fragment) = (text.len(), n, None);
} else {
let q = ends.partition_point(|&e| e < b32);
let b_is_boundary = ends[q] == b32;
let prev_ws = text[..b]
.chars()
.next_back()
.is_some_and(char::is_whitespace);
if b_is_boundary && !prev_ws {
(y, j, right_fragment) = (b, q + 1, None);
} else {
let y0 = if b_is_boundary {
b
} else if q == 0 {
0
} else {
ends[q - 1] as usize
};
let yy = self.snap_down_boundary(scan_back_whitespace(text, y0));
let jj = ends.partition_point(|&e| e <= yy as u32);
(y, j, right_fragment) = (yy, jj, Some((yy, b)));
}
}
if z >= y {
return Ok(encode_content_len(tok, &text[a..b])? + 2);
}
let mut total = left_count;
total += (self.count_prefix[j] - self.count_prefix[i]) as usize;
if let Some((fy, fb)) = right_fragment {
total += encode_content_len(tok, &text[fy..fb])?;
}
Ok(total + 2)
}
pub(crate) fn measure_range_fast(
&self,
tok: &Tokenizer,
text: &str,
a: usize,
b: usize,
lane: &mut FastLane<'_>,
) -> Result<usize> {
if self.direct_only || a >= b {
return self.measure_range(tok, text, a, b);
}
let ends = &self.pretoken_ends;
let p = ends.partition_point(|&e| e <= a as u32);
if p >= ends.len() {
return self.measure_range(tok, text, a, b);
}
let p_start = if p == 0 { 0 } else { ends[p - 1] as usize };
let p_end = ends[p] as usize;
if b > p_end || (a == p_start && b == p_end) {
return self.measure_range(tok, text, a, b);
}
if lane.table.is_none() && b - a <= FastLane::ENGAGE_BYTES {
return self.measure_range(tok, text, a, b);
}
let qualifies = match lane.class_ok.get(&p) {
Some(&q) => q,
None => {
let q = match &lane.tail {
Some(tail) => {
let memo = &mut lane.tail_memo;
text[p_start..p_end]
.chars()
.skip(1)
.all(|c| *memo.entry(c).or_insert_with(|| tail.contains(c)))
}
None => false,
};
lane.class_ok.insert(p, q);
q
}
};
if !qualifies {
return self.measure_range(tok, text, a, b);
}
if lane.table.is_none() {
lane.table = Some(lane.lazy.get());
}
let Some(table) = lane.table.flatten() else {
return self.measure_range(tok, text, a, b);
};
if b - a <= table.max_token_bytes() {
let bytes = &text.as_bytes()[a..b];
if table.whole_word_id(tok, bytes).is_some() {
lane_oracle(tok, text.get(a..b), 1 + 2);
return Ok(1 + 2);
}
return match table.process(bytes) {
Some(run) => {
let answer = run.ends.len() + 2;
lane_oracle(tok, text.get(a..b), answer);
Ok(answer)
}
None => self.measure_range(tok, text, a, b),
};
}
if lane.dead_start == Some(a) {
return self.measure_range(tok, text, a, b);
}
lane.uses += 1;
let now = lane.uses;
let snap = |mut end: usize| -> usize {
end = end.min(p_end);
while end > a && !text.is_char_boundary(end) {
end -= 1;
}
end
};
let slot = match lane.sessions.iter().position(|(s, _)| s.start() == a) {
Some(i) => i,
None => {
if lane.sessions.len() >= FastLane::SESSIONS {
let lru = lane
.sessions
.iter()
.enumerate()
.min_by_key(|(_, (_, used))| *used)
.map_or(0, |(i, _)| i);
lane.sessions.swap_remove(lru);
}
let cap = {
let tokens = (self.count_prefix[p + 1] - self.count_prefix[p]).max(1) as usize;
let bytes = p_end - p_start;
let per_token = bytes.div_ceil(tokens).max(1);
(lane.window.saturating_mul(per_token).saturating_mul(5) / 4).clamp(1024, INITIAL_CAP)
};
let end = snap(a.saturating_add(cap));
let built = if end > a {
Session::build(tok, table, text, a, end)?
} else {
None
};
match built {
Some(s) => {
lane.sessions.push((s, now));
lane.sessions.len() - 1
}
None => {
lane.dead_start = Some(a);
return self.measure_range(tok, text, a, b);
}
}
}
};
lane.sessions[slot].1 = now;
while b > lane.sessions[slot].0.end() {
let have = lane.sessions[slot].0.end();
let end = snap(a.saturating_add((have - a).saturating_mul(2)));
let grown = if end > have {
Session::build(tok, table, text, a, end)?
} else {
None
};
match grown {
Some(s) => lane.sessions[slot].0 = s,
None => {
lane.dead_start = Some(a);
lane.sessions.swap_remove(slot);
return self.measure_range(tok, text, a, b);
}
}
}
let session = &mut lane.sessions[slot].0;
match session.measure_prefix(tok, table, b) {
Prefix::Count(content) => {
lane_oracle(tok, text.get(a..b), content + 2);
Ok(content + 2)
}
Prefix::Direct | Prefix::PastCap => Ok(encode_content_len(tok, &text[a..b])? + 2),
}
}
fn left_resync(
&self,
tok: &Tokenizer,
text: &str,
a: usize,
b: usize,
pp: usize,
) -> Result<Resync> {
let ends = &self.pretoken_ends;
let n = ends.len();
let hi = (ends[(pp + 2).min(n - 1)] as usize).min(b);
if hi <= a {
return Ok(Resync::Direct);
}
#[cfg(test)]
encode_meter::add(hi - a);
let enc = tok.encode(&text[a..hi], false).map_err(Error::Tokenize)?;
let offsets = enc.get_offsets();
let word_ids = enc.get_word_ids();
let m = offsets.len();
for k in 0..m {
let abs = a + offsets[k].1;
let sub_boundary = k + 1 == m || word_ids[k + 1] != word_ids[k];
if sub_boundary && abs > a && self.is_full_boundary(abs as u32) {
if abs == hi && hi < b {
return Ok(Resync::Direct);
}
return Ok(Resync::Boundary(Boundary::new(abs, k + 1)));
}
}
if hi == b {
Ok(Resync::WholeQuery(m))
} else {
Ok(Resync::Direct)
}
}
fn is_full_boundary(&self, pos: u32) -> bool {
pos == 0 || self.pretoken_ends.binary_search(&pos).is_ok()
}
fn snap_down_boundary(&self, pos: usize) -> usize {
let c = self.pretoken_ends.partition_point(|&e| (e as usize) <= pos);
if c == 0 {
0
} else {
self.pretoken_ends[c - 1] as usize
}
}
}
fn scan_back_whitespace(text: &str, from: usize) -> usize {
let mut y = from;
for (idx, ch) in text[..from].char_indices().rev() {
if ch.is_whitespace() {
y = idx;
} else {
break;
}
}
y
}
#[cfg(debug_assertions)]
fn lane_oracle(tok: &Tokenizer, probe: Option<&str>, answer: usize) {
if let Some(s) = probe
&& let Ok(enc) = tok.encode(s, true)
{
debug_assert_eq!(answer, enc.get_ids().len(), "fast-lane answer for {s:?}");
}
}
#[cfg(not(debug_assertions))]
#[inline]
const fn lane_oracle(_: &Tokenizer, _: Option<&str>, _: usize) {}
#[derive(Clone, Copy)]
pub(crate) struct LazyTable<'t> {
cell: &'t OnceLock<Option<MergeTable>>,
tok: &'t Tokenizer,
}
impl<'t> LazyTable<'t> {
pub(crate) const fn new(cell: &'t OnceLock<Option<MergeTable>>, tok: &'t Tokenizer) -> Self {
Self { cell, tok }
}
pub(crate) fn get(&self) -> Option<&'t MergeTable> {
self
.cell
.get_or_init(|| MergeTable::from_tokenizer(self.tok))
.as_ref()
}
}
pub(crate) struct FastLane<'t> {
lazy: LazyTable<'t>,
table: Option<Option<&'t MergeTable>>,
tail: Option<TailClass>,
sessions: Vec<(Session, u64)>,
uses: u64,
dead_start: Option<usize>,
class_ok: HashMap<usize, bool>,
tail_memo: HashMap<char, bool>,
window: usize,
}
impl<'t> FastLane<'t> {
const SESSIONS: usize = 2;
const ENGAGE_BYTES: usize = 128;
fn new(window: usize, lazy: LazyTable<'t>) -> Self {
Self {
lazy,
table: None,
tail: TailClass::new(),
sessions: Vec::new(),
uses: 0,
dead_start: None,
class_ok: HashMap::new(),
tail_memo: HashMap::new(),
window,
}
}
#[cfg(test)]
pub(crate) fn engaged(window: usize, lazy: LazyTable<'t>) -> Self {
let mut lane = Self::new(window, lazy);
lane.table = Some(lazy.get());
lane
}
}
pub(crate) struct IndexMeasure<'a> {
text: &'a str,
index: &'a TokenIndex,
tok: &'a Tokenizer,
lane: RefCell<FastLane<'a>>,
}
impl<'a> IndexMeasure<'a> {
pub(crate) fn new(
text: &'a str,
index: &'a TokenIndex,
tok: &'a Tokenizer,
table: LazyTable<'a>,
window: usize,
) -> Self {
Self {
text,
index,
tok,
lane: RefCell::new(FastLane::new(window, table)),
}
}
}
impl windit::split::MeasureText for IndexMeasure<'_> {
fn measure(&self, s: &str) -> usize {
let base = self.text.as_ptr() as usize;
let sp = s.as_ptr() as usize;
if let Some(off) = sp.checked_sub(base)
&& off <= self.text.len()
&& off + s.len() <= self.text.len()
{
let measured = {
let mut lane = self.lane.borrow_mut();
self
.index
.measure_range_fast(self.tok, self.text, off, off + s.len(), &mut lane)
};
return measured.unwrap_or(usize::MAX);
}
self
.tok
.encode(s, true)
.map(|e| e.get_ids().len())
.unwrap_or(usize::MAX)
}
fn measure_within(&self, s: &str, limit: usize) -> Option<usize> {
let covered_subslice = !self.index.is_direct_only() && {
let base = self.text.as_ptr() as usize;
let sp = s.as_ptr() as usize;
sp.checked_sub(base)
.is_some_and(|off| off <= self.text.len() && off + s.len() <= self.text.len())
};
if covered_subslice {
let content_floor = if s.is_empty() {
0
} else {
self
.lane
.borrow()
.table
.flatten()
.map_or(1, |t| s.len().div_ceil(t.max_token_bytes().max(1)))
};
if 2 + content_floor > limit {
return None;
}
}
let measured = self.measure(s);
(measured <= limit).then_some(measured)
}
}
#[cfg(test)]
mod tests;