use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Mutex, PoisonError};
use kime_tok::layout::CompatSequence;
pub(crate) type Key = [u8; 32];
#[derive(Debug, Clone, PartialEq)]
pub(crate) struct Entry {
pub(crate) logits: Box<[f32]>,
pub(crate) act: [f32; 2],
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CacheMode {
#[default]
Use,
Bypass,
Refresh,
}
impl CacheMode {
#[must_use]
pub fn parse(s: &str) -> Option<CacheMode> {
match s {
"use" => Some(CacheMode::Use),
"bypass" => Some(CacheMode::Bypass),
"refresh" => Some(CacheMode::Refresh),
_ => None,
}
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
pub struct CacheStats {
pub capacity: usize,
pub entries: usize,
pub hits: u64,
pub misses: u64,
}
#[derive(Default)]
struct Generations {
young: HashMap<Key, Entry>,
old: HashMap<Key, Entry>,
}
pub(crate) struct AnswerCache {
half: usize,
maps: Mutex<Generations>,
hits: AtomicU64,
misses: AtomicU64,
}
impl AnswerCache {
pub(crate) fn new(capacity: usize) -> Self {
AnswerCache {
half: capacity.div_ceil(2).max(1),
maps: Mutex::default(),
hits: AtomicU64::new(0),
misses: AtomicU64::new(0),
}
}
pub(crate) fn key(seq: &CompatSequence, qtype: u8) -> Key {
let mut h = blake3::Hasher::new();
h.update(&[qtype]);
h.update(&(seq.ids.len() as u64).to_le_bytes());
for id in &seq.ids {
h.update(&id.to_le_bytes());
}
for m in &seq.markers {
h.update(&m.to_le_bytes());
}
*h.finalize().as_bytes()
}
pub(crate) fn get(&self, keys: &[Key]) -> Vec<Option<Entry>> {
let mut g = self.maps.lock().unwrap_or_else(PoisonError::into_inner);
let out: Vec<Option<Entry>> = keys
.iter()
.map(|k| {
if let Some(e) = g.young.get(k) {
return Some(e.clone());
}
let e = g.old.remove(k)?;
self.put(&mut g, *k, e.clone());
Some(e)
})
.collect();
drop(g);
let hits = out.iter().filter(|e| e.is_some()).count() as u64;
self.hits.fetch_add(hits, Ordering::Relaxed);
self.misses.fetch_add(keys.len() as u64 - hits, Ordering::Relaxed);
out
}
pub(crate) fn all(&self, keys: &[Key]) -> Option<Vec<Entry>> {
let g = self.maps.lock().unwrap_or_else(PoisonError::into_inner);
let out = keys
.iter()
.map(|k| g.young.get(k).or_else(|| g.old.get(k)).cloned())
.collect::<Option<Vec<_>>>()?;
drop(g);
self.hits.fetch_add(keys.len() as u64, Ordering::Relaxed);
Some(out)
}
pub(crate) fn insert(&self, entries: impl IntoIterator<Item = (Key, Entry)>) {
let mut g = self.maps.lock().unwrap_or_else(PoisonError::into_inner);
for (k, e) in entries {
g.old.remove(&k);
self.put(&mut g, k, e);
}
}
fn put(&self, g: &mut Generations, k: Key, e: Entry) {
if g.young.len() >= self.half && !g.young.contains_key(&k) {
g.old = std::mem::take(&mut g.young);
}
g.young.insert(k, e);
}
pub(crate) fn stats(&self) -> CacheStats {
let g = self.maps.lock().unwrap_or_else(PoisonError::into_inner);
CacheStats {
capacity: self.half * 2,
entries: g.young.len() + g.old.len(),
hits: self.hits.load(Ordering::Relaxed),
misses: self.misses.load(Ordering::Relaxed),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn entry(x: f32) -> Entry {
Entry { logits: vec![x].into(), act: [x, 0.0] }
}
#[test]
fn keeps_what_is_used_and_counts() {
let c = AnswerCache::new(4);
let k = |i: u8| [i; 32];
c.insert((0..2).map(|i| (k(i), entry(f32::from(i)))));
c.insert((2..4).map(|i| (k(i), entry(f32::from(i)))));
assert_eq!(c.stats().entries, 4);
assert_eq!(c.get(&[k(0)]), vec![Some(entry(0.0))]);
c.insert([(k(4), entry(4.0))]);
let got = c.get(&[k(0), k(1), k(4), k(9)]);
assert_eq!(got, vec![Some(entry(0.0)), None, Some(entry(4.0)), None]);
let s = c.stats();
assert_eq!((s.hits, s.misses, s.capacity), (3, 2, 4));
assert!(s.entries <= 4, "{s:?}");
}
#[test]
fn insert_overwrites() {
let c = AnswerCache::new(10);
c.insert([([1; 32], entry(1.0))]);
c.insert([([1; 32], entry(2.0))]);
assert_eq!(c.get(&[[1; 32]]), vec![Some(entry(2.0))]);
assert_eq!(c.stats().entries, 1);
}
#[test]
fn modes() {
assert_eq!(CacheMode::parse("bypass"), Some(CacheMode::Bypass));
assert_eq!(CacheMode::parse("refresh"), Some(CacheMode::Refresh));
assert_eq!(CacheMode::parse("use"), Some(CacheMode::Use));
assert_eq!(CacheMode::parse("off"), None);
}
}