use indexmap::IndexMap;
use rustc_hash::FxHasher;
use std::hash::{BuildHasher, BuildHasherDefault, Hash};
use std::sync::Arc;
type FxIndexMap<K, V> = IndexMap<K, V, BuildHasherDefault<FxHasher>>;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
pub struct AdapterId(pub u64);
impl AdapterId {
pub const BASE: Self = Self(0);
pub const fn new(value: u64) -> Self {
Self(value)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct PrefixKey {
pub adapter_id: AdapterId,
pub token_hash: u64,
}
impl PrefixKey {
pub fn from_token_ids(adapter_id: AdapterId, token_ids: &[u32]) -> Self {
Self {
adapter_id,
token_hash: Self::hash_token_ids(token_ids),
}
}
pub fn hash_token_ids(token_ids: &[u32]) -> u64 {
BuildHasherDefault::<FxHasher>::default().hash_one(token_ids)
}
}
#[derive(Debug, Clone)]
pub struct SharedPageRef {
page: Arc<[f32]>,
}
impl SharedPageRef {
pub fn from_vec(page: Vec<f32>) -> Self {
Self {
page: Arc::from(page.into_boxed_slice()),
}
}
pub fn as_slice(&self) -> &[f32] {
self.page.as_ref()
}
pub fn len(&self) -> usize {
self.page.len()
}
pub fn is_empty(&self) -> bool {
self.page.is_empty()
}
pub fn strong_count(&self) -> usize {
Arc::strong_count(&self.page)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PrefixPageCacheConfig {
pub capacity: usize,
pub prefix_page_size: usize,
pub num_layers: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
}
impl PrefixPageCacheConfig {
pub const DEFAULT_CAPACITY: usize = 128;
pub const DEFAULT_PREFIX_PAGE_SIZE: usize = 64;
pub fn kv_dim(&self) -> usize {
self.num_kv_heads * self.head_dim
}
pub fn floats_per_prefix_page(&self) -> usize {
self.num_layers * 2 * self.prefix_page_size * self.kv_dim()
}
}
#[derive(Debug, Clone)]
pub struct PrefixEntry {
pub prefix_len: usize,
pub prefix_page_size: usize,
pub pages: Vec<SharedPageRef>,
pub last_used: u64,
}
impl PrefixEntry {
pub fn new(
prefix_len: usize,
prefix_page_size: usize,
pages: Vec<SharedPageRef>,
last_used: u64,
) -> Self {
Self {
prefix_len,
prefix_page_size,
pages,
last_used,
}
}
pub fn pages_for_tokens(token_count: usize, page_size: usize) -> usize {
assert!(page_size > 0, "page_size must be non-zero");
if token_count == 0 {
0
} else {
((token_count - 1) / page_size) + 1
}
}
}
#[derive(Debug)]
pub struct PrefixPageCache {
config: PrefixPageCacheConfig,
entries: FxIndexMap<PrefixKey, PrefixEntry>,
clock: u64,
}
impl PrefixPageCache {
pub fn new(config: PrefixPageCacheConfig) -> Self {
assert!(
config.prefix_page_size > 0,
"prefix_page_size must be non-zero"
);
Self {
config,
entries: IndexMap::with_capacity_and_hasher(
config.capacity,
BuildHasherDefault::<FxHasher>::default(),
),
clock: 0,
}
}
pub fn config(&self) -> PrefixPageCacheConfig {
self.config
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn contains_key(&self, key: &PrefixKey) -> bool {
self.entries.contains_key(key)
}
pub fn lookup(&mut self, key: &PrefixKey) -> Option<PrefixEntry> {
let last_used = self.next_clock();
let mut entry = self.entries.swap_remove(key)?;
entry.last_used = last_used;
let cloned = entry.clone();
self.entries.insert(*key, entry);
Some(cloned)
}
pub fn insert(
&mut self,
key: PrefixKey,
prefix_len: usize,
pages: Vec<SharedPageRef>,
) -> Option<PrefixEntry> {
let entry = PrefixEntry::new(prefix_len, self.config.prefix_page_size, pages, 0);
self.insert_entry(key, entry)
}
pub fn insert_entry(&mut self, key: PrefixKey, mut entry: PrefixEntry) -> Option<PrefixEntry> {
self.validate_entry(&entry);
let last_used = self.next_clock();
entry.last_used = last_used;
let replaced = self.entries.swap_remove(&key);
self.entries.insert(key, entry);
self.evict_until_within_capacity();
replaced
}
pub fn evict_lru(&mut self) -> usize {
if self.entries.is_empty() {
return 0;
}
let (_key, entry) = self.entries.shift_remove_index(0).expect("non-empty");
entry
.pages
.iter()
.filter(|page| page.strong_count() == 1)
.count()
}
pub fn clear(&mut self) {
self.entries.clear();
self.clock = 0;
}
fn next_clock(&mut self) -> u64 {
self.clock = self.clock.wrapping_add(1);
self.clock
}
fn evict_until_within_capacity(&mut self) {
while self.entries.len() > self.config.capacity {
let before = self.entries.len();
let _ = self.evict_lru();
if self.entries.len() == before {
break;
}
}
}
fn validate_entry(&self, entry: &PrefixEntry) {
assert_eq!(
entry.prefix_page_size, self.config.prefix_page_size,
"prefix entry page size must match prefix cache config"
);
let expected_pages =
PrefixEntry::pages_for_tokens(entry.prefix_len, entry.prefix_page_size);
assert_eq!(
entry.pages.len(),
expected_pages,
"prefix entry page count does not match prefix length"
);
let expected_len = self.config.floats_per_prefix_page();
for page in &entry.pages {
assert_eq!(
page.len(),
expected_len,
"prefix page length does not match prefix cache geometry"
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_config(capacity: usize) -> PrefixPageCacheConfig {
PrefixPageCacheConfig {
capacity,
prefix_page_size: 4,
num_layers: 2,
num_kv_heads: 2,
head_dim: 4,
}
}
fn make_page(config: PrefixPageCacheConfig, marker: f32) -> SharedPageRef {
SharedPageRef::from_vec(vec![marker; config.floats_per_prefix_page()])
}
fn make_pages(
config: PrefixPageCacheConfig,
prefix_len: usize,
marker: f32,
) -> Vec<SharedPageRef> {
let count = PrefixEntry::pages_for_tokens(prefix_len, config.prefix_page_size);
(0..count)
.map(|idx| make_page(config, marker + idx as f32))
.collect()
}
#[test]
fn test_empty_cache_lookup() {
let config = make_config(4);
let mut cache = PrefixPageCache::new(config);
let key = PrefixKey::from_token_ids(AdapterId::BASE, &[1, 2, 3]);
assert!(cache.lookup(&key).is_none());
assert!(cache.is_empty());
}
#[test]
fn test_prefix_lookup_hit() {
let config = make_config(4);
let mut cache = PrefixPageCache::new(config);
let key = PrefixKey::from_token_ids(AdapterId::BASE, &[1, 2, 3, 4]);
let pages = make_pages(config, 4, 7.0);
cache.insert(key, 4, pages);
let entry = cache.lookup(&key).expect("prefix should hit");
assert_eq!(entry.prefix_len, 4);
assert_eq!(entry.pages.len(), 1);
assert_eq!(entry.pages[0].as_slice()[0], 7.0);
assert!(entry.last_used > 0);
}
#[test]
fn test_prefix_cache_miss() {
let config = make_config(4);
let mut cache = PrefixPageCache::new(config);
let key = PrefixKey::from_token_ids(AdapterId::BASE, &[99, 100]);
assert!(cache.lookup(&key).is_none());
}
#[test]
fn test_evict_lru_reclaims_unreferenced() {
let config = make_config(2);
let mut cache = PrefixPageCache::new(config);
let key_a = PrefixKey::from_token_ids(AdapterId::BASE, &[1, 2, 3, 4]);
let key_b = PrefixKey::from_token_ids(AdapterId::BASE, &[5, 6, 7, 8]);
cache.insert(key_a, 4, make_pages(config, 4, 1.0));
cache.insert(key_b, 4, make_pages(config, 4, 2.0));
let _ = cache
.lookup(&key_a)
.expect("key_a should hit and become MRU");
let pages_freed = cache.evict_lru();
assert_eq!(pages_freed, 1);
assert!(cache.contains_key(&key_a));
assert!(!cache.contains_key(&key_b));
}
#[test]
fn test_adapter_keying_separates_entries() {
let config = make_config(4);
let mut cache = PrefixPageCache::new(config);
let tokens = [1, 2, 3, 4];
let base_key = PrefixKey::from_token_ids(AdapterId::BASE, &tokens);
let adapter_key = PrefixKey::from_token_ids(AdapterId::new(42), &tokens);
cache.insert(base_key, 4, make_pages(config, 4, 1.0));
assert!(cache.lookup(&base_key).is_some());
assert!(cache.lookup(&adapter_key).is_none());
}
#[test]
fn test_insert_at_capacity() {
let config = make_config(1);
let mut cache = PrefixPageCache::new(config);
let key_a = PrefixKey::from_token_ids(AdapterId::BASE, &[1, 2, 3, 4]);
let key_b = PrefixKey::from_token_ids(AdapterId::BASE, &[9, 8, 7, 6]);
cache.insert(key_a, 4, make_pages(config, 4, 1.0));
cache.insert(key_b, 4, make_pages(config, 4, 2.0));
assert_eq!(cache.len(), 1);
assert!(!cache.contains_key(&key_a));
assert!(cache.contains_key(&key_b));
}
}