#![allow(dead_code, unused_imports)]
pub(crate) mod pretoken_cache;
pub mod tiktoken;
pub(crate) fn madvise_hugepage(ptr: *mut u8, bytes: usize) {
#[cfg(target_os = "linux")]
{
const PAGE: usize = 4096;
let start = (ptr as usize + PAGE - 1) & !(PAGE - 1);
let end = (ptr as usize).saturating_add(bytes);
if end > start {
unsafe {
libc::madvise(start as *mut libc::c_void, end - start, libc::MADV_HUGEPAGE);
}
}
}
#[cfg(not(target_os = "linux"))]
let _ = (ptr, bytes);
}
use crate::token::TokenId;
use anyhow::{Result, anyhow};
use std::collections::HashMap;
#[derive(Clone)]
pub struct ByteRemapping {
mapping: Vec<TokenId>,
}
fn is_valid_utf8_byte(b: u8) -> bool {
!matches!(b, 0xC0 | 0xC1 | 0xF5..=0xFF)
}
impl ByteRemapping {
pub fn from_byte_vocab(vocab: &[impl AsRef<[u8]>]) -> Result<Option<Self>> {
const UNSET: u32 = u32::MAX;
let mut mapping = vec![TokenId::from(UNSET); 256];
for (id, entry) in vocab.iter().enumerate() {
if let &[b] = entry.as_ref()
&& mapping[b as usize].0 == UNSET
{
mapping[b as usize] = TokenId::from(id as u32);
}
}
if let Some(missing) = mapping
.iter()
.enumerate()
.position(|(b, t)| t.0 == UNSET && is_valid_utf8_byte(b as u8))
{
return Err(anyhow!(
"Byte remapping failed: no single-byte vocab entry for byte {missing:#04x}"
));
}
for t in mapping.iter_mut() {
if t.0 == UNSET {
*t = TokenId::from(0);
}
}
Ok(mapping
.iter()
.enumerate()
.any(|(b, t)| t.0 != b as u32)
.then_some(ByteRemapping { mapping }))
}
}
pub(crate) struct PairRankTable {
dense: Box<[u32]>,
dense_log2: u32,
slots: Box<[u64]>,
mask: usize,
shift: u32,
}
const PAIR_ID_BITS: u32 = 21;
impl PairRankTable {
pub(crate) fn build<S: std::hash::BuildHasher>(
merges: &HashMap<(TokenId, TokenId), TokenId, S>,
byte_remapping: Option<&ByteRemapping>,
vocab_len: usize,
) -> Option<Self> {
if vocab_len > 1 << PAIR_ID_BITS {
return None;
}
let id_limit = 1u32 << PAIR_ID_BITS;
if merges
.iter()
.any(|(&(a, b), &m)| a.0 >= id_limit || b.0 >= id_limit || m.0 >= id_limit)
{
return None;
}
let max_initial =
byte_remapping.map_or(255, |br| br.mapping.iter().map(|t| t.0).max().unwrap_or(0));
let dense_log2 = (32 - max_initial.leading_zeros()).clamp(11, 11);
let mut dense = vec![u32::MAX; 1usize << (2 * dense_log2)].into_boxed_slice();
let n_slots = (merges.len().max(1) * 2).next_power_of_two().max(64);
let shift = 64 - n_slots.trailing_zeros();
let mask = n_slots - 1;
let mut slots = vec![u64::MAX; n_slots].into_boxed_slice();
for (&(a, b), &m) in merges {
if (a.0 | b.0) >> dense_log2 == 0 {
dense[((a.0 as usize) << dense_log2) | b.0 as usize] = m.0;
}
let key = ((a.0 as u64) << PAIR_ID_BITS) | b.0 as u64;
let mut idx = (key.wrapping_mul(0x9E37_79B9_7F4A_7C15) >> shift) as usize;
let mut displacement = 0usize;
while slots[idx] != u64::MAX {
idx = (idx + 1) & mask;
displacement += 1;
if displacement > 64 {
return None;
}
}
slots[idx] = (key << PAIR_ID_BITS) | m.0 as u64;
}
Some(PairRankTable {
dense,
dense_log2,
slots,
mask,
shift,
})
}
#[inline(always)]
pub(crate) fn rank(&self, a: TokenId, b: TokenId) -> u32 {
if (a.0 | b.0) >> self.dense_log2 == 0 {
let idx = ((a.0 as usize) << self.dense_log2) | b.0 as usize;
return unsafe { *self.dense.get_unchecked(idx) };
}
let key = ((a.0 as u64) << PAIR_ID_BITS) | b.0 as u64;
let mut idx = (key.wrapping_mul(0x9E37_79B9_7F4A_7C15) >> self.shift) as usize;
loop {
let slot = unsafe { *self.slots.get_unchecked(idx) };
if slot >> PAIR_ID_BITS == key {
return (slot & ((1 << PAIR_ID_BITS) - 1)) as u32;
}
if slot == u64::MAX {
return u32::MAX;
}
idx = (idx + 1) & self.mask;
}
}
#[inline(always)]
pub(crate) fn prefetch_rank(&self, a: TokenId, b: TokenId) {
#[cfg(target_arch = "x86_64")]
unsafe {
use core::arch::x86_64::{_MM_HINT_T0, _mm_prefetch};
let addr = if (a.0 | b.0) >> self.dense_log2 == 0 {
let idx = ((a.0 as usize) << self.dense_log2) | b.0 as usize;
self.dense.as_ptr().add(idx) as *const i8
} else {
let key = ((a.0 as u64) << PAIR_ID_BITS) | b.0 as u64;
let idx = (key.wrapping_mul(0x9E37_79B9_7F4A_7C15) >> self.shift) as usize;
self.slots.as_ptr().add(idx) as *const i8
};
_mm_prefetch(addr, _MM_HINT_T0);
}
#[cfg(not(target_arch = "x86_64"))]
let _ = (a, b);
}
}
#[derive(Default)]
pub struct MergeScratch {
next: Vec<u32>,
prev: Vec<u32>,
heap: Vec<std::cmp::Reverse<u64>>,
}
#[inline(always)]
fn pack_merge_entry(merged: TokenId, pos: u32) -> u64 {
((merged.0 as u64) << 32) | pos as u64
}
pub fn bpe_merge_symbols<S: std::hash::BuildHasher>(
merges: &HashMap<(TokenId, TokenId), TokenId, S>,
symbols: &mut Vec<TokenId>,
) {
bpe_merge_symbols_with_scratch(merges, symbols, &mut MergeScratch::default());
}
pub fn bpe_merge_symbols_with_scratch<S: std::hash::BuildHasher>(
merges: &HashMap<(TokenId, TokenId), TokenId, S>,
symbols: &mut Vec<TokenId>,
scratch: &mut MergeScratch,
) {
bpe_merge_symbols_by_rank(
&|a, b| merges.get(&(a, b)).map_or(u32::MAX, |m| m.0),
symbols,
scratch,
);
}
#[inline(never)]
pub(crate) fn bpe_merge_symbols_by_rank(
get_rank: &impl Fn(TokenId, TokenId) -> u32,
symbols: &mut Vec<TokenId>,
scratch: &mut MergeScratch,
) {
use std::cmp::Reverse;
use std::collections::BinaryHeap;
let n = symbols.len();
if n < 2 {
return;
}
if n <= SMALL_MERGE_MAX {
bpe_merge_symbols_small(get_rank, symbols);
return;
}
const NONE: u32 = u32::MAX;
let next = &mut scratch.next;
let prev = &mut scratch.prev;
next.clear();
next.extend(1..n as u32);
next.push(NONE);
prev.clear();
prev.push(NONE);
prev.extend(0..n as u32 - 1);
let mut seeds = std::mem::take(&mut scratch.heap);
seeds.clear();
for i in 0..n - 1 {
let m = get_rank(symbols[i], symbols[i + 1]);
if m != u32::MAX {
seeds.push(Reverse(pack_merge_entry(TokenId(m), i as u32)));
}
}
let mut heap: BinaryHeap<Reverse<u64>> = BinaryHeap::from(seeds);
while let Some(Reverse(entry)) = heap.pop() {
let pos = (entry & u32::MAX as u64) as usize;
let expected_merged = (entry >> 32) as u32;
let right = next[pos];
if right == NONE {
continue;
}
let right = right as usize;
let merged = get_rank(symbols[pos], symbols[right]);
if merged == expected_merged {
symbols[pos] = TokenId(merged);
let right_right = next[right];
next[pos] = right_right;
if right_right != NONE {
prev[right_right as usize] = pos as u32;
}
next[right] = NONE;
let left = prev[pos];
if left != NONE {
let m = get_rank(symbols[left as usize], symbols[pos]);
if m != u32::MAX {
heap.push(Reverse(pack_merge_entry(TokenId(m), left)));
}
}
if next[pos] != NONE {
let m = get_rank(symbols[pos], symbols[next[pos] as usize]);
if m != u32::MAX {
heap.push(Reverse(pack_merge_entry(TokenId(m), pos as u32)));
}
}
}
}
scratch.heap = heap.into_vec();
let mut write = 0;
let mut i = 0;
loop {
symbols[write] = symbols[i];
write += 1;
if next[i] == NONE {
break;
}
i = next[i] as usize;
}
symbols.truncate(write);
}
const SMALL_MERGE_MAX: usize = 32;
fn bpe_merge_symbols_small(
get_rank: &impl Fn(TokenId, TokenId) -> u32,
symbols: &mut Vec<TokenId>,
) {
let n = symbols.len();
debug_assert!((2..=SMALL_MERGE_MAX).contains(&n));
let mut next = [0u8; SMALL_MERGE_MAX];
let mut prev = [0u8; SMALL_MERGE_MAX];
for i in 0..n {
next[i] = (i + 1) as u8;
prev[i] = (i as u8).wrapping_sub(1);
}
let mut ranks = [u32::MAX; SMALL_MERGE_MAX];
for i in 0..n - 1 {
ranks[i] = get_rank(symbols[i], symbols[i + 1]);
}
loop {
let mut best = u32::MAX;
let mut best_i = 0;
for (i, &rank) in ranks[..n - 1].iter().enumerate() {
if rank < best {
best = rank;
best_i = i;
}
}
if best == u32::MAX {
break;
}
let i = best_i;
symbols[i] = TokenId::from(best);
let dead = next[i] as usize;
let new_right = next[dead] as usize;
next[i] = new_right as u8;
ranks[dead] = u32::MAX;
if new_right < n {
prev[new_right] = i as u8;
ranks[i] = get_rank(symbols[i], symbols[new_right]);
} else {
ranks[i] = u32::MAX;
}
let left = prev[i] as usize;
if left < n {
ranks[left] = get_rank(symbols[left], symbols[i]);
}
}
let mut write = 0;
let mut i = 0;
while i < n {
symbols[write] = symbols[i];
write += 1;
i = next[i] as usize;
}
symbols.truncate(write);
}
pub(crate) const SHORT_MERGE_MAX: usize = 16;
pub(crate) fn bpe_merge_symbols_short_scalar(
get_rank: impl Fn(TokenId, TokenId) -> u32,
prefetch_rank: impl Fn(TokenId, TokenId),
symbols: &mut [TokenId; SHORT_MERGE_MAX],
n: usize,
) -> usize {
debug_assert!((2..=SHORT_MERGE_MAX - 1).contains(&n));
let mut next = [0u8; SHORT_MERGE_MAX];
let mut prev = [0u8; SHORT_MERGE_MAX];
for i in 0..n {
next[i] = (i + 1) as u8;
prev[i] = (i as u8).wrapping_sub(1);
}
let mut ranks = [u32::MAX; SHORT_MERGE_MAX];
for i in 0..n - 1 {
ranks[i] = get_rank(symbols[i], symbols[i + 1]);
}
loop {
let mut best = u32::MAX;
let mut best_i = 0;
for (i, &rank) in ranks[..n - 1].iter().enumerate() {
if rank < best {
best = rank;
best_i = i;
}
}
if best == u32::MAX {
break;
}
let i = best_i;
let dead = next[i] as usize;
let new_right = next[dead] as usize;
let left = prev[i] as usize;
if new_right < n {
prefetch_rank(TokenId(best), symbols[new_right]);
}
if left < n {
prefetch_rank(symbols[left], TokenId(best));
}
symbols[i] = TokenId(best);
next[i] = new_right as u8;
ranks[dead] = u32::MAX;
if new_right < n {
prev[new_right] = i as u8;
ranks[i] = get_rank(symbols[i], symbols[new_right]);
} else {
ranks[i] = u32::MAX;
}
if left < n {
ranks[left] = get_rank(symbols[left], symbols[i]);
}
}
let mut write = 0;
let mut i = 0;
while i < n {
symbols[write] = symbols[i];
write += 1;
i = next[i] as usize;
}
write
}
#[cfg(target_arch = "aarch64")]
pub(crate) fn bpe_merge_symbols_short_neon(
table: &PairRankTable,
symbols: &mut [TokenId; SHORT_MERGE_MAX],
n: usize,
) -> usize {
use core::arch::aarch64::{vld1q_u32, vminq_u32, vminvq_u32};
debug_assert!((2..=SHORT_MERGE_MAX - 1).contains(&n));
const NO_MERGE_FLOOR: u32 = u32::MAX << 8;
let pack = |rank: u32, i: usize| (rank << 8) | i as u32;
let mut next = [0u8; SHORT_MERGE_MAX];
let mut prev = [0u8; SHORT_MERGE_MAX];
for i in 0..n {
next[i] = (i + 1) as u8;
prev[i] = (i as u8).wrapping_sub(1);
}
let mut pr = [u32::MAX; SHORT_MERGE_MAX];
for i in 0..n - 1 {
pr[i] = pack(table.rank(symbols[i], symbols[i + 1]), i);
}
let narrow = n <= 8;
loop {
let best = unsafe {
let p = pr.as_ptr();
let m01 = vminq_u32(vld1q_u32(p), vld1q_u32(p.add(4)));
let m = if narrow {
m01
} else {
let m23 = vminq_u32(vld1q_u32(p.add(8)), vld1q_u32(p.add(12)));
vminq_u32(m01, m23)
};
vminvq_u32(m)
};
if best >= NO_MERGE_FLOOR {
break;
}
let i = (best & 0xFF) as usize;
symbols[i] = TokenId(best >> 8);
let dead = next[i] as usize;
let new_right = next[dead] as usize;
next[i] = new_right as u8;
pr[dead] = u32::MAX;
if new_right < n {
prev[new_right] = i as u8;
pr[i] = pack(table.rank(symbols[i], symbols[new_right]), i);
} else {
pr[i] = u32::MAX;
}
let left = prev[i] as usize;
if left < n {
pr[left] = pack(table.rank(symbols[left], symbols[i]), left);
}
}
let mut write = 0;
let mut i = 0;
while i < n {
symbols[write] = symbols[i];
write += 1;
i = next[i] as usize;
}
write
}
#[cfg(target_arch = "x86_64")]
#[cfg_attr(not(test), allow(dead_code))]
#[target_feature(enable = "avx512f")]
fn bpe_merge_symbols_short_avx512(
table: &PairRankTable,
symbols: &mut [TokenId; SHORT_MERGE_MAX],
n: usize,
) -> usize {
use core::arch::x86_64::{_mm512_loadu_si512, _mm512_reduce_min_epu32};
debug_assert!((2..=SHORT_MERGE_MAX - 1).contains(&n));
const NO_MERGE_FLOOR: u32 = u32::MAX << 8;
let pack = |rank: u32, i: usize| (rank << 8) | i as u32;
let mut next = [0u8; SHORT_MERGE_MAX];
let mut prev = [0u8; SHORT_MERGE_MAX];
for i in 0..n {
next[i] = (i + 1) as u8;
prev[i] = (i as u8).wrapping_sub(1);
}
let mut pr = [u32::MAX; SHORT_MERGE_MAX];
for i in 0..n - 1 {
pr[i] = pack(table.rank(symbols[i], symbols[i + 1]), i);
}
loop {
let best = unsafe { _mm512_reduce_min_epu32(_mm512_loadu_si512(pr.as_ptr() as *const _)) };
if best >= NO_MERGE_FLOOR {
break;
}
let i = (best & 0xFF) as usize;
symbols[i] = TokenId(best >> 8);
let dead = next[i] as usize;
let new_right = next[dead] as usize;
next[i] = new_right as u8;
pr[dead] = u32::MAX;
if new_right < n {
prev[new_right] = i as u8;
pr[i] = pack(table.rank(symbols[i], symbols[new_right]), i);
} else {
pr[i] = u32::MAX;
}
let left = prev[i] as usize;
if left < n {
pr[left] = pack(table.rank(symbols[left], symbols[i]), left);
}
}
let mut write = 0;
let mut i = 0;
while i < n {
symbols[write] = symbols[i];
write += 1;
i = next[i] as usize;
}
write
}
#[cfg(target_arch = "x86_64")]
#[cfg_attr(not(test), allow(dead_code))]
#[target_feature(enable = "avx2")]
fn bpe_merge_symbols_short_avx2(
table: &PairRankTable,
symbols: &mut [TokenId; SHORT_MERGE_MAX],
n: usize,
) -> usize {
use core::arch::x86_64::*;
debug_assert!((2..=SHORT_MERGE_MAX - 1).contains(&n));
const NO_MERGE_FLOOR: u32 = u32::MAX << 8;
let pack = |rank: u32, i: usize| (rank << 8) | i as u32;
let mut next = [0u8; SHORT_MERGE_MAX];
let mut prev = [0u8; SHORT_MERGE_MAX];
for i in 0..n {
next[i] = (i + 1) as u8;
prev[i] = (i as u8).wrapping_sub(1);
}
let mut pr = [u32::MAX; SHORT_MERGE_MAX];
for i in 0..n - 1 {
pr[i] = pack(table.rank(symbols[i], symbols[i + 1]), i);
}
let narrow = n <= 8;
loop {
let best = unsafe {
let p = pr.as_ptr() as *const __m256i;
let m = if narrow {
_mm256_loadu_si256(p)
} else {
_mm256_min_epu32(_mm256_loadu_si256(p), _mm256_loadu_si256(p.add(1)))
};
let m128 = _mm_min_epu32(_mm256_castsi256_si128(m), _mm256_extracti128_si256::<1>(m));
let m128 = _mm_min_epu32(m128, _mm_shuffle_epi32::<0b01_00_11_10>(m128));
let m128 = _mm_min_epu32(m128, _mm_shuffle_epi32::<0b00_00_00_01>(m128));
_mm_cvtsi128_si32(m128) as u32
};
if best >= NO_MERGE_FLOOR {
break;
}
let i = (best & 0xFF) as usize;
symbols[i] = TokenId(best >> 8);
let dead = next[i] as usize;
let new_right = next[dead] as usize;
next[i] = new_right as u8;
pr[dead] = u32::MAX;
if new_right < n {
prev[new_right] = i as u8;
pr[i] = pack(table.rank(symbols[i], symbols[new_right]), i);
} else {
pr[i] = u32::MAX;
}
let left = prev[i] as usize;
if left < n {
pr[left] = pack(table.rank(symbols[left], symbols[i]), left);
}
}
let mut write = 0;
let mut i = 0;
while i < n {
symbols[write] = symbols[i];
write += 1;
i = next[i] as usize;
}
write
}
pub(crate) fn vocab_entries(vocab: &[std::sync::Arc<[u8]>]) -> impl Iterator<Item = (u32, &[u8])> {
vocab
.iter()
.enumerate()
.filter(|(_, bytes)| !bytes.is_empty())
.map(|(id, bytes)| (id as u32, bytes.as_ref()))
}
#[inline(always)]
pub fn ranked_merge_key(a: TokenId, b: TokenId) -> u64 {
((a.0 as u64) << 32) | b.0 as u64
}
pub(crate) fn bpe_merge_symbols_ranked_slice<S: std::hash::BuildHasher>(
merges: &HashMap<u64, (TokenId, u32), S>,
symbols: &mut [TokenId],
) -> usize {
let get = |a: TokenId, b: TokenId| -> (TokenId, u32) {
merges
.get(&ranked_merge_key(a, b))
.map_or((TokenId::from(0u32), u32::MAX), |&m| m)
};
let n = symbols.len();
debug_assert!((2..=SMALL_MERGE_MAX).contains(&n));
let mut next = [0u8; SMALL_MERGE_MAX];
let mut prev = [0u8; SMALL_MERGE_MAX];
for i in 0..n {
next[i] = (i + 1) as u8;
prev[i] = (i as u8).wrapping_sub(1);
}
let mut ranks = [u32::MAX; SMALL_MERGE_MAX];
let mut merged = [TokenId::from(0u32); SMALL_MERGE_MAX];
for i in 0..n - 1 {
(merged[i], ranks[i]) = get(symbols[i], symbols[i + 1]);
}
loop {
let mut best = u32::MAX;
let mut best_i = 0;
for (i, &rank) in ranks[..n - 1].iter().enumerate() {
if rank < best {
best = rank;
best_i = i;
}
}
if best == u32::MAX {
break;
}
let i = best_i;
symbols[i] = merged[i];
let dead = next[i] as usize;
let new_right = next[dead] as usize;
next[i] = new_right as u8;
ranks[dead] = u32::MAX;
if new_right < n {
prev[new_right] = i as u8;
(merged[i], ranks[i]) = get(symbols[i], symbols[new_right]);
} else {
ranks[i] = u32::MAX;
}
let left = prev[i] as usize;
if left < n {
(merged[left], ranks[left]) = get(symbols[left], symbols[i]);
}
}
let mut write = 0;
let mut i = 0;
while i < n {
symbols[write] = symbols[i];
write += 1;
i = next[i] as usize;
}
write
}
pub fn bpe_merge_symbols_ranked<S: std::hash::BuildHasher>(
merges: &HashMap<u64, (TokenId, u32), S>,
symbols: &mut Vec<TokenId>,
) {
use std::cmp::Reverse;
use std::collections::BinaryHeap;
let n = symbols.len();
if n < 2 {
return;
}
if n <= SMALL_MERGE_MAX {
let new_len = bpe_merge_symbols_ranked_slice(merges, symbols);
symbols.truncate(new_len);
return;
}
const NONE: usize = usize::MAX;
let mut next: Vec<usize> = (1..n).chain(std::iter::once(NONE)).collect();
let mut prev: Vec<usize> = std::iter::once(NONE).chain(0..n - 1).collect();
let mut token = symbols.clone();
let mut heap: BinaryHeap<Reverse<(u32, usize)>> = BinaryHeap::new();
let mut i = 0;
while i < n {
let j = next[i];
if j == NONE {
break;
}
if let Some(&(_, rank)) = merges.get(&ranked_merge_key(token[i], token[j])) {
heap.push(Reverse((rank, i)));
}
i = j;
}
while let Some(Reverse((rank, pos))) = heap.pop() {
let right = next[pos];
if right == NONE {
continue;
}
let pair = ranked_merge_key(token[pos], token[right]);
match merges.get(&pair) {
Some(&(merged_token, r)) if r == rank => {
token[pos] = merged_token;
let right_right = next[right];
next[pos] = right_right;
if right_right != NONE {
prev[right_right] = pos;
}
next[right] = NONE;
prev[right] = NONE;
let left = prev[pos];
if left != NONE
&& let Some(&(_, rank)) = merges.get(&ranked_merge_key(token[left], token[pos]))
{
heap.push(Reverse((rank, left)));
}
if next[pos] != NONE
&& let Some(&(_, rank)) =
merges.get(&ranked_merge_key(token[pos], token[next[pos]]))
{
heap.push(Reverse((rank, pos)));
}
}
_ => continue, }
}
symbols.clear();
let mut i = 0;
loop {
symbols.push(token[i]);
if next[i] == NONE {
break;
}
i = next[i];
}
}
pub fn simple_bpe_merge<S: std::hash::BuildHasher>(
merges: &HashMap<(TokenId, TokenId), TokenId, S>,
pre_token: &[u8],
) -> Vec<TokenId> {
let mut symbols: Vec<TokenId> = pre_token.iter().map(|&b| TokenId::from(b as u32)).collect();
bpe_merge_symbols(merges, &mut symbols);
symbols
}
pub use tiktoken::Tokenizer;