use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
#[cfg(test)]
use super::prefix::PrefixPageCacheConfig;
use super::prefix::{AdapterId, PrefixEntry, PrefixKey, PrefixPageCache, SharedPageRef};
use crate::error::InferenceError;
#[derive(Debug, Clone)]
pub struct PagedKVCacheConfig {
pub page_size: usize,
pub max_pages: usize,
pub num_layers: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
pub eviction: EvictionPolicy,
}
impl PagedKVCacheConfig {
#[inline]
pub fn kv_dim(&self) -> usize {
self.num_kv_heads * self.head_dim
}
#[inline]
fn floats_per_page(&self) -> usize {
self.num_layers * 2 * self.page_size * self.kv_dim()
}
pub fn bytes_per_page(&self) -> usize {
self.floats_per_page() * std::mem::size_of::<f32>()
}
pub fn total_bytes(&self) -> usize {
self.max_pages * self.bytes_per_page()
}
pub fn max_tokens(&self) -> usize {
self.max_pages * self.page_size
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EvictionPolicy {
None,
Lru,
}
#[derive(Debug)]
pub struct PagePool {
data: Vec<f32>,
free_list: Vec<usize>,
max_pages: usize,
floats_per_page: usize,
}
impl PagePool {
pub fn new(max_pages: usize, floats_per_page: usize) -> Self {
let data = vec![0.0f32; max_pages * floats_per_page];
let free_list: Vec<usize> = (0..max_pages).rev().collect();
Self {
data,
free_list,
max_pages,
floats_per_page,
}
}
pub fn alloc(&mut self) -> Option<usize> {
self.free_list.pop()
}
pub fn free(&mut self, page_idx: usize) {
debug_assert!(page_idx < self.max_pages);
self.free_list.push(page_idx);
}
pub fn free_count(&self) -> usize {
self.free_list.len()
}
pub fn allocated_count(&self) -> usize {
self.max_pages - self.free_list.len()
}
#[inline]
pub fn page_data(&self, page_idx: usize) -> &[f32] {
let start = page_idx * self.floats_per_page;
&self.data[start..start + self.floats_per_page]
}
#[inline]
pub fn page_data_mut(&mut self, page_idx: usize) -> &mut [f32] {
let start = page_idx * self.floats_per_page;
&mut self.data[start..start + self.floats_per_page]
}
}
#[derive(Debug, Clone)]
pub struct PageTable {
entries: Vec<usize>,
seq_len: usize,
page_size: usize,
}
impl PageTable {
pub fn new(page_size: usize) -> Self {
assert!(page_size > 0, "PageTable page_size must be non-zero");
Self {
entries: Vec::new(),
seq_len: 0,
page_size,
}
}
pub fn num_pages(&self) -> usize {
self.entries.len()
}
pub fn seq_len(&self) -> usize {
self.seq_len
}
#[inline]
pub fn resolve(&self, token_pos: usize) -> (usize, usize) {
let logical = token_pos / self.page_size;
let offset = token_pos % self.page_size;
debug_assert!(logical < self.entries.len(), "token_pos out of range");
(self.entries[logical], offset)
}
pub fn push_page(&mut self, physical_idx: usize) {
self.entries.push(physical_idx);
}
pub fn pop_page(&mut self) -> Option<usize> {
let phys = self.entries.pop()?;
let max = self.entries.len() * self.page_size;
if self.seq_len > max {
self.seq_len = max;
}
Some(phys)
}
pub fn pop_front_page(&mut self) -> Option<usize> {
if self.entries.is_empty() {
return None;
}
let phys = self.entries.remove(0);
let max = self.entries.len() * self.page_size;
if self.seq_len > max {
self.seq_len = max;
}
Some(phys)
}
pub fn set_seq_len(&mut self, len: usize) {
debug_assert!(len <= self.entries.len() * self.page_size);
self.seq_len = len;
}
pub fn physical_pages(&self) -> &[usize] {
&self.entries
}
pub fn clear(&mut self) {
self.entries.clear();
self.seq_len = 0;
}
}
#[derive(Debug)]
pub struct PagedKVCache {
pool: PagePool,
table: PageTable,
config: PagedKVCacheConfig,
prefix_cache: Option<Arc<Mutex<PrefixPageCache>>>,
lru_order: VecDeque<usize>,
}
impl PagedKVCache {
pub fn new(config: PagedKVCacheConfig) -> Self {
Self::with_prefix_cache(config, None)
}
pub fn with_prefix_cache(
config: PagedKVCacheConfig,
prefix_cache: Option<Arc<Mutex<PrefixPageCache>>>,
) -> Self {
assert!(
config.page_size > 0,
"PagedKVCacheConfig.page_size must be non-zero"
);
let fpp = config.floats_per_page();
let pool = PagePool::new(config.max_pages, fpp);
let table = PageTable::new(config.page_size);
Self {
pool,
table,
config,
prefix_cache,
lru_order: VecDeque::new(),
}
}
pub fn restore_prefix(
&mut self,
adapter_id: AdapterId,
token_ids: &[u32],
) -> Result<Option<usize>, InferenceError> {
if self.seq_len() != 0 {
return Err(InferenceError::PrefixCache(
"restore_prefix requires an empty PagedKVCache".into(),
));
}
let key = PrefixKey::from_token_ids(adapter_id, token_ids);
let entry = match &self.prefix_cache {
Some(prefix_cache) => {
let mut guard = prefix_cache.lock().map_err(|_| {
InferenceError::PrefixCache("prefix cache lock poisoned".into())
})?;
guard.lookup(&key)
}
None => None,
};
let Some(entry) = entry else {
return Ok(None);
};
if entry.prefix_len != token_ids.len() {
return Err(InferenceError::PrefixCache(format!(
"prefix hash collision or invalid entry length: key length {}, entry length {}",
token_ids.len(),
entry.prefix_len
)));
}
let restored = self.restore_prefix_entry(&entry)?;
Ok(Some(restored))
}
pub fn promote_to_prefix(
&mut self,
adapter_id: AdapterId,
token_ids: &[u32],
) -> Result<Option<usize>, InferenceError> {
let prefix_len = self.seq_len();
if prefix_len == 0 {
return Ok(None);
}
if token_ids.len() != prefix_len {
return Err(InferenceError::PrefixCache(format!(
"promote_to_prefix token length {} does not match seq_len {}",
token_ids.len(),
prefix_len
)));
}
let Some(prefix_cache) = self.prefix_cache.as_ref().cloned() else {
return Ok(None);
};
let prefix_page_size = {
let guard = prefix_cache
.lock()
.map_err(|_| InferenceError::PrefixCache("prefix cache lock poisoned".into()))?;
guard.config().prefix_page_size
};
let pages = self.copy_owned_pages_to_shared(prefix_page_size)?;
let page_count = pages.len();
let key = PrefixKey::from_token_ids(adapter_id, token_ids);
let mut guard = prefix_cache
.lock()
.map_err(|_| InferenceError::PrefixCache("prefix cache lock poisoned".into()))?;
guard.insert(key, prefix_len, pages);
Ok(Some(page_count))
}
fn restore_prefix_entry(&mut self, entry: &PrefixEntry) -> Result<usize, InferenceError> {
self.validate_prefix_entry(entry)?;
let live_page_count =
PrefixEntry::pages_for_tokens(entry.prefix_len, self.config.page_size);
let mut owned_pages = Vec::with_capacity(live_page_count);
for _ in 0..live_page_count {
let Some(phys) = self.pool.alloc() else {
for allocated in owned_pages {
self.pool.free(allocated);
}
return Err(InferenceError::PrefixCache(format!(
"not enough free pages to restore prefix: needed {}, free {}",
live_page_count,
self.pool.free_count()
)));
};
self.pool.page_data_mut(phys).fill(0.0);
owned_pages.push(phys);
}
for token_pos in 0..entry.prefix_len {
let src_page_idx = token_pos / entry.prefix_page_size;
let src_offset = token_pos % entry.prefix_page_size;
let dst_page_idx = token_pos / self.config.page_size;
let dst_offset = token_pos % self.config.page_size;
let src_page = entry.pages[src_page_idx].as_slice();
let dst_page = self.pool.page_data_mut(owned_pages[dst_page_idx]);
Self::copy_token_between_page_layouts(
src_page,
entry.prefix_page_size,
src_offset,
dst_page,
self.config.page_size,
dst_offset,
self.config.num_layers,
self.config.kv_dim(),
);
}
for phys in owned_pages.iter().copied() {
self.table.push_page(phys);
self.touch_page(phys);
}
self.table.set_seq_len(entry.prefix_len);
Ok(entry.prefix_len)
}
fn copy_owned_pages_to_shared(
&self,
prefix_page_size: usize,
) -> Result<Vec<SharedPageRef>, InferenceError> {
if prefix_page_size == 0 {
return Err(InferenceError::PrefixCache(
"prefix_page_size must be non-zero".into(),
));
}
let prefix_len = self.seq_len();
let prefix_page_count = PrefixEntry::pages_for_tokens(prefix_len, prefix_page_size);
let kv_dim = self.config.kv_dim();
let floats_per_prefix_page = self.config.num_layers * 2 * prefix_page_size * kv_dim;
let mut pages = Vec::with_capacity(prefix_page_count);
for prefix_page_idx in 0..prefix_page_count {
let mut page = vec![0.0f32; floats_per_prefix_page];
let start = prefix_page_idx * prefix_page_size;
let end = (start + prefix_page_size).min(prefix_len);
for token_pos in start..end {
let (src_phys, src_offset) = self.table.resolve(token_pos);
let src_page = self.pool.page_data(src_phys);
let dst_offset = token_pos - start;
Self::copy_token_between_page_layouts(
src_page,
self.config.page_size,
src_offset,
&mut page,
prefix_page_size,
dst_offset,
self.config.num_layers,
kv_dim,
);
}
pages.push(SharedPageRef::from_vec(page));
}
Ok(pages)
}
fn validate_prefix_entry(&self, entry: &PrefixEntry) -> Result<(), InferenceError> {
if entry.prefix_page_size == 0 {
return Err(InferenceError::PrefixCache(
"prefix entry page size must be non-zero".into(),
));
}
let expected_pages =
PrefixEntry::pages_for_tokens(entry.prefix_len, entry.prefix_page_size);
if entry.pages.len() != expected_pages {
return Err(InferenceError::PrefixCache(format!(
"prefix entry page count {} does not match expected {}",
entry.pages.len(),
expected_pages
)));
}
let expected_page_len =
self.config.num_layers * 2 * entry.prefix_page_size * self.config.kv_dim();
for page in &entry.pages {
if page.len() != expected_page_len {
return Err(InferenceError::PrefixCache(format!(
"prefix page has {} floats, expected {}",
page.len(),
expected_page_len
)));
}
}
Ok(())
}
fn copy_token_between_page_layouts(
src_page: &[f32],
src_page_size: usize,
src_offset: usize,
dst_page: &mut [f32],
dst_page_size: usize,
dst_offset: usize,
num_layers: usize,
kv_dim: usize,
) {
let src_layer_stride = 2 * src_page_size * kv_dim;
let dst_layer_stride = 2 * dst_page_size * kv_dim;
for layer in 0..num_layers {
let src_k_base = layer * src_layer_stride + src_offset * kv_dim;
let src_v_base =
layer * src_layer_stride + src_page_size * kv_dim + src_offset * kv_dim;
let dst_k_base = layer * dst_layer_stride + dst_offset * kv_dim;
let dst_v_base =
layer * dst_layer_stride + dst_page_size * kv_dim + dst_offset * kv_dim;
dst_page[dst_k_base..dst_k_base + kv_dim]
.copy_from_slice(&src_page[src_k_base..src_k_base + kv_dim]);
dst_page[dst_v_base..dst_v_base + kv_dim]
.copy_from_slice(&src_page[src_v_base..src_v_base + kv_dim]);
}
}
pub fn seq_len(&self) -> usize {
self.table.seq_len()
}
pub fn max_tokens(&self) -> usize {
self.config.max_tokens()
}
pub fn num_pages(&self) -> usize {
self.table.num_pages()
}
pub fn free_pages(&self) -> usize {
self.pool.free_count()
}
pub fn append_kv_layer(&mut self, layer: usize, k_token: &[f32], v_token: &[f32]) {
let kv_dim = self.config.kv_dim();
assert_eq!(k_token.len(), kv_dim);
assert_eq!(v_token.len(), kv_dim);
assert!(layer < self.config.num_layers);
let pos = self.table.seq_len();
let page_size = self.config.page_size;
let needed_pages = (pos / page_size) + 1;
while self.table.num_pages() < needed_pages {
let phys = self.alloc_page();
self.table.push_page(phys);
}
let (phys_page, offset) = self.table.resolve(pos);
self.touch_page(phys_page);
let page_data = self.pool.page_data_mut(phys_page);
let layer_stride = 2 * page_size * kv_dim;
let k_base = layer * layer_stride + offset * kv_dim;
let v_base = layer * layer_stride + page_size * kv_dim + offset * kv_dim;
page_data[k_base..k_base + kv_dim].copy_from_slice(k_token);
page_data[v_base..v_base + kv_dim].copy_from_slice(v_token);
}
pub fn advance(&mut self) {
let new_len = self.table.seq_len() + 1;
self.table.set_seq_len(new_len);
}
pub fn gather_k(&self, layer: usize, dst: &mut [f32]) {
let seq_len = self.table.seq_len();
let kv_dim = self.config.kv_dim();
let page_size = self.config.page_size;
assert_eq!(dst.len(), seq_len * kv_dim);
let layer_stride = 2 * page_size * kv_dim;
let mut pos = 0usize;
while pos < seq_len {
let (phys_page, offset) = self.table.resolve(pos);
let run_len = (page_size - offset).min(seq_len - pos);
let len = run_len * kv_dim;
let page_data = self.pool.page_data(phys_page);
let src_base = layer * layer_stride + offset * kv_dim;
let dst_base = pos * kv_dim;
dst[dst_base..dst_base + len].copy_from_slice(&page_data[src_base..src_base + len]);
pos += run_len;
}
}
pub fn gather_v(&self, layer: usize, dst: &mut [f32]) {
let seq_len = self.table.seq_len();
let kv_dim = self.config.kv_dim();
let page_size = self.config.page_size;
assert_eq!(dst.len(), seq_len * kv_dim);
let layer_stride = 2 * page_size * kv_dim;
let mut pos = 0usize;
while pos < seq_len {
let (phys_page, offset) = self.table.resolve(pos);
let run_len = (page_size - offset).min(seq_len - pos);
let len = run_len * kv_dim;
let page_data = self.pool.page_data(phys_page);
let src_base = layer * layer_stride + page_size * kv_dim + offset * kv_dim;
let dst_base = pos * kv_dim;
dst[dst_base..dst_base + len].copy_from_slice(&page_data[src_base..src_base + len]);
pos += run_len;
}
}
pub fn reset(&mut self) {
for &phys in self.table.physical_pages() {
self.pool.free(phys);
}
self.table.clear();
self.lru_order.clear();
}
pub fn total_memory_bytes(&self) -> usize {
self.config.total_bytes()
}
pub fn used_memory_bytes(&self) -> usize {
self.pool.allocated_count() * self.config.bytes_per_page()
}
fn alloc_page(&mut self) -> usize {
if let Some(phys) = self.pool.alloc() {
return phys;
}
match self.config.eviction {
EvictionPolicy::None => {
panic!(
"PagePool exhausted ({} pages allocated, eviction=None)",
self.pool.allocated_count()
);
}
EvictionPolicy::Lru => self.evict_lru(),
}
}
fn evict_lru(&mut self) -> usize {
let evicted = self
.lru_order
.pop_front()
.expect("LRU order empty but pool exhausted");
let removed = self
.table
.pop_front_page()
.expect("invariant: page table has an LRU page when pool is exhausted");
debug_assert_eq!(removed, evicted);
let fpp = self.config.floats_per_page();
self.pool.page_data_mut(evicted)[..fpp].fill(0.0);
evicted
}
fn touch_page(&mut self, phys: usize) {
if let Some(pos) = self.lru_order.iter().position(|&p| p == phys) {
self.lru_order.remove(pos);
}
self.lru_order.push_back(phys);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_prefix_cache(capacity: usize) -> Arc<Mutex<PrefixPageCache>> {
Arc::new(Mutex::new(PrefixPageCache::new(PrefixPageCacheConfig {
capacity,
prefix_page_size: 4,
num_layers: 2,
num_kv_heads: 2,
head_dim: 4,
})))
}
#[test]
fn test_prefix_cache_miss_fallthrough() {
let config = make_config(4);
let kv_dim = config.kv_dim();
let prefix_cache = make_prefix_cache(4);
let mut cache = PagedKVCache::with_prefix_cache(config, Some(prefix_cache));
let restored = cache
.restore_prefix(AdapterId::BASE, &[1, 2, 3])
.expect("restore miss should not fail");
assert_eq!(restored, None);
let k = vec![1.0; kv_dim];
let v = vec![2.0; kv_dim];
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
assert_eq!(cache.seq_len(), 1);
assert_eq!(cache.num_pages(), 1);
}
#[test]
fn test_restore_prefix_hit_fast_forwards_seq_len() {
let config = make_config(4);
let kv_dim = config.kv_dim();
let prefix_cache = make_prefix_cache(4);
let tokens: [u32; 3] = [1, 2, 3];
let mut source =
PagedKVCache::with_prefix_cache(config.clone(), Some(Arc::clone(&prefix_cache)));
for step in 0..tokens.len() {
for layer in 0..2 {
let marker = (step * 10 + layer) as f32;
let k = vec![marker; kv_dim];
let v = vec![marker + 0.5; kv_dim];
source.append_kv_layer(layer, &k, &v);
}
source.advance();
}
assert_eq!(
source
.promote_to_prefix(AdapterId::BASE, &tokens)
.expect("promotion should succeed"),
Some(1)
);
let mut restored = PagedKVCache::with_prefix_cache(config, Some(prefix_cache));
assert_eq!(
restored
.restore_prefix(AdapterId::BASE, &tokens)
.expect("restore should succeed"),
Some(tokens.len())
);
assert_eq!(restored.seq_len(), tokens.len());
let mut k_buf = vec![0.0f32; tokens.len() * kv_dim];
restored.gather_k(0, &mut k_buf);
assert_eq!(k_buf[0], 0.0);
assert_eq!(k_buf[kv_dim], 10.0);
assert_eq!(k_buf[2 * kv_dim], 20.0);
}
#[test]
fn test_restore_prefix_with_different_page_sizes() {
let config = PagedKVCacheConfig {
page_size: 8,
max_pages: 4,
num_layers: 2,
num_kv_heads: 2,
head_dim: 4,
eviction: EvictionPolicy::None,
};
let kv_dim = config.kv_dim(); let prefix_cache = Arc::new(Mutex::new(PrefixPageCache::new(PrefixPageCacheConfig {
capacity: 4,
prefix_page_size: 2, num_layers: 2,
num_kv_heads: 2,
head_dim: 4,
})));
let tokens: [u32; 5] = [10, 20, 30, 40, 50];
let mut source =
PagedKVCache::with_prefix_cache(config.clone(), Some(Arc::clone(&prefix_cache)));
for (step, _) in tokens.iter().enumerate() {
for layer in 0..2 {
let k_val = (step * 100 + layer * 10) as f32;
let k = vec![k_val; kv_dim];
let v = vec![k_val + 0.5; kv_dim];
source.append_kv_layer(layer, &k, &v);
}
source.advance();
}
assert_eq!(source.seq_len(), 5);
let page_count = source
.promote_to_prefix(AdapterId::BASE, &tokens)
.expect("promote should succeed")
.expect("promote should insert pages");
assert_eq!(page_count, 3);
let mut restored = PagedKVCache::with_prefix_cache(config.clone(), Some(prefix_cache));
let prefix_len = restored
.restore_prefix(AdapterId::BASE, &tokens)
.expect("restore should succeed")
.expect("restore should hit");
assert_eq!(prefix_len, 5);
assert_eq!(restored.seq_len(), 5);
for layer in 0..2 {
let mut k_buf = vec![0.0f32; tokens.len() * kv_dim];
let mut v_buf = vec![0.0f32; tokens.len() * kv_dim];
restored.gather_k(layer, &mut k_buf);
restored.gather_v(layer, &mut v_buf);
for (step, _) in tokens.iter().enumerate() {
let k_expected = (step * 100 + layer * 10) as f32;
let v_expected = k_expected + 0.5;
for i in 0..kv_dim {
assert_eq!(
k_buf[step * kv_dim + i],
k_expected,
"K mismatch at step={step}, layer={layer}, i={i}"
);
assert_eq!(
v_buf[step * kv_dim + i],
v_expected,
"V mismatch at step={step}, layer={layer}, i={i}"
);
}
}
}
}
fn make_config(max_pages: usize) -> PagedKVCacheConfig {
PagedKVCacheConfig {
page_size: 4, max_pages,
num_layers: 2,
num_kv_heads: 2,
head_dim: 4,
eviction: EvictionPolicy::None,
}
}
#[test]
fn paged_append_gather_roundtrip() {
let config = make_config(4);
let kv_dim = config.kv_dim(); let mut cache = PagedKVCache::new(config);
for step in 0..3u32 {
for layer in 0..2 {
let marker = (step * 10 + layer as u32) as f32;
let k = vec![marker; kv_dim];
let v = vec![marker + 0.5; kv_dim];
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
assert_eq!(cache.seq_len(), 3);
let mut k_buf = vec![0.0f32; 3 * kv_dim];
let mut v_buf = vec![0.0f32; 3 * kv_dim];
cache.gather_k(0, &mut k_buf);
cache.gather_v(0, &mut v_buf);
assert_eq!(k_buf[0], 0.0);
assert_eq!(v_buf[0], 0.5);
assert_eq!(k_buf[kv_dim], 10.0);
assert_eq!(v_buf[kv_dim], 10.5);
assert_eq!(k_buf[2 * kv_dim], 20.0);
assert_eq!(v_buf[2 * kv_dim], 20.5);
cache.gather_k(1, &mut k_buf);
cache.gather_v(1, &mut v_buf);
assert_eq!(k_buf[0], 1.0); assert_eq!(v_buf[0], 1.5);
}
#[test]
fn paged_page_allocation() {
let config = make_config(8);
let kv_dim = config.kv_dim();
let mut cache = PagedKVCache::new(config);
assert_eq!(cache.num_pages(), 0);
assert_eq!(cache.free_pages(), 8);
let k = vec![1.0; kv_dim];
let v = vec![2.0; kv_dim];
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
assert_eq!(cache.num_pages(), 1);
assert_eq!(cache.free_pages(), 7);
for _ in 0..3 {
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
assert_eq!(cache.num_pages(), 1);
assert_eq!(cache.seq_len(), 4);
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
assert_eq!(cache.num_pages(), 2);
assert_eq!(cache.free_pages(), 6);
}
#[test]
fn paged_reset() {
let config = make_config(4);
let kv_dim = config.kv_dim();
let mut cache = PagedKVCache::new(config);
let k = vec![1.0; kv_dim];
let v = vec![2.0; kv_dim];
for _ in 0..6 {
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
assert!(cache.num_pages() > 0);
cache.reset();
assert_eq!(cache.seq_len(), 0);
assert_eq!(cache.num_pages(), 0);
assert_eq!(cache.free_pages(), 4);
}
#[test]
fn paged_cross_page_boundary() {
let config = make_config(4);
let kv_dim = config.kv_dim(); let mut cache = PagedKVCache::new(config);
for step in 0..6u32 {
for layer in 0..2 {
let k: Vec<f32> = (0..kv_dim)
.map(|i| step as f32 * 100.0 + i as f32)
.collect();
let v: Vec<f32> = (0..kv_dim)
.map(|i| step as f32 * 100.0 + i as f32 + 0.5)
.collect();
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
assert_eq!(cache.num_pages(), 2);
let mut k_buf = vec![0.0f32; 6 * kv_dim];
let mut v_buf = vec![0.0f32; 6 * kv_dim];
cache.gather_k(0, &mut k_buf);
cache.gather_v(0, &mut v_buf);
for step in 0..6u32 {
for i in 0..kv_dim {
let k_expected = step as f32 * 100.0 + i as f32;
let v_expected = k_expected + 0.5;
let k_got = k_buf[step as usize * kv_dim + i];
let v_got = v_buf[step as usize * kv_dim + i];
assert!(
(k_got - k_expected).abs() < 1e-6,
"K step={step}, i={i}: expected {k_expected}, got {k_got}"
);
assert!(
(v_got - v_expected).abs() < 1e-6,
"V step={step}, i={i}: expected {v_expected}, got {v_got}"
);
}
}
}
#[test]
fn paged_non_multiple_gather_kv() {
let config = make_config(8);
let kv_dim = config.kv_dim(); let page_size = config.page_size; let seq_len = page_size * 2 + 3; let mut cache = PagedKVCache::new(config);
for step in 0..seq_len as u32 {
for layer in 0..2 {
let k: Vec<f32> = (0..kv_dim)
.map(|i| step as f32 * 10.0 + layer as f32 + i as f32 * 0.1)
.collect();
let v: Vec<f32> = k.iter().map(|&x| x + 0.5).collect();
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
assert_eq!(cache.seq_len(), seq_len);
let mut k_buf = vec![0.0f32; seq_len * kv_dim];
let mut v_buf = vec![0.0f32; seq_len * kv_dim];
cache.gather_k(0, &mut k_buf);
cache.gather_v(0, &mut v_buf);
for step in 0..seq_len as u32 {
for i in 0..kv_dim {
let k_expected = step as f32 * 10.0 + 0.0 + i as f32 * 0.1;
let v_expected = k_expected + 0.5;
let k_got = k_buf[step as usize * kv_dim + i];
let v_got = v_buf[step as usize * kv_dim + i];
assert!(
(k_got - k_expected).abs() < 1e-5,
"K step={step}, i={i}: expected {k_expected}, got {k_got}"
);
assert!(
(v_got - v_expected).abs() < 1e-5,
"V step={step}, i={i}: expected {v_expected}, got {v_got}"
);
}
}
}
#[test]
fn paged_memory_accounting() {
let config = PagedKVCacheConfig {
page_size: 256,
max_pages: 16,
num_layers: 28,
num_kv_heads: 8,
head_dim: 128,
eviction: EvictionPolicy::None,
};
let expected_per_page = 28 * 2 * 256 * 1024 * 4;
assert_eq!(config.bytes_per_page(), expected_per_page);
let cache = PagedKVCache::new(config);
assert_eq!(cache.total_memory_bytes(), 16 * expected_per_page);
assert_eq!(cache.used_memory_bytes(), 0);
}
#[test]
#[should_panic(expected = "PagePool exhausted")]
fn paged_no_eviction_panics_on_exhaustion() {
let config = make_config(1); let kv_dim = config.kv_dim();
let mut cache = PagedKVCache::new(config);
let k = vec![1.0; kv_dim];
let v = vec![2.0; kv_dim];
for _ in 0..4 {
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
for layer in 0..2 {
cache.append_kv_layer(layer, &k, &v);
}
}
#[test]
#[should_panic(expected = "page_size must be non-zero")]
fn paged_zero_page_size_panics_at_construction() {
let mut config = make_config(4);
config.page_size = 0;
let _ = PagedKVCache::new(config);
}
#[test]
#[should_panic(expected = "page_size must be non-zero")]
fn page_table_zero_page_size_panics() {
let _ = PageTable::new(0);
}
#[test]
fn page_pool_alloc_free_cycle() {
let mut pool = PagePool::new(4, 16);
assert_eq!(pool.free_count(), 4);
let p0 = pool.alloc().unwrap();
let _p1 = pool.alloc().unwrap();
assert_eq!(pool.free_count(), 2);
assert_eq!(pool.allocated_count(), 2);
pool.free(p0);
assert_eq!(pool.free_count(), 3);
let p2 = pool.alloc().unwrap();
assert_eq!(p2, p0);
}
#[test]
fn page_table_resolve() {
let mut table = PageTable::new(4);
table.push_page(10); table.push_page(5); table.set_seq_len(6);
assert_eq!(table.resolve(0), (10, 0));
assert_eq!(table.resolve(3), (10, 3));
assert_eq!(table.resolve(4), (5, 0));
assert_eq!(table.resolve(5), (5, 1));
}
}