use std::collections::HashMap;
#[derive(Debug)]
pub(crate) struct BanditEmbedCache {
map: HashMap<u64, Vec<f32>>,
order: std::collections::VecDeque<u64>,
capacity: usize,
}
impl BanditEmbedCache {
pub(crate) fn new(capacity: usize) -> Self {
Self {
map: HashMap::with_capacity(capacity),
order: std::collections::VecDeque::with_capacity(capacity),
capacity,
}
}
pub(crate) fn get(&self, key: u64) -> Option<&Vec<f32>> {
self.map.get(&key)
}
pub(crate) fn insert(&mut self, key: u64, value: Vec<f32>) {
if self.map.contains_key(&key) {
return;
}
if self.map.len() >= self.capacity
&& let Some(evict) = self.order.pop_front()
{
self.map.remove(&evict);
}
self.map.insert(key, value);
self.order.push_back(key);
}
}
impl Default for BanditEmbedCache {
fn default() -> Self {
Self::new(512)
}
}
#[derive(Debug, Default)]
pub(crate) struct TurnEmbedCache {
entries: HashMap<String, Vec<f32>>,
}
impl TurnEmbedCache {
pub(crate) fn get(&self, text: &str) -> Option<&Vec<f32>> {
self.entries.get(text)
}
pub(crate) fn insert(&mut self, text: impl Into<String>, embedding: Vec<f32>) {
self.entries.insert(text.into(), embedding);
}
}