1use burn::tensor::ops::AttentionModuleOptions;
16use burn::tensor::{Bool, Device, Int, Tensor, TensorData, activation::softmax, backend::Backend};
17
18use crate::matmul::safe_matmul;
19
20fn flash_enabled() -> bool {
24 static ENABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
25 *ENABLED.get_or_init(|| {
26 std::env::var("COMBS_ATTN").map(|v| v != "manual").unwrap_or(true)
27 })
28}
29
30#[derive(Debug, Clone, Copy, PartialEq, Eq)]
32pub enum CacheKind {
33 Contiguous,
35 Paged,
37}
38
39#[derive(Debug, Clone, Copy)]
41pub struct CacheConfig {
42 pub max_seq_len: usize,
44 pub page_size: usize,
46 pub kind: CacheKind,
48}
49
50impl CacheConfig {
51 pub const DEFAULT_PAGE_SIZE: usize = 16;
53
54 pub fn paged(max_seq_len: usize) -> Self {
56 CacheConfig {
57 max_seq_len,
58 page_size: Self::DEFAULT_PAGE_SIZE,
59 kind: CacheKind::Paged,
60 }
61 }
62
63 pub fn contiguous(max_seq_len: usize) -> Self {
65 CacheConfig {
66 max_seq_len,
67 page_size: Self::DEFAULT_PAGE_SIZE,
68 kind: CacheKind::Contiguous,
69 }
70 }
71
72 pub fn num_pages(&self) -> usize {
74 self.max_seq_len.div_ceil(self.page_size)
75 }
76}
77
78pub trait KVCache<B: Backend>: Send {
84 fn attention(
92 &mut self,
93 layer: usize,
94 q: Tensor<B, 4>,
95 k: Tensor<B, 4>,
96 v: Tensor<B, 4>,
97 pos: usize,
98 scale: f64,
99 ) -> Tensor<B, 4>;
100
101 fn seq_len(&self) -> usize;
103
104 fn popn(&mut self, n: usize) -> usize {
108 let _ = n;
109 0
110 }
111
112 fn reset(&mut self);
114
115 fn pages_used(&self) -> Option<usize> {
117 None
118 }
119}
120
121fn repeat_kv<B: Backend>(x: Tensor<B, 4>, n_rep: usize) -> Tensor<B, 4> {
124 if n_rep == 1 {
125 return x;
126 }
127 let [b, nkv, s, d] = x.dims();
128 x.unsqueeze_dim::<5>(2)
129 .expand([b, nkv, n_rep, s, d])
130 .reshape([b, nkv * n_rep, s, d])
131}
132
133fn attend<B: Backend>(
146 q: Tensor<B, 4>,
147 k: Tensor<B, 4>,
148 v: Tensor<B, 4>,
149 pos: usize,
150 scale: f64,
151) -> Tensor<B, 4> {
152 let device = q.device();
153 let [_, n_q, seq, d] = q.dims();
154 let [_, n_kv, total, _] = k.dims();
155 let n_rep = n_q / n_kv;
156 let k = repeat_kv(k, n_rep);
157 let v = repeat_kv(v, n_rep);
158
159 let default_scale = 1.0 / (d as f64).sqrt();
160 if flash_enabled() && (scale - default_scale).abs() < 1e-12 {
161 return burn::tensor::module::attention(
162 q,
163 k,
164 v,
165 None,
166 None,
167 AttentionModuleOptions {
168 scale: None,
169 softcap: None,
170 is_causal: seq > 1,
173 },
174 );
175 }
176
177 let scores = q.matmul(k.transpose()).mul_scalar(scale);
179 let scores = if seq > 1 {
180 let q_pos =
182 Tensor::<B, 1, Int>::arange((pos as i64)..((pos + seq) as i64), &device)
183 .reshape([seq, 1]);
184 let k_pos = Tensor::<B, 1, Int>::arange(0..(total as i64), &device).reshape([1, total]);
185 let forbidden: Tensor<B, 2, Bool> = k_pos.greater(q_pos);
186 let mask = forbidden
187 .unsqueeze_dims::<4>(&[0, 1])
188 .expand([1, n_q, seq, total]);
189 scores.mask_fill(mask, -1e30f32)
190 } else {
191 scores };
193
194 safe_matmul(softmax(scores, 3), v)
197}
198
199pub struct ContiguousKVCache<B: Backend> {
205 layers: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
206 seq_len: usize,
207}
208
209impl<B: Backend> ContiguousKVCache<B> {
210 pub fn new(num_layers: usize) -> Self {
212 ContiguousKVCache {
213 layers: (0..num_layers).map(|_| None).collect(),
214 seq_len: 0,
215 }
216 }
217}
218
219impl<B: Backend> KVCache<B> for ContiguousKVCache<B> {
220 fn attention(
221 &mut self,
222 layer: usize,
223 q: Tensor<B, 4>,
224 k: Tensor<B, 4>,
225 v: Tensor<B, 4>,
226 pos: usize,
227 scale: f64,
228 ) -> Tensor<B, 4> {
229 let slot = &mut self.layers[layer];
230 let (k_full, v_full) = match slot.take() {
231 Some((k_old, v_old)) => (
232 Tensor::cat(vec![k_old, k], 2),
233 Tensor::cat(vec![v_old, v], 2),
234 ),
235 None => (k, v),
236 };
237 self.seq_len = k_full.dims()[2];
238 let out = attend(q, k_full.clone(), v_full.clone(), pos, scale);
239 *slot = Some((k_full, v_full));
240 out
241 }
242
243 fn seq_len(&self) -> usize {
244 self.seq_len
245 }
246
247 fn reset(&mut self) {
248 for slot in &mut self.layers {
249 *slot = None;
250 }
251 self.seq_len = 0;
252 }
253}
254
255#[derive(Debug)]
257struct PageAllocator {
258 free: Vec<usize>,
259}
260
261impl PageAllocator {
262 fn new(num_pages: usize) -> Self {
263 PageAllocator {
265 free: (0..num_pages).rev().collect(),
266 }
267 }
268
269 fn alloc(&mut self) -> Option<usize> {
270 self.free.pop()
271 }
272
273 fn free_page(&mut self, id: usize) {
274 self.free.push(id);
275 }
276
277 fn num_free(&self) -> usize {
278 self.free.len()
279 }
280
281 fn reset(&mut self, num_pages: usize) {
282 *self = PageAllocator::new(num_pages);
283 }
284}
285
286pub struct PagedKVCache<B: Backend> {
301 config: CacheConfig,
302 allocator: PageAllocator,
303 table: Vec<usize>,
305 seq_len: usize,
306 arenas: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
307 device: Option<Device<B>>,
308}
309
310impl<B: Backend> PagedKVCache<B> {
311 pub fn new(num_layers: usize, config: CacheConfig) -> Self {
314 PagedKVCache {
315 allocator: PageAllocator::new(config.num_pages()),
316 config,
317 table: Vec::new(),
318 seq_len: 0,
319 arenas: (0..num_layers).map(|_| None).collect(),
320 device: None,
321 }
322 }
323
324 pub fn num_free_pages(&self) -> usize {
326 self.allocator.num_free()
327 }
328
329 fn ensure_pages(&mut self, total: usize) -> usize {
331 let pages_needed = total.div_ceil(self.config.page_size);
332 while self.table.len() < pages_needed {
333 let page = self
334 .allocator
335 .alloc()
336 .expect("page allocator exhausted (max_seq_len exceeded)");
337 self.table.push(page);
338 }
339 pages_needed
340 }
341
342 fn gather_window(
346 &self,
347 arena: Tensor<B, 4>,
348 pages: usize,
349 total: usize,
350 ) -> Tensor<B, 4> {
351 let [_, n_kv, page_size, head_dim] = arena.dims();
352 let ids: Vec<i32> = self.table[..pages].iter().map(|&p| p as i32).collect();
353 let device = self
354 .device
355 .as_ref()
356 .expect("device set on first attention call");
357 let indices = Tensor::<B, 1, Int>::from_data(TensorData::new(ids, [pages]), device);
358 arena
359 .select(0, indices) .swap_dims(0, 1) .reshape([1, n_kv, pages * page_size, head_dim])
362 .narrow(2, 0, total)
363 }
364}
365
366impl<B: Backend> KVCache<B> for PagedKVCache<B> {
367 fn attention(
368 &mut self,
369 layer: usize,
370 q: Tensor<B, 4>,
371 k: Tensor<B, 4>,
372 v: Tensor<B, 4>,
373 pos: usize,
374 scale: f64,
375 ) -> Tensor<B, 4> {
376 let [_, n_kv, seq, head_dim] = k.dims();
377 let total = pos + seq;
378 if layer == 0 {
382 assert_eq!(
383 pos, self.seq_len,
384 "paged cache expects dense contiguous appends (pos == seq_len)"
385 );
386 self.seq_len = total;
387 } else {
388 debug_assert_eq!(total, self.seq_len);
389 }
390 assert!(
391 total <= self.config.max_seq_len,
392 "paged cache capacity exceeded: {total} > {}",
393 self.config.max_seq_len
394 );
395
396 if self.device.is_none() {
397 self.device = Some(k.device());
398 }
399 if self.arenas[layer].is_none() {
400 let device = k.device();
401 let shape = [self.config.num_pages(), n_kv, self.config.page_size, head_dim];
402 self.arenas[layer] = Some((
403 Tensor::zeros(shape, &device),
404 Tensor::zeros(shape, &device),
405 ));
406 }
407
408 let pages = self.ensure_pages(total);
409 let page_size = self.config.page_size;
410
411 let (mut arena_k, mut arena_v) = self.arenas[layer].take().expect("arena initialized");
414 let mut written = 0;
415 while written < seq {
416 let global = pos + written;
417 let slot = global % page_size;
418 let run = (page_size - slot).min(seq - written);
419 let phys = self.table[global / page_size];
420 let range = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..head_dim];
421 arena_k = arena_k.slice_assign(range.clone(), k.clone().narrow(2, written, run));
422 arena_v = arena_v.slice_assign(range, v.clone().narrow(2, written, run));
423 written += run;
424 }
425
426 let k_full = self.gather_window(arena_k.clone(), pages, total);
427 let v_full = self.gather_window(arena_v.clone(), pages, total);
428 self.arenas[layer] = Some((arena_k, arena_v));
429
430 attend(q, k_full, v_full, pos, scale)
431 }
432
433 fn seq_len(&self) -> usize {
434 self.seq_len
435 }
436
437 fn popn(&mut self, n: usize) -> usize {
441 let n = n.min(self.seq_len);
442 self.seq_len -= n;
443 let keep = self.seq_len.div_ceil(self.config.page_size);
444 while self.table.len() > keep {
445 let page = self.table.pop().expect("table nonempty");
446 self.allocator.free_page(page);
447 }
448 n
449 }
450
451 fn reset(&mut self) {
452 self.table.clear();
453 self.allocator.reset(self.config.num_pages());
454 self.seq_len = 0;
455 }
458
459 fn pages_used(&self) -> Option<usize> {
460 Some(self.table.len())
461 }
462}
463
464#[cfg(test)]
465mod tests {
466 use super::*;
467
468 #[test]
469 fn allocator_alloc_in_order_and_exhaust() {
470 let mut a = PageAllocator::new(3);
471 assert_eq!(a.num_free(), 3);
472 assert_eq!(a.alloc(), Some(0));
473 assert_eq!(a.alloc(), Some(1));
474 assert_eq!(a.alloc(), Some(2));
475 assert_eq!(a.alloc(), None);
476 assert_eq!(a.num_free(), 0);
477 }
478
479 #[test]
480 fn allocator_free_and_realloc_lifo() {
481 let mut a = PageAllocator::new(2);
482 let p0 = a.alloc().unwrap();
483 let p1 = a.alloc().unwrap();
484 a.free_page(p1);
485 a.free_page(p0);
486 assert_eq!(a.num_free(), 2);
487 assert_eq!(a.alloc(), Some(p0));
489 assert_eq!(a.alloc(), Some(p1));
490 }
491
492 #[test]
493 fn allocator_reset_restores_all_pages() {
494 let mut a = PageAllocator::new(4);
495 a.alloc();
496 a.alloc();
497 a.reset(4);
498 assert_eq!(a.num_free(), 4);
499 assert_eq!(a.alloc(), Some(0));
500 }
501
502 #[test]
503 fn cache_config_num_pages_rounds_up() {
504 assert_eq!(CacheConfig::paged(16).num_pages(), 1);
505 assert_eq!(CacheConfig::paged(17).num_pages(), 2);
506 assert_eq!(CacheConfig::paged(1).num_pages(), 1);
507 }
508
509 type TestBackend = burn::backend::NdArray<f32>;
513
514 fn cache(max_seq_len: usize, page_size: usize) -> PagedKVCache<TestBackend> {
515 PagedKVCache::new(
516 2,
517 CacheConfig {
518 max_seq_len,
519 page_size,
520 kind: CacheKind::Paged,
521 },
522 )
523 }
524
525 fn grow(c: &mut PagedKVCache<TestBackend>, total: usize) {
527 c.ensure_pages(total);
528 c.seq_len = total;
529 }
530
531 #[test]
532 fn popn_frees_only_fully_unused_pages() {
533 let mut c = cache(64, 16);
534 grow(&mut c, 40); assert_eq!(c.pages_used(), Some(3));
536 assert_eq!(c.num_free_pages(), 1);
537
538 c.popn(9); assert_eq!(c.seq_len(), 31);
540 assert_eq!(c.pages_used(), Some(2));
541 assert_eq!(c.num_free_pages(), 2);
542
543 c.popn(15); assert_eq!(c.pages_used(), Some(1));
545 c.popn(1); assert_eq!(c.pages_used(), Some(1));
547
548 c.popn(1000); assert_eq!(c.seq_len(), 0);
550 assert_eq!(c.pages_used(), Some(0));
551 assert_eq!(c.num_free_pages(), 4);
552 }
553
554 #[test]
555 fn popn_boundary_exact_page_edge() {
556 let mut c = cache(64, 16);
557 grow(&mut c, 32); c.popn(16); assert_eq!(c.pages_used(), Some(1));
560 assert_eq!(c.num_free_pages(), 3);
561 c.popn(16);
562 assert_eq!(c.pages_used(), Some(0));
563 assert_eq!(c.num_free_pages(), 4);
564 }
565
566 #[test]
567 fn regrowth_after_popn_reuses_freed_pages() {
568 let mut c = cache(64, 16);
569 grow(&mut c, 40);
570 c.popn(9); grow(&mut c, 33); assert_eq!(c.pages_used(), Some(3));
573 assert_eq!(c.num_free_pages(), 1);
574 }
575
576 #[test]
577 fn reset_releases_all_pages() {
578 let mut c = cache(64, 16);
579 grow(&mut c, 40);
580 c.reset();
581 assert_eq!(c.seq_len(), 0);
582 assert_eq!(c.pages_used(), Some(0));
583 assert_eq!(c.num_free_pages(), 4);
584 }
585}