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, Default, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum CacheType {
#[default]
F32,
Q8,
}
#[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
}
pub(crate) fn try_kv_dim(&self) -> Result<usize, InferenceError> {
self.num_kv_heads.checked_mul(self.head_dim).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"num_kv_heads ({}) * head_dim ({}) overflows 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(crate) fn try_floats_per_page(&self) -> Result<usize, InferenceError> {
let kv_dim = self.try_kv_dim()?;
self.num_layers
.checked_mul(2)
.and_then(|n| n.checked_mul(self.page_size))
.and_then(|n| n.checked_mul(kv_dim))
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"num_layers ({}) * 2 * page_size ({}) * kv_dim ({kv_dim}) overflows usize",
self.num_layers, self.page_size
))
})
}
#[inline]
fn q8_scales_per_page(&self) -> usize {
self.num_layers * 2 * self.page_size
}
fn try_q8_scales_per_page(&self) -> Result<usize, InferenceError> {
self.num_layers
.checked_mul(2)
.and_then(|n| n.checked_mul(self.page_size))
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"num_layers ({}) * 2 * page_size ({}) overflows usize",
self.num_layers, self.page_size
))
})
}
pub fn bytes_per_page(&self) -> usize {
self.floats_per_page() * std::mem::size_of::<f32>()
}
pub fn try_bytes_per_page(&self) -> Result<usize, InferenceError> {
let floats_per_page = self.try_floats_per_page()?;
floats_per_page
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"floats_per_page ({floats_per_page}) * size_of::<f32>() ({}) overflows usize",
std::mem::size_of::<f32>()
))
})
}
pub fn total_bytes(&self) -> usize {
self.max_pages * self.bytes_per_page()
}
pub fn try_total_bytes(&self) -> Result<usize, InferenceError> {
let bytes_per_page = self.try_bytes_per_page()?;
self.max_pages.checked_mul(bytes_per_page).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({}) * bytes_per_page ({bytes_per_page}) overflows usize",
self.max_pages
))
})
}
pub fn max_tokens(&self) -> usize {
self.max_pages * self.page_size
}
pub fn try_max_tokens(&self) -> Result<usize, InferenceError> {
self.max_pages.checked_mul(self.page_size).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({}) * page_size ({}) overflows usize",
self.max_pages, self.page_size
))
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
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 try_new(max_pages: usize, floats_per_page: usize) -> Result<Self, InferenceError> {
let len = max_pages.checked_mul(floats_per_page).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({max_pages}) * floats_per_page ({floats_per_page}) overflows usize"
))
})?;
let byte_len = len.checked_mul(std::mem::size_of::<f32>()).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"page pool byte size ({len} * {}) overflows usize",
std::mem::size_of::<f32>()
))
})?;
if byte_len > isize::MAX as usize {
return Err(InferenceError::InvalidInput(format!(
"page pool byte size ({byte_len}) exceeds isize::MAX — allocation would panic"
)));
}
let free_list_bytes = max_pages
.checked_mul(std::mem::size_of::<usize>())
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({max_pages}) * size_of::<usize>() overflows usize"
))
})?;
if free_list_bytes > isize::MAX as usize {
return Err(InferenceError::InvalidInput(format!(
"free-list allocation ({free_list_bytes} bytes) exceeds isize::MAX — allocation would panic"
)));
}
let data = vec![0.0f32; len];
let free_list: Vec<usize> = (0..max_pages).rev().collect();
Ok(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)]
struct Q8PagePool {
data: Vec<i8>,
scales: Vec<f32>,
free_list: Vec<usize>,
max_pages: usize,
values_per_page: usize,
scales_per_page: usize,
}
impl Q8PagePool {
fn new(max_pages: usize, values_per_page: usize, scales_per_page: usize) -> Self {
Self {
data: vec![0; max_pages * values_per_page],
scales: vec![1.0; max_pages * scales_per_page],
free_list: (0..max_pages).rev().collect(),
max_pages,
values_per_page,
scales_per_page,
}
}
fn try_new(
max_pages: usize,
values_per_page: usize,
scales_per_page: usize,
) -> Result<Self, InferenceError> {
let data_len = max_pages.checked_mul(values_per_page).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({max_pages}) * Q8 values_per_page ({values_per_page}) overflows usize"
))
})?;
if data_len > isize::MAX as usize {
return Err(InferenceError::InvalidInput(format!(
"Q8 page data allocation ({data_len} bytes) exceeds isize::MAX"
)));
}
let scales_len = max_pages.checked_mul(scales_per_page).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({max_pages}) * Q8 scales_per_page ({scales_per_page}) overflows usize"
))
})?;
let scales_bytes = scales_len
.checked_mul(std::mem::size_of::<f32>())
.filter(|&bytes| bytes <= isize::MAX as usize)
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"Q8 page scale allocation ({scales_len} * {}) exceeds isize::MAX",
std::mem::size_of::<f32>()
))
})?;
let _ = scales_bytes;
let free_list_bytes = max_pages
.checked_mul(std::mem::size_of::<usize>())
.filter(|&bytes| bytes <= isize::MAX as usize)
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"Q8 free-list allocation ({max_pages} * {}) exceeds isize::MAX",
std::mem::size_of::<usize>()
))
})?;
let _ = free_list_bytes;
let scale_bytes_per_page = scales_per_page
.checked_mul(std::mem::size_of::<f32>())
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"Q8 scales_per_page ({scales_per_page}) * {} overflows usize",
std::mem::size_of::<f32>()
))
})?;
let bytes_per_page = values_per_page
.checked_add(scale_bytes_per_page)
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"Q8 values_per_page ({values_per_page}) + scale bytes \
({scale_bytes_per_page}) overflows usize"
))
})?;
max_pages.checked_mul(bytes_per_page).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"max_pages ({max_pages}) * Q8 bytes_per_page ({bytes_per_page}) overflows usize"
))
})?;
Ok(Self {
data: vec![0; data_len],
scales: vec![1.0; scales_len],
free_list: (0..max_pages).rev().collect(),
max_pages,
values_per_page,
scales_per_page,
})
}
fn alloc(&mut self) -> Option<usize> {
self.free_list.pop()
}
fn free(&mut self, page_idx: usize) {
debug_assert!(page_idx < self.max_pages);
self.free_list.push(page_idx);
}
fn free_count(&self) -> usize {
self.free_list.len()
}
fn allocated_count(&self) -> usize {
self.max_pages - self.free_list.len()
}
fn store_vector(
&mut self,
page_idx: usize,
value_offset: usize,
scale_offset: usize,
values: &[f32],
) {
let abs_max = values.iter().fold(0.0f32, |max, &value| {
assert!(value.is_finite(), "Q8 KV input must be finite");
max.max(value.abs())
});
let scale = if abs_max == 0.0 {
1.0
} else {
let scale = abs_max / 127.0;
if scale == 0.0 { abs_max } else { scale }
};
let value_base = page_idx * self.values_per_page + value_offset;
let scale_idx = page_idx * self.scales_per_page + scale_offset;
self.scales[scale_idx] = scale;
for (&value, quantized) in values
.iter()
.zip(&mut self.data[value_base..value_base + values.len()])
{
*quantized = (value / scale).round().clamp(-127.0, 127.0) as i8;
}
}
fn gather_vector(
&self,
page_idx: usize,
value_offset: usize,
scale_offset: usize,
dst: &mut [f32],
) {
let value_base = page_idx * self.values_per_page + value_offset;
let scale = self.scales[page_idx * self.scales_per_page + scale_offset];
for (&quantized, value) in self.data[value_base..value_base + dst.len()]
.iter()
.zip(dst)
{
*value = f32::from(quantized) * scale;
}
}
fn clear_page(&mut self, page_idx: usize) {
let value_base = page_idx * self.values_per_page;
self.data[value_base..value_base + self.values_per_page].fill(0);
let scale_base = page_idx * self.scales_per_page;
self.scales[scale_base..scale_base + self.scales_per_page].fill(1.0);
}
fn bytes_per_page(&self) -> usize {
self.values_per_page + self.scales_per_page * std::mem::size_of::<f32>()
}
}
#[derive(Debug)]
enum PagedPagePool {
F32(PagePool),
Q8(Q8PagePool),
}
impl PagedPagePool {
fn cache_type(&self) -> CacheType {
match self {
Self::F32(_) => CacheType::F32,
Self::Q8(_) => CacheType::Q8,
}
}
fn alloc(&mut self) -> Option<usize> {
match self {
Self::F32(pool) => pool.alloc(),
Self::Q8(pool) => pool.alloc(),
}
}
fn free(&mut self, page_idx: usize) {
match self {
Self::F32(pool) => pool.free(page_idx),
Self::Q8(pool) => pool.free(page_idx),
}
}
fn free_count(&self) -> usize {
match self {
Self::F32(pool) => pool.free_count(),
Self::Q8(pool) => pool.free_count(),
}
}
fn allocated_count(&self) -> usize {
match self {
Self::F32(pool) => pool.allocated_count(),
Self::Q8(pool) => pool.allocated_count(),
}
}
fn f32_page(&self, page_idx: usize) -> Result<&[f32], InferenceError> {
match self {
Self::F32(pool) => Ok(pool.page_data(page_idx)),
Self::Q8(_) => Err(InferenceError::PrefixCache(
"prefix sharing requires f32 paged KV storage".into(),
)),
}
}
fn f32_page_mut(&mut self, page_idx: usize) -> Result<&mut [f32], InferenceError> {
match self {
Self::F32(pool) => Ok(pool.page_data_mut(page_idx)),
Self::Q8(_) => Err(InferenceError::PrefixCache(
"prefix sharing requires f32 paged KV storage".into(),
)),
}
}
fn clear_page(&mut self, page_idx: usize) {
match self {
Self::F32(pool) => pool.page_data_mut(page_idx).fill(0.0),
Self::Q8(pool) => pool.clear_page(page_idx),
}
}
fn bytes_per_page(&self) -> usize {
match self {
Self::F32(pool) => pool.floats_per_page * std::mem::size_of::<f32>(),
Self::Q8(pool) => pool.bytes_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: PagedPagePool,
table: PageTable,
config: PagedKVCacheConfig,
prefix_cache: Option<Arc<Mutex<PrefixPageCache>>>,
lru_order: VecDeque<usize>,
}
impl PagedKVCache {
pub fn new(config: PagedKVCacheConfig) -> Self {
Self::with_cache_type(config, CacheType::F32)
}
pub fn try_new(config: PagedKVCacheConfig) -> Result<Self, InferenceError> {
Self::try_with_cache_type(config, CacheType::F32)
}
pub fn with_cache_type(config: PagedKVCacheConfig, cache_type: CacheType) -> Self {
match cache_type {
CacheType::F32 => Self::with_prefix_cache(config, None),
CacheType::Q8 => {
assert!(
config.page_size > 0,
"PagedKVCacheConfig.page_size must be non-zero"
);
let values_per_page = config.floats_per_page();
let scales_per_page = config.q8_scales_per_page();
let pool = PagedPagePool::Q8(Q8PagePool::new(
config.max_pages,
values_per_page,
scales_per_page,
));
let table = PageTable::new(config.page_size);
Self {
pool,
table,
config,
prefix_cache: None,
lru_order: VecDeque::new(),
}
}
}
}
pub fn try_with_cache_type(
config: PagedKVCacheConfig,
cache_type: CacheType,
) -> Result<Self, InferenceError> {
match cache_type {
CacheType::F32 => Self::try_with_prefix_cache(config, None),
CacheType::Q8 => {
if config.page_size == 0 {
return Err(InferenceError::InvalidInput(
"PagedKVCacheConfig.page_size must be non-zero".into(),
));
}
let values_per_page = config.try_floats_per_page()?;
let scales_per_page = config.try_q8_scales_per_page()?;
let pool = PagedPagePool::Q8(Q8PagePool::try_new(
config.max_pages,
values_per_page,
scales_per_page,
)?);
let table = PageTable::new(config.page_size);
Ok(Self {
pool,
table,
config,
prefix_cache: None,
lru_order: VecDeque::new(),
})
}
}
}
pub fn try_with_prefix_cache(
config: PagedKVCacheConfig,
prefix_cache: Option<Arc<Mutex<PrefixPageCache>>>,
) -> Result<Self, InferenceError> {
if config.page_size == 0 {
return Err(InferenceError::InvalidInput(
"PagedKVCacheConfig.page_size must be non-zero".into(),
));
}
let fpp = config.try_floats_per_page()?;
let _total_bytes = config.try_total_bytes()?;
if let Some(ref cache_arc) = prefix_cache {
let guard = cache_arc
.lock()
.map_err(|_| InferenceError::PrefixCache("prefix cache lock poisoned".into()))?;
let pc = guard.config();
let kv_dim = config
.num_kv_heads
.checked_mul(config.head_dim)
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"num_kv_heads ({}) * head_dim ({}) overflows usize",
config.num_kv_heads, config.head_dim
))
})?;
let floats = config
.num_layers
.checked_mul(2)
.and_then(|n| n.checked_mul(pc.prefix_page_size))
.and_then(|n| n.checked_mul(kv_dim))
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"prefix page float count (num_layers={} * 2 * prefix_page_size={} * kv_dim={kv_dim}) overflows usize",
config.num_layers, pc.prefix_page_size
))
})?;
let prefix_page_bytes = floats
.checked_mul(std::mem::size_of::<f32>())
.filter(|&b| b <= isize::MAX as usize)
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"prefix page byte size ({floats} * {}) exceeds isize::MAX",
std::mem::size_of::<f32>()
))
})?;
let _ = prefix_page_bytes;
}
let pool = PagedPagePool::F32(PagePool::try_new(config.max_pages, fpp)?);
let table = PageTable::new(config.page_size);
Ok(Self {
pool,
table,
config,
prefix_cache,
lru_order: VecDeque::new(),
})
}
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 = PagedPagePool::F32(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 || self.num_pages() != 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.f32_page_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.f32_page_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
.checked_mul(2)
.and_then(|n| n.checked_mul(prefix_page_size))
.and_then(|n| n.checked_mul(kv_dim))
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"prefix page float count (num_layers={} * 2 * prefix_page_size={prefix_page_size} * kv_dim={kv_dim}) overflows usize",
self.config.num_layers
))
})?;
let prefix_page_bytes = floats_per_prefix_page
.checked_mul(std::mem::size_of::<f32>())
.filter(|&b| b <= isize::MAX as usize)
.ok_or_else(|| {
InferenceError::InvalidInput(format!(
"prefix page byte size ({floats_per_prefix_page} * {}) exceeds isize::MAX",
std::mem::size_of::<f32>()
))
})?;
let _ = prefix_page_bytes; 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.f32_page(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 cache_type(&self) -> CacheType {
self.pool.cache_type()
}
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;
if self.config.eviction == EvictionPolicy::Lru {
assert!(
needed_pages <= self.config.max_pages,
"PagedKVCache::append_kv_layer: position {pos} needs {needed_pages} pages but \
max_pages is {}; EvictionPolicy::Lru does not yet support sequences beyond \
max_tokens ({}) (see issue #337)",
self.config.max_pages,
self.config.max_tokens(),
);
}
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 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;
match &mut self.pool {
PagedPagePool::F32(pool) => {
let page_data = pool.page_data_mut(phys_page);
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);
}
PagedPagePool::Q8(pool) => {
let scale_layer_stride = 2 * page_size;
let k_scale = layer * scale_layer_stride + offset;
let v_scale = layer * scale_layer_stride + page_size + offset;
pool.store_vector(phys_page, k_base, k_scale, k_token);
pool.store_vector(phys_page, v_base, v_scale, 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]) {
self.gather(layer, false, dst);
}
pub fn gather_v(&self, layer: usize, dst: &mut [f32]) {
self.gather(layer, true, dst);
}
fn gather(&self, layer: usize, value_cache: bool, 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 value_offset = usize::from(value_cache) * page_size;
match &self.pool {
PagedPagePool::F32(pool) => {
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 = pool.page_data(phys_page);
let src_base = layer * layer_stride + (value_offset + 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;
}
}
PagedPagePool::Q8(pool) => {
let layer_stride = 2 * page_size * kv_dim;
let scale_layer_stride = 2 * page_size;
for pos in 0..seq_len {
let (phys_page, offset) = self.table.resolve(pos);
let src_base = layer * layer_stride + (value_offset + offset) * kv_dim;
let scale_offset = layer * scale_layer_stride + value_offset + offset;
let dst_base = pos * kv_dim;
pool.gather_vector(
phys_page,
src_base,
scale_offset,
&mut dst[dst_base..dst_base + kv_dim],
);
}
}
}
}
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.max_pages * self.pool.bytes_per_page()
}
pub fn used_memory_bytes(&self) -> usize {
self.pool.allocated_count() * self.pool.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);
self.pool.clear_page(evicted);
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 restore_prefix_rejects_append_before_advance() {
let config = make_config(4);
let kv_dim = config.kv_dim();
let mut cache = PagedKVCache::new(config);
cache.append_kv_layer(0, &vec![1.0; kv_dim], &vec![2.0; kv_dim]);
assert_eq!(cache.seq_len(), 0);
assert_eq!(cache.num_pages(), 1);
let err = cache
.restore_prefix(AdapterId::BASE, &[1, 2, 3])
.expect_err("partially initialized cache must be rejected");
assert!(
matches!(err, InferenceError::PrefixCache(_)),
"expected PrefixCache error, got {err:?}"
);
}
#[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);
assert_eq!(cache.cache_type(), CacheType::F32);
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_q8_append_gather_roundtrip() {
let config = make_config(4);
let kv_dim = config.kv_dim();
let values_per_page = config.num_layers * 2 * config.page_size * kv_dim;
let scales_per_page = config.num_layers * 2 * config.page_size;
let expected_bytes_per_page = values_per_page + scales_per_page * size_of::<f32>();
let expected_total_bytes = config.max_pages * expected_bytes_per_page;
assert!(expected_total_bytes < config.total_bytes());
let mut cache = PagedKVCache::try_with_cache_type(config, CacheType::Q8)
.expect("valid Q8 config must succeed");
let mut expected_k = vec![Vec::new(); 2];
let mut expected_v = vec![Vec::new(); 2];
for step in 0..6 {
for layer in 0..2 {
let k: Vec<f32> = (0..kv_dim)
.map(|i| {
let magnitude =
0.37 * (step + 1) as f32 + 0.11 * (layer + 1) as f32 + 0.073 * i as f32;
if (step + layer + i) % 2 == 0 {
magnitude
} else {
-magnitude
}
})
.collect();
let v: Vec<f32> = k.iter().map(|value| value * -0.61 + 0.19).collect();
expected_k[layer].extend_from_slice(&k);
expected_v[layer].extend_from_slice(&v);
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
assert_eq!(cache.cache_type(), CacheType::Q8);
assert_eq!(cache.total_memory_bytes(), expected_total_bytes);
assert_eq!(cache.num_pages(), 2);
assert_eq!(cache.used_memory_bytes(), 2 * expected_bytes_per_page);
let mut observed_quantization_error = false;
for layer in 0..2 {
let mut gathered_k = vec![0.0; 6 * kv_dim];
let mut gathered_v = vec![0.0; 6 * kv_dim];
cache.gather_k(layer, &mut gathered_k);
cache.gather_v(layer, &mut gathered_v);
for (expected, actual) in [
(&expected_k[layer], &gathered_k),
(&expected_v[layer], &gathered_v),
] {
for (expected_token, actual_token) in expected
.chunks_exact(kv_dim)
.zip(actual.chunks_exact(kv_dim))
{
let abs_max = expected_token
.iter()
.fold(0.0f32, |max, value| max.max(value.abs()));
let tolerance = abs_max / 254.0 + f32::EPSILON * abs_max;
for (&expected_value, &actual_value) in expected_token.iter().zip(actual_token)
{
let error = (expected_value - actual_value).abs();
assert!(
error <= tolerance,
"Q8 round-trip error {error} exceeds {tolerance}"
);
observed_quantization_error |= error > f32::EPSILON;
}
}
}
}
assert!(observed_quantization_error);
}
#[test]
fn paged_q8_preserves_subnormal_vectors() {
let config = make_config(2);
let kv_dim = config.kv_dim();
let mut cache = PagedKVCache::try_with_cache_type(config, CacheType::Q8)
.expect("valid Q8 config must succeed");
let smallest = f32::from_bits(1);
let mut k = vec![0.0f32; kv_dim];
k[0] = smallest;
k[1] = -smallest;
let v: Vec<f32> = k.iter().map(|value| value * 4.0).collect();
assert!(k.iter().all(|value| value.is_finite()));
assert!(k[0] > 0.0 && k[1] < 0.0);
cache.append_kv_layer(0, &k, &v);
cache.append_kv_layer(1, &k, &v);
cache.advance();
let mut gathered_k = vec![0.0; kv_dim];
let mut gathered_v = vec![0.0; kv_dim];
cache.gather_k(0, &mut gathered_k);
cache.gather_v(0, &mut gathered_v);
assert!(
gathered_k[0] > 0.0,
"subnormal K vanished: {:?}",
&gathered_k[..2]
);
assert!(
gathered_k[1] < 0.0,
"subnormal K sign lost: {:?}",
&gathered_k[..2]
);
assert!(
gathered_v[0] > 0.0,
"subnormal V vanished: {:?}",
&gathered_v[..2]
);
assert!(gathered_k.iter().all(|value| value.is_finite()));
for (expected, actual) in k.iter().zip(&gathered_k) {
assert!((expected - actual).abs() <= smallest);
}
}
#[test]
fn paged_q8_try_rejects_overflowing_page_bytes() {
let config = PagedKVCacheConfig {
page_size: 1,
max_pages: 0,
num_layers: 1,
num_kv_heads: (usize::MAX - 1) / 2,
head_dim: 1,
eviction: EvictionPolicy::None,
};
assert_eq!(config.try_floats_per_page().expect("fits"), usize::MAX - 1);
let err = PagedKVCache::try_with_cache_type(config.clone(), CacheType::Q8)
.expect_err("overflowing per-page byte count must be rejected");
assert!(
matches!(err, InferenceError::InvalidInput(_)),
"expected InvalidInput, got {err:?}"
);
assert!(matches!(
PagedKVCache::try_with_cache_type(config, CacheType::F32),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
#[should_panic(expected = "Q8 KV input must be finite")]
fn paged_q8_rejects_non_finite_input() {
let config = make_config(2);
let kv_dim = config.kv_dim();
let mut cache = PagedKVCache::try_with_cache_type(config, CacheType::Q8)
.expect("valid Q8 config must succeed");
let mut k = vec![0.5f32; kv_dim];
k[kv_dim - 1] = f32::NAN;
cache.append_kv_layer(0, &k, &vec![0.5; kv_dim]);
}
#[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 = "does not yet support sequences beyond max_tokens")]
fn paged_lru_eviction_fails_closed_beyond_max_tokens() {
let mut config = make_config(2); config.eviction = EvictionPolicy::Lru;
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..8 {
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]
fn paged_lru_overflow_leaves_live_pages_unchanged() {
let mut config = make_config(2); config.eviction = EvictionPolicy::Lru;
let kv_dim = config.kv_dim();
let max_tokens = config.max_tokens();
let mut cache = PagedKVCache::new(config);
for step in 0..max_tokens {
for layer in 0..2 {
let k = vec![step as f32; kv_dim];
let v = vec![1000.0 + step as f32; kv_dim];
cache.append_kv_layer(layer, &k, &v);
}
cache.advance();
}
let overflow = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
for layer in 0..2 {
let k = vec![999.0; kv_dim];
let v = vec![1999.0; kv_dim];
cache.append_kv_layer(layer, &k, &v);
}
}));
assert!(overflow.is_err(), "LRU overflow should fail closed");
assert_eq!(cache.seq_len(), max_tokens);
assert_eq!(cache.num_pages(), 2);
let mut k_buf = vec![0.0f32; max_tokens * kv_dim];
let mut v_buf = vec![0.0f32; max_tokens * kv_dim];
cache.gather_k(0, &mut k_buf);
cache.gather_v(0, &mut v_buf);
for step in 0..max_tokens {
for i in 0..kv_dim {
assert_eq!(k_buf[step * kv_dim + i], step as f32);
assert_eq!(v_buf[step * kv_dim + i], 1000.0 + step as f32);
}
}
}
#[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));
}
#[test]
fn paged_try_new_overflow_kv_dim_returns_invalid_input() {
let config = PagedKVCacheConfig {
page_size: 1,
max_pages: 1,
num_layers: 1,
num_kv_heads: usize::MAX,
head_dim: 2,
eviction: EvictionPolicy::None,
};
let r = PagedKVCache::try_new(config);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput on kv_dim overflow, got {r:?}"
);
}
#[test]
fn paged_try_new_overflow_floats_per_page_returns_invalid_input() {
let config = PagedKVCacheConfig {
page_size: 8,
max_pages: 1,
num_layers: usize::MAX / 16 + 2,
num_kv_heads: 1,
head_dim: 1,
eviction: EvictionPolicy::None,
};
let r = PagedKVCache::try_new(config);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput on floats_per_page overflow, got {r:?}"
);
}
#[test]
fn paged_try_new_overflow_total_bytes_returns_invalid_input() {
let config = PagedKVCacheConfig {
page_size: 1,
max_pages: usize::MAX / 8 + 2,
num_layers: 1,
num_kv_heads: 1,
head_dim: 1,
eviction: EvictionPolicy::None,
};
let r = PagedKVCache::try_new(config);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput on total_bytes overflow, got {r:?}"
);
}
#[test]
fn page_pool_try_new_overflow_capacity_returns_invalid_input() {
let r = PagePool::try_new(usize::MAX, 2);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput on page-pool capacity overflow, got {r:?}"
);
}
#[test]
fn page_pool_try_new_overflow_bytes_returns_invalid_input() {
let r = PagePool::try_new(1, usize::MAX / 2);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput on page-pool byte overflow, got {r:?}"
);
}
#[test]
fn paged_try_new_valid_config_succeeds() {
let config = make_config(2);
let cache = PagedKVCache::try_new(config).expect("valid config must succeed");
assert_eq!(cache.seq_len(), 0);
assert_eq!(cache.free_pages(), 2);
}
#[test]
fn page_pool_try_new_isize_max_byte_bound_returns_invalid_input() {
let floats_per_page = (isize::MAX as usize / std::mem::size_of::<f32>()) + 1;
let r = PagePool::try_new(1, floats_per_page);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput when byte size exceeds isize::MAX, got {r:?}"
);
}
#[test]
fn page_pool_try_new_free_list_capacity_bound_returns_invalid_input() {
let max_pages = (isize::MAX as usize / std::mem::size_of::<usize>()) + 1;
let r = PagePool::try_new(max_pages, 0);
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput when free_list byte size exceeds isize::MAX, got {r:?}"
);
}
#[test]
fn try_with_prefix_cache_oversized_prefix_geometry_returns_invalid_input() {
let kv_dim = 1usize; let bad_prefix_page_size = (isize::MAX as usize / (2 * std::mem::size_of::<f32>())) + 1;
let prefix_cache = Arc::new(Mutex::new(PrefixPageCache::new(PrefixPageCacheConfig {
capacity: 1,
prefix_page_size: bad_prefix_page_size,
num_layers: 1,
num_kv_heads: 1,
head_dim: kv_dim,
})));
let config = PagedKVCacheConfig {
page_size: 4,
max_pages: 1,
num_layers: 1,
num_kv_heads: 1,
head_dim: kv_dim,
eviction: EvictionPolicy::None,
};
let r = PagedKVCache::try_with_prefix_cache(config, Some(prefix_cache));
assert!(
matches!(r, Err(InferenceError::InvalidInput(_))),
"expected InvalidInput on oversized prefix geometry, got {r:?}"
);
}
}