use burn::tensor::ops::AttentionModuleOptions;
use burn::tensor::{Bool, Device, Int, Tensor, TensorData, activation::softmax, backend::Backend};
use crate::matmul::safe_matmul;
use crate::precision::{to_f32, to_float};
fn flash_enabled() -> bool {
static ENABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ENABLED.get_or_init(|| {
std::env::var("COMBS_ATTN").map(|v| v != "manual").unwrap_or(true)
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CacheKind {
Contiguous,
Paged,
}
#[derive(Debug, Clone, Copy)]
pub struct CacheConfig {
pub max_seq_len: usize,
pub page_size: usize,
pub kind: CacheKind,
pub quantize_kv: bool,
}
impl CacheConfig {
pub const DEFAULT_PAGE_SIZE: usize = 16;
pub fn paged(max_seq_len: usize) -> Self {
CacheConfig {
max_seq_len,
page_size: Self::DEFAULT_PAGE_SIZE,
kind: CacheKind::Paged,
quantize_kv: false,
}
}
pub fn contiguous(max_seq_len: usize) -> Self {
CacheConfig {
max_seq_len,
page_size: Self::DEFAULT_PAGE_SIZE,
kind: CacheKind::Contiguous,
quantize_kv: false,
}
}
pub fn num_pages(&self) -> usize {
self.max_seq_len.div_ceil(self.page_size)
}
}
pub trait KVCache<B: Backend>: Send {
fn attention(
&mut self,
layer: usize,
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
) -> Tensor<B, 4> {
self.attention_opts(layer, q, k, v, pos, scale, None)
}
fn attention_opts(
&mut self,
layer: usize,
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
window: Option<usize>,
) -> Tensor<B, 4>;
fn seq_len(&self) -> usize;
fn popn(&mut self, n: usize) -> usize {
let _ = n;
0
}
fn reset(&mut self);
fn pages_used(&self) -> Option<usize> {
None
}
fn page_stats(&self) -> Option<PageStats> {
None
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PageStats {
pub pages_used: usize,
pub pages_free: usize,
pub num_pages: usize,
pub page_size: usize,
pub seq_len: usize,
pub layers_materialized: usize,
pub layers_total: usize,
pub layers_sliding: usize,
}
fn repeat_kv<B: Backend>(x: Tensor<B, 4>, n_rep: usize) -> Tensor<B, 4> {
if n_rep == 1 {
return x;
}
let [b, nkv, s, d] = x.dims();
x.unsqueeze_dim::<5>(2)
.expand([b, nkv, n_rep, s, d])
.reshape([b, nkv * n_rep, s, d])
}
fn attend<B: Backend>(
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
window: Option<usize>,
) -> Tensor<B, 4> {
let device = q.device();
let out_dtype = q.dtype();
let q = to_f32(q);
let k = to_f32(k);
let v = to_f32(v);
let [_, n_q, seq, d] = q.dims();
let [_, n_kv, total, _] = k.dims();
let n_rep = n_q / n_kv;
let k = repeat_kv(k, n_rep);
let v = repeat_kv(v, n_rep);
let default_scale = 1.0 / (d as f64).sqrt();
if flash_enabled() && window.is_none() && (scale - default_scale).abs() < 1e-12 {
let out = burn::tensor::module::attention(
q,
k,
v,
None,
None,
AttentionModuleOptions {
scale: None,
softcap: None,
is_causal: seq > 1,
},
);
return to_float(out, out_dtype);
}
let scores = q.matmul(k.transpose()).mul_scalar(scale);
let scores = if seq > 1 || window.is_some() {
let q_pos =
Tensor::<B, 1, Int>::arange((pos as i64)..((pos + seq) as i64), &device)
.reshape([seq, 1]);
let k_pos = Tensor::<B, 1, Int>::arange(0..(total as i64), &device).reshape([1, total]);
let mut forbidden: Tensor<B, 2, Bool> = k_pos.clone().greater(q_pos.clone());
if let Some(w) = window {
let too_old = k_pos
.add_scalar(w as i64 - 1)
.lower(q_pos);
forbidden = forbidden.bool_or(too_old);
}
let mask = forbidden
.unsqueeze_dims::<4>(&[0, 1])
.expand([1, n_q, seq, total]);
scores.mask_fill(mask, -1e30f32)
} else {
scores };
to_float(safe_matmul(softmax(scores, 3), v), out_dtype)
}
const KV_QUANT_GROUP: usize = 32;
fn kv_quantize<B: Backend>(x: Tensor<B, 4>) -> (Tensor<B, 4, Int>, Tensor<B, 4>) {
let [b, h, s, d] = x.dims();
debug_assert_eq!(d % KV_QUANT_GROUP, 0);
let groups = d / KV_QUANT_GROUP;
let native = x.dtype();
let g = to_f32(x).reshape([b, h, s, groups, KV_QUANT_GROUP]);
let scale = g
.clone()
.abs()
.max_dim(4) .div_scalar(127.0)
.clamp_min(1e-8);
let q = g
.div(scale.clone().expand([b, h, s, groups, KV_QUANT_GROUP]))
.round()
.clamp(-127.0, 127.0)
.int()
.reshape([b, h, s, d / 4, 4]);
let lane = |i: usize| q.clone().narrow(4, i, 1).reshape([b, h, s, d / 4]);
let packed = lane(0).add_scalar(128)
+ lane(1).add_scalar(128).mul_scalar(256)
+ lane(2).add_scalar(128).mul_scalar(65536)
+ lane(3).mul_scalar(16777216);
(packed, to_float(scale.reshape([b, h, s, groups]), native))
}
fn kv_dequantize<B: Backend>(
packed: Tensor<B, 4, Int>,
scales: Tensor<B, 4>,
d: usize,
) -> Tensor<B, 4> {
let [b, h, s, _] = packed.dims();
let groups = d / KV_QUANT_GROUP;
let t = packed.clone().div_scalar(16777216);
let r3 = packed - t.clone().mul_scalar(16777216);
let neg = r3.clone().lower_elem(0);
let q3 = t.clone().mask_where(neg.clone(), t.sub_scalar(1));
let r = r3.clone().mask_where(neg, r3.add_scalar(16777216));
let l2 = r.clone().div_scalar(65536);
let r = r - l2.clone().mul_scalar(65536);
let l1 = r.clone().div_scalar(256);
let l0 = r - l1.clone().mul_scalar(256);
let q = Tensor::stack::<5>(
vec![
l0.sub_scalar(128),
l1.sub_scalar(128),
l2.sub_scalar(128),
q3,
],
4,
)
.reshape([b, h, s, d]);
let g = q.float().reshape([b, h, s, groups, KV_QUANT_GROUP]);
let scales = to_float(scales, g.dtype());
g.mul(
scales
.reshape([b, h, s, groups, 1])
.expand([b, h, s, groups, KV_QUANT_GROUP]),
)
.reshape([b, h, s, d])
}
pub struct ContiguousKVCache<B: Backend> {
layers: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
seq_len: usize,
}
impl<B: Backend> ContiguousKVCache<B> {
pub fn new(num_layers: usize) -> Self {
ContiguousKVCache {
layers: (0..num_layers).map(|_| None).collect(),
seq_len: 0,
}
}
}
impl<B: Backend> KVCache<B> for ContiguousKVCache<B> {
fn attention_opts(
&mut self,
layer: usize,
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
window: Option<usize>,
) -> Tensor<B, 4> {
let slot = &mut self.layers[layer];
let (k_full, v_full) = match slot.take() {
Some((k_old, v_old)) => (
Tensor::cat(vec![k_old, k], 2),
Tensor::cat(vec![v_old, v], 2),
),
None => (k, v),
};
self.seq_len = k_full.dims()[2];
let out = attend(q, k_full.clone(), v_full.clone(), pos, scale, window);
*slot = Some((k_full, v_full));
out
}
fn seq_len(&self) -> usize {
self.seq_len
}
fn reset(&mut self) {
for slot in &mut self.layers {
*slot = None;
}
self.seq_len = 0;
}
}
#[derive(Debug)]
struct PageAllocator {
free: Vec<usize>,
}
impl PageAllocator {
fn new(num_pages: usize) -> Self {
PageAllocator {
free: (0..num_pages).rev().collect(),
}
}
fn alloc(&mut self) -> Option<usize> {
self.free.pop()
}
fn free_page(&mut self, id: usize) {
self.free.push(id);
}
fn num_free(&self) -> usize {
self.free.len()
}
fn reset(&mut self, num_pages: usize) {
*self = PageAllocator::new(num_pages);
}
}
enum Arena<B: Backend> {
Fp {
k: Tensor<B, 4>,
v: Tensor<B, 4>,
},
Quant {
k_packed: Tensor<B, 4, Int>,
k_scales: Tensor<B, 4>,
v_packed: Tensor<B, 4, Int>,
v_scales: Tensor<B, 4>,
},
}
pub struct PagedKVCache<B: Backend> {
config: CacheConfig,
allocator: PageAllocator,
table: Vec<usize>,
seq_len: usize,
arenas: Vec<Option<Arena<B>>>,
layer_windows: Vec<Option<usize>>,
sliding: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
device: Option<Device<B>>,
}
impl<B: Backend> PagedKVCache<B> {
pub fn new(num_layers: usize, config: CacheConfig) -> Self {
Self::new_with_windows(num_layers, config, vec![None; num_layers])
}
pub fn new_with_windows(
num_layers: usize,
config: CacheConfig,
windows: Vec<Option<usize>>,
) -> Self {
assert_eq!(windows.len(), num_layers, "one window entry per layer");
for w in windows.iter().flatten() {
assert!(*w >= 2, "sliding window must be >= 2, got {w}");
}
PagedKVCache {
allocator: PageAllocator::new(config.num_pages()),
config,
table: Vec::new(),
seq_len: 0,
arenas: (0..num_layers).map(|_| None).collect(),
layer_windows: windows,
sliding: (0..num_layers).map(|_| None).collect(),
device: None,
}
}
pub fn num_free_pages(&self) -> usize {
self.allocator.num_free()
}
pub fn page_stats_inner(&self) -> PageStats {
PageStats {
pages_used: self.table.len(),
pages_free: self.allocator.num_free(),
num_pages: self.config.num_pages(),
page_size: self.config.page_size,
seq_len: self.seq_len,
layers_materialized: self.arenas.iter().filter(|a| a.is_some()).count(),
layers_total: self.arenas.len(),
layers_sliding: self.sliding.iter().filter(|s| s.is_some()).count(),
}
}
fn ensure_pages(&mut self, total: usize) -> usize {
let pages_needed = total.div_ceil(self.config.page_size);
while self.table.len() < pages_needed {
let page = self
.allocator
.alloc()
.expect("page allocator exhausted (max_seq_len exceeded)");
self.table.push(page);
}
pages_needed
}
fn sliding_attention(
&mut self,
layer: usize,
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
w: usize,
) -> Tensor<B, 4> {
let seq = k.dims()[2];
let slot = &mut self.sliding[layer];
let (k_full, v_full) = match slot.take() {
Some((k_old, v_old)) => (
Tensor::cat(vec![k_old, k], 2),
Tensor::cat(vec![v_old, v], 2),
),
None => (k, v),
};
let full_len = k_full.dims()[2];
let kv_offset = pos + seq - full_len;
let out = attend(
q,
k_full.clone(),
v_full.clone(),
pos - kv_offset,
scale,
Some(w),
);
let keep = full_len.min(w - 1);
*slot = Some((
k_full.narrow(2, full_len - keep, keep),
v_full.narrow(2, full_len - keep, keep),
));
out
}
fn page_indices(&self, pages: usize) -> Tensor<B, 1, Int> {
let ids: Vec<i32> = self.table[..pages].iter().map(|&p| p as i32).collect();
let device = self
.device
.as_ref()
.expect("device set on first attention call");
Tensor::<B, 1, Int>::from_data(TensorData::new(ids, [pages]), device)
}
fn gather_window(
&self,
arena: Tensor<B, 4>,
pages: usize,
total: usize,
) -> Tensor<B, 4> {
let [_, n_kv, page_size, last] = arena.dims();
arena
.select(0, self.page_indices(pages)) .swap_dims(0, 1) .reshape([1, n_kv, pages * page_size, last])
.narrow(2, 0, total)
}
fn gather_window_int(
&self,
arena: Tensor<B, 4, Int>,
pages: usize,
total: usize,
) -> Tensor<B, 4, Int> {
let [_, n_kv, page_size, last] = arena.dims();
arena
.select(0, self.page_indices(pages))
.swap_dims(0, 1)
.reshape([1, n_kv, pages * page_size, last])
.narrow(2, 0, total)
}
}
impl<B: Backend> KVCache<B> for PagedKVCache<B> {
fn attention_opts(
&mut self,
layer: usize,
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
window: Option<usize>,
) -> Tensor<B, 4> {
let [_, n_kv, seq, head_dim] = k.dims();
let total = pos + seq;
if layer == 0 {
assert_eq!(
pos, self.seq_len,
"paged cache expects dense contiguous appends (pos == seq_len)"
);
self.seq_len = total;
} else {
debug_assert_eq!(total, self.seq_len);
}
assert!(
total <= self.config.max_seq_len,
"paged cache capacity exceeded: {total} > {}",
self.config.max_seq_len
);
if self.device.is_none() {
self.device = Some(k.device());
}
if let Some(w) = self.layer_windows.get(layer).copied().flatten() {
return self.sliding_attention(layer, q, k, v, pos, scale, w);
}
let quant = self.config.quantize_kv && head_dim % KV_QUANT_GROUP == 0;
if self.arenas[layer].is_none() {
let device = k.device();
let np = self.config.num_pages();
let ps = self.config.page_size;
self.arenas[layer] = Some(if quant {
Arena::Quant {
k_packed: Tensor::zeros([np, n_kv, ps, head_dim / 4], &device),
k_scales: Tensor::zeros([np, n_kv, ps, head_dim / KV_QUANT_GROUP], &device),
v_packed: Tensor::zeros([np, n_kv, ps, head_dim / 4], &device),
v_scales: Tensor::zeros([np, n_kv, ps, head_dim / KV_QUANT_GROUP], &device),
}
} else {
let shape = [np, n_kv, ps, head_dim];
Arena::Fp {
k: Tensor::zeros(shape, &device),
v: Tensor::zeros(shape, &device),
}
});
}
let pages = self.ensure_pages(total);
let page_size = self.config.page_size;
let (k_full, v_full) = match self.arenas[layer].take().expect("arena initialized") {
Arena::Fp { k: mut arena_k, v: mut arena_v } => {
let mut written = 0;
while written < seq {
let global = pos + written;
let slot = global % page_size;
let run = (page_size - slot).min(seq - written);
let phys = self.table[global / page_size];
let range = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..head_dim];
arena_k =
arena_k.slice_assign(range.clone(), k.clone().narrow(2, written, run));
arena_v = arena_v.slice_assign(range, v.clone().narrow(2, written, run));
written += run;
}
let k_full = self.gather_window(arena_k.clone(), pages, total);
let v_full = self.gather_window(arena_v.clone(), pages, total);
self.arenas[layer] = Some(Arena::Fp { k: arena_k, v: arena_v });
(k_full, v_full)
}
Arena::Quant {
mut k_packed,
mut k_scales,
mut v_packed,
mut v_scales,
} => {
let (kq, ks) = kv_quantize(k);
let (vq, vs) = kv_quantize(v);
let dp = head_dim / 4;
let dg = head_dim / KV_QUANT_GROUP;
let mut written = 0;
while written < seq {
let global = pos + written;
let slot = global % page_size;
let run = (page_size - slot).min(seq - written);
let phys = self.table[global / page_size];
let rp = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..dp];
let rs = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..dg];
k_packed =
k_packed.slice_assign(rp.clone(), kq.clone().narrow(2, written, run));
k_scales =
k_scales.slice_assign(rs.clone(), ks.clone().narrow(2, written, run));
v_packed = v_packed.slice_assign(rp, vq.clone().narrow(2, written, run));
v_scales = v_scales.slice_assign(rs, vs.clone().narrow(2, written, run));
written += run;
}
let k_full = kv_dequantize(
self.gather_window_int(k_packed.clone(), pages, total),
self.gather_window(k_scales.clone(), pages, total),
head_dim,
);
let v_full = kv_dequantize(
self.gather_window_int(v_packed.clone(), pages, total),
self.gather_window(v_scales.clone(), pages, total),
head_dim,
);
self.arenas[layer] = Some(Arena::Quant {
k_packed,
k_scales,
v_packed,
v_scales,
});
(k_full, v_full)
}
};
attend(q, k_full, v_full, pos, scale, window)
}
fn seq_len(&self) -> usize {
self.seq_len
}
fn popn(&mut self, n: usize) -> usize {
let n = n.min(self.seq_len);
if n == 0 {
return 0;
}
for w in self.layer_windows.iter().flatten() {
if self.seq_len > w - 1 {
return 0; }
}
for slot in self.sliding.iter_mut() {
if let Some((k, v)) = slot.take() {
let len = k.dims()[2];
let keep = len.saturating_sub(n);
if keep > 0 {
*slot = Some((k.narrow(2, 0, keep), v.narrow(2, 0, keep)));
}
}
}
self.seq_len -= n;
let keep = self.seq_len.div_ceil(self.config.page_size);
while self.table.len() > keep {
let page = self.table.pop().expect("table nonempty");
self.allocator.free_page(page);
}
n
}
fn reset(&mut self) {
self.table.clear();
self.allocator.reset(self.config.num_pages());
self.seq_len = 0;
for slot in &mut self.sliding {
*slot = None;
}
}
fn pages_used(&self) -> Option<usize> {
Some(self.table.len())
}
fn page_stats(&self) -> Option<PageStats> {
Some(self.page_stats_inner())
}
}
#[cfg(test)]
mod tests {
use super::*;
type TB = burn::backend::NdArray<f32>;
fn kv_tok(i: usize, n_kv: usize, d: usize) -> (Tensor<TB, 4>, Tensor<TB, 4>) {
let dev = Default::default();
let mk = |salt: usize| {
let data: Vec<f32> = (0..n_kv * d)
.map(|j| ((i * 7 + j * 3 + salt) % 13) as f32 / 13.0 - 0.5)
.collect();
Tensor::<TB, 4>::from_data(TensorData::new(data, [1, n_kv, 1, d]), &dev)
};
(mk(0), mk(5))
}
fn q_tok(i: usize, n_q: usize, d: usize) -> Tensor<TB, 4> {
let dev = Default::default();
let data: Vec<f32> = (0..n_q * d)
.map(|j| ((i * 11 + j * 5) % 17) as f32 / 17.0 - 0.5)
.collect();
Tensor::<TB, 4>::from_data(TensorData::new(data, [1, n_q, 1, d]), &dev)
}
fn assert_close4(a: Tensor<TB, 4>, b: Tensor<TB, 4>, what: &str) {
let av: Vec<f32> = a.into_data().to_vec().unwrap();
let bv: Vec<f32> = b.into_data().to_vec().unwrap();
assert_eq!(av.len(), bv.len(), "{what}: shape");
for (i, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
assert!((x - y).abs() < 1e-5, "{what}[{i}]: {x} vs {y}");
}
}
#[test]
fn kv_quant_roundtrip_exact_on_grid_values() {
let dev = Default::default();
let d = 64usize;
let data: Vec<f32> = (0..2 * 3 * d)
.map(|i| {
let q = ((i * 37) % 255) as i64 - 127; q as f32 * 0.5
})
.collect();
let mut data = data;
for g in 0..(2 * 3 * d) / 32 {
data[g * 32] = 63.5;
}
let x = Tensor::<TB, 4>::from_data(TensorData::new(data.clone(), [1, 2, 3, d]), &dev);
let (packed, scales) = kv_quantize(x);
let back: Vec<f32> = kv_dequantize(packed, scales, d)
.into_data()
.to_vec()
.unwrap();
for (i, (a, b)) in data.iter().zip(back.iter()).enumerate() {
assert!((a - b).abs() < 1e-6, "[{i}]: {a} vs {b} (must be exact)");
}
}
#[test]
fn kv_quant_error_bounded_by_half_step() {
let dev = Default::default();
let d = 64usize;
let data: Vec<f32> = (0..1 * 2 * 5 * d)
.map(|i| ((i * 7919) % 1000) as f32 / 250.0 - 2.0)
.collect();
let x = Tensor::<TB, 4>::from_data(TensorData::new(data.clone(), [1, 2, 5, d]), &dev);
let (packed, scales) = kv_quantize(x);
let back: Vec<f32> = kv_dequantize(packed, scales, d)
.into_data()
.to_vec()
.unwrap();
for (g, chunk) in data.chunks(32).enumerate() {
let absmax = chunk.iter().fold(0f32, |m, v| m.max(v.abs()));
let half_step = absmax / 254.0 + 1e-6;
for (j, (a, b)) in chunk
.iter()
.zip(back[g * 32..g * 32 + 32].iter())
.enumerate()
{
assert!(
(a - b).abs() <= half_step,
"group {g} elem {j}: |{a} - {b}| > {half_step}"
);
}
}
}
#[test]
fn quantized_paged_cache_matches_fp_within_tolerance() {
let (n_kv, n_q, d) = (2usize, 4usize, 32usize);
let mut cfg_q = CacheConfig::paged(64);
cfg_q.quantize_kv = true;
let cfg_f = CacheConfig::paged(64);
let scale = 1.0 / (d as f64).sqrt() * 0.9;
let mut fp = PagedKVCache::<TB>::new(1, cfg_f);
let mut qn = PagedKVCache::<TB>::new(1, cfg_q);
let ks: Vec<_> = (0..20).map(|i| kv_tok(i, n_kv, d)).collect();
let k20 = Tensor::cat(ks.iter().map(|(k, _)| k.clone()).collect(), 2);
let v20 = Tensor::cat(ks.iter().map(|(_, v)| v.clone()).collect(), 2);
let q20 = Tensor::cat((0..20).map(|i| q_tok(i, n_q, d)).collect(), 2);
let a = fp.attention_opts(0, q20.clone(), k20.clone(), v20.clone(), 0, scale, None);
let b = qn.attention_opts(0, q20, k20, v20, 0, scale, None);
let av: Vec<f32> = a.into_data().to_vec().unwrap();
let bv: Vec<f32> = b.into_data().to_vec().unwrap();
for (i, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
assert!((x - y).abs() < 1e-2, "prefill[{i}]: {x} vs {y}");
}
for i in 20..30 {
let (k, v) = kv_tok(i, n_kv, d);
let q = q_tok(i, n_q, d);
let a = fp.attention_opts(0, q.clone(), k.clone(), v.clone(), i, scale, None);
let b = qn.attention_opts(0, q, k, v, i, scale, None);
let av: Vec<f32> = a.into_data().to_vec().unwrap();
let bv: Vec<f32> = b.into_data().to_vec().unwrap();
for (j, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
assert!((x - y).abs() < 1e-2, "decode {i}[{j}]: {x} vs {y}");
}
}
assert_eq!(qn.pages_used(), fp.pages_used());
assert_eq!(qn.popn(5), 5, "quantized rollback works (no sliding)");
assert_eq!(qn.seq_len(), 25);
}
#[test]
fn sliding_layer_matches_masked_global() {
let (n_kv, n_q, d, w) = (2usize, 4usize, 4usize, 5usize);
let cfg = CacheConfig::paged(64);
let scale = 1.0 / (d as f64).sqrt() * 0.9; let mut global = PagedKVCache::<TB>::new(1, cfg);
let mut sliding = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
let ks: Vec<_> = (0..7).map(|i| kv_tok(i, n_kv, d)).collect();
let k7 = Tensor::cat(ks.iter().map(|(k, _)| k.clone()).collect(), 2);
let v7 = Tensor::cat(ks.iter().map(|(_, v)| v.clone()).collect(), 2);
let q7 = Tensor::cat((0..7).map(|i| q_tok(i, n_q, d)).collect(), 2);
let a = global.attention_opts(0, q7.clone(), k7.clone(), v7.clone(), 0, scale, Some(w));
let b = sliding.attention_opts(0, q7, k7, v7, 0, scale, Some(w));
assert_close4(a, b, "prefill chunk");
for i in 7..14 {
let (k, v) = kv_tok(i, n_kv, d);
let q = q_tok(i, n_q, d);
let a = global.attention_opts(0, q.clone(), k.clone(), v.clone(), i, scale, Some(w));
let b = sliding.attention_opts(0, q, k, v, i, scale, Some(w));
assert_close4(a, b, &format!("decode step {i}"));
}
assert_eq!(sliding.pages_used(), Some(0), "sliding layers use no pages");
assert!(global.pages_used().unwrap() > 0);
}
#[test]
fn sliding_popn_before_eviction_matches_replay() {
let (n_kv, n_q, d, w) = (2usize, 2usize, 4usize, 8usize);
let cfg = CacheConfig::paged(64);
let scale = 0.4;
let mut cache = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
for i in 0..4 {
let (k, v) = kv_tok(i, n_kv, d);
cache.attention_opts(0, q_tok(i, n_q, d), k, v, i, scale, Some(w));
}
assert_eq!(cache.popn(2), 2, "un-evicted rollback succeeds");
assert_eq!(cache.seq_len(), 2);
let mut fresh = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
for i in 0..2 {
let (k, v) = kv_tok(i, n_kv, d);
fresh.attention_opts(0, q_tok(i, n_q, d), k, v, i, scale, Some(w));
}
let (k, v) = kv_tok(9, n_kv, d);
let a = cache.attention_opts(0, q_tok(9, n_q, d), k.clone(), v.clone(), 2, scale, Some(w));
let b = fresh.attention_opts(0, q_tok(9, n_q, d), k, v, 2, scale, Some(w));
assert_close4(a, b, "post-rollback step");
}
#[test]
fn sliding_popn_after_eviction_refuses() {
let (n_kv, n_q, d, w) = (1usize, 1usize, 4usize, 4usize);
let cfg = CacheConfig::paged(64);
let mut cache = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
for i in 0..6 {
let (k, v) = kv_tok(i, n_kv, d);
cache.attention_opts(0, q_tok(i, n_q, d), k, v, i, 0.5, Some(w));
}
assert_eq!(cache.popn(1), 0, "evicted sliding layer refuses rollback");
assert_eq!(cache.seq_len(), 6, "refused rollback leaves state intact");
assert_eq!(cache.popn(0), 0);
}
#[test]
fn allocator_alloc_in_order_and_exhaust() {
let mut a = PageAllocator::new(3);
assert_eq!(a.num_free(), 3);
assert_eq!(a.alloc(), Some(0));
assert_eq!(a.alloc(), Some(1));
assert_eq!(a.alloc(), Some(2));
assert_eq!(a.alloc(), None);
assert_eq!(a.num_free(), 0);
}
#[test]
fn allocator_free_and_realloc_lifo() {
let mut a = PageAllocator::new(2);
let p0 = a.alloc().unwrap();
let p1 = a.alloc().unwrap();
a.free_page(p1);
a.free_page(p0);
assert_eq!(a.num_free(), 2);
assert_eq!(a.alloc(), Some(p0));
assert_eq!(a.alloc(), Some(p1));
}
#[test]
fn allocator_reset_restores_all_pages() {
let mut a = PageAllocator::new(4);
a.alloc();
a.alloc();
a.reset(4);
assert_eq!(a.num_free(), 4);
assert_eq!(a.alloc(), Some(0));
}
#[test]
fn cache_config_num_pages_rounds_up() {
assert_eq!(CacheConfig::paged(16).num_pages(), 1);
assert_eq!(CacheConfig::paged(17).num_pages(), 2);
assert_eq!(CacheConfig::paged(1).num_pages(), 1);
}
type TestBackend = burn::backend::NdArray<f32>;
fn cache(max_seq_len: usize, page_size: usize) -> PagedKVCache<TestBackend> {
PagedKVCache::new(
2,
CacheConfig {
max_seq_len,
page_size,
kind: CacheKind::Paged,
quantize_kv: false,
},
)
}
fn grow(c: &mut PagedKVCache<TestBackend>, total: usize) {
c.ensure_pages(total);
c.seq_len = total;
}
#[test]
fn popn_frees_only_fully_unused_pages() {
let mut c = cache(64, 16);
grow(&mut c, 40); assert_eq!(c.pages_used(), Some(3));
assert_eq!(c.num_free_pages(), 1);
c.popn(9); assert_eq!(c.seq_len(), 31);
assert_eq!(c.pages_used(), Some(2));
assert_eq!(c.num_free_pages(), 2);
c.popn(15); assert_eq!(c.pages_used(), Some(1));
c.popn(1); assert_eq!(c.pages_used(), Some(1));
c.popn(1000); assert_eq!(c.seq_len(), 0);
assert_eq!(c.pages_used(), Some(0));
assert_eq!(c.num_free_pages(), 4);
}
#[test]
fn popn_boundary_exact_page_edge() {
let mut c = cache(64, 16);
grow(&mut c, 32); c.popn(16); assert_eq!(c.pages_used(), Some(1));
assert_eq!(c.num_free_pages(), 3);
c.popn(16);
assert_eq!(c.pages_used(), Some(0));
assert_eq!(c.num_free_pages(), 4);
}
#[test]
fn regrowth_after_popn_reuses_freed_pages() {
let mut c = cache(64, 16);
grow(&mut c, 40);
c.popn(9); grow(&mut c, 33); assert_eq!(c.pages_used(), Some(3));
assert_eq!(c.num_free_pages(), 1);
}
#[test]
fn reset_releases_all_pages() {
let mut c = cache(64, 16);
grow(&mut c, 40);
c.reset();
assert_eq!(c.seq_len(), 0);
assert_eq!(c.pages_used(), Some(0));
assert_eq!(c.num_free_pages(), 4);
}
fn kv_tok_on<B: Backend>(
dev: &B::Device,
i: usize,
n_kv: usize,
d: usize,
) -> (Tensor<B, 4>, Tensor<B, 4>) {
let mk = |salt: usize| {
let data: Vec<f32> = (0..n_kv * d)
.map(|j| ((i * 7 + j * 3 + salt) % 13) as f32 / 13.0 - 0.5)
.collect();
Tensor::<B, 4>::from_data(TensorData::new(data, [1, n_kv, 1, d]), dev)
};
(mk(0), mk(5))
}
fn q_tok_on<B: Backend>(dev: &B::Device, i: usize, n_q: usize, d: usize) -> Tensor<B, 4> {
let data: Vec<f32> = (0..n_q * d)
.map(|j| ((i * 11 + j * 5) % 17) as f32 / 17.0 - 0.5)
.collect();
Tensor::<B, 4>::from_data(TensorData::new(data, [1, n_q, 1, d]), dev)
}
fn quant_roundtrip_on<B: Backend>(dev: &B::Device) {
let d = 64usize;
let data: Vec<f32> = (0..2 * d)
.map(|j| ((j * 5) % 251) as f32 / 251.0 - 0.5)
.collect();
let x = Tensor::<B, 4>::from_data(TensorData::new(data.clone(), [1, 2, 1, d]), dev);
let (packed, scales) = kv_quantize(x);
let y = kv_dequantize(packed, scales, d);
let yv: Vec<f32> = y.into_data().convert::<f32>().to_vec().unwrap();
for (i, (orig, got)) in data.iter().zip(yv.iter()).enumerate() {
assert!(got.is_finite(), "dequant[{i}] not finite: {got}");
assert!(
(orig - got).abs() < 0.01,
"dequant[{i}]: {orig} vs {got}"
);
}
}
fn quant_parity_on<B: Backend>(dev: &B::Device, tol: f32) {
let (n_kv, n_q, d) = (2usize, 4usize, 32usize);
let mut cfg_q = CacheConfig::paged(64);
cfg_q.quantize_kv = true;
let cfg_f = CacheConfig::paged(64);
let scale = 1.0 / (d as f64).sqrt() * 0.9;
let mut fp = PagedKVCache::<B>::new(1, cfg_f);
let mut qn = PagedKVCache::<B>::new(1, cfg_q);
for i in 0..24 {
let (k, v) = kv_tok_on::<B>(dev, i, n_kv, d);
let q = q_tok_on::<B>(dev, i, n_q, d);
let a = fp.attention_opts(0, q.clone(), k.clone(), v.clone(), i, scale, None);
let b = qn.attention_opts(0, q, k, v, i, scale, None);
let av: Vec<f32> = a.into_data().convert::<f32>().to_vec().unwrap();
let bv: Vec<f32> = b.into_data().convert::<f32>().to_vec().unwrap();
for (j, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
assert!(y.is_finite(), "step {i}[{j}] not finite: {y}");
assert!((x - y).abs() < tol, "step {i}[{j}]: {x} vs {y}");
}
}
}
#[test]
#[ignore = "gpu"]
fn kv_quant_roundtrip_on_production_backend() {
quant_roundtrip_on::<combs_core::CombsBackend>(&Default::default());
}
#[test]
#[ignore = "gpu"]
fn quantized_paged_cache_matches_fp_on_production_backend() {
quant_parity_on::<combs_core::CombsBackend>(&Default::default(), 5e-2);
}
}