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)]
50pub struct LayerKvCache {
51 pub mode: KvMode,
52 k: Vec<Vec<f32>>,
54 v: Vec<Vec<f32>>,
56 kq: Vec<Vec<i8>>,
58 ks: Vec<Vec<f32>>,
59 vq: Vec<Vec<i8>>,
60 vs: Vec<Vec<f32>>,
61 kcol: Vec<Vec<f32>>,
63 vcol: Vec<Vec<f32>>,
64 imp: Vec<f32>,
67 pub seq_len: usize,
69 pub num_kv_heads: usize,
70 pub head_dim: usize,
71 pub linear_state: Vec<f64>,
73 pub linear_scratch: Vec<f64>,
75}
76
77impl LayerKvCache {
78 pub fn new(num_kv_heads: usize, head_dim: usize) -> Self {
79 Self {
80 mode: KvMode::from_env(),
81 k: vec![Vec::new(); num_kv_heads],
82 v: vec![Vec::new(); num_kv_heads],
83 kq: vec![Vec::new(); num_kv_heads],
84 ks: vec![Vec::new(); num_kv_heads],
85 vq: vec![Vec::new(); num_kv_heads],
86 vs: vec![Vec::new(); num_kv_heads],
87 kcol: vec![Vec::new(); num_kv_heads],
88 vcol: vec![Vec::new(); num_kv_heads],
89 imp: Vec::new(),
90 seq_len: 0,
91 num_kv_heads,
92 head_dim,
93 linear_state: Vec::new(),
94 linear_scratch: Vec::new(),
95 }
96 }
97
98 fn quant_row(row: &[f32], col: &[f32], q: &mut Vec<i8>, sc: &mut Vec<f32>,
101 group: usize) {
102 let mut resid = vec![0.0f32; row.len()];
103 for (d, &x) in row.iter().enumerate() {
104 resid[d] = if col.is_empty() { x } else { x / col[d] };
105 }
106 for g0 in (0..row.len()).step_by(group) {
107 let g1 = (g0 + group).min(row.len());
108 let mut absmax = 0.0f32;
109 for &r in &resid[g0..g1] {
110 absmax = absmax.max(r.abs());
111 }
112 let s = (absmax / 127.0).max(1e-12);
113 sc.push(s);
114 for &r in &resid[g0..g1] {
115 q.push((r / s).round().clamp(-127.0, 127.0) as i8);
116 }
117 }
118 }
119
120 fn freeze_cols(&mut self) {
123 let hd = self.head_dim;
124 let ngk = hd.div_ceil(KV_K_GROUP);
125 for h in 0..self.num_kv_heads {
126 for (qv, sv, colv, group) in [
127 (&mut self.kq[h], &mut self.ks[h], &mut self.kcol[h], KV_K_GROUP),
128 (&mut self.vq[h], &mut self.vs[h], &mut self.vcol[h], hd),
129 ] {
130 let spp = if group == hd { 1 } else { ngk }; let n = sv.len() / spp;
132 if n == 0 {
133 continue;
134 }
135 let mut rows = vec![0.0f32; n * hd];
137 for p in 0..n {
138 for d in 0..hd {
139 rows[p * hd + d] =
140 qv[p * hd + d] as f32 * sv[p * spp + d / group];
141 }
142 }
143 let mut col = vec![0.0f32; hd];
144 for p in 0..n {
145 for d in 0..hd {
146 col[d] += rows[p * hd + d] * rows[p * hd + d];
147 }
148 }
149 for c in col.iter_mut() {
150 *c = (*c / n as f32).sqrt().max(1e-6);
151 }
152 qv.clear();
153 sv.clear();
154 for p in 0..n {
155 Self::quant_row(&rows[p * hd..(p + 1) * hd], &col, qv, sv, group);
156 }
157 *colv = col;
158 }
159 }
160 }
161
162 pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
166 debug_assert_eq!(k_new.len(), self.num_kv_heads * self.head_dim);
167 debug_assert_eq!(v_new.len(), self.num_kv_heads * self.head_dim);
168 if matches!(self.mode, KvMode::Q8 { .. })
174 && self.seq_len >= KV_COL_WARMUP
175 && self.kcol.iter().all(Vec::is_empty)
176 && self.vcol.iter().all(Vec::is_empty)
177 {
178 self.freeze_cols();
179 }
180 for h in 0..self.num_kv_heads {
181 if !alive.get(h).copied().unwrap_or(true) {
182 continue;
183 }
184 let s = h * self.head_dim;
185 if self.mode.quant_k() {
186 Self::quant_row(&k_new[s..s + self.head_dim],
187 &self.kcol[h], &mut self.kq[h], &mut self.ks[h],
188 KV_K_GROUP);
189 } else {
190 self.k[h].extend_from_slice(&k_new[s..s + self.head_dim]);
191 }
192 if self.mode.quant_v() {
193 Self::quant_row(&v_new[s..s + self.head_dim],
194 &self.vcol[h], &mut self.vq[h], &mut self.vs[h],
195 self.head_dim);
196 } else {
197 self.v[h].extend_from_slice(&v_new[s..s + self.head_dim]);
198 }
199 }
200 self.imp.push(0.0);
201 self.seq_len += 1;
202 }
203
204 pub fn attend(&self, q: &[f32], kv_head: usize) -> (Vec<f32>, Vec<f32>) {
209 let hd = self.head_dim;
210 if self.mode == KvMode::F32 {
211 let stored = self.k[kv_head].len() / hd;
212 return crate::attention::attention_head(
213 q, &self.k[kv_head], &self.v[kv_head], hd, stored);
214 }
215 let stored = self.head_len(kv_head);
216 let scale = 1.0 / (hd as f32).sqrt();
217 let mut scores = vec![0.0f32; stored];
218 if self.mode.quant_k() {
219 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
220 let kcol = &self.kcol[kv_head];
222 let mut qc = vec![0.0f32; hd];
223 for d in 0..hd {
224 qc[d] = if kcol.is_empty() { q[d] } else { q[d] * kcol[d] };
225 }
226 let ng = hd.div_ceil(KV_K_GROUP);
227 for p in 0..stored {
228 let row = &kq[p * hd..(p + 1) * hd];
229 let mut dot = 0.0f32;
230 for g in 0..ng {
231 let g0 = g * KV_K_GROUP;
232 let g1 = (g0 + KV_K_GROUP).min(hd);
233 let mut gd = 0.0f32;
234 for d in g0..g1 {
235 gd += qc[d] * row[d] as f32;
236 }
237 dot += gd * ks[p * ng + g];
238 }
239 scores[p] = dot * scale;
240 }
241 } else {
242 let k = &self.k[kv_head];
243 for p in 0..stored {
244 let row = &k[p * hd..(p + 1) * hd];
245 let mut dot = 0.0f32;
246 for d in 0..hd {
247 dot += q[d] * row[d];
248 }
249 scores[p] = dot * scale;
250 }
251 }
252 let max_score = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
253 let mut sum = 0.0f32;
254 for s in scores.iter_mut() {
255 *s = (*s - max_score).exp();
256 sum += *s;
257 }
258 if sum > 0.0 {
259 for s in scores.iter_mut() {
260 *s /= sum;
261 }
262 }
263 let mut acc = vec![0.0f32; hd];
264 if self.mode.quant_v() {
265 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
266 for p in 0..stored {
267 let w = scores[p] * vs[p];
268 if w.abs() < 1e-12 {
269 continue;
270 }
271 let row = &vq[p * hd..(p + 1) * hd];
272 for d in 0..hd {
273 acc[d] += w * row[d] as f32;
274 }
275 }
276 let vcol = &self.vcol[kv_head];
277 if !vcol.is_empty() {
278 for d in 0..hd {
279 acc[d] *= vcol[d];
280 }
281 }
282 } else {
283 let v = &self.v[kv_head];
284 for p in 0..stored {
285 let w = scores[p];
286 if w.abs() < 1e-12 {
287 continue;
288 }
289 let row = &v[p * hd..(p + 1) * hd];
290 for d in 0..hd {
291 acc[d] += w * row[d];
292 }
293 }
294 }
295 (acc, scores)
296 }
297
298 pub fn truncate_last(&mut self, n_drop: usize) {
300 let d = n_drop.min(self.seq_len);
301 for h in 0..self.num_kv_heads {
302 let keep = self.k[h].len().saturating_sub(d * self.head_dim);
303 self.k[h].truncate(keep);
304 self.v[h].truncate(keep);
305 let ngk = self.head_dim.div_ceil(KV_K_GROUP);
306 let keep_q = self.kq[h].len().saturating_sub(d * self.head_dim);
307 self.kq[h].truncate(keep_q);
308 let keep_vq = self.vq[h].len().saturating_sub(d * self.head_dim);
309 self.vq[h].truncate(keep_vq);
310 let keep_ks = self.ks[h].len().saturating_sub(d * ngk);
311 self.ks[h].truncate(keep_ks);
312 let keep_vs = self.vs[h].len().saturating_sub(d);
313 self.vs[h].truncate(keep_vs);
314 }
315 self.imp.truncate(self.imp.len().saturating_sub(d));
316 self.seq_len -= d;
317 }
318
319 pub fn accumulate_imp(&mut self, probs: &[f32]) {
321 for (dst, &p) in self.imp.iter_mut().zip(probs) {
322 *dst += p;
323 }
324 }
325
326 pub fn head_keys(&self, kv_head: usize) -> &[f32] {
328 &self.k[kv_head]
329 }
330
331 pub fn head_values(&self, kv_head: usize) -> &[f32] {
332 &self.v[kv_head]
333 }
334
335 pub fn head_len(&self, kv_head: usize) -> usize {
337 let ng = self.head_dim.div_ceil(KV_K_GROUP);
338 (self.k[kv_head].len() / self.head_dim)
339 .max(self.ks[kv_head].len() / ng)
340 .max(self.vs[kv_head].len())
341 }
342
343 pub fn clear(&mut self) {
345 for h in 0..self.num_kv_heads {
346 self.k[h].clear();
347 self.v[h].clear();
348 self.kq[h].clear();
349 self.ks[h].clear();
350 self.vq[h].clear();
351 self.vs[h].clear();
352 self.kcol[h].clear();
353 self.vcol[h].clear();
354 }
355 self.imp.clear();
356 self.linear_state.clear();
357 self.linear_scratch.clear();
358 self.seq_len = 0;
359 }
360
361 pub fn memory_bytes(&self) -> usize {
363 let floats: usize = self.k.iter().map(Vec::len).sum::<usize>()
364 + self.v.iter().map(Vec::len).sum::<usize>()
365 + self.ks.iter().map(Vec::len).sum::<usize>()
366 + self.vs.iter().map(Vec::len).sum::<usize>()
367 + self.kcol.iter().map(Vec::len).sum::<usize>()
368 + self.vcol.iter().map(Vec::len).sum::<usize>();
369 let bytes: usize = self.kq.iter().map(Vec::len).sum::<usize>()
370 + self.vq.iter().map(Vec::len).sum::<usize>();
371 floats * std::mem::size_of::<f32>() + bytes
372 }
373
374 fn evict(&mut self, keep_last: usize) {
376 if self.seq_len <= keep_last {
377 return;
378 }
379 let drop = self.seq_len - keep_last;
380 for h in 0..self.num_kv_heads {
381 let stored = self.head_len(h);
383 let d = drop.min(stored);
384 let hd = self.head_dim;
385 fn drop_front<T>(v: &mut Vec<T>, n: usize) {
386 let n = n.min(v.len());
387 v.drain(..n);
388 }
389 drop_front(&mut self.k[h], d * hd);
390 drop_front(&mut self.v[h], d * hd);
391 drop_front(&mut self.kq[h], d * hd);
392 drop_front(&mut self.vq[h], d * hd);
393 drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
394 drop_front(&mut self.vs[h], d);
395 }
396 let d = drop.min(self.imp.len());
397 self.imp.drain(..d);
398 self.seq_len = keep_last;
399 }
400
401 fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
406 let stored = self.imp.len();
407 if stored <= keep_last {
408 return;
409 }
410 let sink_n = sink.min(keep_last);
413 let recent_n = recent.min(keep_last - sink_n);
414 let mut keep = vec![false; stored];
415 for k in keep.iter_mut().take(sink_n) {
416 *k = true;
417 }
418 for k in keep.iter_mut().skip(stored.saturating_sub(recent_n)) {
419 *k = true;
420 }
421 let mut budget = keep_last.saturating_sub(keep.iter().filter(|&&x| x).count());
422 let mut order: Vec<usize> = (0..stored).filter(|&i| !keep[i]).collect();
424 order.sort_by(|&a, &b| {
425 self.imp[b].partial_cmp(&self.imp[a]).unwrap_or(std::cmp::Ordering::Equal)
426 });
427 for i in order {
428 if budget == 0 {
429 break;
430 }
431 keep[i] = true;
432 budget -= 1;
433 }
434
435 let kept: Vec<usize> = (0..stored).filter(|&i| keep[i]).collect();
436 let hd = self.head_dim;
437 fn gather<T: Copy>(src: &[T], kept: &[usize], step: usize) -> Vec<T> {
438 let mut out = Vec::with_capacity(kept.len() * step);
439 for &i in kept {
440 out.extend_from_slice(&src[i * step..(i + 1) * step]);
441 }
442 out
443 }
444 for h in 0..self.num_kv_heads {
449 if !self.k[h].is_empty() {
450 self.k[h] = gather(&self.k[h], &kept, hd);
451 }
452 if !self.v[h].is_empty() {
453 self.v[h] = gather(&self.v[h], &kept, hd);
454 }
455 if !self.kq[h].is_empty() {
456 self.kq[h] = gather(&self.kq[h], &kept, hd);
457 self.ks[h] = gather(&self.ks[h], &kept, hd.div_ceil(KV_K_GROUP));
458 }
459 if !self.vq[h].is_empty() {
460 self.vq[h] = gather(&self.vq[h], &kept, hd);
461 self.vs[h] = gather(&self.vs[h], &kept, 1);
462 }
463 }
464 self.imp = kept.iter().map(|&i| self.imp[i]).collect();
465 self.seq_len = kept.len();
466 }
467}
468
469#[derive(Debug, Clone, Copy, PartialEq, Eq)]
471pub enum EvictionPolicy {
472 Recent,
474 Born { sink: usize },
476}
477
478#[derive(Debug)]
480pub struct KvCache {
481 pub layers: Vec<LayerKvCache>,
482 pub max_seq_len: usize,
483 pub policy: EvictionPolicy,
484}
485
486impl KvCache {
487 pub fn new(num_layers: usize, num_kv_heads: usize, head_dim: usize, max_seq_len: usize) -> Self {
488 let layers = (0..num_layers)
489 .map(|_| LayerKvCache::new(num_kv_heads, head_dim))
490 .collect();
491 Self {
492 layers,
493 max_seq_len,
494 policy: EvictionPolicy::Born { sink: 4 },
495 }
496 }
497
498 pub fn clear(&mut self) {
499 for layer in &mut self.layers {
500 layer.clear();
501 }
502 }
503
504 pub fn total_memory_bytes(&self) -> usize {
505 self.layers.iter().map(|l| l.memory_bytes()).sum()
506 }
507
508 pub fn seq_len(&self) -> usize {
510 self.layers.iter().map(|l| l.seq_len).max().unwrap_or(0)
511 }
512
513 pub fn needs_eviction(&self) -> bool {
514 self.seq_len() >= self.max_seq_len
515 }
516
517 pub fn evict(&mut self, keep_last: usize) {
519 match self.policy {
520 EvictionPolicy::Recent => {
521 for layer in &mut self.layers {
522 layer.evict(keep_last);
523 }
524 }
525 EvictionPolicy::Born { sink } => {
526 let recent = (keep_last / 2).max(1);
527 for layer in &mut self.layers {
528 layer.evict_born(keep_last, sink, recent);
529 }
530 }
531 }
532 }
533}
534
535#[cfg(test)]
536mod tests {
537 use super::*;
538
539 #[test]
540 fn append_tracks_seq_len_and_layout() {
541 let mut cache = LayerKvCache::new(4, 8);
542 cache.mode = KvMode::F32;
543 assert_eq!(cache.seq_len, 0);
544
545 let k: Vec<f32> = (0..32).map(|i| i as f32).collect();
546 let v = vec![2.0f32; 32];
547 cache.append(&k, &v, &[true; 4]);
548
549 assert_eq!(cache.seq_len, 1);
550 assert_eq!(cache.head_len(0), 1);
551 assert_eq!(cache.head_keys(1), &k[8..16]);
553 assert_eq!(cache.memory_bytes(), 256);
554 }
555
556 #[test]
557 fn dead_head_stores_nothing() {
558 let mut cache = LayerKvCache::new(2, 4);
559 cache.mode = KvMode::F32;
560 let k = vec![1.0f32; 8];
561 let v = vec![2.0f32; 8];
562 cache.append(&k, &v, &[true, false]);
563 cache.append(&k, &v, &[true, false]);
564
565 assert_eq!(cache.seq_len, 2);
566 assert_eq!(cache.head_len(0), 2);
567 assert_eq!(cache.head_len(1), 0, "dead head must not store KV");
568 assert_eq!(cache.memory_bytes(), 2 * 2 * 4 * 4);
569 }
570
571 #[test]
572 fn eviction_keeps_recent() {
573 let mut cache = KvCache::new(2, 4, 8, 10);
574 cache.policy = EvictionPolicy::Recent;
575 for l in &mut cache.layers { l.mode = KvMode::F32; }
576 let k = vec![1.0f32; 32];
577 let v = vec![2.0f32; 32];
578 for _ in 0..8 {
579 for layer in &mut cache.layers {
580 layer.append(&k, &v, &[true; 4]);
581 }
582 }
583 assert_eq!(cache.seq_len(), 8);
584 assert!(!cache.needs_eviction());
585
586 cache.evict(4);
587 assert_eq!(cache.seq_len(), 4);
588 assert_eq!(cache.layers[0].head_len(0), 4);
589 }
590
591 #[test]
592 fn truncate_rolls_back_speculative_positions() {
593 let mut cache = LayerKvCache::new(2, 4);
594 cache.mode = KvMode::F32;
595 for pos in 0..5 {
596 let k = vec![pos as f32; 8];
597 let v = vec![pos as f32; 8];
598 cache.append(&k, &v, &[true; 2]);
599 }
600 cache.truncate_last(2);
601 assert_eq!(cache.seq_len, 3);
602 assert_eq!(cache.head_len(0), 3);
603 assert_eq!(cache.head_keys(0)[2 * 4], 2.0, "position 2 survives");
604 }
605
606 #[test]
610 fn q8_attend_matches_f32_within_grid() {
611 let (heads, hd) = (2, 32);
612 let mut f = LayerKvCache::new(heads, hd);
613 f.mode = KvMode::F32;
614 let mut q8 = LayerKvCache::new(heads, hd);
615 q8.mode = KvMode::Q8 { k: true, v: true };
616
617 let synth = |p: usize, salt: usize| -> Vec<f32> {
618 (0..heads * hd)
619 .map(|i| {
620 let x = ((i * 31 + p * 17 + salt * 7 + 3) % 97) as f32 / 97.0 - 0.5;
621 if i % 2 == 0 { x * 4.0 } else { x * 0.25 }
623 })
624 .collect()
625 };
626 for p in 0..100 {
627 let k = synth(p, 1);
628 let v = synth(p, 2);
629 f.append(&k, &v, &[true; 2]);
630 q8.append(&k, &v, &[true; 2]);
631 }
632 let q: Vec<f32> = (0..hd).map(|i| ((i * 13 + 5) % 89) as f32 / 89.0 - 0.5).collect();
633 for g in 0..heads {
634 let (of, pf) = f.attend(&q, g);
635 let (o8, p8) = q8.attend(&q, g);
636 let scale = of.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-6);
637 for d in 0..hd {
638 assert!(
639 (of[d] - o8[d]).abs() <= scale * 0.03 + 1e-3,
640 "g{g} d{d}: f32 {} vs q8 {}", of[d], o8[d]
641 );
642 }
643 for p in 0..100 {
644 assert!((pf[p] - p8[p]).abs() < 0.02, "prob p{p}");
645 }
646 }
647 q8.truncate_last(30);
649 assert_eq!(q8.head_len(0), 70);
650 let imp: Vec<f32> = (0..70).map(|i| i as f32).collect();
651 q8.accumulate_imp(&imp);
652 q8.evict_born(20, 2, 8);
653 assert_eq!(q8.head_len(0), 20);
654 let (o, _) = q8.attend(&q, 0);
655 assert!(o.iter().all(|x| x.is_finite()));
656 assert!(q8.memory_bytes() * 3 < f.memory_bytes());
658 }
659
660 #[test]
664 fn born_eviction_mixed_modes_stay_consistent() {
665 for (mk, mv) in [(false, true), (true, false)] {
666 let mut c = LayerKvCache::new(1, 4);
667 c.mode = KvMode::Q8 { k: mk, v: mv };
668 for p in 0..80 {
669 let k = vec![p as f32 * 0.01; 4];
670 let v = vec![p as f32; 4];
671 c.append(&k, &v, &[true]);
672 }
673 let imp: Vec<f32> = (0..80).map(|i| i as f32).collect();
674 c.accumulate_imp(&imp);
675 let before = c.memory_bytes();
676 c.evict_born(20, 4, 8); assert_eq!(c.head_len(0), 20, "k={mk} v={mv}");
678 assert!(c.memory_bytes() < before / 2,
679 "memory must shrink (k={mk} v={mv})");
680 let (out, _) = c.attend(&[1.0, 1.0, 1.0, 1.0], 0);
683 assert!(out[0] > 30.0,
684 "V from the kept tail, not the stale head (k={mk} v={mv}, out {})",
685 out[0]);
686 }
687 }
688
689 #[test]
690 fn born_eviction_keeps_high_mass_position() {
691 let mut cache = KvCache::new(1, 1, 2, 16);
692 cache.policy = EvictionPolicy::Born { sink: 1 };
693 for l in &mut cache.layers { l.mode = KvMode::F32; }
694 let layer = &mut cache.layers[0];
695 for pos in 0..8 {
698 let k = vec![pos as f32; 2];
699 let v = vec![pos as f32 + 100.0; 2];
700 layer.append(&k, &v, &[true]);
701 }
702 let mut imp = vec![0.05f32; 8];
704 imp[3] = 5.0;
705 layer.accumulate_imp(&imp);
706
707 cache.evict(4); let layer = &cache.layers[0];
709 assert_eq!(layer.seq_len, 4);
710 let kept_keys: Vec<f32> = (0..4).map(|i| layer.head_keys(0)[i * 2]).collect();
711 assert_eq!(
712 kept_keys,
713 vec![0.0, 3.0, 6.0, 7.0],
714 "kept = sink(0) + Born-top(3) + recent(6,7)"
715 );
716 assert_eq!(layer.head_len(0), 4);
718 }
719}