use kopitiam_core::{Error, Result};
use kopitiam_tensor::Tensor;
struct LayerCache {
k: Option<Tensor>,
v: Option<Tensor>,
}
pub struct KvCache {
layers: Vec<LayerCache>,
max_context: usize,
}
impl KvCache {
pub fn new(n_layers: usize, max_context: usize) -> Self {
let layers = (0..n_layers).map(|_| LayerCache { k: None, v: None }).collect();
Self { layers, max_context }
}
pub fn len(&self) -> usize {
self.layers
.first()
.and_then(|l| l.k.as_ref())
.map(|k| k.shape().dims()[1])
.unwrap_or(0)
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn max_context(&self) -> usize {
self.max_context
}
pub fn clear(&mut self) {
for layer in &mut self.layers {
layer.k = None;
layer.v = None;
}
}
pub(crate) fn append(&mut self, layer: usize, new_k: Tensor, new_v: Tensor) -> Result<(Tensor, Tensor)> {
let new_len = new_k.shape().dims()[1];
let existing_len = self.layers[layer].k.as_ref().map(|k| k.shape().dims()[1]).unwrap_or(0);
let total_len = existing_len + new_len;
if total_len > self.max_context {
return Err(Error::IndexOutOfBounds { dim: 1, index: total_len, len: self.max_context });
}
let (full_k, full_v) = match (&self.layers[layer].k, &self.layers[layer].v) {
(Some(prev_k), Some(prev_v)) => {
(Tensor::concat(&[prev_k.clone(), new_k], 1)?, Tensor::concat(&[prev_v.clone(), new_v], 1)?)
}
_ => (new_k, new_v),
};
self.layers[layer].k = Some(full_k.clone());
self.layers[layer].v = Some(full_v.clone());
Ok((full_k, full_v))
}
}
#[cfg(test)]
mod tests {
use super::*;
fn kv(seq: usize, fill: f32) -> Tensor {
Tensor::from_f32(vec![fill; seq], [1, seq, 1]).unwrap()
}
#[test]
fn a_fresh_cache_is_empty() {
let cache = KvCache::new(2, 128);
assert_eq!(cache.len(), 0);
assert!(cache.is_empty());
}
#[test]
fn append_accumulates_across_calls_and_reports_the_new_length() {
let mut cache = KvCache::new(1, 128);
let (k, _v) = cache.append(0, kv(3, 1.0), kv(3, 2.0)).unwrap();
assert_eq!(k.shape().dims(), &[1, 3, 1]);
assert_eq!(cache.len(), 3);
let (k, v) = cache.append(0, kv(1, 9.0), kv(1, 9.0)).unwrap();
assert_eq!(k.shape().dims(), &[1, 4, 1]);
assert_eq!(cache.len(), 4);
assert_eq!(k.to_vec_f32().unwrap(), vec![1.0, 1.0, 1.0, 9.0]);
assert_eq!(v.to_vec_f32().unwrap(), vec![2.0, 2.0, 2.0, 9.0]);
}
#[test]
fn exceeding_max_context_is_rejected() {
let mut cache = KvCache::new(1, 4);
cache.append(0, kv(4, 1.0), kv(4, 1.0)).unwrap();
assert!(matches!(
cache.append(0, kv(1, 1.0), kv(1, 1.0)),
Err(Error::IndexOutOfBounds { .. })
));
}
#[test]
fn clear_resets_every_layer_to_empty() {
let mut cache = KvCache::new(2, 128);
cache.append(0, kv(3, 1.0), kv(3, 1.0)).unwrap();
cache.append(1, kv(3, 1.0), kv(3, 1.0)).unwrap();
cache.clear();
assert_eq!(cache.len(), 0);
let (k, _) = cache.append(0, kv(2, 1.0), kv(2, 1.0)).unwrap();
assert_eq!(k.shape().dims(), &[1, 2, 1]);
}
#[test]
fn different_layers_are_independent() {
let mut cache = KvCache::new(2, 128);
cache.append(0, kv(3, 1.0), kv(3, 1.0)).unwrap();
assert_eq!(cache.len(), 3); let (k, _) = cache.append(1, kv(2, 5.0), kv(2, 5.0)).unwrap();
assert_eq!(k.shape().dims(), &[1, 2, 1]);
}
}