use std::collections::VecDeque;
#[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 {
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,
lru_order: VecDeque<usize>,
}
impl PagedKVCache {
pub fn new(config: PagedKVCacheConfig) -> Self {
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,
lru_order: VecDeque::new(),
}
}
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_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]
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));
}
}