1#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum KvMode {
14 F32,
15 Q8 { k: bool, v: bool },
17}
18
19impl KvMode {
20 pub fn from_env() -> Self {
21 match std::env::var("CMF_KV").as_deref() {
22 Ok("q8") | Ok("q8_2f") => KvMode::Q8 { k: true, v: true },
23 Ok("q8k") => KvMode::Q8 { k: true, v: false },
24 Ok("q8v") => KvMode::Q8 { k: false, v: true },
25 _ => KvMode::F32,
26 }
27 }
28
29 fn quant_k(self) -> bool {
30 matches!(self, KvMode::Q8 { k: true, .. })
31 }
32
33 fn quant_v(self) -> bool {
34 matches!(self, KvMode::Q8 { v: true, .. })
35 }
36}
37
38const KV_COL_WARMUP: usize = 64;
41
42const KV_K_GROUP: usize = 32;
47
48#[derive(Debug, Clone)]
59pub enum O1State {
60 Collecting {
61 m: usize,
62 w: usize,
63 sink: usize,
64 rect: crate::nystrom::O1Rect,
65 q_buf: Vec<f32>,
67 },
68 Sealed { groups: Vec<crate::nystrom::NystromState> },
74}
75
76#[derive(Debug, Clone)]
78pub struct LayerKvCache {
79 pub mode: KvMode,
80 k: Vec<Vec<f32>>,
82 v: Vec<Vec<f32>>,
84 kq: Vec<Vec<i8>>,
86 ks: Vec<Vec<f32>>,
87 vq: Vec<Vec<i8>>,
88 vs: Vec<Vec<f32>>,
89 kcol: Vec<Vec<f32>>,
91 vcol: Vec<Vec<f32>>,
92 imp: Vec<f32>,
95 pub seq_len: usize,
97 pub num_kv_heads: usize,
98 pub head_dim: usize,
99 pub linear_state: Vec<f32>,
101 pub linear_scratch: Vec<f32>,
103 pub o1: Option<O1State>,
105}
106
107impl LayerKvCache {
108 pub fn new(num_kv_heads: usize, head_dim: usize) -> Self {
109 Self {
110 mode: KvMode::from_env(),
111 k: vec![Vec::new(); num_kv_heads],
112 v: vec![Vec::new(); num_kv_heads],
113 kq: vec![Vec::new(); num_kv_heads],
114 ks: vec![Vec::new(); num_kv_heads],
115 vq: vec![Vec::new(); num_kv_heads],
116 vs: vec![Vec::new(); num_kv_heads],
117 kcol: vec![Vec::new(); num_kv_heads],
118 vcol: vec![Vec::new(); num_kv_heads],
119 imp: Vec::new(),
120 seq_len: 0,
121 num_kv_heads,
122 head_dim,
123 linear_state: Vec::new(),
124 linear_scratch: Vec::new(),
125 o1: None,
126 }
127 }
128
129 pub fn o1_begin(&mut self, m: usize, w: usize, sink: usize, rect: crate::nystrom::O1Rect) {
133 self.o1 = Some(O1State::Collecting { m, w, sink, rect, q_buf: Vec::new() });
134 }
135
136 pub fn o1_push_q(&mut self, q_all: &[f32]) {
141 if let Some(O1State::Collecting { q_buf, .. }) = &mut self.o1 {
142 q_buf.extend_from_slice(q_all);
143 }
144 }
145
146 pub fn o1_sealed(&self) -> bool {
147 matches!(self.o1, Some(O1State::Sealed { .. }))
148 }
149
150 pub fn o1_seal(&mut self, num_heads: usize) -> bool {
156 if !matches!(self.o1, Some(O1State::Collecting { .. })) {
159 return self.o1_sealed();
160 }
161 let Some(O1State::Collecting { m, w, sink, rect, q_buf }) = self.o1.take() else {
162 unreachable!("checked above");
163 };
164 let (hd, t) = (self.head_dim, self.seq_len);
165 let nkv = self.num_kv_heads.max(1);
166 let hpk = num_heads / nkv;
167 let ok = t > 0
168 && self.mode == KvMode::F32
169 && q_buf.len() == t * num_heads * hd
170 && hpk * nkv == num_heads
171 && (0..self.num_kv_heads).all(|g| self.head_len(g) == t);
172 if !ok {
173 tracing::warn!(
174 "o1: cannot seal (needs f32 KV mode, dense heads, full query \
175 trace, num_heads divisible by num_kv_heads) — layer keeps \
176 exact attention"
177 );
178 return false;
179 }
180 let mut groups = Vec::with_capacity(self.num_kv_heads);
181 let mut qh = vec![0.0f32; hpk * t * hd];
184 for g in 0..self.num_kv_heads {
185 for hh in 0..hpk {
186 let h = g * hpk + hh;
187 for p in 0..t {
188 let src = (p * num_heads + h) * hd;
189 let dst = (hh * t + p) * hd;
190 qh[dst..dst + hd].copy_from_slice(&q_buf[src..src + hd]);
191 }
192 }
193 let qs: Vec<&[f32]> = (0..hpk).map(|hh| &qh[hh * t * hd..(hh + 1) * t * hd]).collect();
194 let mut st = crate::nystrom::NystromState::new_group(m, w, sink, hpk).with_rect(rect);
195 st.prefill_group(&qs, &self.k[g], &self.v[g], t, hd, hd);
196 groups.push(st);
197 }
198 for h in 0..self.num_kv_heads {
201 self.k[h] = Vec::new();
202 self.v[h] = Vec::new();
203 }
204 self.imp = Vec::new();
205 self.o1 = Some(O1State::Sealed { groups });
206 true
207 }
208
209 pub fn o1_step(
215 &mut self,
216 q_all: &[f32],
217 k_new: &[f32],
218 v_new: &[f32],
219 num_heads: usize,
220 ) -> Vec<f32> {
221 let hd = self.head_dim;
222 let hpk = num_heads / self.num_kv_heads.max(1);
223 let mut out = vec![0.0f32; num_heads * hd];
224 let Some(O1State::Sealed { groups }) = &mut self.o1 else {
225 debug_assert!(false, "o1_step on an unsealed layer");
226 return out;
227 };
228 for (g, st) in groups.iter_mut().enumerate() {
229 let (lo, hi) = (g * hpk * hd, (g + 1) * hpk * hd);
230 st.step_group(
231 &q_all[lo..hi],
232 &k_new[g * hd..(g + 1) * hd],
233 &v_new[g * hd..(g + 1) * hd],
234 &mut out[lo..hi],
235 );
236 }
237 self.seq_len += 1;
240 out
241 }
242
243 pub fn o1_memory_bytes(&self) -> usize {
246 match &self.o1 {
247 Some(O1State::Collecting { q_buf, .. }) => {
248 q_buf.len() * std::mem::size_of::<f32>()
249 }
250 Some(O1State::Sealed { groups }) => {
251 groups.iter().map(|s| s.memory_bytes()).sum()
252 }
253 None => 0,
254 }
255 }
256
257 fn quant_row(row: &[f32], col: &[f32], q: &mut Vec<i8>, sc: &mut Vec<f32>,
260 group: usize) {
261 let mut resid = vec![0.0f32; row.len()];
262 for (d, &x) in row.iter().enumerate() {
263 resid[d] = if col.is_empty() { x } else { x / col[d] };
264 }
265 for g0 in (0..row.len()).step_by(group) {
266 let g1 = (g0 + group).min(row.len());
267 let mut absmax = 0.0f32;
268 for &r in &resid[g0..g1] {
269 absmax = absmax.max(r.abs());
270 }
271 let s = (absmax / 127.0).max(1e-12);
272 sc.push(s);
273 for &r in &resid[g0..g1] {
274 q.push((r / s).round().clamp(-127.0, 127.0) as i8);
275 }
276 }
277 }
278
279 fn freeze_cols(&mut self) {
282 let hd = self.head_dim;
283 let ngk = hd.div_ceil(KV_K_GROUP);
284 for h in 0..self.num_kv_heads {
285 for (qv, sv, colv, group) in [
286 (&mut self.kq[h], &mut self.ks[h], &mut self.kcol[h], KV_K_GROUP),
287 (&mut self.vq[h], &mut self.vs[h], &mut self.vcol[h], hd),
288 ] {
289 let spp = if group == hd { 1 } else { ngk }; let n = sv.len() / spp;
291 if n == 0 {
292 continue;
293 }
294 let mut rows = vec![0.0f32; n * hd];
296 for p in 0..n {
297 for d in 0..hd {
298 rows[p * hd + d] =
299 qv[p * hd + d] as f32 * sv[p * spp + d / group];
300 }
301 }
302 let mut col = vec![0.0f32; hd];
303 for p in 0..n {
304 for d in 0..hd {
305 col[d] += rows[p * hd + d] * rows[p * hd + d];
306 }
307 }
308 for c in col.iter_mut() {
309 *c = (*c / n as f32).sqrt().max(1e-6);
310 }
311 qv.clear();
312 sv.clear();
313 for p in 0..n {
314 Self::quant_row(&rows[p * hd..(p + 1) * hd], &col, qv, sv, group);
315 }
316 *colv = col;
317 }
318 }
319 }
320
321 pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
325 debug_assert_eq!(k_new.len(), self.num_kv_heads * self.head_dim);
326 debug_assert_eq!(v_new.len(), self.num_kv_heads * self.head_dim);
327 if matches!(self.mode, KvMode::Q8 { .. })
333 && self.seq_len >= KV_COL_WARMUP
334 && self.kcol.iter().all(Vec::is_empty)
335 && self.vcol.iter().all(Vec::is_empty)
336 {
337 self.freeze_cols();
338 }
339 for h in 0..self.num_kv_heads {
340 if !alive.get(h).copied().unwrap_or(true) {
341 continue;
342 }
343 let s = h * self.head_dim;
344 if self.mode.quant_k() {
345 Self::quant_row(&k_new[s..s + self.head_dim],
346 &self.kcol[h], &mut self.kq[h], &mut self.ks[h],
347 KV_K_GROUP);
348 } else {
349 self.k[h].extend_from_slice(&k_new[s..s + self.head_dim]);
350 }
351 if self.mode.quant_v() {
352 Self::quant_row(&v_new[s..s + self.head_dim],
353 &self.vcol[h], &mut self.vq[h], &mut self.vs[h],
354 self.head_dim);
355 } else {
356 self.v[h].extend_from_slice(&v_new[s..s + self.head_dim]);
357 }
358 }
359 self.imp.push(0.0);
360 self.seq_len += 1;
361 }
362
363 pub fn attend(&self, q: &[f32], kv_head: usize) -> (Vec<f32>, Vec<f32>) {
368 let hd = self.head_dim;
369 if self.mode == KvMode::F32 {
370 let stored = self.k[kv_head].len() / hd;
371 return crate::attention::attention_head(
372 q, &self.k[kv_head], &self.v[kv_head], hd, stored);
373 }
374 let stored = self.head_len(kv_head);
375 let scale = 1.0 / (hd as f32).sqrt();
376 let mut scores = vec![0.0f32; stored];
377 if self.mode.quant_k() {
378 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
379 let kcol = &self.kcol[kv_head];
381 let mut qc = vec![0.0f32; hd];
382 for d in 0..hd {
383 qc[d] = if kcol.is_empty() { q[d] } else { q[d] * kcol[d] };
384 }
385 let ng = hd.div_ceil(KV_K_GROUP);
386 for p in 0..stored {
387 let row = &kq[p * hd..(p + 1) * hd];
388 let row_u8 = unsafe {
391 std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len())
392 };
393 let mut dot = 0.0f32;
394 for g in 0..ng {
395 let g0 = g * KV_K_GROUP;
396 let g1 = (g0 + KV_K_GROUP).min(hd);
397 dot += crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qc[g0..g1])
398 * ks[p * ng + g];
399 }
400 scores[p] = dot * scale;
401 }
402 } else {
403 let k = &self.k[kv_head];
404 for p in 0..stored {
405 let row = &k[p * hd..(p + 1) * hd];
406 scores[p] = crate::attention::dot_f32(q, row) * scale;
407 }
408 }
409 let max_score = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
410 let mut sum = 0.0f32;
411 for s in scores.iter_mut() {
412 *s = (*s - max_score).exp();
413 sum += *s;
414 }
415 if sum > 0.0 {
416 for s in scores.iter_mut() {
417 *s /= sum;
418 }
419 }
420 let mut acc = vec![0.0f32; hd];
421 if self.mode.quant_v() {
422 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
423 for p in 0..stored {
424 let w = scores[p] * vs[p];
425 if w.abs() < 1e-12 {
426 continue;
427 }
428 crate::qtensor::axpy_i8_f32(&mut acc, &vq[p * hd..(p + 1) * hd], w);
429 }
430 let vcol = &self.vcol[kv_head];
431 if !vcol.is_empty() {
432 for d in 0..hd {
433 acc[d] *= vcol[d];
434 }
435 }
436 } else {
437 let v = &self.v[kv_head];
438 for p in 0..stored {
439 let w = scores[p];
440 if w.abs() < 1e-12 {
441 continue;
442 }
443 crate::attention::axpy_f32(&mut acc, &v[p * hd..(p + 1) * hd], w);
444 }
445 }
446 (acc, scores)
447 }
448
449 #[allow(clippy::too_many_arguments)]
463 pub fn attend_group(
464 &self,
465 q_group: &[f32],
466 kv_head: usize,
467 out: &mut [f32],
468 imp_acc: &mut [f32],
469 scale: f32,
470 first: usize,
471 ) {
472 let hd = self.head_dim;
473 let nheads = q_group.len() / hd;
474 debug_assert_eq!(out.len(), nheads * hd);
475 let stored = if self.mode == KvMode::F32 {
476 self.k[kv_head].len() / hd
477 } else {
478 self.head_len(kv_head)
479 };
480 if stored == 0 {
481 out.fill(0.0);
482 return;
483 }
484 let first = first.min(stored.saturating_sub(1));
485
486 thread_local! {
487 static GQA_SCORES: std::cell::RefCell<Vec<f32>> =
489 const { std::cell::RefCell::new(Vec::new()) };
490 static GQA_QC: std::cell::RefCell<Vec<f32>> =
492 const { std::cell::RefCell::new(Vec::new()) };
493 }
494
495 GQA_SCORES.with(|sc| {
496 let mut scores = sc.borrow_mut();
497 if first > 0 {
498 scores.clear();
501 scores.resize(nheads * stored, f32::NEG_INFINITY);
502 } else {
503 scores.resize(nheads * stored, 0.0);
504 }
505
506 if self.mode.quant_k() {
508 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
509 let kcol = &self.kcol[kv_head];
510 let ng = hd.div_ceil(KV_K_GROUP);
511 GQA_QC.with(|qc| {
512 let mut qcb = qc.borrow_mut();
513 qcb.resize(nheads * hd, 0.0);
514 for h in 0..nheads {
515 for d in 0..hd {
516 let qv = q_group[h * hd + d];
517 qcb[h * hd + d] = if kcol.is_empty() { qv } else { qv * kcol[d] };
518 }
519 }
520 for p in first..stored {
521 let row = &kq[p * hd..(p + 1) * hd];
522 let row_u8 = unsafe {
525 std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len())
526 };
527 for h in 0..nheads {
528 let qch = &qcb[h * hd..(h + 1) * hd];
529 let mut dot = 0.0f32;
530 for g in 0..ng {
531 let g0 = g * KV_K_GROUP;
532 let g1 = (g0 + KV_K_GROUP).min(hd);
533 dot += crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qch[g0..g1])
534 * ks[p * ng + g];
535 }
536 scores[h * stored + p] = dot * scale;
537 }
538 }
539 });
540 } else {
541 let k = &self.k[kv_head];
542 for p in first..stored {
543 let row = &k[p * hd..(p + 1) * hd];
544 for h in 0..nheads {
545 scores[h * stored + p] =
546 crate::attention::dot_f32(&q_group[h * hd..(h + 1) * hd], row)
547 * scale;
548 }
549 }
550 }
551
552 for h in 0..nheads {
554 let s = &mut scores[h * stored..(h + 1) * stored];
555 let max_score = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
556 let mut sum = 0.0f32;
557 for v in s.iter_mut() {
558 *v = (*v - max_score).exp();
559 sum += *v;
560 }
561 if sum > 0.0 {
562 for v in s.iter_mut() {
563 *v /= sum;
564 }
565 }
566 }
567
568 out.fill(0.0);
570 if self.mode.quant_v() {
571 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
572 for p in first..stored {
573 let row = &vq[p * hd..(p + 1) * hd];
574 for h in 0..nheads {
575 let w = scores[h * stored + p] * vs[p];
576 if w.abs() < 1e-12 {
577 continue;
578 }
579 crate::qtensor::axpy_i8_f32(&mut out[h * hd..(h + 1) * hd], row, w);
580 }
581 }
582 let vcol = &self.vcol[kv_head];
583 if !vcol.is_empty() {
584 for h in 0..nheads {
585 for d in 0..hd {
586 out[h * hd + d] *= vcol[d];
587 }
588 }
589 }
590 } else {
591 let v = &self.v[kv_head];
592 for p in first..stored {
593 let row = &v[p * hd..(p + 1) * hd];
594 for h in 0..nheads {
595 let w = scores[h * stored + p];
596 if w.abs() < 1e-12 {
597 continue;
598 }
599 crate::attention::axpy_f32(&mut out[h * hd..(h + 1) * hd], row, w);
600 }
601 }
602 }
603
604 let n = imp_acc.len().min(stored);
607 for h in 0..nheads {
608 let s = &scores[h * stored..(h + 1) * stored];
609 for (dst, &p) in imp_acc[..n].iter_mut().zip(s) {
610 *dst += p;
611 }
612 }
613 });
614 }
615
616 #[cfg(target_arch = "aarch64")]
624 #[allow(clippy::too_many_arguments)]
625 pub fn attend_chunk(
626 &mut self,
627 q_all: &[f32],
628 b: usize,
629 s0: usize,
630 nh: usize,
631 heads_per_kv: usize,
632 hd: usize,
633 out: &mut [f32],
634 pool: Option<&crate::pool::Pool>,
635 scale: f32,
636 window: Option<usize>,
637 ) {
638 let n = s0 + b;
639 struct SendPtr(*mut f32);
640 unsafe impl Send for SendPtr {}
641 unsafe impl Sync for SendPtr {}
642 impl SendPtr {
643 fn at(&self, i: usize) -> *mut f32 {
644 unsafe { self.0.add(i) }
647 }
648 }
649 thread_local! {
650 static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>)> =
651 const { std::cell::RefCell::new((Vec::new(), Vec::new(), Vec::new(), Vec::new())) };
652 }
653 let neon_gemm = cfg!(not(target_os = "macos"))
658 || std::env::var("CMF_FORCE_NEON_GEMM").map(|v| v == "1").unwrap_or(false);
659 SCRATCH.with(|s| {
660 let mut s = s.borrow_mut();
661 let (qpanel, scores, aopanel, ktpack) = &mut *s;
662 let m = heads_per_kv * b;
667 qpanel.resize(m * hd, 0.0);
668 scores.resize(m * n, 0.0);
669 aopanel.resize(m * hd, 0.0);
670 for g in 0..self.num_kv_heads {
671 let kmat = &self.k[g];
672 let vmat = &self.v[g];
673 debug_assert_eq!(kmat.len(), n * hd);
674 for hl in 0..heads_per_kv {
675 let hh = g * heads_per_kv + hl;
676 for bi in 0..b {
677 qpanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]
678 .copy_from_slice(&q_all[bi * nh * hd + hh * hd..][..hd]);
679 }
680 }
681 if neon_gemm {
682 ktpack.resize(hd * n, 0.0);
683 for p in 0..n {
684 let row = &kmat[p * hd..(p + 1) * hd];
685 for (d, &v) in row.iter().enumerate() {
686 ktpack[d * n + p] = v;
687 }
688 }
689 let sp_q = SendPtr(qpanel.as_ptr() as *mut f32);
692 let sp_s = SendPtr(scores.as_mut_ptr());
693 let kt = &*ktpack;
694 let run = |start: usize, end: usize| {
695 if end > start {
696 let a = unsafe {
698 std::slice::from_raw_parts(sp_q.at(start * hd), (end - start) * hd)
699 };
700 let c = unsafe {
701 std::slice::from_raw_parts_mut(sp_s.at(start * n), (end - start) * n)
702 };
703 crate::qtensor::neon_gemm_rm(
704 end - start, n, hd, scale, a, hd, kt, n, false, c, n,
705 );
706 }
707 };
708 match pool {
709 Some(p) if m >= 64 => p.run_rows(m, &run),
710 _ => run(0, m),
711 }
712 } else {
713 crate::qtensor::sgemm_rm(m, n, hd, scale, qpanel, hd, kmat, hd, true, scores, n);
714 }
715 let sp = SendPtr(scores.as_mut_ptr());
717 let run = |start: usize, end: usize| {
718 for r in start..end {
719 let allowed = s0 + (r % b) + 1;
720 let lo = window.map(|w| allowed.saturating_sub(w)).unwrap_or(0);
724 let row = unsafe { std::slice::from_raw_parts_mut(sp.at(r * n), n) };
726 crate::attention::softmax_row(&mut row[lo..allowed]);
727 row[..lo].fill(0.0);
728 row[allowed..].fill(0.0);
729 }
730 };
731 match pool {
732 Some(p) if m >= 64 => p.run_rows(m, &run),
733 _ => run(0, m),
734 }
735 let ni = self.imp.len().min(n);
739 for r in 0..m {
740 let al = (s0 + (r % b) + 1).min(ni);
741 for (dst, &p) in self.imp[..al].iter_mut().zip(&scores[r * n..r * n + al]) {
742 *dst += p;
743 }
744 }
745 if neon_gemm {
746 let sp_s = SendPtr(scores.as_mut_ptr());
747 let sp_o = SendPtr(aopanel.as_mut_ptr());
748 let run = |start: usize, end: usize| {
749 if end > start {
750 let a = unsafe {
752 std::slice::from_raw_parts(sp_s.at(start * n), (end - start) * n)
753 };
754 let c = unsafe {
755 std::slice::from_raw_parts_mut(
756 sp_o.at(start * hd),
757 (end - start) * hd,
758 )
759 };
760 crate::qtensor::neon_gemm_rm(
761 end - start, hd, n, 1.0, a, n, vmat, hd, false, c, hd,
762 );
763 }
764 };
765 match pool {
766 Some(p) if m >= 64 => p.run_rows(m, &run),
767 _ => run(0, m),
768 }
769 } else {
770 crate::qtensor::sgemm_rm(m, hd, n, 1.0, scores, n, vmat, hd, false, aopanel, hd);
771 }
772 for hl in 0..heads_per_kv {
773 let hh = g * heads_per_kv + hl;
774 for bi in 0..b {
775 out[bi * nh * hd + hh * hd..][..hd]
776 .copy_from_slice(&aopanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]);
777 }
778 }
779 }
780 });
781 }
782
783 pub fn truncate_last(&mut self, n_drop: usize) {
785 let d = n_drop.min(self.seq_len);
786 for h in 0..self.num_kv_heads {
787 let keep = self.k[h].len().saturating_sub(d * self.head_dim);
788 self.k[h].truncate(keep);
789 self.v[h].truncate(keep);
790 let ngk = self.head_dim.div_ceil(KV_K_GROUP);
791 let keep_q = self.kq[h].len().saturating_sub(d * self.head_dim);
792 self.kq[h].truncate(keep_q);
793 let keep_vq = self.vq[h].len().saturating_sub(d * self.head_dim);
794 self.vq[h].truncate(keep_vq);
795 let keep_ks = self.ks[h].len().saturating_sub(d * ngk);
796 self.ks[h].truncate(keep_ks);
797 let keep_vs = self.vs[h].len().saturating_sub(d);
798 self.vs[h].truncate(keep_vs);
799 }
800 self.imp.truncate(self.imp.len().saturating_sub(d));
801 self.seq_len -= d;
802 }
803
804 pub fn accumulate_imp(&mut self, probs: &[f32]) {
806 for (dst, &p) in self.imp.iter_mut().zip(probs) {
807 *dst += p;
808 }
809 }
810
811 pub fn head_keys(&self, kv_head: usize) -> &[f32] {
813 &self.k[kv_head]
814 }
815
816 pub fn head_values(&self, kv_head: usize) -> &[f32] {
817 &self.v[kv_head]
818 }
819
820 pub fn head_len(&self, kv_head: usize) -> usize {
822 let ng = self.head_dim.div_ceil(KV_K_GROUP);
823 (self.k[kv_head].len() / self.head_dim)
824 .max(self.ks[kv_head].len() / ng)
825 .max(self.vs[kv_head].len())
826 }
827
828 pub fn clear(&mut self) {
830 for h in 0..self.num_kv_heads {
831 self.k[h].clear();
832 self.v[h].clear();
833 self.kq[h].clear();
834 self.ks[h].clear();
835 self.vq[h].clear();
836 self.vs[h].clear();
837 self.kcol[h].clear();
838 self.vcol[h].clear();
839 }
840 self.imp.clear();
841 self.linear_state.clear();
842 self.linear_scratch.clear();
843 self.o1 = None;
846 self.seq_len = 0;
847 }
848
849 pub fn memory_bytes(&self) -> usize {
851 let floats: usize = self.k.iter().map(Vec::len).sum::<usize>()
852 + self.v.iter().map(Vec::len).sum::<usize>()
853 + self.ks.iter().map(Vec::len).sum::<usize>()
854 + self.vs.iter().map(Vec::len).sum::<usize>()
855 + self.kcol.iter().map(Vec::len).sum::<usize>()
856 + self.vcol.iter().map(Vec::len).sum::<usize>();
857 let bytes: usize = self.kq.iter().map(Vec::len).sum::<usize>()
858 + self.vq.iter().map(Vec::len).sum::<usize>();
859 floats * std::mem::size_of::<f32>()
860 + bytes
861 + self.linear_state.len() * std::mem::size_of::<f32>()
865 + self.o1_memory_bytes()
868 }
869
870 fn evict(&mut self, keep_last: usize) {
872 if self.o1_sealed() || self.seq_len <= keep_last {
876 return;
877 }
878 let drop = self.seq_len - keep_last;
879 for h in 0..self.num_kv_heads {
880 let stored = self.head_len(h);
882 let d = drop.min(stored);
883 let hd = self.head_dim;
884 fn drop_front<T>(v: &mut Vec<T>, n: usize) {
885 let n = n.min(v.len());
886 v.drain(..n);
887 }
888 drop_front(&mut self.k[h], d * hd);
889 drop_front(&mut self.v[h], d * hd);
890 drop_front(&mut self.kq[h], d * hd);
891 drop_front(&mut self.vq[h], d * hd);
892 drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
893 drop_front(&mut self.vs[h], d);
894 }
895 let d = drop.min(self.imp.len());
896 self.imp.drain(..d);
897 self.seq_len = keep_last;
898 }
899
900 fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
905 if self.o1_sealed() {
906 return; }
908 let stored = self.imp.len();
909 if stored <= keep_last {
910 return;
911 }
912 let sink_n = sink.min(keep_last);
915 let recent_n = recent.min(keep_last - sink_n);
916 let mut keep = vec![false; stored];
917 for k in keep.iter_mut().take(sink_n) {
918 *k = true;
919 }
920 for k in keep.iter_mut().skip(stored.saturating_sub(recent_n)) {
921 *k = true;
922 }
923 let mut budget = keep_last.saturating_sub(keep.iter().filter(|&&x| x).count());
924 let mut order: Vec<usize> = (0..stored).filter(|&i| !keep[i]).collect();
926 order.sort_by(|&a, &b| {
927 self.imp[b].partial_cmp(&self.imp[a]).unwrap_or(std::cmp::Ordering::Equal)
928 });
929 for i in order {
930 if budget == 0 {
931 break;
932 }
933 keep[i] = true;
934 budget -= 1;
935 }
936
937 let kept: Vec<usize> = (0..stored).filter(|&i| keep[i]).collect();
938 let hd = self.head_dim;
939 fn gather<T: Copy>(src: &[T], kept: &[usize], step: usize) -> Vec<T> {
940 let mut out = Vec::with_capacity(kept.len() * step);
941 for &i in kept {
942 out.extend_from_slice(&src[i * step..(i + 1) * step]);
943 }
944 out
945 }
946 for h in 0..self.num_kv_heads {
951 if !self.k[h].is_empty() {
952 self.k[h] = gather(&self.k[h], &kept, hd);
953 }
954 if !self.v[h].is_empty() {
955 self.v[h] = gather(&self.v[h], &kept, hd);
956 }
957 if !self.kq[h].is_empty() {
958 self.kq[h] = gather(&self.kq[h], &kept, hd);
959 self.ks[h] = gather(&self.ks[h], &kept, hd.div_ceil(KV_K_GROUP));
960 }
961 if !self.vq[h].is_empty() {
962 self.vq[h] = gather(&self.vq[h], &kept, hd);
963 self.vs[h] = gather(&self.vs[h], &kept, 1);
964 }
965 }
966 self.imp = kept.iter().map(|&i| self.imp[i]).collect();
967 self.seq_len = kept.len();
968 }
969}
970
971#[derive(Debug, Clone, Copy, PartialEq, Eq)]
973pub enum EvictionPolicy {
974 Recent,
976 Born { sink: usize },
978}
979
980#[derive(Debug)]
982pub struct KvCache {
983 pub layers: Vec<LayerKvCache>,
984 pub max_seq_len: usize,
985 pub policy: EvictionPolicy,
986}
987
988impl KvCache {
989 pub fn new(num_layers: usize, num_kv_heads: usize, head_dim: usize, max_seq_len: usize) -> Self {
990 let layers = (0..num_layers)
991 .map(|_| LayerKvCache::new(num_kv_heads, head_dim))
992 .collect();
993 Self {
994 layers,
995 max_seq_len,
996 policy: EvictionPolicy::Born { sink: 4 },
997 }
998 }
999
1000 pub fn clear(&mut self) {
1001 for layer in &mut self.layers {
1002 layer.clear();
1003 }
1004 }
1005
1006 pub fn total_memory_bytes(&self) -> usize {
1007 self.layers.iter().map(|l| l.memory_bytes()).sum()
1008 }
1009
1010 pub fn seq_len(&self) -> usize {
1012 self.layers.iter().map(|l| l.seq_len).max().unwrap_or(0)
1013 }
1014
1015 pub fn needs_eviction(&self) -> bool {
1016 self.seq_len() >= self.max_seq_len
1017 }
1018
1019 pub fn evict(&mut self, keep_last: usize) {
1021 match self.policy {
1022 EvictionPolicy::Recent => {
1023 for layer in &mut self.layers {
1024 layer.evict(keep_last);
1025 }
1026 }
1027 EvictionPolicy::Born { sink } => {
1028 let recent = (keep_last / 2).max(1);
1029 for layer in &mut self.layers {
1030 layer.evict_born(keep_last, sink, recent);
1031 }
1032 }
1033 }
1034 }
1035}
1036
1037#[cfg(test)]
1038mod tests {
1039 use super::*;
1040
1041 #[test]
1042 fn append_tracks_seq_len_and_layout() {
1043 let mut cache = LayerKvCache::new(4, 8);
1044 cache.mode = KvMode::F32;
1045 assert_eq!(cache.seq_len, 0);
1046
1047 let k: Vec<f32> = (0..32).map(|i| i as f32).collect();
1048 let v = vec![2.0f32; 32];
1049 cache.append(&k, &v, &[true; 4]);
1050
1051 assert_eq!(cache.seq_len, 1);
1052 assert_eq!(cache.head_len(0), 1);
1053 assert_eq!(cache.head_keys(1), &k[8..16]);
1055 assert_eq!(cache.memory_bytes(), 256);
1056 }
1057
1058 #[test]
1059 fn dead_head_stores_nothing() {
1060 let mut cache = LayerKvCache::new(2, 4);
1061 cache.mode = KvMode::F32;
1062 let k = vec![1.0f32; 8];
1063 let v = vec![2.0f32; 8];
1064 cache.append(&k, &v, &[true, false]);
1065 cache.append(&k, &v, &[true, false]);
1066
1067 assert_eq!(cache.seq_len, 2);
1068 assert_eq!(cache.head_len(0), 2);
1069 assert_eq!(cache.head_len(1), 0, "dead head must not store KV");
1070 assert_eq!(cache.memory_bytes(), 2 * 2 * 4 * 4);
1071 }
1072
1073 #[test]
1074 fn eviction_keeps_recent() {
1075 let mut cache = KvCache::new(2, 4, 8, 10);
1076 cache.policy = EvictionPolicy::Recent;
1077 for l in &mut cache.layers { l.mode = KvMode::F32; }
1078 let k = vec![1.0f32; 32];
1079 let v = vec![2.0f32; 32];
1080 for _ in 0..8 {
1081 for layer in &mut cache.layers {
1082 layer.append(&k, &v, &[true; 4]);
1083 }
1084 }
1085 assert_eq!(cache.seq_len(), 8);
1086 assert!(!cache.needs_eviction());
1087
1088 cache.evict(4);
1089 assert_eq!(cache.seq_len(), 4);
1090 assert_eq!(cache.layers[0].head_len(0), 4);
1091 }
1092
1093 #[test]
1094 fn truncate_rolls_back_speculative_positions() {
1095 let mut cache = LayerKvCache::new(2, 4);
1096 cache.mode = KvMode::F32;
1097 for pos in 0..5 {
1098 let k = vec![pos as f32; 8];
1099 let v = vec![pos as f32; 8];
1100 cache.append(&k, &v, &[true; 2]);
1101 }
1102 cache.truncate_last(2);
1103 assert_eq!(cache.seq_len, 3);
1104 assert_eq!(cache.head_len(0), 3);
1105 assert_eq!(cache.head_keys(0)[2 * 4], 2.0, "position 2 survives");
1106 }
1107
1108 #[test]
1112 fn q8_attend_matches_f32_within_grid() {
1113 let (heads, hd) = (2, 32);
1114 let mut f = LayerKvCache::new(heads, hd);
1115 f.mode = KvMode::F32;
1116 let mut q8 = LayerKvCache::new(heads, hd);
1117 q8.mode = KvMode::Q8 { k: true, v: true };
1118
1119 let synth = |p: usize, salt: usize| -> Vec<f32> {
1120 (0..heads * hd)
1121 .map(|i| {
1122 let x = ((i * 31 + p * 17 + salt * 7 + 3) % 97) as f32 / 97.0 - 0.5;
1123 if i % 2 == 0 { x * 4.0 } else { x * 0.25 }
1125 })
1126 .collect()
1127 };
1128 for p in 0..100 {
1129 let k = synth(p, 1);
1130 let v = synth(p, 2);
1131 f.append(&k, &v, &[true; 2]);
1132 q8.append(&k, &v, &[true; 2]);
1133 }
1134 let q: Vec<f32> = (0..hd).map(|i| ((i * 13 + 5) % 89) as f32 / 89.0 - 0.5).collect();
1135 for g in 0..heads {
1136 let (of, pf) = f.attend(&q, g);
1137 let (o8, p8) = q8.attend(&q, g);
1138 let scale = of.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-6);
1139 for d in 0..hd {
1140 assert!(
1141 (of[d] - o8[d]).abs() <= scale * 0.03 + 1e-3,
1142 "g{g} d{d}: f32 {} vs q8 {}", of[d], o8[d]
1143 );
1144 }
1145 for p in 0..100 {
1146 assert!((pf[p] - p8[p]).abs() < 0.02, "prob p{p}");
1147 }
1148 }
1149 q8.truncate_last(30);
1151 assert_eq!(q8.head_len(0), 70);
1152 let imp: Vec<f32> = (0..70).map(|i| i as f32).collect();
1153 q8.accumulate_imp(&imp);
1154 q8.evict_born(20, 2, 8);
1155 assert_eq!(q8.head_len(0), 20);
1156 let (o, _) = q8.attend(&q, 0);
1157 assert!(o.iter().all(|x| x.is_finite()));
1158 assert!(q8.memory_bytes() * 3 < f.memory_bytes());
1160 }
1161
1162 #[test]
1165 fn attend_group_equals_per_head_attend_bitexact() {
1166 let (kv_heads, hd, hpk) = (2usize, 32usize, 3usize); for mode in [KvMode::F32, KvMode::Q8 { k: true, v: true }] {
1168 let mut c = LayerKvCache::new(kv_heads, hd);
1169 c.mode = mode;
1170 for p in 0..70 {
1171 let k: Vec<f32> = (0..kv_heads * hd)
1172 .map(|i| ((i * 31 + p * 17 + 3) % 97) as f32 / 97.0 - 0.5)
1173 .collect();
1174 let v: Vec<f32> = (0..kv_heads * hd)
1175 .map(|i| ((i * 13 + p * 29 + 7) % 89) as f32 / 89.0 - 0.5)
1176 .collect();
1177 c.append(&k, &v, &[true; 2]);
1178 }
1179 let q: Vec<f32> = (0..kv_heads * hpk * hd)
1180 .map(|i| ((i * 11 + 5) % 83) as f32 / 83.0 - 0.5)
1181 .collect();
1182 for g in 0..kv_heads {
1183 let span = g * hpk * hd..(g + 1) * hpk * hd;
1184 let mut out = vec![0f32; hpk * hd];
1185 let mut imp = vec![0f32; 70];
1186 c.attend_group(
1187 &q[span.clone()],
1188 g,
1189 &mut out,
1190 &mut imp,
1191 1.0 / (hd as f32).sqrt(),
1192 0,
1193 );
1194 let mut imp_ref = vec![0f32; 70];
1195 for h in 0..hpk {
1196 let qh = &q[span.start + h * hd..span.start + (h + 1) * hd];
1197 let (o, probs) = c.attend(qh, g);
1198 assert_eq!(
1199 &out[h * hd..(h + 1) * hd],
1200 &o[..],
1201 "mode {mode:?} g{g} h{h}: grouped attend must be bit-identical"
1202 );
1203 for (dst, &p) in imp_ref.iter_mut().zip(&probs) {
1204 *dst += p;
1205 }
1206 }
1207 assert_eq!(imp, imp_ref, "mode {mode:?} g{g}: Born mass must match");
1208 }
1209 }
1210 }
1211
1212 #[test]
1216 fn born_eviction_mixed_modes_stay_consistent() {
1217 for (mk, mv) in [(false, true), (true, false)] {
1218 let mut c = LayerKvCache::new(1, 4);
1219 c.mode = KvMode::Q8 { k: mk, v: mv };
1220 for p in 0..80 {
1221 let k = vec![p as f32 * 0.01; 4];
1222 let v = vec![p as f32; 4];
1223 c.append(&k, &v, &[true]);
1224 }
1225 let imp: Vec<f32> = (0..80).map(|i| i as f32).collect();
1226 c.accumulate_imp(&imp);
1227 let before = c.memory_bytes();
1228 c.evict_born(20, 4, 8); assert_eq!(c.head_len(0), 20, "k={mk} v={mv}");
1230 assert!(c.memory_bytes() < before / 2,
1231 "memory must shrink (k={mk} v={mv})");
1232 let (out, _) = c.attend(&[1.0, 1.0, 1.0, 1.0], 0);
1235 assert!(out[0] > 30.0,
1236 "V from the kept tail, not the stale head (k={mk} v={mv}, out {})",
1237 out[0]);
1238 }
1239 }
1240
1241 #[test]
1242 fn born_eviction_keeps_high_mass_position() {
1243 let mut cache = KvCache::new(1, 1, 2, 16);
1244 cache.policy = EvictionPolicy::Born { sink: 1 };
1245 for l in &mut cache.layers { l.mode = KvMode::F32; }
1246 let layer = &mut cache.layers[0];
1247 for pos in 0..8 {
1250 let k = vec![pos as f32; 2];
1251 let v = vec![pos as f32 + 100.0; 2];
1252 layer.append(&k, &v, &[true]);
1253 }
1254 let mut imp = vec![0.05f32; 8];
1256 imp[3] = 5.0;
1257 layer.accumulate_imp(&imp);
1258
1259 cache.evict(4); let layer = &cache.layers[0];
1261 assert_eq!(layer.seq_len, 4);
1262 let kept_keys: Vec<f32> = (0..4).map(|i| layer.head_keys(0)[i * 2]).collect();
1263 assert_eq!(
1264 kept_keys,
1265 vec![0.0, 3.0, 6.0, 7.0],
1266 "kept = sink(0) + Born-top(3) + recent(6,7)"
1267 );
1268 assert_eq!(layer.head_len(0), 4);
1270 }
1271}