use crate::core::encoder::Encoder;
pub(crate) fn merge_ranks<'a>(
merged: Vec<String>,
vocab_in_id_order: impl Iterator<Item = &'a str>,
) -> Encoder {
let mut merge_set: rustc_hash::FxHashSet<&str> =
rustc_hash::FxHashSet::with_capacity_and_hasher(merged.len(), rustc_hash::FxBuildHasher);
merge_set.extend(merged.iter().map(String::as_str));
let mut ranks: Encoder = Encoder::with_capacity(merged.len() + 512);
for token in vocab_in_id_order.filter(|t| !merge_set.contains(t)) {
let next = ranks.len() as u32;
ranks.insert_if_absent(token.as_bytes(), next);
}
drop(merge_set);
let base_count = ranks.len() as u32;
for (i, token) in merged.into_iter().enumerate() {
ranks.insert_if_absent(token.as_bytes(), base_count + i as u32);
}
ranks
}
const SHORT_MIN: usize = 3;
const SHORT_MAX: usize = 4;
const SHORT_MAX_RANK: u32 = u32::MAX >> 1;
pub(crate) struct BytePairRanks {
ranks: Box<[u32]>,
short: Option<ShortRanks>,
}
impl BytePairRanks {
pub(crate) fn build(map: &Encoder) -> Self {
let mut ranks = vec![u32::MAX; 256 * 256];
for (key, rank) in map {
if let [hi, lo] = key[..] {
ranks[(hi as usize) << 8 | lo as usize] = rank;
}
}
Self {
ranks: ranks.into_boxed_slice(),
short: ShortRanks::build(map),
}
}
#[inline]
fn get(&self, hi: u8, lo: u8) -> u32 {
self.ranks[(hi as usize) << 8 | lo as usize]
}
}
struct ShortRanks {
slots: Box<[u64]>,
mask: usize,
}
impl ShortRanks {
const EMPTY: u64 = u64::MAX;
#[inline]
fn pack_key(key: &[u8]) -> u64 {
let mut bytes = [0u8; 4];
bytes[..key.len()].copy_from_slice(key);
u32::from_le_bytes(bytes) as u64 | ((key.len() - SHORT_MIN) as u64) << 32
}
#[inline]
fn slot_of(&self, packed_key: u64) -> usize {
(packed_key.wrapping_mul(0x9E37_79B9_7F4A_7C15) >> 32) as usize & self.mask
}
fn build(map: &Encoder) -> Option<Self> {
let count = map
.keys()
.filter(|k| (SHORT_MIN..=SHORT_MAX).contains(&k.len()))
.count();
let capacity = (count * 2).next_power_of_two().max(16);
let mut table = Self {
slots: vec![Self::EMPTY; capacity].into_boxed_slice(),
mask: capacity - 1,
};
for (key, rank) in map {
if !(SHORT_MIN..=SHORT_MAX).contains(&key.len()) {
continue;
}
if rank > SHORT_MAX_RANK {
return None;
}
let packed_key = Self::pack_key(key);
let mut slot = table.slot_of(packed_key);
while table.slots[slot] != Self::EMPTY {
slot = (slot + 1) & table.mask;
}
table.slots[slot] = packed_key | (rank as u64) << 33;
}
Some(table)
}
#[inline]
fn get(&self, key: &[u8]) -> u32 {
let packed_key = Self::pack_key(key);
let mut slot = self.slot_of(packed_key);
loop {
let entry = self.slots[slot];
if entry == Self::EMPTY {
return u32::MAX;
}
if entry & 0x1_FFFF_FFFF == packed_key {
return (entry >> 33) as u32;
}
slot = (slot + 1) & self.mask;
}
}
}
#[derive(Clone, Copy)]
pub(crate) struct RankLookup<'a> {
map: &'a Encoder,
pairs: Option<&'a BytePairRanks>,
short: Option<&'a ShortRanks>,
}
impl<'a> RankLookup<'a> {
pub(crate) fn new(map: &'a Encoder) -> Self {
Self {
map,
pairs: None,
short: None,
}
}
pub(crate) fn with_pairs(map: &'a Encoder, pairs: &'a BytePairRanks) -> Self {
Self {
map,
pairs: Some(pairs),
short: pairs.short.as_ref(),
}
}
#[inline]
pub(crate) fn get(&self, key: &[u8]) -> u32 {
if let Some(pairs) = self.pairs {
if let [hi, lo] = key {
return pairs.get(*hi, *lo);
}
}
if let Some(short) = self.short {
if (SHORT_MIN..=SHORT_MAX).contains(&key.len()) {
return short.get(key);
}
}
self.map.get(key).unwrap_or(u32::MAX)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn encoder(entries: &[(&[u8], u32)]) -> Encoder {
entries.iter().copied().collect()
}
#[test]
fn a_trailing_zero_byte_is_not_the_shorter_key() {
let map = encoder(&[(b"abc", 7), (b"abc\0", 9)]);
let short = ShortRanks::build(&map).unwrap();
assert_eq!(short.get(b"abc"), 7);
assert_eq!(short.get(b"abc\0"), 9);
}
#[test]
fn the_table_answers_for_every_short_key_and_only_those() {
let map = encoder(&[(b"ab", 1), (b"xyz", 2), (b"wxyz", 3), (b"abcde", 4)]);
let short = ShortRanks::build(&map).unwrap();
assert_eq!(short.get(b"xyz"), 2);
assert_eq!(short.get(b"wxyz"), 3);
assert_eq!(short.get(b"qqq"), u32::MAX);
}
#[test]
fn an_unpackable_rank_gives_up_on_the_whole_table() {
assert!(ShortRanks::build(&encoder(&[(b"abc", u32::MAX)])).is_none());
}
#[test]
fn the_fronted_lookup_agrees_with_the_map() {
let entries: &[(&[u8], u32)] = &[
(b"ab", 1),
(b"abc", 2),
(b"abcd", 3),
(b"abcde", 4),
(b"abcdef", 5),
];
let map = encoder(entries);
let pairs = BytePairRanks::build(&map);
let fronted = RankLookup::with_pairs(&map, &pairs);
let plain = RankLookup::new(&map);
for (key, rank) in entries {
assert_eq!(fronted.get(key), *rank);
assert_eq!(plain.get(key), *rank);
}
for miss in [&b"zz"[..], b"zzz", b"zzzz", b"zzzzz"] {
assert_eq!(fronted.get(miss), u32::MAX);
assert_eq!(plain.get(miss), u32::MAX);
}
}
}