use crate::core::token_bytes::Encoder;
use rustc_hash::FxHashMap;
pub(crate) fn merge_ranks<'a>(
merged: &[String],
vocab_in_id_order: impl Iterator<Item = &'a str>,
) -> FxHashMap<Vec<u8>, u32> {
let merge_set: std::collections::HashSet<&str> = merged.iter().map(String::as_str).collect();
let mut ranks: FxHashMap<Vec<u8>, u32> = FxHashMap::default();
for token in vocab_in_id_order.filter(|t| !merge_set.contains(t)) {
let next = ranks.len() as u32;
ranks.entry(token.as_bytes().to_vec()).or_insert(next);
}
let base_count = ranks.len() as u32;
for (i, token) in merged.iter().enumerate() {
ranks
.entry(token.as_bytes().to_vec())
.or_insert(base_count + i as u32);
}
ranks
}
pub(crate) struct BytePairRanks {
ranks: Box<[u32]>,
}
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(),
}
}
#[inline]
fn get(&self, hi: u8, lo: u8) -> u32 {
self.ranks[(hi as usize) << 8 | lo as usize]
}
}
#[derive(Clone, Copy)]
pub(crate) struct RankLookup<'a> {
map: &'a Encoder,
pairs: Option<&'a BytePairRanks>,
}
impl<'a> RankLookup<'a> {
pub(crate) fn new(map: &'a Encoder) -> Self {
Self { map, pairs: None }
}
pub(crate) fn with_pairs(map: &'a Encoder, pairs: &'a BytePairRanks) -> Self {
Self {
map,
pairs: Some(pairs),
}
}
#[inline]
pub(crate) fn get(&self, key: &[u8]) -> u32 {
if let (Some(pairs), [hi, lo]) = (self.pairs, key) {
return pairs.get(*hi, *lo);
}
self.map.get(key).copied().unwrap_or(u32::MAX)
}
}