use std::sync::{Arc, Mutex, RwLock, RwLockReadGuard, RwLockWriteGuard};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct KvPoolExhausted;
impl std::fmt::Display for KvPoolExhausted {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "KV cache block pool exhausted: no free blocks remain")
}
}
impl std::error::Error for KvPoolExhausted {}
pub struct KvBlockPool {
block_size: usize,
total_blocks: usize,
free_blocks: usize,
}
impl KvBlockPool {
pub fn new(block_size: usize, total_blocks: usize) -> Self {
assert!(block_size > 0, "block_size must be positive");
KvBlockPool {
block_size,
total_blocks,
free_blocks: total_blocks,
}
}
pub fn block_size(&self) -> usize {
self.block_size
}
pub fn total_blocks(&self) -> usize {
self.total_blocks
}
pub fn free_blocks(&self) -> usize {
self.free_blocks
}
pub fn resize(&mut self, total_blocks: usize) -> Result<(), usize> {
let in_use = self.total_blocks - self.free_blocks;
if total_blocks < in_use {
return Err(in_use);
}
self.free_blocks = total_blocks - in_use;
self.total_blocks = total_blocks;
Ok(())
}
fn try_acquire(&mut self, n: usize) -> bool {
if n <= self.free_blocks {
self.free_blocks -= n;
true
} else {
false
}
}
fn release(&mut self, n: usize) {
self.free_blocks = (self.free_blocks + n).min(self.total_blocks);
}
}
struct PooledState {
pool: Arc<Mutex<KvBlockPool>>,
block_size: usize,
blocks_held: usize,
}
pub struct KvCache {
pub n_kv_heads: usize,
pub head_dim: usize,
pub k: Vec<f32>, pub v: Vec<f32>,
pub seq_len: usize,
planned_capacity: Option<usize>,
pool_state: Option<PooledState>,
}
impl Clone for KvCache {
fn clone(&self) -> Self {
KvCache {
n_kv_heads: self.n_kv_heads,
head_dim: self.head_dim,
k: self.k.clone(),
v: self.v.clone(),
seq_len: self.seq_len,
planned_capacity: self.planned_capacity,
pool_state: None,
}
}
}
impl Drop for KvCache {
fn drop(&mut self) {
if let Some(state) = &self.pool_state {
if let Ok(mut pool) = state.pool.lock() {
pool.release(state.blocks_held);
}
}
}
}
impl KvCache {
pub fn new(n_kv_heads: usize, head_dim: usize) -> Self {
KvCache {
n_kv_heads,
head_dim,
k: Vec::new(),
v: Vec::new(),
seq_len: 0,
planned_capacity: None,
pool_state: None,
}
}
pub fn with_capacity(n_kv_heads: usize, head_dim: usize, max_seq_len: usize) -> Self {
let elems_per_position = n_kv_heads * head_dim;
KvCache {
n_kv_heads,
head_dim,
k: Vec::with_capacity(max_seq_len * elems_per_position),
v: Vec::with_capacity(max_seq_len * elems_per_position),
seq_len: 0,
planned_capacity: Some(max_seq_len),
pool_state: None,
}
}
pub fn with_pool(
n_kv_heads: usize,
head_dim: usize,
pool: Arc<Mutex<KvBlockPool>>,
max_seq_len: usize,
) -> Result<Self, KvPoolExhausted> {
let block_size = pool.lock().unwrap().block_size();
let blocks_needed = max_seq_len.div_ceil(block_size).max(1);
if !pool.lock().unwrap().try_acquire(blocks_needed) {
return Err(KvPoolExhausted);
}
let elems_per_position = n_kv_heads * head_dim;
Ok(KvCache {
n_kv_heads,
head_dim,
k: Vec::with_capacity(blocks_needed * block_size * elems_per_position),
v: Vec::with_capacity(blocks_needed * block_size * elems_per_position),
seq_len: 0,
planned_capacity: None,
pool_state: Some(PooledState {
pool,
block_size,
blocks_held: blocks_needed,
}),
})
}
pub fn push(&mut self, k_step: &[f32], v_step: &[f32]) -> Result<(), KvPoolExhausted> {
assert_eq!(k_step.len(), self.n_kv_heads * self.head_dim);
assert_eq!(v_step.len(), self.n_kv_heads * self.head_dim);
let elems_per_position = self.n_kv_heads * self.head_dim;
if let Some(state) = &mut self.pool_state {
let capacity_positions = self.k.capacity() / elems_per_position;
if self.seq_len == capacity_positions {
if !state.pool.lock().unwrap().try_acquire(1) {
return Err(KvPoolExhausted);
}
state.blocks_held += 1;
self.k.reserve_exact(state.block_size * elems_per_position);
self.v.reserve_exact(state.block_size * elems_per_position);
}
}
self.k.extend_from_slice(k_step);
self.v.extend_from_slice(v_step);
self.seq_len += 1;
Ok(())
}
pub fn advance_len(&mut self, n: usize) -> Result<(), KvPoolExhausted> {
if n == 0 {
return Ok(());
}
let elems_per_position = self.n_kv_heads * self.head_dim;
let zeros = vec![0f32; elems_per_position];
for _ in 0..n {
self.push(&zeros, &zeros)?;
}
Ok(())
}
pub fn release_to_pool(&mut self) {
if let Some(state) = self.pool_state.take() {
if let Ok(mut pool) = state.pool.lock() {
pool.release(state.blocks_held);
}
}
}
pub fn clear(&mut self) {
self.k.clear();
self.v.clear();
self.seq_len = 0;
}
pub fn truncate(&mut self, new_seq_len: usize) {
assert!(
new_seq_len <= self.seq_len,
"truncate target {new_seq_len} must not exceed current seq_len {}",
self.seq_len
);
let elems_per_position = self.n_kv_heads * self.head_dim;
self.k.truncate(new_seq_len * elems_per_position);
self.v.truncate(new_seq_len * elems_per_position);
self.seq_len = new_seq_len;
}
pub fn allocated_bytes(&self) -> usize {
(self.k.capacity() + self.v.capacity()) * std::mem::size_of::<f32>()
}
pub fn is_within_planned_capacity(&self) -> bool {
match self.planned_capacity {
Some(cap) => {
self.seq_len <= cap
&& self.k.capacity() >= self.seq_len * self.n_kv_heads * self.head_dim
}
None => false,
}
}
}
pub struct PagedKvStore {
block_size: usize,
n_kv_heads: usize,
head_dim: usize,
k: Vec<f32>, v: Vec<f32>,
free_block_ids: Vec<usize>,
}
impl PagedKvStore {
pub fn new(block_size: usize, total_blocks: usize, n_kv_heads: usize, head_dim: usize) -> Self {
assert!(block_size > 0, "block_size must be positive");
let elems_per_block = block_size * n_kv_heads * head_dim;
PagedKvStore {
block_size,
n_kv_heads,
head_dim,
k: vec![0.0; total_blocks * elems_per_block],
v: vec![0.0; total_blocks * elems_per_block],
free_block_ids: (0..total_blocks).rev().collect(),
}
}
pub fn block_size(&self) -> usize {
self.block_size
}
pub fn free_block_count(&self) -> usize {
self.free_block_ids.len()
}
pub fn n_kv_heads(&self) -> usize {
self.n_kv_heads
}
pub fn head_dim(&self) -> usize {
self.head_dim
}
fn acquire_block(&mut self) -> Option<usize> {
self.free_block_ids.pop()
}
fn release_block(&mut self, id: usize) {
self.free_block_ids.push(id);
}
fn elems_per_block(&self) -> usize {
self.block_size * self.n_kv_heads * self.head_dim
}
pub fn k_row(&self, id: usize, offset: usize) -> &[f32] {
let elems_per_position = self.n_kv_heads * self.head_dim;
let start = id * self.elems_per_block() + offset * elems_per_position;
&self.k[start..start + elems_per_position]
}
pub fn v_row(&self, id: usize, offset: usize) -> &[f32] {
let elems_per_position = self.n_kv_heads * self.head_dim;
let start = id * self.elems_per_block() + offset * elems_per_position;
&self.v[start..start + elems_per_position]
}
fn k_row_mut(&mut self, id: usize, offset: usize) -> &mut [f32] {
let elems_per_position = self.n_kv_heads * self.head_dim;
let start = id * self.elems_per_block() + offset * elems_per_position;
&mut self.k[start..start + elems_per_position]
}
fn v_row_mut(&mut self, id: usize, offset: usize) -> &mut [f32] {
let elems_per_position = self.n_kv_heads * self.head_dim;
let start = id * self.elems_per_block() + offset * elems_per_position;
&mut self.v[start..start + elems_per_position]
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PagedStoreExhausted;
impl std::fmt::Display for PagedStoreExhausted {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "paged KV store exhausted: no free blocks remain")
}
}
impl std::error::Error for PagedStoreExhausted {}
#[derive(Debug, Clone, Default)]
pub struct PagedKvCache {
block_table: Vec<usize>,
seq_len: usize,
}
impl PagedKvCache {
pub fn new() -> Self {
PagedKvCache {
block_table: Vec::new(),
seq_len: 0,
}
}
pub fn seq_len(&self) -> usize {
self.seq_len
}
pub fn block_table(&self) -> &[usize] {
&self.block_table
}
pub fn push(
&mut self,
store: &mut PagedKvStore,
k_step: &[f32],
v_step: &[f32],
) -> Result<(), PagedStoreExhausted> {
let block_size = store.block_size();
let offset_in_block = self.seq_len % block_size;
let block_index = self.seq_len / block_size;
if block_index >= self.block_table.len() {
let id = store.acquire_block().ok_or(PagedStoreExhausted)?;
self.block_table.push(id);
}
let block_id = self.block_table[block_index];
store
.k_row_mut(block_id, offset_in_block)
.copy_from_slice(k_step);
store
.v_row_mut(block_id, offset_in_block)
.copy_from_slice(v_step);
self.seq_len += 1;
Ok(())
}
pub fn append_block(&mut self, block_id: usize) {
self.block_table.push(block_id);
}
pub fn release(&mut self, store: &mut PagedKvStore) {
let mut seen: Vec<usize> = Vec::new();
for id in self.block_table.drain(..) {
if !seen.contains(&id) {
seen.push(id);
store.release_block(id);
}
}
self.seq_len = 0;
}
pub fn blocks_needed_for(&self, store: &PagedKvStore, n_new: usize) -> usize {
let held_capacity = self.block_table.len() * store.block_size();
let unused = held_capacity.saturating_sub(self.seq_len);
n_new.saturating_sub(unused).div_ceil(store.block_size())
}
pub fn reserve(
&mut self,
store: &mut PagedKvStore,
n_new: usize,
) -> Result<(), PagedStoreExhausted> {
let need = self.blocks_needed_for(store, n_new);
if need > store.free_block_count() {
return Err(PagedStoreExhausted);
}
for _ in 0..need {
let id = store
.acquire_block()
.expect("checked against free_block_count immediately above");
self.block_table.push(id);
}
Ok(())
}
pub fn adopt_blocks(&mut self, block_table: Vec<usize>, seq_len: usize, block_size: usize) {
assert_eq!(
seq_len % block_size,
0,
"an adopted prefix must end on a block boundary, or the first \
append writes into a block another sequence is reading"
);
assert!(
seq_len / block_size <= block_table.len(),
"block table too short for the adopted length"
);
self.block_table = block_table;
self.seq_len = seq_len;
}
pub fn to_contiguous(&self, store: &PagedKvStore) -> KvCache {
let elems_per_position = store.n_kv_heads * store.head_dim;
let mut cache = KvCache::with_capacity(store.n_kv_heads, store.head_dim, self.seq_len);
cache.k.reserve_exact(self.seq_len * elems_per_position);
cache.v.reserve_exact(self.seq_len * elems_per_position);
for pos in 0..self.seq_len {
let block_id = self.block_table[pos / store.block_size];
let offset = pos % store.block_size;
cache.k.extend_from_slice(store.k_row(block_id, offset));
cache.v.extend_from_slice(store.v_row(block_id, offset));
}
cache.seq_len = self.seq_len;
cache
}
pub fn append_contiguous(
&mut self,
store: &mut PagedKvStore,
k: &[f32],
v: &[f32],
count: usize,
) -> Result<(), PagedStoreExhausted> {
let elems_per_position = store.n_kv_heads * store.head_dim;
assert_eq!(k.len(), count * elems_per_position, "k row count");
assert_eq!(v.len(), count * elems_per_position, "v row count");
if self.blocks_needed_for(store, count) > store.free_block_count() {
return Err(PagedStoreExhausted);
}
for i in 0..count {
let lo = i * elems_per_position;
let hi = lo + elems_per_position;
self.push(store, &k[lo..hi], &v[lo..hi])
.expect("blocks reserved above, so no push here can exhaust the store");
}
Ok(())
}
}
pub struct SharedPagedKv {
layers: Vec<RwLock<PagedKvStore>>,
groups: Mutex<GroupTable>,
}
impl SharedPagedKv {
pub fn new(
n_layers: usize,
block_size: usize,
blocks_per_layer: usize,
n_kv_heads: usize,
head_dim: usize,
) -> Self {
SharedPagedKv {
layers: (0..n_layers)
.map(|_| {
RwLock::new(PagedKvStore::new(
block_size,
blocks_per_layer,
n_kv_heads,
head_dim,
))
})
.collect(),
groups: Mutex::new(GroupTable::default()),
}
}
pub fn from_stores(stores: Vec<PagedKvStore>) -> Self {
SharedPagedKv {
layers: stores.into_iter().map(RwLock::new).collect(),
groups: Mutex::new(GroupTable::default()),
}
}
pub fn layer_count(&self) -> usize {
self.layers.len()
}
pub fn read(&self, layer: usize) -> RwLockReadGuard<'_, PagedKvStore> {
self.layers[layer]
.read()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn write(&self, layer: usize) -> RwLockWriteGuard<'_, PagedKvStore> {
self.layers[layer]
.write()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub fn write_all(&self) -> Vec<RwLockWriteGuard<'_, PagedKvStore>> {
self.layers
.iter()
.map(|l| l.write().unwrap_or_else(|poisoned| poisoned.into_inner()))
.collect()
}
pub fn free_blocks(&self, layer: usize) -> usize {
self.read(layer).free_block_count()
}
pub fn acquire_group(&self) -> Option<PageGroup> {
let mut guards = self.write_all();
if guards.iter().any(|s| s.free_block_count() == 0) {
return None;
}
let blocks: Vec<usize> = guards
.iter_mut()
.map(|s| {
s.acquire_block()
.expect("checked every layer under these same guards")
})
.collect();
let mut groups = self
.groups
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
Some(PageGroup(groups.insert(blocks)))
}
pub fn retain_group(&self, group: PageGroup) {
let mut groups = self
.groups
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
groups.retain(group.0);
}
pub fn release_group(&self, group: PageGroup) -> bool {
let blocks = {
let mut groups = self
.groups
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
match groups.release(group.0) {
Some(blocks) => blocks,
None => return false,
}
};
let mut guards = self.write_all();
for (store, block) in guards.iter_mut().zip(blocks) {
store.release_block(block);
}
true
}
pub fn group_blocks(&self, group: PageGroup) -> Vec<usize> {
let groups = self
.groups
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
groups.blocks(group.0).to_vec()
}
pub fn group_refs(&self, group: PageGroup) -> u32 {
let groups = self
.groups
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
groups.refs(group.0)
}
pub fn free_groups(&self) -> usize {
(0..self.layers.len())
.map(|l| self.free_blocks(l))
.min()
.unwrap_or(0)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct PageGroup(pub u32);
#[derive(Debug, Default)]
struct GroupTable {
blocks: Vec<Option<Vec<usize>>>,
refs: Vec<u32>,
free_ids: Vec<u32>,
}
impl GroupTable {
fn insert(&mut self, blocks: Vec<usize>) -> u32 {
if let Some(id) = self.free_ids.pop() {
self.blocks[id as usize] = Some(blocks);
self.refs[id as usize] = 1;
return id;
}
self.blocks.push(Some(blocks));
self.refs.push(1);
(self.blocks.len() - 1) as u32
}
fn retain(&mut self, id: u32) {
let refs = &mut self.refs[id as usize];
assert!(*refs > 0, "cannot retain group {id}, which has no holders");
*refs += 1;
}
fn release(&mut self, id: u32) -> Option<Vec<usize>> {
let refs = &mut self.refs[id as usize];
assert!(*refs > 0, "double free of group {id}");
*refs -= 1;
if *refs > 0 {
return None;
}
let blocks = self.blocks[id as usize]
.take()
.expect("a group with holders always has blocks");
self.free_ids.push(id);
Some(blocks)
}
fn blocks(&self, id: u32) -> &[usize] {
self.blocks[id as usize]
.as_deref()
.expect("group has no blocks; it was already released")
}
fn refs(&self, id: u32) -> u32 {
self.refs.get(id as usize).copied().unwrap_or(0)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn blocks_needed_for_accounts_for_the_part_full_tail_block() {
let mut store = PagedKvStore::new( 4, 64, 1, 1);
let mut cache = PagedKvCache::new();
let row = [1.0f32];
let advance = |cache: &mut PagedKvCache, store: &mut PagedKvStore, n: usize| {
for _ in 0..n {
cache.push(store, &row, &row).unwrap();
}
};
assert_eq!(cache.blocks_needed_for(&store, 0), 0);
assert_eq!(cache.blocks_needed_for(&store, 1), 1);
assert_eq!(cache.blocks_needed_for(&store, 4), 1);
assert_eq!(cache.blocks_needed_for(&store, 5), 2, "5 into 4s needs 2");
advance(&mut cache, &mut store, 1);
assert_eq!(cache.blocks_needed_for(&store, 3), 0, "fits in the tail");
assert_eq!(cache.blocks_needed_for(&store, 4), 1);
assert_eq!(cache.blocks_needed_for(&store, 8), 2);
advance(&mut cache, &mut store, 2); assert_eq!(cache.blocks_needed_for(&store, 6), 2);
assert_eq!(cache.blocks_needed_for(&store, 5), 1);
advance(&mut cache, &mut store, 1); assert_eq!(cache.blocks_needed_for(&store, 1), 1);
assert_eq!(cache.blocks_needed_for(&store, 4), 1);
cache.reserve(&mut store, 4).unwrap();
assert_eq!(
cache.blocks_needed_for(&store, 4),
0,
"a reserved block is already held"
);
assert_eq!(cache.blocks_needed_for(&store, 5), 1);
}
#[test]
fn a_group_takes_one_block_from_every_layer_and_returns_them_together() {
let kv = SharedPagedKv::new(3, 2, 4, 1, 1);
assert_eq!(kv.free_groups(), 4);
let g = kv.acquire_group().expect("4 groups available");
let blocks = kv.group_blocks(g);
assert_eq!(blocks.len(), 3, "one block per layer");
for l in 0..3 {
assert_eq!(kv.free_blocks(l), 3, "layer {l} gave up exactly one");
}
assert_eq!(kv.free_groups(), 3);
assert!(kv.release_group(g), "sole holder, so this frees it");
for l in 0..3 {
assert_eq!(kv.free_blocks(l), 4, "layer {l} got its block back");
}
assert_eq!(kv.free_groups(), 4);
}
#[test]
fn a_group_shared_by_two_holders_survives_the_first_release() {
let kv = SharedPagedKv::new(2, 2, 2, 1, 1);
let g = kv.acquire_group().unwrap();
let blocks = kv.group_blocks(g);
kv.retain_group(g);
assert_eq!(kv.group_refs(g), 2);
assert!(
!kv.release_group(g),
"one holder remains, so nothing is freed"
);
assert_eq!(kv.group_refs(g), 1);
assert_eq!(kv.free_blocks(0), 1, "the blocks are still held");
assert_eq!(kv.group_blocks(g), blocks, "and still name the same blocks");
assert!(kv.release_group(g), "last holder frees it");
assert_eq!(kv.group_refs(g), 0);
assert_eq!(kv.free_blocks(0), 2);
}
#[test]
fn group_capacity_is_bounded_by_the_layer_with_the_fewest_blocks() {
let kv = SharedPagedKv::from_stores(vec![
PagedKvStore::new(2, 5, 1, 1),
PagedKvStore::new(2, 1, 1, 1),
]);
assert_eq!(kv.free_groups(), 1, "layer 1 has only one block");
let g = kv.acquire_group().expect("one group fits");
assert_eq!(kv.free_groups(), 0);
assert!(
kv.acquire_group().is_none(),
"layer 1 is empty, so no group can be formed"
);
assert_eq!(kv.free_blocks(0), 4, "a refused group leaks nothing");
kv.release_group(g);
assert_eq!(kv.free_blocks(0), 5);
}
#[test]
fn a_released_group_id_is_reused_with_a_fresh_refcount() {
let kv = SharedPagedKv::new(1, 2, 2, 1, 1);
let first = kv.acquire_group().unwrap();
kv.retain_group(first);
assert_eq!(kv.group_refs(first), 2);
kv.release_group(first);
kv.release_group(first);
assert_eq!(kv.group_refs(first), 0, "gone, not merely decremented");
let second = kv.acquire_group().unwrap();
assert_eq!(second, first, "the id is reused");
assert_eq!(
kv.group_refs(second),
1,
"a reused id must not inherit the old count"
);
assert_eq!(kv.group_blocks(second).len(), 1);
assert_eq!(kv.free_blocks(0), 1);
}
#[test]
#[should_panic(expected = "already released")]
fn reading_a_released_group_panics_rather_than_returning_stale_blocks() {
let kv = SharedPagedKv::new(2, 2, 2, 1, 1);
let g = kv.acquire_group().unwrap();
assert!(kv.release_group(g));
let _ = kv.group_blocks(g);
}
#[test]
#[should_panic(expected = "double free of group")]
fn releasing_a_group_twice_panics_rather_than_freeing_it_twice() {
let kv = SharedPagedKv::new(1, 2, 2, 1, 1);
let g = kv.acquire_group().unwrap();
assert!(kv.release_group(g));
kv.release_group(g);
}
#[test]
fn a_recycled_block_backs_a_later_position_without_touching_the_store() {
let mut store = PagedKvStore::new(2, 2, 1, 2);
let mut cache = PagedKvCache::new();
cache.push(&mut store, &[1.0, 1.0], &[1.0, 1.0]).unwrap();
cache.push(&mut store, &[2.0, 2.0], &[2.0, 2.0]).unwrap();
assert_eq!(store.free_block_count(), 1, "one block per position pair");
let recycled = cache.block_table()[0];
cache.append_block(recycled);
assert_eq!(
store.free_block_count(),
1,
"recycling must not take a block from the store"
);
cache.push(&mut store, &[3.0, 3.0], &[3.0, 3.0]).unwrap();
assert_eq!(cache.seq_len(), 3);
assert_eq!(
cache.block_table(),
&[recycled, recycled],
"the same block at the stale index and the live one"
);
let flat = cache.to_contiguous(&store);
assert_eq!(&flat.k[4..6], &[3.0, 3.0], "position 2 reads its own row");
assert_eq!(
&flat.k[0..2],
&[3.0, 3.0],
"position 0 now reads the recycled row, and nothing may read it"
);
}
#[test]
fn releasing_a_recycled_table_gives_each_block_back_once() {
let mut store = PagedKvStore::new(2, 1, 1, 2);
let mut cache = PagedKvCache::new();
cache.push(&mut store, &[1.0, 1.0], &[1.0, 1.0]).unwrap();
let held = cache.block_table()[0];
cache.append_block(held);
cache.append_block(held);
let free_before = store.free_block_count();
cache.release(&mut store);
assert_eq!(
store.free_block_count(),
free_before + 1,
"three table entries naming one block are one block back"
);
assert!(store.acquire_block().is_some());
assert!(store.acquire_block().is_none());
}
#[test]
fn a_gathered_sequence_round_trips_through_the_store() {
let mut store = PagedKvStore::new(2, 8, 2, 2);
let mut cache = PagedKvCache::new();
let rows: Vec<[f32; 4]> = (0..5)
.map(|i| {
let b = i as f32 * 10.0;
[b + 1.0, b + 2.0, b + 3.0, b + 4.0]
})
.collect();
for r in &rows {
cache.push(&mut store, r, r).unwrap();
}
let flat = cache.to_contiguous(&store);
assert_eq!(flat.seq_len, 5);
assert_eq!(flat.k.len(), 5 * 4);
for (i, r) in rows.iter().enumerate() {
assert_eq!(&flat.k[i * 4..(i + 1) * 4], r, "position {i} k");
assert_eq!(&flat.v[i * 4..(i + 1) * 4], r, "position {i} v");
}
let mut rebuilt = PagedKvCache::new();
let mut store2 = PagedKvStore::new(2, 8, 2, 2);
rebuilt
.append_contiguous(&mut store2, &flat.k, &flat.v, 5)
.unwrap();
let again = rebuilt.to_contiguous(&store2);
assert_eq!(again.k, flat.k);
assert_eq!(again.v, flat.v);
assert_eq!(again.seq_len, flat.seq_len);
}
#[test]
fn push_grows_seq_len_and_stores_values() {
let mut cache = KvCache::new(2, 2);
cache
.push(&[1.0, 2.0, 3.0, 4.0], &[5.0, 6.0, 7.0, 8.0])
.unwrap();
assert_eq!(cache.seq_len, 1);
cache
.push(&[9.0, 10.0, 11.0, 12.0], &[13.0, 14.0, 15.0, 16.0])
.unwrap();
assert_eq!(cache.seq_len, 2);
assert_eq!(cache.k.len(), 2 * 2 * 2);
assert_eq!(cache.k[4], 9.0);
}
#[test]
#[should_panic]
fn push_wrong_size_panics() {
let mut cache = KvCache::new(2, 2);
let _ = cache.push(&[1.0, 2.0], &[1.0, 2.0]); }
#[test]
fn clear_resets_state() {
let mut cache = KvCache::new(1, 1);
cache.push(&[1.0], &[2.0]).unwrap();
cache.clear();
assert_eq!(cache.seq_len, 0);
assert!(cache.k.is_empty());
}
#[test]
fn truncate_rolls_back_to_exact_length_preserving_earlier_data() {
let mut cache = KvCache::new(2, 2);
cache
.push(&[1.0, 2.0, 3.0, 4.0], &[10.0, 20.0, 30.0, 40.0])
.unwrap();
cache
.push(&[5.0, 6.0, 7.0, 8.0], &[50.0, 60.0, 70.0, 80.0])
.unwrap();
cache
.push(&[9.0, 9.0, 9.0, 9.0], &[90.0, 90.0, 90.0, 90.0])
.unwrap();
assert_eq!(cache.seq_len, 3);
cache.truncate(1);
assert_eq!(cache.seq_len, 1);
assert_eq!(cache.k, vec![1.0, 2.0, 3.0, 4.0]);
assert_eq!(cache.v, vec![10.0, 20.0, 30.0, 40.0]);
}
#[test]
fn truncate_to_current_length_is_a_no_op() {
let mut cache = KvCache::new(1, 2);
cache.push(&[1.0, 2.0], &[3.0, 4.0]).unwrap();
cache.truncate(1);
assert_eq!(cache.seq_len, 1);
assert_eq!(cache.k, vec![1.0, 2.0]);
}
#[test]
fn truncate_to_zero_empties_the_cache() {
let mut cache = KvCache::new(1, 2);
cache.push(&[1.0, 2.0], &[3.0, 4.0]).unwrap();
cache.truncate(0);
assert_eq!(cache.seq_len, 0);
assert!(cache.k.is_empty());
assert!(cache.v.is_empty());
}
#[test]
#[should_panic]
fn truncate_beyond_current_length_panics() {
let mut cache = KvCache::new(1, 2);
cache.push(&[1.0, 2.0], &[3.0, 4.0]).unwrap();
cache.truncate(5);
}
#[test]
fn push_after_truncate_continues_correctly() {
let mut cache = KvCache::new(1, 1);
cache.push(&[1.0], &[10.0]).unwrap();
cache.push(&[2.0], &[20.0]).unwrap();
cache.push(&[3.0], &[30.0]).unwrap(); cache.truncate(2);
cache.push(&[99.0], &[990.0]).unwrap(); assert_eq!(cache.seq_len, 3);
assert_eq!(cache.k, vec![1.0, 2.0, 99.0]);
assert_eq!(cache.v, vec![10.0, 20.0, 990.0]);
}
#[test]
fn with_capacity_preallocates_and_never_reallocates_within_plan() {
let n_kv_heads = 4;
let head_dim = 8;
let max_seq_len = 16;
let mut cache = KvCache::with_capacity(n_kv_heads, head_dim, max_seq_len);
let expected_elems = max_seq_len * n_kv_heads * head_dim;
assert!(cache.k.capacity() >= expected_elems);
assert!(cache.v.capacity() >= expected_elems);
let step = vec![0.5f32; n_kv_heads * head_dim];
let k_ptr_before = cache.k.as_ptr();
for _ in 0..max_seq_len {
cache.push(&step, &step).unwrap();
}
let k_ptr_after = cache.k.as_ptr();
assert_eq!(
k_ptr_before, k_ptr_after,
"pushing exactly up to the planned capacity must not reallocate"
);
assert!(cache.is_within_planned_capacity());
}
#[test]
fn allocated_bytes_reflects_preallocated_capacity_not_just_used_length() {
let cache = KvCache::with_capacity(4, 8, 100);
let expected_min = 100 * 4 * 8 * 2 * 4;
assert!(
cache.allocated_bytes() >= expected_min,
"allocated_bytes={} expected_min={expected_min}",
cache.allocated_bytes()
);
assert_eq!(cache.seq_len, 0);
}
#[test]
fn grow_as_you_go_cache_reports_not_within_planned_capacity() {
let mut cache = KvCache::new(2, 2);
cache
.push(&[1.0, 2.0, 3.0, 4.0], &[5.0, 6.0, 7.0, 8.0])
.unwrap();
assert!(
!cache.is_within_planned_capacity(),
"a cache built with `new` has no plan to be within"
);
}
#[test]
fn with_pool_acquires_one_block_and_reports_it_in_free_blocks() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(4, 10)));
let cache = KvCache::with_pool(2, 2, pool.clone(), 0).unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 9);
assert_eq!(cache.seq_len, 0);
}
#[test]
fn with_pool_fails_without_mutating_the_pool_when_exhausted() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(4, 0)));
let result = KvCache::with_pool(2, 2, pool.clone(), 0);
assert!(result.is_err());
assert_eq!(pool.lock().unwrap().free_blocks(), 0);
}
#[test]
fn push_acquires_additional_blocks_as_the_cache_crosses_block_boundaries() {
let block_size = 2;
let pool = Arc::new(Mutex::new(KvBlockPool::new(block_size, 10)));
let mut cache = KvCache::with_pool(1, 1, pool.clone(), 0).unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 9);
cache.push(&[1.0], &[1.0]).unwrap();
cache.push(&[2.0], &[2.0]).unwrap();
assert_eq!(
pool.lock().unwrap().free_blocks(),
9,
"filling exactly the first block must not acquire a second one"
);
cache.push(&[3.0], &[3.0]).unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 8);
assert_eq!(cache.seq_len, 3);
assert_eq!(cache.k, vec![1.0, 2.0, 3.0]);
}
#[test]
fn push_returns_pool_exhausted_and_leaves_state_unchanged_when_no_blocks_remain() {
let block_size = 1;
let pool = Arc::new(Mutex::new(KvBlockPool::new(block_size, 1)));
let mut cache = KvCache::with_pool(1, 1, pool.clone(), 0).unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 0);
cache.push(&[1.0], &[1.0]).unwrap();
let before_k = cache.k.clone();
let result = cache.push(&[2.0], &[2.0]);
assert_eq!(result, Err(KvPoolExhausted));
assert_eq!(cache.seq_len, 1, "a failed push must not change seq_len");
assert_eq!(cache.k, before_k, "a failed push must not append data");
}
#[test]
fn dropping_a_pooled_cache_returns_its_blocks_to_the_pool() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(1, 2)));
{
let mut cache = KvCache::with_pool(1, 1, pool.clone(), 0).unwrap();
cache.push(&[1.0], &[1.0]).unwrap(); cache.push(&[2.0], &[2.0]).unwrap(); assert_eq!(pool.lock().unwrap().free_blocks(), 0);
}
assert_eq!(
pool.lock().unwrap().free_blocks(),
2,
"both blocks held by the dropped cache must return to the pool"
);
}
#[test]
fn release_to_pool_is_explicit_and_idempotent() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(4, 5)));
let mut cache = KvCache::with_pool(1, 1, pool.clone(), 0).unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 4);
cache.release_to_pool();
assert_eq!(pool.lock().unwrap().free_blocks(), 5);
cache.release_to_pool(); assert_eq!(pool.lock().unwrap().free_blocks(), 5);
drop(cache); assert_eq!(pool.lock().unwrap().free_blocks(), 5);
}
#[test]
fn two_pooled_caches_share_one_bounded_budget() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(1, 1)));
let cache_a = KvCache::with_pool(1, 1, pool.clone(), 0).unwrap();
let cache_b = KvCache::with_pool(1, 1, pool.clone(), 0);
assert!(
cache_b.is_err(),
"a second concurrent request must not be admitted when the shared budget is full"
);
drop(cache_a);
let cache_c = KvCache::with_pool(1, 1, pool, 0);
assert!(
cache_c.is_ok(),
"once the first request's cache is dropped, its budget must become available again"
);
}
#[test]
fn cloning_a_pooled_cache_detaches_the_clone_from_pool_accounting() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(4, 3)));
let original = KvCache::with_pool(1, 1, pool.clone(), 0).unwrap();
assert_eq!(pool.lock().unwrap().free_blocks(), 2);
let clone = original.clone();
assert_eq!(
pool.lock().unwrap().free_blocks(),
2,
"cloning must not acquire additional blocks"
);
assert_eq!(clone.k, original.k);
drop(clone);
assert_eq!(
pool.lock().unwrap().free_blocks(),
2,
"dropping a detached clone must not release the original's blocks"
);
drop(original);
assert_eq!(
pool.lock().unwrap().free_blocks(),
3,
"dropping the original must release its blocks exactly once"
);
}
#[test]
fn a_pool_refuses_to_shrink_below_what_is_already_held() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(4, 10)));
let held = KvCache::with_pool(2, 4, Arc::clone(&pool), 24).expect("blocks");
let in_use = {
let p = pool.lock().unwrap();
p.total_blocks() - p.free_blocks()
};
assert!(in_use > 0, "the fixture must actually hold blocks");
let mut p = pool.lock().unwrap();
assert_eq!(p.resize(in_use - 1), Err(in_use));
assert_eq!(p.total_blocks(), 10, "a refused resize changes nothing");
assert_eq!(p.free_blocks(), 10 - in_use);
assert_eq!(p.resize(in_use), Ok(()));
assert_eq!(p.free_blocks(), 0);
drop(p);
drop(held);
}
#[test]
fn growing_a_pool_adds_to_what_is_free_and_not_to_what_is_held() {
let pool = Arc::new(Mutex::new(KvBlockPool::new(4, 8)));
let held = KvCache::with_pool(2, 4, Arc::clone(&pool), 16).expect("blocks");
let mut p = pool.lock().unwrap();
let in_use = p.total_blocks() - p.free_blocks();
assert_eq!(p.resize(32), Ok(()));
assert_eq!(p.total_blocks(), 32);
assert_eq!(p.free_blocks(), 32 - in_use);
drop(p);
drop(held);
}
}