1use burn::tensor::ops::AttentionModuleOptions;
16use burn::tensor::{Bool, Device, Int, Tensor, TensorData, activation::softmax, backend::Backend};
17
18use crate::matmul::safe_matmul;
19use crate::precision::{to_f32, to_float};
20
21fn flash_enabled() -> bool {
25 static ENABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
26 *ENABLED.get_or_init(|| {
27 std::env::var("COMBS_ATTN").map(|v| v != "manual").unwrap_or(true)
28 })
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum CacheKind {
34 Contiguous,
36 Paged,
38}
39
40#[derive(Debug, Clone, Copy)]
42pub struct CacheConfig {
43 pub max_seq_len: usize,
45 pub page_size: usize,
47 pub kind: CacheKind,
49 pub quantize_kv: bool,
55}
56
57impl CacheConfig {
58 pub const DEFAULT_PAGE_SIZE: usize = 16;
60
61 pub fn paged(max_seq_len: usize) -> Self {
63 CacheConfig {
64 max_seq_len,
65 page_size: Self::DEFAULT_PAGE_SIZE,
66 kind: CacheKind::Paged,
67 quantize_kv: false,
68 }
69 }
70
71 pub fn contiguous(max_seq_len: usize) -> Self {
73 CacheConfig {
74 max_seq_len,
75 page_size: Self::DEFAULT_PAGE_SIZE,
76 kind: CacheKind::Contiguous,
77 quantize_kv: false,
78 }
79 }
80
81 pub fn num_pages(&self) -> usize {
83 self.max_seq_len.div_ceil(self.page_size)
84 }
85}
86
87pub trait KVCache<B: Backend>: Send {
93 fn attention(
101 &mut self,
102 layer: usize,
103 q: Tensor<B, 4>,
104 k: Tensor<B, 4>,
105 v: Tensor<B, 4>,
106 pos: usize,
107 scale: f64,
108 ) -> Tensor<B, 4> {
109 self.attention_opts(layer, q, k, v, pos, scale, None)
110 }
111
112 fn attention_opts(
117 &mut self,
118 layer: usize,
119 q: Tensor<B, 4>,
120 k: Tensor<B, 4>,
121 v: Tensor<B, 4>,
122 pos: usize,
123 scale: f64,
124 window: Option<usize>,
125 ) -> Tensor<B, 4>;
126
127 fn seq_len(&self) -> usize;
129
130 fn popn(&mut self, n: usize) -> usize {
134 let _ = n;
135 0
136 }
137
138 fn reset(&mut self);
140
141 fn pages_used(&self) -> Option<usize> {
143 None
144 }
145
146 fn page_stats(&self) -> Option<PageStats> {
149 None
150 }
151}
152
153#[derive(Debug, Clone, Copy, PartialEq, Eq)]
155pub struct PageStats {
156 pub pages_used: usize,
158 pub pages_free: usize,
160 pub num_pages: usize,
162 pub page_size: usize,
164 pub seq_len: usize,
166 pub layers_materialized: usize,
171 pub layers_total: usize,
173 pub layers_sliding: usize,
175}
176
177fn repeat_kv<B: Backend>(x: Tensor<B, 4>, n_rep: usize) -> Tensor<B, 4> {
180 if n_rep == 1 {
181 return x;
182 }
183 let [b, nkv, s, d] = x.dims();
184 x.unsqueeze_dim::<5>(2)
185 .expand([b, nkv, n_rep, s, d])
186 .reshape([b, nkv * n_rep, s, d])
187}
188
189fn attend<B: Backend>(
202 q: Tensor<B, 4>,
203 k: Tensor<B, 4>,
204 v: Tensor<B, 4>,
205 pos: usize,
206 scale: f64,
207 window: Option<usize>,
208) -> Tensor<B, 4> {
209 let device = q.device();
210 let out_dtype = q.dtype();
213 let q = to_f32(q);
214 let k = to_f32(k);
215 let v = to_f32(v);
216 let [_, n_q, seq, d] = q.dims();
217 let [_, n_kv, total, _] = k.dims();
218 let n_rep = n_q / n_kv;
219 let k = repeat_kv(k, n_rep);
220 let v = repeat_kv(v, n_rep);
221
222 let default_scale = 1.0 / (d as f64).sqrt();
223 if flash_enabled() && window.is_none() && (scale - default_scale).abs() < 1e-12 {
224 let out = burn::tensor::module::attention(
225 q,
226 k,
227 v,
228 None,
229 None,
230 AttentionModuleOptions {
231 scale: None,
232 softcap: None,
233 is_causal: seq > 1,
236 },
237 );
238 return to_float(out, out_dtype);
239 }
240
241 let scores = q.matmul(k.transpose()).mul_scalar(scale);
244 let scores = if seq > 1 || window.is_some() {
245 let q_pos =
248 Tensor::<B, 1, Int>::arange((pos as i64)..((pos + seq) as i64), &device)
249 .reshape([seq, 1]);
250 let k_pos = Tensor::<B, 1, Int>::arange(0..(total as i64), &device).reshape([1, total]);
251 let mut forbidden: Tensor<B, 2, Bool> = k_pos.clone().greater(q_pos.clone());
252 if let Some(w) = window {
253 let too_old = k_pos
254 .add_scalar(w as i64 - 1)
255 .lower(q_pos);
256 forbidden = forbidden.bool_or(too_old);
257 }
258 let mask = forbidden
259 .unsqueeze_dims::<4>(&[0, 1])
260 .expand([1, n_q, seq, total]);
261 scores.mask_fill(mask, -1e30f32)
262 } else {
263 scores };
265
266 to_float(safe_matmul(softmax(scores, 3), v), out_dtype)
269}
270
271const KV_QUANT_GROUP: usize = 32;
273
274fn kv_quantize<B: Backend>(x: Tensor<B, 4>) -> (Tensor<B, 4, Int>, Tensor<B, 4>) {
286 let [b, h, s, d] = x.dims();
287 debug_assert_eq!(d % KV_QUANT_GROUP, 0);
288 let groups = d / KV_QUANT_GROUP;
289 let native = x.dtype();
294 let g = to_f32(x).reshape([b, h, s, groups, KV_QUANT_GROUP]);
295 let scale = g
296 .clone()
297 .abs()
298 .max_dim(4) .div_scalar(127.0)
300 .clamp_min(1e-8);
301 let q = g
302 .div(scale.clone().expand([b, h, s, groups, KV_QUANT_GROUP]))
303 .round()
304 .clamp(-127.0, 127.0)
305 .int()
306 .reshape([b, h, s, d / 4, 4]);
307 let lane = |i: usize| q.clone().narrow(4, i, 1).reshape([b, h, s, d / 4]);
308 let packed = lane(0).add_scalar(128)
309 + lane(1).add_scalar(128).mul_scalar(256)
310 + lane(2).add_scalar(128).mul_scalar(65536)
311 + lane(3).mul_scalar(16777216);
312 (packed, to_float(scale.reshape([b, h, s, groups]), native))
313}
314
315fn kv_dequantize<B: Backend>(
320 packed: Tensor<B, 4, Int>,
321 scales: Tensor<B, 4>,
322 d: usize,
323) -> Tensor<B, 4> {
324 let [b, h, s, _] = packed.dims();
325 let groups = d / KV_QUANT_GROUP;
326 let t = packed.clone().div_scalar(16777216);
327 let r3 = packed - t.clone().mul_scalar(16777216);
328 let neg = r3.clone().lower_elem(0);
329 let q3 = t.clone().mask_where(neg.clone(), t.sub_scalar(1));
330 let r = r3.clone().mask_where(neg, r3.add_scalar(16777216));
331 let l2 = r.clone().div_scalar(65536);
332 let r = r - l2.clone().mul_scalar(65536);
333 let l1 = r.clone().div_scalar(256);
334 let l0 = r - l1.clone().mul_scalar(256);
335 let q = Tensor::stack::<5>(
339 vec![
340 l0.sub_scalar(128),
341 l1.sub_scalar(128),
342 l2.sub_scalar(128),
343 q3,
344 ],
345 4,
346 )
347 .reshape([b, h, s, d]);
348 let g = q.float().reshape([b, h, s, groups, KV_QUANT_GROUP]);
349 let scales = to_float(scales, g.dtype());
352 g.mul(
353 scales
354 .reshape([b, h, s, groups, 1])
355 .expand([b, h, s, groups, KV_QUANT_GROUP]),
356 )
357 .reshape([b, h, s, d])
358}
359
360pub struct ContiguousKVCache<B: Backend> {
366 layers: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
367 seq_len: usize,
368}
369
370impl<B: Backend> ContiguousKVCache<B> {
371 pub fn new(num_layers: usize) -> Self {
373 ContiguousKVCache {
374 layers: (0..num_layers).map(|_| None).collect(),
375 seq_len: 0,
376 }
377 }
378}
379
380impl<B: Backend> KVCache<B> for ContiguousKVCache<B> {
381 fn attention_opts(
382 &mut self,
383 layer: usize,
384 q: Tensor<B, 4>,
385 k: Tensor<B, 4>,
386 v: Tensor<B, 4>,
387 pos: usize,
388 scale: f64,
389 window: Option<usize>,
390 ) -> Tensor<B, 4> {
391 let slot = &mut self.layers[layer];
392 let (k_full, v_full) = match slot.take() {
393 Some((k_old, v_old)) => (
394 Tensor::cat(vec![k_old, k], 2),
395 Tensor::cat(vec![v_old, v], 2),
396 ),
397 None => (k, v),
398 };
399 self.seq_len = k_full.dims()[2];
400 let out = attend(q, k_full.clone(), v_full.clone(), pos, scale, window);
401 *slot = Some((k_full, v_full));
402 out
403 }
404
405 fn seq_len(&self) -> usize {
406 self.seq_len
407 }
408
409 fn reset(&mut self) {
410 for slot in &mut self.layers {
411 *slot = None;
412 }
413 self.seq_len = 0;
414 }
415}
416
417#[derive(Debug)]
419struct PageAllocator {
420 free: Vec<usize>,
421}
422
423impl PageAllocator {
424 fn new(num_pages: usize) -> Self {
425 PageAllocator {
427 free: (0..num_pages).rev().collect(),
428 }
429 }
430
431 fn alloc(&mut self) -> Option<usize> {
432 self.free.pop()
433 }
434
435 fn free_page(&mut self, id: usize) {
436 self.free.push(id);
437 }
438
439 fn num_free(&self) -> usize {
440 self.free.len()
441 }
442
443 fn reset(&mut self, num_pages: usize) {
444 *self = PageAllocator::new(num_pages);
445 }
446}
447
448enum Arena<B: Backend> {
465 Fp {
466 k: Tensor<B, 4>,
467 v: Tensor<B, 4>,
468 },
469 Quant {
470 k_packed: Tensor<B, 4, Int>,
471 k_scales: Tensor<B, 4>,
472 v_packed: Tensor<B, 4, Int>,
473 v_scales: Tensor<B, 4>,
474 },
475}
476
477pub struct PagedKVCache<B: Backend> {
478 config: CacheConfig,
479 allocator: PageAllocator,
480 table: Vec<usize>,
482 seq_len: usize,
483 arenas: Vec<Option<Arena<B>>>,
484 layer_windows: Vec<Option<usize>>,
490 sliding: Vec<Option<(Tensor<B, 4>, Tensor<B, 4>)>>,
494 device: Option<Device<B>>,
495}
496
497impl<B: Backend> PagedKVCache<B> {
498 pub fn new(num_layers: usize, config: CacheConfig) -> Self {
501 Self::new_with_windows(num_layers, config, vec![None; num_layers])
502 }
503
504 pub fn new_with_windows(
509 num_layers: usize,
510 config: CacheConfig,
511 windows: Vec<Option<usize>>,
512 ) -> Self {
513 assert_eq!(windows.len(), num_layers, "one window entry per layer");
514 for w in windows.iter().flatten() {
515 assert!(*w >= 2, "sliding window must be >= 2, got {w}");
516 }
517 PagedKVCache {
518 allocator: PageAllocator::new(config.num_pages()),
519 config,
520 table: Vec::new(),
521 seq_len: 0,
522 arenas: (0..num_layers).map(|_| None).collect(),
523 layer_windows: windows,
524 sliding: (0..num_layers).map(|_| None).collect(),
525 device: None,
526 }
527 }
528
529 pub fn num_free_pages(&self) -> usize {
531 self.allocator.num_free()
532 }
533
534 pub fn page_stats_inner(&self) -> PageStats {
536 PageStats {
537 pages_used: self.table.len(),
538 pages_free: self.allocator.num_free(),
539 num_pages: self.config.num_pages(),
540 page_size: self.config.page_size,
541 seq_len: self.seq_len,
542 layers_materialized: self.arenas.iter().filter(|a| a.is_some()).count(),
543 layers_total: self.arenas.len(),
544 layers_sliding: self.sliding.iter().filter(|s| s.is_some()).count(),
545 }
546 }
547
548 fn ensure_pages(&mut self, total: usize) -> usize {
550 let pages_needed = total.div_ceil(self.config.page_size);
551 while self.table.len() < pages_needed {
552 let page = self
553 .allocator
554 .alloc()
555 .expect("page allocator exhausted (max_seq_len exceeded)");
556 self.table.push(page);
557 }
558 pages_needed
559 }
560
561 fn sliding_attention(
571 &mut self,
572 layer: usize,
573 q: Tensor<B, 4>,
574 k: Tensor<B, 4>,
575 v: Tensor<B, 4>,
576 pos: usize,
577 scale: f64,
578 w: usize,
579 ) -> Tensor<B, 4> {
580 let seq = k.dims()[2];
581 let slot = &mut self.sliding[layer];
582 let (k_full, v_full) = match slot.take() {
583 Some((k_old, v_old)) => (
584 Tensor::cat(vec![k_old, k], 2),
585 Tensor::cat(vec![v_old, v], 2),
586 ),
587 None => (k, v),
588 };
589 let full_len = k_full.dims()[2];
590 let kv_offset = pos + seq - full_len;
591 let out = attend(
592 q,
593 k_full.clone(),
594 v_full.clone(),
595 pos - kv_offset,
596 scale,
597 Some(w),
598 );
599 let keep = full_len.min(w - 1);
604 *slot = Some((
605 k_full.narrow(2, full_len - keep, keep),
606 v_full.narrow(2, full_len - keep, keep),
607 ));
608 out
609 }
610
611 fn page_indices(&self, pages: usize) -> Tensor<B, 1, Int> {
613 let ids: Vec<i32> = self.table[..pages].iter().map(|&p| p as i32).collect();
614 let device = self
615 .device
616 .as_ref()
617 .expect("device set on first attention call");
618 Tensor::<B, 1, Int>::from_data(TensorData::new(ids, [pages]), device)
619 }
620
621 fn gather_window(
625 &self,
626 arena: Tensor<B, 4>,
627 pages: usize,
628 total: usize,
629 ) -> Tensor<B, 4> {
630 let [_, n_kv, page_size, last] = arena.dims();
631 arena
632 .select(0, self.page_indices(pages)) .swap_dims(0, 1) .reshape([1, n_kv, pages * page_size, last])
635 .narrow(2, 0, total)
636 }
637
638 fn gather_window_int(
640 &self,
641 arena: Tensor<B, 4, Int>,
642 pages: usize,
643 total: usize,
644 ) -> Tensor<B, 4, Int> {
645 let [_, n_kv, page_size, last] = arena.dims();
646 arena
647 .select(0, self.page_indices(pages))
648 .swap_dims(0, 1)
649 .reshape([1, n_kv, pages * page_size, last])
650 .narrow(2, 0, total)
651 }
652}
653
654impl<B: Backend> KVCache<B> for PagedKVCache<B> {
655 fn attention_opts(
656 &mut self,
657 layer: usize,
658 q: Tensor<B, 4>,
659 k: Tensor<B, 4>,
660 v: Tensor<B, 4>,
661 pos: usize,
662 scale: f64,
663 window: Option<usize>,
664 ) -> Tensor<B, 4> {
665 let [_, n_kv, seq, head_dim] = k.dims();
666 let total = pos + seq;
667 if layer == 0 {
671 assert_eq!(
672 pos, self.seq_len,
673 "paged cache expects dense contiguous appends (pos == seq_len)"
674 );
675 self.seq_len = total;
676 } else {
677 debug_assert_eq!(total, self.seq_len);
678 }
679 assert!(
680 total <= self.config.max_seq_len,
681 "paged cache capacity exceeded: {total} > {}",
682 self.config.max_seq_len
683 );
684
685 if self.device.is_none() {
686 self.device = Some(k.device());
687 }
688 if let Some(w) = self.layer_windows.get(layer).copied().flatten() {
692 return self.sliding_attention(layer, q, k, v, pos, scale, w);
693 }
694 let quant = self.config.quantize_kv && head_dim % KV_QUANT_GROUP == 0;
695 if self.arenas[layer].is_none() {
696 let device = k.device();
697 let np = self.config.num_pages();
698 let ps = self.config.page_size;
699 self.arenas[layer] = Some(if quant {
700 Arena::Quant {
701 k_packed: Tensor::zeros([np, n_kv, ps, head_dim / 4], &device),
702 k_scales: Tensor::zeros([np, n_kv, ps, head_dim / KV_QUANT_GROUP], &device),
703 v_packed: Tensor::zeros([np, n_kv, ps, head_dim / 4], &device),
704 v_scales: Tensor::zeros([np, n_kv, ps, head_dim / KV_QUANT_GROUP], &device),
705 }
706 } else {
707 let shape = [np, n_kv, ps, head_dim];
708 Arena::Fp {
709 k: Tensor::zeros(shape, &device),
710 v: Tensor::zeros(shape, &device),
711 }
712 });
713 }
714
715 let pages = self.ensure_pages(total);
716 let page_size = self.config.page_size;
717
718 let (k_full, v_full) = match self.arenas[layer].take().expect("arena initialized") {
722 Arena::Fp { k: mut arena_k, v: mut arena_v } => {
723 let mut written = 0;
724 while written < seq {
725 let global = pos + written;
726 let slot = global % page_size;
727 let run = (page_size - slot).min(seq - written);
728 let phys = self.table[global / page_size];
729 let range = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..head_dim];
730 arena_k =
731 arena_k.slice_assign(range.clone(), k.clone().narrow(2, written, run));
732 arena_v = arena_v.slice_assign(range, v.clone().narrow(2, written, run));
733 written += run;
734 }
735 let k_full = self.gather_window(arena_k.clone(), pages, total);
736 let v_full = self.gather_window(arena_v.clone(), pages, total);
737 self.arenas[layer] = Some(Arena::Fp { k: arena_k, v: arena_v });
738 (k_full, v_full)
739 }
740 Arena::Quant {
741 mut k_packed,
742 mut k_scales,
743 mut v_packed,
744 mut v_scales,
745 } => {
746 let (kq, ks) = kv_quantize(k);
751 let (vq, vs) = kv_quantize(v);
752 let dp = head_dim / 4;
753 let dg = head_dim / KV_QUANT_GROUP;
754 let mut written = 0;
755 while written < seq {
756 let global = pos + written;
757 let slot = global % page_size;
758 let run = (page_size - slot).min(seq - written);
759 let phys = self.table[global / page_size];
760 let rp = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..dp];
761 let rs = [phys..phys + 1, 0..n_kv, slot..slot + run, 0..dg];
762 k_packed =
763 k_packed.slice_assign(rp.clone(), kq.clone().narrow(2, written, run));
764 k_scales =
765 k_scales.slice_assign(rs.clone(), ks.clone().narrow(2, written, run));
766 v_packed = v_packed.slice_assign(rp, vq.clone().narrow(2, written, run));
767 v_scales = v_scales.slice_assign(rs, vs.clone().narrow(2, written, run));
768 written += run;
769 }
770 let k_full = kv_dequantize(
771 self.gather_window_int(k_packed.clone(), pages, total),
772 self.gather_window(k_scales.clone(), pages, total),
773 head_dim,
774 );
775 let v_full = kv_dequantize(
776 self.gather_window_int(v_packed.clone(), pages, total),
777 self.gather_window(v_scales.clone(), pages, total),
778 head_dim,
779 );
780 self.arenas[layer] = Some(Arena::Quant {
781 k_packed,
782 k_scales,
783 v_packed,
784 v_scales,
785 });
786 (k_full, v_full)
787 }
788 };
789
790 attend(q, k_full, v_full, pos, scale, window)
791 }
792
793 fn seq_len(&self) -> usize {
794 self.seq_len
795 }
796
797 fn popn(&mut self, n: usize) -> usize {
808 let n = n.min(self.seq_len);
809 if n == 0 {
810 return 0;
811 }
812 for w in self.layer_windows.iter().flatten() {
813 if self.seq_len > w - 1 {
814 return 0; }
816 }
817 for slot in self.sliding.iter_mut() {
818 if let Some((k, v)) = slot.take() {
819 let len = k.dims()[2];
822 let keep = len.saturating_sub(n);
823 if keep > 0 {
824 *slot = Some((k.narrow(2, 0, keep), v.narrow(2, 0, keep)));
825 }
826 }
827 }
828 self.seq_len -= n;
829 let keep = self.seq_len.div_ceil(self.config.page_size);
830 while self.table.len() > keep {
831 let page = self.table.pop().expect("table nonempty");
832 self.allocator.free_page(page);
833 }
834 n
835 }
836
837 fn reset(&mut self) {
838 self.table.clear();
839 self.allocator.reset(self.config.num_pages());
840 self.seq_len = 0;
841 for slot in &mut self.sliding {
842 *slot = None;
843 }
844 }
847
848 fn pages_used(&self) -> Option<usize> {
849 Some(self.table.len())
850 }
851
852 fn page_stats(&self) -> Option<PageStats> {
853 Some(self.page_stats_inner())
854 }
855}
856
857#[cfg(test)]
858mod tests {
859 use super::*;
860
861 type TB = burn::backend::NdArray<f32>;
862
863 fn kv_tok(i: usize, n_kv: usize, d: usize) -> (Tensor<TB, 4>, Tensor<TB, 4>) {
865 let dev = Default::default();
866 let mk = |salt: usize| {
867 let data: Vec<f32> = (0..n_kv * d)
868 .map(|j| ((i * 7 + j * 3 + salt) % 13) as f32 / 13.0 - 0.5)
869 .collect();
870 Tensor::<TB, 4>::from_data(TensorData::new(data, [1, n_kv, 1, d]), &dev)
871 };
872 (mk(0), mk(5))
873 }
874
875 fn q_tok(i: usize, n_q: usize, d: usize) -> Tensor<TB, 4> {
877 let dev = Default::default();
878 let data: Vec<f32> = (0..n_q * d)
879 .map(|j| ((i * 11 + j * 5) % 17) as f32 / 17.0 - 0.5)
880 .collect();
881 Tensor::<TB, 4>::from_data(TensorData::new(data, [1, n_q, 1, d]), &dev)
882 }
883
884 fn assert_close4(a: Tensor<TB, 4>, b: Tensor<TB, 4>, what: &str) {
885 let av: Vec<f32> = a.into_data().to_vec().unwrap();
886 let bv: Vec<f32> = b.into_data().to_vec().unwrap();
887 assert_eq!(av.len(), bv.len(), "{what}: shape");
888 for (i, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
889 assert!((x - y).abs() < 1e-5, "{what}[{i}]: {x} vs {y}");
890 }
891 }
892
893 #[test]
898 fn kv_quant_roundtrip_exact_on_grid_values() {
899 let dev = Default::default();
900 let d = 64usize;
901 let data: Vec<f32> = (0..2 * 3 * d)
903 .map(|i| {
904 let q = ((i * 37) % 255) as i64 - 127; q as f32 * 0.5
906 })
907 .collect();
908 let mut data = data;
910 for g in 0..(2 * 3 * d) / 32 {
911 data[g * 32] = 63.5;
912 }
913 let x = Tensor::<TB, 4>::from_data(TensorData::new(data.clone(), [1, 2, 3, d]), &dev);
914 let (packed, scales) = kv_quantize(x);
915 let back: Vec<f32> = kv_dequantize(packed, scales, d)
916 .into_data()
917 .to_vec()
918 .unwrap();
919 for (i, (a, b)) in data.iter().zip(back.iter()).enumerate() {
920 assert!((a - b).abs() < 1e-6, "[{i}]: {a} vs {b} (must be exact)");
921 }
922 }
923
924 #[test]
927 fn kv_quant_error_bounded_by_half_step() {
928 let dev = Default::default();
929 let d = 64usize;
930 let data: Vec<f32> = (0..1 * 2 * 5 * d)
931 .map(|i| ((i * 7919) % 1000) as f32 / 250.0 - 2.0)
932 .collect();
933 let x = Tensor::<TB, 4>::from_data(TensorData::new(data.clone(), [1, 2, 5, d]), &dev);
934 let (packed, scales) = kv_quantize(x);
935 let back: Vec<f32> = kv_dequantize(packed, scales, d)
936 .into_data()
937 .to_vec()
938 .unwrap();
939 for (g, chunk) in data.chunks(32).enumerate() {
940 let absmax = chunk.iter().fold(0f32, |m, v| m.max(v.abs()));
941 let half_step = absmax / 254.0 + 1e-6;
942 for (j, (a, b)) in chunk
943 .iter()
944 .zip(back[g * 32..g * 32 + 32].iter())
945 .enumerate()
946 {
947 assert!(
948 (a - b).abs() <= half_step,
949 "group {g} elem {j}: |{a} - {b}| > {half_step}"
950 );
951 }
952 }
953 }
954
955 #[test]
959 fn quantized_paged_cache_matches_fp_within_tolerance() {
960 let (n_kv, n_q, d) = (2usize, 4usize, 32usize);
961 let mut cfg_q = CacheConfig::paged(64);
962 cfg_q.quantize_kv = true;
963 let cfg_f = CacheConfig::paged(64);
964 let scale = 1.0 / (d as f64).sqrt() * 0.9;
965 let mut fp = PagedKVCache::<TB>::new(1, cfg_f);
966 let mut qn = PagedKVCache::<TB>::new(1, cfg_q);
967
968 let ks: Vec<_> = (0..20).map(|i| kv_tok(i, n_kv, d)).collect();
970 let k20 = Tensor::cat(ks.iter().map(|(k, _)| k.clone()).collect(), 2);
971 let v20 = Tensor::cat(ks.iter().map(|(_, v)| v.clone()).collect(), 2);
972 let q20 = Tensor::cat((0..20).map(|i| q_tok(i, n_q, d)).collect(), 2);
973 let a = fp.attention_opts(0, q20.clone(), k20.clone(), v20.clone(), 0, scale, None);
974 let b = qn.attention_opts(0, q20, k20, v20, 0, scale, None);
975 let av: Vec<f32> = a.into_data().to_vec().unwrap();
976 let bv: Vec<f32> = b.into_data().to_vec().unwrap();
977 for (i, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
978 assert!((x - y).abs() < 1e-2, "prefill[{i}]: {x} vs {y}");
979 }
980 for i in 20..30 {
981 let (k, v) = kv_tok(i, n_kv, d);
982 let q = q_tok(i, n_q, d);
983 let a = fp.attention_opts(0, q.clone(), k.clone(), v.clone(), i, scale, None);
984 let b = qn.attention_opts(0, q, k, v, i, scale, None);
985 let av: Vec<f32> = a.into_data().to_vec().unwrap();
986 let bv: Vec<f32> = b.into_data().to_vec().unwrap();
987 for (j, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
988 assert!((x - y).abs() < 1e-2, "decode {i}[{j}]: {x} vs {y}");
989 }
990 }
991 assert_eq!(qn.pages_used(), fp.pages_used());
992 assert_eq!(qn.popn(5), 5, "quantized rollback works (no sliding)");
993 assert_eq!(qn.seq_len(), 25);
994 }
995
996 #[test]
1002 fn sliding_layer_matches_masked_global() {
1003 let (n_kv, n_q, d, w) = (2usize, 4usize, 4usize, 5usize);
1004 let cfg = CacheConfig::paged(64);
1005 let scale = 1.0 / (d as f64).sqrt() * 0.9; let mut global = PagedKVCache::<TB>::new(1, cfg);
1007 let mut sliding = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
1008
1009 let ks: Vec<_> = (0..7).map(|i| kv_tok(i, n_kv, d)).collect();
1011 let k7 = Tensor::cat(ks.iter().map(|(k, _)| k.clone()).collect(), 2);
1012 let v7 = Tensor::cat(ks.iter().map(|(_, v)| v.clone()).collect(), 2);
1013 let q7 = Tensor::cat((0..7).map(|i| q_tok(i, n_q, d)).collect(), 2);
1014 let a = global.attention_opts(0, q7.clone(), k7.clone(), v7.clone(), 0, scale, Some(w));
1015 let b = sliding.attention_opts(0, q7, k7, v7, 0, scale, Some(w));
1016 assert_close4(a, b, "prefill chunk");
1017
1018 for i in 7..14 {
1020 let (k, v) = kv_tok(i, n_kv, d);
1021 let q = q_tok(i, n_q, d);
1022 let a = global.attention_opts(0, q.clone(), k.clone(), v.clone(), i, scale, Some(w));
1023 let b = sliding.attention_opts(0, q, k, v, i, scale, Some(w));
1024 assert_close4(a, b, &format!("decode step {i}"));
1025 }
1026
1027 assert_eq!(sliding.pages_used(), Some(0), "sliding layers use no pages");
1029 assert!(global.pages_used().unwrap() > 0);
1030 }
1031
1032 #[test]
1035 fn sliding_popn_before_eviction_matches_replay() {
1036 let (n_kv, n_q, d, w) = (2usize, 2usize, 4usize, 8usize);
1037 let cfg = CacheConfig::paged(64);
1038 let scale = 0.4;
1039 let mut cache = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
1040 for i in 0..4 {
1041 let (k, v) = kv_tok(i, n_kv, d);
1042 cache.attention_opts(0, q_tok(i, n_q, d), k, v, i, scale, Some(w));
1043 }
1044 assert_eq!(cache.popn(2), 2, "un-evicted rollback succeeds");
1045 assert_eq!(cache.seq_len(), 2);
1046
1047 let mut fresh = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
1048 for i in 0..2 {
1049 let (k, v) = kv_tok(i, n_kv, d);
1050 fresh.attention_opts(0, q_tok(i, n_q, d), k, v, i, scale, Some(w));
1051 }
1052 let (k, v) = kv_tok(9, n_kv, d);
1054 let a = cache.attention_opts(0, q_tok(9, n_q, d), k.clone(), v.clone(), 2, scale, Some(w));
1055 let b = fresh.attention_opts(0, q_tok(9, n_q, d), k, v, 2, scale, Some(w));
1056 assert_close4(a, b, "post-rollback step");
1057 }
1058
1059 #[test]
1062 fn sliding_popn_after_eviction_refuses() {
1063 let (n_kv, n_q, d, w) = (1usize, 1usize, 4usize, 4usize);
1064 let cfg = CacheConfig::paged(64);
1065 let mut cache = PagedKVCache::<TB>::new_with_windows(1, cfg, vec![Some(w)]);
1066 for i in 0..6 {
1067 let (k, v) = kv_tok(i, n_kv, d);
1068 cache.attention_opts(0, q_tok(i, n_q, d), k, v, i, 0.5, Some(w));
1069 }
1070 assert_eq!(cache.popn(1), 0, "evicted sliding layer refuses rollback");
1071 assert_eq!(cache.seq_len(), 6, "refused rollback leaves state intact");
1072 assert_eq!(cache.popn(0), 0);
1073 }
1074
1075 #[test]
1076 fn allocator_alloc_in_order_and_exhaust() {
1077 let mut a = PageAllocator::new(3);
1078 assert_eq!(a.num_free(), 3);
1079 assert_eq!(a.alloc(), Some(0));
1080 assert_eq!(a.alloc(), Some(1));
1081 assert_eq!(a.alloc(), Some(2));
1082 assert_eq!(a.alloc(), None);
1083 assert_eq!(a.num_free(), 0);
1084 }
1085
1086 #[test]
1087 fn allocator_free_and_realloc_lifo() {
1088 let mut a = PageAllocator::new(2);
1089 let p0 = a.alloc().unwrap();
1090 let p1 = a.alloc().unwrap();
1091 a.free_page(p1);
1092 a.free_page(p0);
1093 assert_eq!(a.num_free(), 2);
1094 assert_eq!(a.alloc(), Some(p0));
1096 assert_eq!(a.alloc(), Some(p1));
1097 }
1098
1099 #[test]
1100 fn allocator_reset_restores_all_pages() {
1101 let mut a = PageAllocator::new(4);
1102 a.alloc();
1103 a.alloc();
1104 a.reset(4);
1105 assert_eq!(a.num_free(), 4);
1106 assert_eq!(a.alloc(), Some(0));
1107 }
1108
1109 #[test]
1110 fn cache_config_num_pages_rounds_up() {
1111 assert_eq!(CacheConfig::paged(16).num_pages(), 1);
1112 assert_eq!(CacheConfig::paged(17).num_pages(), 2);
1113 assert_eq!(CacheConfig::paged(1).num_pages(), 1);
1114 }
1115
1116 type TestBackend = burn::backend::NdArray<f32>;
1120
1121 fn cache(max_seq_len: usize, page_size: usize) -> PagedKVCache<TestBackend> {
1122 PagedKVCache::new(
1123 2,
1124 CacheConfig {
1125 max_seq_len,
1126 page_size,
1127 kind: CacheKind::Paged,
1128 quantize_kv: false,
1129 },
1130 )
1131 }
1132
1133 fn grow(c: &mut PagedKVCache<TestBackend>, total: usize) {
1135 c.ensure_pages(total);
1136 c.seq_len = total;
1137 }
1138
1139 #[test]
1140 fn popn_frees_only_fully_unused_pages() {
1141 let mut c = cache(64, 16);
1142 grow(&mut c, 40); assert_eq!(c.pages_used(), Some(3));
1144 assert_eq!(c.num_free_pages(), 1);
1145
1146 c.popn(9); assert_eq!(c.seq_len(), 31);
1148 assert_eq!(c.pages_used(), Some(2));
1149 assert_eq!(c.num_free_pages(), 2);
1150
1151 c.popn(15); assert_eq!(c.pages_used(), Some(1));
1153 c.popn(1); assert_eq!(c.pages_used(), Some(1));
1155
1156 c.popn(1000); assert_eq!(c.seq_len(), 0);
1158 assert_eq!(c.pages_used(), Some(0));
1159 assert_eq!(c.num_free_pages(), 4);
1160 }
1161
1162 #[test]
1163 fn popn_boundary_exact_page_edge() {
1164 let mut c = cache(64, 16);
1165 grow(&mut c, 32); c.popn(16); assert_eq!(c.pages_used(), Some(1));
1168 assert_eq!(c.num_free_pages(), 3);
1169 c.popn(16);
1170 assert_eq!(c.pages_used(), Some(0));
1171 assert_eq!(c.num_free_pages(), 4);
1172 }
1173
1174 #[test]
1175 fn regrowth_after_popn_reuses_freed_pages() {
1176 let mut c = cache(64, 16);
1177 grow(&mut c, 40);
1178 c.popn(9); grow(&mut c, 33); assert_eq!(c.pages_used(), Some(3));
1181 assert_eq!(c.num_free_pages(), 1);
1182 }
1183
1184 #[test]
1185 fn reset_releases_all_pages() {
1186 let mut c = cache(64, 16);
1187 grow(&mut c, 40);
1188 c.reset();
1189 assert_eq!(c.seq_len(), 0);
1190 assert_eq!(c.pages_used(), Some(0));
1191 assert_eq!(c.num_free_pages(), 4);
1192 }
1193
1194 fn kv_tok_on<B: Backend>(
1203 dev: &B::Device,
1204 i: usize,
1205 n_kv: usize,
1206 d: usize,
1207 ) -> (Tensor<B, 4>, Tensor<B, 4>) {
1208 let mk = |salt: usize| {
1209 let data: Vec<f32> = (0..n_kv * d)
1210 .map(|j| ((i * 7 + j * 3 + salt) % 13) as f32 / 13.0 - 0.5)
1211 .collect();
1212 Tensor::<B, 4>::from_data(TensorData::new(data, [1, n_kv, 1, d]), dev)
1213 };
1214 (mk(0), mk(5))
1215 }
1216
1217 fn q_tok_on<B: Backend>(dev: &B::Device, i: usize, n_q: usize, d: usize) -> Tensor<B, 4> {
1218 let data: Vec<f32> = (0..n_q * d)
1219 .map(|j| ((i * 11 + j * 5) % 17) as f32 / 17.0 - 0.5)
1220 .collect();
1221 Tensor::<B, 4>::from_data(TensorData::new(data, [1, n_q, 1, d]), dev)
1222 }
1223
1224 fn quant_roundtrip_on<B: Backend>(dev: &B::Device) {
1227 let d = 64usize;
1228 let data: Vec<f32> = (0..2 * d)
1229 .map(|j| ((j * 5) % 251) as f32 / 251.0 - 0.5)
1230 .collect();
1231 let x = Tensor::<B, 4>::from_data(TensorData::new(data.clone(), [1, 2, 1, d]), dev);
1232 let (packed, scales) = kv_quantize(x);
1233 let y = kv_dequantize(packed, scales, d);
1234 let yv: Vec<f32> = y.into_data().convert::<f32>().to_vec().unwrap();
1235 for (i, (orig, got)) in data.iter().zip(yv.iter()).enumerate() {
1236 assert!(got.is_finite(), "dequant[{i}] not finite: {got}");
1237 assert!(
1238 (orig - got).abs() < 0.01,
1239 "dequant[{i}]: {orig} vs {got}"
1240 );
1241 }
1242 }
1243
1244 fn quant_parity_on<B: Backend>(dev: &B::Device, tol: f32) {
1246 let (n_kv, n_q, d) = (2usize, 4usize, 32usize);
1247 let mut cfg_q = CacheConfig::paged(64);
1248 cfg_q.quantize_kv = true;
1249 let cfg_f = CacheConfig::paged(64);
1250 let scale = 1.0 / (d as f64).sqrt() * 0.9;
1251 let mut fp = PagedKVCache::<B>::new(1, cfg_f);
1252 let mut qn = PagedKVCache::<B>::new(1, cfg_q);
1253 for i in 0..24 {
1254 let (k, v) = kv_tok_on::<B>(dev, i, n_kv, d);
1255 let q = q_tok_on::<B>(dev, i, n_q, d);
1256 let a = fp.attention_opts(0, q.clone(), k.clone(), v.clone(), i, scale, None);
1257 let b = qn.attention_opts(0, q, k, v, i, scale, None);
1258 let av: Vec<f32> = a.into_data().convert::<f32>().to_vec().unwrap();
1259 let bv: Vec<f32> = b.into_data().convert::<f32>().to_vec().unwrap();
1260 for (j, (x, y)) in av.iter().zip(bv.iter()).enumerate() {
1261 assert!(y.is_finite(), "step {i}[{j}] not finite: {y}");
1262 assert!((x - y).abs() < tol, "step {i}[{j}]: {x} vs {y}");
1263 }
1264 }
1265 }
1266
1267 #[test]
1268 #[ignore = "gpu"]
1269 fn kv_quant_roundtrip_on_production_backend() {
1270 quant_roundtrip_on::<combs_core::CombsBackend>(&Default::default());
1271 }
1272
1273 #[test]
1274 #[ignore = "gpu"]
1275 fn quantized_paged_cache_matches_fp_on_production_backend() {
1276 quant_parity_on::<combs_core::CombsBackend>(&Default::default(), 5e-2);
1277 }
1278}