use burn::tensor::ops::AttentionModuleOptions;
use burn::tensor::{Bool, Device, Int, Tensor, TensorData, activation::softmax, backend::Backend};
use crate::matmul::safe_matmul;
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,
}
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,
}
}
pub fn contiguous(max_seq_len: usize) -> Self {
CacheConfig {
max_seq_len,
page_size: Self::DEFAULT_PAGE_SIZE,
kind: CacheKind::Contiguous,
}
}
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>;
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 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,
) -> Tensor<B, 4> {
let device = q.device();
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() && (scale - default_scale).abs() < 1e-12 {
return burn::tensor::module::attention(
q,
k,
v,
None,
None,
AttentionModuleOptions {
scale: None,
softcap: None,
is_causal: seq > 1,
},
);
}
let scores = q.matmul(k.transpose()).mul_scalar(scale);
let scores = if seq > 1 {
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 forbidden: Tensor<B, 2, Bool> = k_pos.greater(q_pos);
let mask = forbidden
.unsqueeze_dims::<4>(&[0, 1])
.expand([1, n_q, seq, total]);
scores.mask_fill(mask, -1e30f32)
} else {
scores };
safe_matmul(softmax(scores, 3), v)
}
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(
&mut self,
layer: usize,
q: Tensor<B, 4>,
k: Tensor<B, 4>,
v: Tensor<B, 4>,
pos: usize,
scale: f64,
) -> 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);
*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);
}
}
pub struct PagedKVCache<B: Backend> {
config: CacheConfig,
allocator: PageAllocator,
table: Vec<usize>,
seq_len: usize,
arenas: 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 {
PagedKVCache {
allocator: PageAllocator::new(config.num_pages()),
config,
table: Vec::new(),
seq_len: 0,
arenas: (0..num_layers).map(|_| None).collect(),
device: None,
}
}
pub fn num_free_pages(&self) -> usize {
self.allocator.num_free()
}
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 gather_window(
&self,
arena: Tensor<B, 4>,
pages: usize,
total: usize,
) -> Tensor<B, 4> {
let [_, n_kv, page_size, head_dim] = arena.dims();
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");
let indices = Tensor::<B, 1, Int>::from_data(TensorData::new(ids, [pages]), device);
arena
.select(0, indices) .swap_dims(0, 1) .reshape([1, n_kv, pages * page_size, head_dim])
.narrow(2, 0, total)
}
}
impl<B: Backend> KVCache<B> for PagedKVCache<B> {
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> {
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 self.arenas[layer].is_none() {
let device = k.device();
let shape = [self.config.num_pages(), n_kv, self.config.page_size, head_dim];
self.arenas[layer] = Some((
Tensor::zeros(shape, &device),
Tensor::zeros(shape, &device),
));
}
let pages = self.ensure_pages(total);
let page_size = self.config.page_size;
let (mut arena_k, mut arena_v) = self.arenas[layer].take().expect("arena initialized");
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_k, arena_v));
attend(q, k_full, v_full, pos, scale)
}
fn seq_len(&self) -> usize {
self.seq_len
}
fn popn(&mut self, n: usize) -> usize {
let n = n.min(self.seq_len);
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;
}
fn pages_used(&self) -> Option<usize> {
Some(self.table.len())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[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,
},
)
}
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);
}
}