1#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum KvMode {
14 F32,
15 Q8 {
17 k: bool,
18 v: bool,
19 },
20}
21
22impl KvMode {
23 pub fn from_env() -> Self {
24 match std::env::var("CMF_KV").as_deref() {
25 Ok("q8") | Ok("q8_2f") => KvMode::Q8 { k: true, v: true },
26 Ok("q8k") => KvMode::Q8 { k: true, v: false },
27 Ok("q8v") => KvMode::Q8 { k: false, v: true },
28 _ => KvMode::F32,
29 }
30 }
31
32 fn quant_k(self) -> bool {
33 matches!(self, KvMode::Q8 { k: true, .. })
34 }
35
36 fn quant_v(self) -> bool {
37 matches!(self, KvMode::Q8 { v: true, .. })
38 }
39}
40
41const KV_COL_WARMUP: usize = 64;
44
45const KV_K_GROUP: usize = 32;
50
51static KV_GEN: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
56
57fn next_gen() -> u64 {
58 KV_GEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
59}
60
61#[derive(Debug, Clone)]
72pub enum O1State {
73 Collecting {
74 m: usize,
75 w: usize,
76 sink: usize,
77 rect: crate::nystrom::O1Rect,
78 seal_at: Option<usize>,
82 q_buf: Vec<f32>,
84 },
85 Sealed {
91 groups: Vec<crate::nystrom::NystromState>,
92 },
93}
94
95#[derive(Debug, Clone)]
97pub struct LayerKvCache {
98 pub mode: KvMode,
99 k: Vec<Vec<f32>>,
101 v: Vec<Vec<f32>>,
103 kq: Vec<Vec<i8>>,
105 ks: Vec<Vec<f32>>,
106 vq: Vec<Vec<i8>>,
107 vs: Vec<Vec<f32>>,
108 kcol: Vec<Vec<f32>>,
110 vcol: Vec<Vec<f32>>,
111 imp: Vec<f32>,
114 pub seq_len: usize,
120 base: usize,
123 tail: Option<usize>,
128 generation: u64,
136 pub num_kv_heads: usize,
137 pub head_dim: usize,
138 pub linear_state: Vec<f32>,
140 pub linear_scratch: Vec<f32>,
142 linear_wire_allowed: bool,
146 pub o1: Option<O1State>,
148 o1_error: Option<String>,
152 o1_transitioned: bool,
155 pub sinks: Option<Vec<f32>>,
162 pub bounded: Option<crate::bounded::BoundedState>,
166 pub wire_kind: WireKind,
168 pub wire_layer: u32,
170 pub wire_identity: u64,
173 kv_amax: [f32; 2],
177 kv_amax_rows: usize,
178}
179
180#[derive(Debug, Clone, Copy, PartialEq, Eq)]
182#[repr(u8)]
183pub enum WireKind {
184 Full = 0,
186 Linear = 1,
188 Bounded = 2,
190 FullTail = 3,
196}
197
198impl WireKind {
199 fn from_u8(v: u8) -> Option<Self> {
200 match v {
201 0 => Some(WireKind::Full),
202 1 => Some(WireKind::Linear),
203 2 => Some(WireKind::Bounded),
204 3 => Some(WireKind::FullTail),
205 _ => None,
206 }
207 }
208}
209
210pub const WIRE_MAGIC: &[u8; 4] = b"CMFS";
212pub const WIRE_VERSION: u32 = 2;
214
215impl LayerKvCache {
216 pub fn new(num_kv_heads: usize, head_dim: usize) -> Self {
217 Self {
218 sinks: None,
219 mode: KvMode::from_env(),
220 k: vec![Vec::new(); num_kv_heads],
221 v: vec![Vec::new(); num_kv_heads],
222 kq: vec![Vec::new(); num_kv_heads],
223 ks: vec![Vec::new(); num_kv_heads],
224 vq: vec![Vec::new(); num_kv_heads],
225 vs: vec![Vec::new(); num_kv_heads],
226 kcol: vec![Vec::new(); num_kv_heads],
227 vcol: vec![Vec::new(); num_kv_heads],
228 imp: Vec::new(),
229 seq_len: 0,
230 base: 0,
231 tail: None,
232 generation: next_gen(),
233 num_kv_heads,
234 head_dim,
235 linear_state: Vec::new(),
236 linear_scratch: Vec::new(),
237 linear_wire_allowed: true,
238 o1: None,
239 o1_error: None,
240 o1_transitioned: false,
241 bounded: None,
242 wire_kind: WireKind::Full,
243 wire_layer: 0,
244 wire_identity: 0,
245 kv_amax: [0.0; 2],
246 kv_amax_rows: 0,
247 }
248 }
249
250 pub(crate) fn kv_abs_max(&mut self) -> (f32, f32) {
258 let hd = self.head_dim.max(1);
259 let rows = self.kv_rows();
260 if rows < self.kv_amax_rows {
261 self.kv_amax_rows = 0;
262 }
263 if self.kv_amax_rows == 0 {
264 self.kv_amax = [0.0; 2];
265 }
266 let from = self.kv_amax_rows * hd;
267 for (m, heads) in self.kv_amax.iter_mut().zip([&self.k, &self.v]) {
268 for x in heads.iter().filter(|x| x.len() > from) {
269 *m = m.max(crate::gpu::abs_max_or_inf(&x[from..]));
270 }
271 }
272 self.kv_amax_rows = rows;
273 (self.kv_amax[0], self.kv_amax[1])
274 }
275
276 pub(crate) fn kv_abs_max_after(&mut self, chunk: (usize, usize, f32, f32)) -> (f32, f32) {
282 let (from, upto, km, vm) = chunk;
283 if self.kv_amax_rows != from || self.kv_rows() != upto || from > upto {
284 return self.kv_abs_max();
285 }
286 if from == 0 {
287 self.kv_amax = [0.0; 2];
288 }
289 self.kv_amax = [self.kv_amax[0].max(km), self.kv_amax[1].max(vm)];
290 self.kv_amax_rows = upto;
291 (self.kv_amax[0], self.kv_amax[1])
292 }
293
294 fn kv_rows(&self) -> usize {
296 let hd = self.head_dim.max(1);
297 self.k
298 .iter()
299 .chain(&self.v)
300 .map(|x| x.len() / hd)
301 .max()
302 .unwrap_or(0)
303 }
304
305 fn kv_rows_changed(&mut self, first: usize) {
307 self.kv_amax_rows = self.kv_amax_rows.min(first);
308 }
309
310 pub fn install_bounded(&mut self, window: usize) {
316 self.bounded = Some(crate::bounded::BoundedState::new(
317 self.num_kv_heads,
318 self.head_dim,
319 window,
320 ));
321 self.wire_kind = WireKind::Bounded;
322 }
323
324 #[allow(clippy::too_many_arguments)]
329 pub fn bounded_step(
330 &mut self,
331 q: &[f32],
332 k: &[f32],
333 v: &[f32],
334 w: &crate::bounded::BoundedWeights,
335 rope: &crate::bounded::BoundedRope,
336 scale: f32,
337 num_heads: usize,
338 out: &mut [f32],
339 ) {
340 let st = self
341 .bounded
342 .as_mut()
343 .expect("bounded_step on a layer without an installed ring");
344 st.insert(k, v);
345 st.attend(q, num_heads, &w.sink_k, &w.sink_v, w.sink, rope, scale, out);
346 self.seq_len += 1;
349 }
350
351 pub fn bounded_state_bytes(&self) -> usize {
353 self.bounded.as_ref().map(|b| b.state_bytes()).unwrap_or(0)
354 }
355
356 pub fn bounded_snapshot(&self) -> Option<crate::bounded::BoundedSnapshot> {
358 self.bounded.as_ref().map(|b| b.snapshot())
359 }
360
361 pub fn bounded_restore(&mut self, s: &crate::bounded::BoundedSnapshot) {
364 if let Some(b) = self.bounded.as_mut() {
365 b.restore(s);
366 self.seq_len = b.seen;
367 self.generation = next_gen();
368 }
369 }
370
371 pub fn set_linear_wire_allowed(&mut self, allowed: bool) {
375 self.linear_wire_allowed = allowed;
376 }
377
378 pub fn discard_linear_scratch(&mut self) {
381 self.linear_scratch.clear();
382 }
383
384 pub fn k_heads(&self) -> &[Vec<f32>] {
387 &self.k
388 }
389 pub fn v_heads(&self) -> &[Vec<f32>] {
391 &self.v
392 }
393
394 pub fn base(&self) -> usize {
399 self.base
400 }
401
402 pub fn pos_len(&self) -> usize {
404 self.base + self.seq_len
405 }
406
407 pub fn generation(&self) -> u64 {
410 self.generation
411 }
412
413 pub fn tail_window(&self) -> Option<usize> {
415 self.tail
416 }
417
418 pub fn trim_window(&mut self, w: usize, slack: usize, align: usize) -> usize {
442 if self.o1.is_some() || self.bounded.is_some() || w == 0 {
443 return 0;
444 }
445 self.tail = Some(w);
446 if self.seq_len <= 2 * w {
447 return 0;
448 }
449 if matches!(self.mode, KvMode::Q8 { .. })
450 && self.kcol.iter().all(Vec::is_empty)
451 && self.vcol.iter().all(Vec::is_empty)
452 {
453 return 0;
454 }
455 let align = align.max(1);
456 let keep_from = self.pos_len().saturating_sub(w + slack);
457 let new_base = keep_from / align * align;
458 if new_base <= self.base {
459 return 0;
460 }
461 let d = new_base - self.base;
462 let hd = self.head_dim;
463 fn drop_front<T>(v: &mut Vec<T>, n: usize) {
464 let n = n.min(v.len());
465 v.drain(..n);
466 }
467 for h in 0..self.num_kv_heads {
468 drop_front(&mut self.k[h], d * hd);
469 drop_front(&mut self.v[h], d * hd);
470 drop_front(&mut self.kq[h], d * hd);
471 drop_front(&mut self.vq[h], d * hd);
472 drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
473 drop_front(&mut self.vs[h], d);
474 }
475 drop_front(&mut self.imp, d);
476 self.base = new_base;
477 self.seq_len -= d;
478 self.generation = next_gen();
479 self.kv_amax_rows = self.kv_amax_rows.saturating_sub(d);
485 d
486 }
487
488 pub fn o1_begin(&mut self, m: usize, w: usize, sink: usize, rect: crate::nystrom::O1Rect) {
492 self.o1_begin_with_boundary(m, w, sink, rect, None);
493 }
494
495 pub(crate) fn o1_begin_with_boundary(
499 &mut self,
500 m: usize,
501 w: usize,
502 sink: usize,
503 rect: crate::nystrom::O1Rect,
504 seal_at: Option<usize>,
505 ) {
506 self.o1 = Some(O1State::Collecting {
507 m,
508 w,
509 sink,
510 rect,
511 seal_at,
512 q_buf: Vec::new(),
513 });
514 self.o1_error = None;
515 self.o1_transitioned = false;
516 }
517
518 pub fn o1_push_q(&mut self, q_all: &[f32]) {
523 if let Some(O1State::Collecting { q_buf, .. }) = &mut self.o1 {
524 q_buf.extend_from_slice(q_all);
525 }
526 }
527
528 pub fn o1_sealed(&self) -> bool {
529 matches!(self.o1, Some(O1State::Sealed { .. }))
530 }
531
532 pub(crate) fn o1_pending_boundary(&self) -> Option<usize> {
535 match &self.o1 {
536 Some(O1State::Collecting { seal_at, .. }) => *seal_at,
537 _ => None,
538 }
539 }
540
541 pub(crate) fn o1_boundary_crossed_by(&self, count: usize) -> bool {
545 let Some(target) = self.o1_pending_boundary() else {
546 return false;
547 };
548 target <= self.seq_len
549 || self
550 .seq_len
551 .checked_add(count)
552 .map_or(true, |next| next >= target)
553 }
554
555 pub(crate) fn take_o1_transition(&mut self) -> bool {
556 std::mem::take(&mut self.o1_transitioned)
557 }
558
559 pub(crate) fn take_o1_error(&self) -> Option<String> {
560 self.o1_error.clone()
566 }
567
568 pub(crate) fn o1_abort(&mut self, err: String) {
573 self.k.iter_mut().for_each(Vec::clear);
574 self.v.iter_mut().for_each(Vec::clear);
575 self.kv_rows_changed(0);
576 self.kq.iter_mut().for_each(Vec::clear);
577 self.ks.iter_mut().for_each(Vec::clear);
578 self.vq.iter_mut().for_each(Vec::clear);
579 self.vs.iter_mut().for_each(Vec::clear);
580 self.kcol.iter_mut().for_each(Vec::clear);
581 self.vcol.iter_mut().for_each(Vec::clear);
582 self.imp.clear();
583 self.o1 = None;
584 self.seq_len = 0;
585 self.base = 0;
586 self.generation = next_gen();
587 self.o1_transitioned = false;
588 self.o1_error = Some(err);
589 }
590
591 pub fn o1_seal(&mut self, num_heads: usize) -> bool {
598 match self.o1_seal_checked(num_heads) {
599 Ok(sealed) => sealed,
600 Err(err) => {
601 tracing::error!("o1: seal aborted: {err}");
602 self.o1_abort(err);
603 false
604 }
605 }
606 }
607
608 pub(crate) fn o1_seal_checked(&mut self, num_heads: usize) -> Result<bool, String> {
612 if let Some(err) = self.o1_error.clone() {
613 return Err(err);
614 }
615 if !matches!(self.o1, Some(O1State::Collecting { .. })) {
618 return Ok(self.o1_sealed());
619 }
620 let (m, w, sink, requested_boundary, q_len) = match &self.o1 {
621 Some(O1State::Collecting {
622 m,
623 w,
624 sink,
625 rect: _,
626 seal_at,
627 q_buf,
628 }) => (*m, *w, *sink, *seal_at, q_buf.len()),
629 _ => unreachable!("checked above"),
630 };
631 let floor = crate::nystrom::o1_deferred_boundary(w, sink)
632 .ok_or_else(|| "o1 seal: w + sink + slack + 1 overflow".to_string())?;
633 let target = requested_boundary.unwrap_or(floor).max(floor);
634 let t = self.seq_len;
635 if t < target {
636 if let Some(O1State::Collecting { seal_at, .. }) = &mut self.o1 {
637 if *seal_at != Some(target) {
638 *seal_at = Some(target);
639 tracing::info!(
640 "o1 deferred seal: current rows={t}, boundary={target} (floor={floor})"
641 );
642 }
643 }
644 return Ok(false);
645 }
646
647 let hd = self.head_dim;
648 if t == 0 {
649 return Err("o1 seal: cannot seal an empty layer".into());
650 }
651 if self.mode != KvMode::F32 {
652 return Err("o1 seal: requires dense F32 KV storage".into());
653 }
654 if self.num_kv_heads == 0 || num_heads == 0 || num_heads % self.num_kv_heads != 0 {
655 return Err(format!(
656 "o1 seal: invalid GQA geometry num_heads={num_heads} num_kv_heads={}",
657 self.num_kv_heads
658 ));
659 }
660 let hpk = num_heads / self.num_kv_heads;
661 let expected_k = t
662 .checked_mul(hd)
663 .ok_or_else(|| "o1 seal: KV row length overflow".to_string())?;
664 let expected_q = expected_k
665 .checked_mul(num_heads)
666 .ok_or_else(|| "o1 seal: query trace length overflow".to_string())?;
667 if q_len != expected_q {
668 return Err(format!(
669 "o1 seal: query trace has {q_len} values, expected {expected_q}"
670 ));
671 }
672 if (0..self.num_kv_heads)
673 .any(|g| self.k[g].len() != expected_k || self.v[g].len() != expected_k)
674 {
675 return Err("o1 seal: KV heads are not densely populated".into());
676 }
677 if m < 4 || w == 0 {
678 return Err(format!("o1 seal: invalid geometry m={m} w={w}"));
679 }
680
681 let Some(O1State::Collecting {
682 m,
683 w,
684 sink,
685 rect,
686 q_buf,
687 ..
688 }) = self.o1.take()
689 else {
690 unreachable!("collecting state disappeared after validation");
691 };
692 let mut groups = Vec::with_capacity(self.num_kv_heads);
693 let mut qh = vec![0.0f32; hpk * t * hd];
696 for g in 0..self.num_kv_heads {
697 for hh in 0..hpk {
698 let h = g * hpk + hh;
699 for p in 0..t {
700 let src = (p * num_heads + h) * hd;
701 let dst = (hh * t + p) * hd;
702 qh[dst..dst + hd].copy_from_slice(&q_buf[src..src + hd]);
703 }
704 }
705 let qs: Vec<&[f32]> = (0..hpk)
706 .map(|hh| &qh[hh * t * hd..(hh + 1) * t * hd])
707 .collect();
708 let mut st = crate::nystrom::NystromState::new_group(m, w, sink, hpk).with_rect(rect);
709 st.prefill_group(&qs, &self.k[g], &self.v[g], t, hd, hd);
710 groups.push(st);
711 }
712 for h in 0..self.num_kv_heads {
715 self.k[h] = Vec::new();
716 self.v[h] = Vec::new();
717 }
718 self.kv_rows_changed(0);
719 self.imp = Vec::new();
720 self.o1 = Some(O1State::Sealed { groups });
721 self.o1_transitioned = true;
722 Ok(true)
723 }
724
725 pub fn o1_views(&self) -> Option<Vec<crate::nystrom::O1DeviceView<'_>>> {
734 let Some(O1State::Sealed { groups }) = &self.o1 else {
735 return None;
736 };
737 let views: Vec<_> = groups.iter().map(|g| g.device_view()).collect();
738 if views.iter().any(|v| v.exact_only) {
739 return None;
740 }
741 Some(views)
742 }
743
744 pub fn o1_step(
745 &mut self,
746 q_all: &[f32],
747 k_new: &[f32],
748 v_new: &[f32],
749 num_heads: usize,
750 ) -> Vec<f32> {
751 let hd = self.head_dim;
752 let hpk = num_heads / self.num_kv_heads.max(1);
753 let mut out = vec![0.0f32; num_heads * hd];
754 let Some(O1State::Sealed { groups }) = &mut self.o1 else {
755 debug_assert!(false, "o1_step on an unsealed layer");
756 return out;
757 };
758 for (g, st) in groups.iter_mut().enumerate() {
759 let (lo, hi) = (g * hpk * hd, (g + 1) * hpk * hd);
760 st.step_group(
761 &q_all[lo..hi],
762 &k_new[g * hd..(g + 1) * hd],
763 &v_new[g * hd..(g + 1) * hd],
764 &mut out[lo..hi],
765 );
766 }
767 self.seq_len += 1;
770 out
771 }
772
773 pub fn o1_memory_bytes(&self) -> usize {
776 match &self.o1 {
777 Some(O1State::Collecting { q_buf, .. }) => q_buf.len() * std::mem::size_of::<f32>(),
778 Some(O1State::Sealed { groups }) => groups.iter().map(|s| s.memory_bytes()).sum(),
779 None => 0,
780 }
781 }
782
783 fn quant_row(row: &[f32], col: &[f32], q: &mut Vec<i8>, sc: &mut Vec<f32>, group: usize) {
786 let mut resid = vec![0.0f32; row.len()];
787 for (d, &x) in row.iter().enumerate() {
788 resid[d] = if col.is_empty() { x } else { x / col[d] };
789 }
790 for g0 in (0..row.len()).step_by(group) {
791 let g1 = (g0 + group).min(row.len());
792 let mut absmax = 0.0f32;
793 for &r in &resid[g0..g1] {
794 absmax = absmax.max(r.abs());
795 }
796 let s = (absmax / 127.0).max(1e-12);
797 sc.push(s);
798 for &r in &resid[g0..g1] {
799 q.push((r / s).round().clamp(-127.0, 127.0) as i8);
800 }
801 }
802 }
803
804 fn freeze_cols(&mut self) {
807 let hd = self.head_dim;
808 let ngk = hd.div_ceil(KV_K_GROUP);
809 for h in 0..self.num_kv_heads {
810 for (qv, sv, colv, group) in [
811 (
812 &mut self.kq[h],
813 &mut self.ks[h],
814 &mut self.kcol[h],
815 KV_K_GROUP,
816 ),
817 (&mut self.vq[h], &mut self.vs[h], &mut self.vcol[h], hd),
818 ] {
819 let spp = if group == hd { 1 } else { ngk }; let n = sv.len() / spp;
821 if n == 0 {
822 continue;
823 }
824 let mut rows = vec![0.0f32; n * hd];
826 for p in 0..n {
827 for d in 0..hd {
828 rows[p * hd + d] = qv[p * hd + d] as f32 * sv[p * spp + d / group];
829 }
830 }
831 let mut col = vec![0.0f32; hd];
832 for p in 0..n {
833 for d in 0..hd {
834 col[d] += rows[p * hd + d] * rows[p * hd + d];
835 }
836 }
837 for c in col.iter_mut() {
838 *c = (*c / n as f32).sqrt().max(1e-6);
839 }
840 qv.clear();
841 sv.clear();
842 for p in 0..n {
843 Self::quant_row(&rows[p * hd..(p + 1) * hd], &col, qv, sv, group);
844 }
845 *colv = col;
846 }
847 }
848 }
849
850 pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
854 if self.o1_error.is_some() {
858 return;
859 }
860 debug_assert_eq!(k_new.len(), self.num_kv_heads * self.head_dim);
861 debug_assert_eq!(v_new.len(), self.num_kv_heads * self.head_dim);
862 if matches!(self.mode, KvMode::Q8 { .. })
868 && self.seq_len >= KV_COL_WARMUP
869 && self.kcol.iter().all(Vec::is_empty)
870 && self.vcol.iter().all(Vec::is_empty)
871 {
872 self.freeze_cols();
873 }
874 for h in 0..self.num_kv_heads {
875 if !alive.get(h).copied().unwrap_or(true) {
876 continue;
877 }
878 let s = h * self.head_dim;
879 if self.mode.quant_k() {
880 Self::quant_row(
881 &k_new[s..s + self.head_dim],
882 &self.kcol[h],
883 &mut self.kq[h],
884 &mut self.ks[h],
885 KV_K_GROUP,
886 );
887 } else {
888 self.k[h].extend_from_slice(&k_new[s..s + self.head_dim]);
889 }
890 if self.mode.quant_v() {
891 Self::quant_row(
892 &v_new[s..s + self.head_dim],
893 &self.vcol[h],
894 &mut self.vq[h],
895 &mut self.vs[h],
896 self.head_dim,
897 );
898 } else {
899 self.v[h].extend_from_slice(&v_new[s..s + self.head_dim]);
900 }
901 }
902 self.imp.push(0.0);
903 self.seq_len += 1;
904 }
905
906 pub fn attend(&self, q: &[f32], kv_head: usize) -> (Vec<f32>, Vec<f32>) {
911 debug_assert!(self.tail.is_none(), "attend() on a sliding-window tail");
913 let hd = self.head_dim;
914 if self.mode == KvMode::F32 {
915 let stored = self.k[kv_head].len() / hd;
916 return crate::attention::attention_head(
917 q,
918 &self.k[kv_head],
919 &self.v[kv_head],
920 hd,
921 stored,
922 );
923 }
924 let stored = self.head_len(kv_head);
925 let scale = 1.0 / (hd as f32).sqrt();
926 let mut scores = vec![0.0f32; stored];
927 if self.mode.quant_k() {
928 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
929 let kcol = &self.kcol[kv_head];
931 let mut qc = vec![0.0f32; hd];
932 for d in 0..hd {
933 qc[d] = if kcol.is_empty() {
934 q[d]
935 } else {
936 q[d] * kcol[d]
937 };
938 }
939 let ng = hd.div_ceil(KV_K_GROUP);
940 for p in 0..stored {
941 let row = &kq[p * hd..(p + 1) * hd];
942 let row_u8 =
945 unsafe { std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len()) };
946 let mut dot = 0.0f32;
947 for g in 0..ng {
948 let g0 = g * KV_K_GROUP;
949 let g1 = (g0 + KV_K_GROUP).min(hd);
950 dot +=
951 crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qc[g0..g1]) * ks[p * ng + g];
952 }
953 scores[p] = dot * scale;
954 }
955 } else {
956 let k = &self.k[kv_head];
957 for p in 0..stored {
958 let row = &k[p * hd..(p + 1) * hd];
959 scores[p] = crate::attention::dot_f32(q, row) * scale;
960 }
961 }
962 let max_score = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
963 let mut sum = 0.0f32;
964 for s in scores.iter_mut() {
965 *s = (*s - max_score).exp();
966 sum += *s;
967 }
968 if sum > 0.0 {
969 for s in scores.iter_mut() {
970 *s /= sum;
971 }
972 }
973 let mut acc = vec![0.0f32; hd];
974 if self.mode.quant_v() {
975 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
976 for p in 0..stored {
977 let w = scores[p] * vs[p];
978 if w.abs() < 1e-12 {
979 continue;
980 }
981 crate::qtensor::axpy_i8_f32(&mut acc, &vq[p * hd..(p + 1) * hd], w);
982 }
983 let vcol = &self.vcol[kv_head];
984 if !vcol.is_empty() {
985 for d in 0..hd {
986 acc[d] *= vcol[d];
987 }
988 }
989 } else {
990 let v = &self.v[kv_head];
991 for p in 0..stored {
992 let w = scores[p];
993 if w.abs() < 1e-12 {
994 continue;
995 }
996 crate::attention::axpy_f32(&mut acc, &v[p * hd..(p + 1) * hd], w);
997 }
998 }
999 (acc, scores)
1000 }
1001
1002 #[allow(clippy::too_many_arguments)]
1019 pub fn attend_group(
1020 &self,
1021 q_group: &[f32],
1022 kv_head: usize,
1023 out: &mut [f32],
1024 imp_acc: &mut [f32],
1025 scale: f32,
1026 first: usize,
1027 softcap: f32,
1028 sinks: &[f32],
1029 ) {
1030 self.attend_group_upto(
1031 q_group,
1032 kv_head,
1033 out,
1034 imp_acc,
1035 scale,
1036 first,
1037 softcap,
1038 usize::MAX,
1039 sinks,
1040 )
1041 }
1042
1043 #[allow(clippy::too_many_arguments)]
1061 pub fn attend_group_upto(
1062 &self,
1063 q_group: &[f32],
1064 kv_head: usize,
1065 out: &mut [f32],
1066 imp_acc: &mut [f32],
1067 scale: f32,
1068 first: usize,
1069 softcap: f32,
1070 upto: usize,
1071 sinks: &[f32],
1072 ) {
1073 let hd = self.head_dim;
1074 let nheads = q_group.len() / hd;
1075 debug_assert_eq!(out.len(), nheads * hd);
1076 assert!(
1077 sinks.is_empty() || sinks.len() == nheads,
1078 "attend_group: {} sink logits for {nheads} heads",
1079 sinks.len()
1080 );
1081 let stored = if self.mode == KvMode::F32 {
1082 self.k[kv_head].len() / hd
1083 } else {
1084 self.head_len(kv_head)
1085 }
1086 .min(upto);
1087 if stored == 0 {
1088 out.fill(0.0);
1089 return;
1090 }
1091 let first = first.min(stored.saturating_sub(1));
1092 let span = stored - first;
1094
1095 thread_local! {
1096 static GQA_SCORES: std::cell::RefCell<Vec<f32>> =
1098 const { std::cell::RefCell::new(Vec::new()) };
1099 static GQA_QC: std::cell::RefCell<Vec<f32>> =
1101 const { std::cell::RefCell::new(Vec::new()) };
1102 }
1103
1104 GQA_SCORES.with(|sc| {
1105 let mut scores = sc.borrow_mut();
1106 scores.resize(nheads * span, 0.0);
1107
1108 if self.mode.quant_k() {
1110 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
1111 let kcol = &self.kcol[kv_head];
1112 let ng = hd.div_ceil(KV_K_GROUP);
1113 GQA_QC.with(|qc| {
1114 let mut qcb = qc.borrow_mut();
1115 qcb.resize(nheads * hd, 0.0);
1116 for h in 0..nheads {
1117 for d in 0..hd {
1118 let qv = q_group[h * hd + d];
1119 qcb[h * hd + d] = if kcol.is_empty() { qv } else { qv * kcol[d] };
1120 }
1121 }
1122 for p in first..stored {
1123 let row = &kq[p * hd..(p + 1) * hd];
1124 let row_u8 = unsafe {
1127 std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len())
1128 };
1129 for h in 0..nheads {
1130 let qch = &qcb[h * hd..(h + 1) * hd];
1131 let mut dot = 0.0f32;
1132 for g in 0..ng {
1133 let g0 = g * KV_K_GROUP;
1134 let g1 = (g0 + KV_K_GROUP).min(hd);
1135 dot += crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qch[g0..g1])
1136 * ks[p * ng + g];
1137 }
1138 scores[h * span + (p - first)] = dot * scale;
1139 }
1140 }
1141 });
1142 } else {
1143 let k = &self.k[kv_head];
1144 for p in first..stored {
1145 let row = &k[p * hd..(p + 1) * hd];
1146 for h in 0..nheads {
1147 scores[h * span + (p - first)] =
1148 crate::attention::dot_f32(&q_group[h * hd..(h + 1) * hd], row) * scale;
1149 }
1150 }
1151 }
1152
1153 if softcap > 0.0 {
1157 for v in scores.iter_mut() {
1158 *v = softcap * (*v / softcap).tanh();
1159 }
1160 }
1161
1162 for h in 0..nheads {
1165 let s = &mut scores[h * span..(h + 1) * span];
1166 let row_max = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
1167 let sink = sinks.get(h).copied();
1168 let max_score = match sink {
1169 Some(z) => row_max.max(z),
1170 None => row_max,
1171 };
1172 let mut sum = 0.0f32;
1173 for v in s.iter_mut() {
1174 *v = (*v - max_score).exp();
1175 sum += *v;
1176 }
1177 if let Some(z) = sink {
1178 sum += (z - max_score).exp();
1179 }
1180 if sum > 0.0 {
1181 for v in s.iter_mut() {
1182 *v /= sum;
1183 }
1184 }
1185 }
1186
1187 out.fill(0.0);
1189 if self.mode.quant_v() {
1190 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
1191 for p in first..stored {
1192 let row = &vq[p * hd..(p + 1) * hd];
1193 for h in 0..nheads {
1194 let w = scores[h * span + (p - first)] * vs[p];
1195 if w.abs() < 1e-12 {
1196 continue;
1197 }
1198 crate::qtensor::axpy_i8_f32(&mut out[h * hd..(h + 1) * hd], row, w);
1199 }
1200 }
1201 let vcol = &self.vcol[kv_head];
1202 if !vcol.is_empty() {
1203 for h in 0..nheads {
1204 for d in 0..hd {
1205 out[h * hd + d] *= vcol[d];
1206 }
1207 }
1208 }
1209 } else {
1210 let v = &self.v[kv_head];
1211 for p in first..stored {
1212 let row = &v[p * hd..(p + 1) * hd];
1213 for h in 0..nheads {
1214 let w = scores[h * span + (p - first)];
1215 if w.abs() < 1e-12 {
1216 continue;
1217 }
1218 crate::attention::axpy_f32(&mut out[h * hd..(h + 1) * hd], row, w);
1219 }
1220 }
1221 }
1222
1223 let n = imp_acc.len().min(stored);
1227 if n > first {
1228 for h in 0..nheads {
1229 let s = &scores[h * span..(h + 1) * span];
1230 for (dst, &p) in imp_acc[first..n].iter_mut().zip(s) {
1231 *dst += p;
1232 }
1233 }
1234 }
1235 });
1236 }
1237
1238 #[cfg(target_arch = "aarch64")]
1246 #[allow(clippy::too_many_arguments)]
1247 pub fn attend_chunk(
1248 &mut self,
1249 q_all: &[f32],
1250 b: usize,
1251 s0: usize,
1252 nh: usize,
1253 heads_per_kv: usize,
1254 hd: usize,
1255 out: &mut [f32],
1256 pool: Option<&crate::pool::Pool>,
1257 scale: f32,
1258 window: Option<usize>,
1259 ) {
1260 let n = s0 + b;
1261 struct SendPtr(*mut f32);
1262 unsafe impl Send for SendPtr {}
1263 unsafe impl Sync for SendPtr {}
1264 impl SendPtr {
1265 fn at(&self, i: usize) -> *mut f32 {
1266 unsafe { self.0.add(i) }
1269 }
1270 }
1271 thread_local! {
1272 static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>)> =
1273 const { std::cell::RefCell::new((Vec::new(), Vec::new(), Vec::new(), Vec::new())) };
1274 }
1275 let neon_gemm = cfg!(not(target_os = "macos"))
1280 || std::env::var("CMF_FORCE_NEON_GEMM")
1281 .map(|v| v == "1")
1282 .unwrap_or(false);
1283 SCRATCH.with(|s| {
1284 let mut s = s.borrow_mut();
1285 let (qpanel, scores, aopanel, ktpack) = &mut *s;
1286 let m = heads_per_kv * b;
1291 qpanel.resize(m * hd, 0.0);
1292 scores.resize(m * n, 0.0);
1293 aopanel.resize(m * hd, 0.0);
1294 for g in 0..self.num_kv_heads {
1295 let kmat = &self.k[g];
1296 let vmat = &self.v[g];
1297 debug_assert_eq!(kmat.len(), n * hd);
1298 for hl in 0..heads_per_kv {
1299 let hh = g * heads_per_kv + hl;
1300 for bi in 0..b {
1301 qpanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]
1302 .copy_from_slice(&q_all[bi * nh * hd + hh * hd..][..hd]);
1303 }
1304 }
1305 if neon_gemm {
1306 ktpack.resize(hd * n, 0.0);
1307 for p in 0..n {
1308 let row = &kmat[p * hd..(p + 1) * hd];
1309 for (d, &v) in row.iter().enumerate() {
1310 ktpack[d * n + p] = v;
1311 }
1312 }
1313 let sp_q = SendPtr(qpanel.as_ptr() as *mut f32);
1316 let sp_s = SendPtr(scores.as_mut_ptr());
1317 let kt = &*ktpack;
1318 let run = |start: usize, end: usize| {
1319 if end > start {
1320 let a = unsafe {
1322 std::slice::from_raw_parts(sp_q.at(start * hd), (end - start) * hd)
1323 };
1324 let c = unsafe {
1325 std::slice::from_raw_parts_mut(
1326 sp_s.at(start * n),
1327 (end - start) * n,
1328 )
1329 };
1330 crate::qtensor::neon_gemm_rm(
1331 end - start,
1332 n,
1333 hd,
1334 scale,
1335 a,
1336 hd,
1337 kt,
1338 n,
1339 false,
1340 c,
1341 n,
1342 );
1343 }
1344 };
1345 match pool {
1346 Some(p) if m >= 64 => p.run_rows(m, &run),
1347 _ => run(0, m),
1348 }
1349 } else {
1350 crate::qtensor::sgemm_rm(
1351 m, n, hd, scale, qpanel, hd, kmat, hd, true, scores, n,
1352 );
1353 }
1354 let sp = SendPtr(scores.as_mut_ptr());
1356 let run = |start: usize, end: usize| {
1357 for r in start..end {
1358 let allowed = s0 + (r % b) + 1;
1359 let lo = window.map(|w| allowed.saturating_sub(w)).unwrap_or(0);
1363 let row = unsafe { std::slice::from_raw_parts_mut(sp.at(r * n), n) };
1365 crate::attention::softmax_row(&mut row[lo..allowed]);
1366 row[..lo].fill(0.0);
1367 row[allowed..].fill(0.0);
1368 }
1369 };
1370 match pool {
1371 Some(p) if m >= 64 => p.run_rows(m, &run),
1372 _ => run(0, m),
1373 }
1374 let ni = self.imp.len().min(n);
1378 for r in 0..m {
1379 let al = (s0 + (r % b) + 1).min(ni);
1380 for (dst, &p) in self.imp[..al].iter_mut().zip(&scores[r * n..r * n + al]) {
1381 *dst += p;
1382 }
1383 }
1384 if neon_gemm {
1385 let sp_s = SendPtr(scores.as_mut_ptr());
1386 let sp_o = SendPtr(aopanel.as_mut_ptr());
1387 let run = |start: usize, end: usize| {
1388 if end > start {
1389 let a = unsafe {
1391 std::slice::from_raw_parts(sp_s.at(start * n), (end - start) * n)
1392 };
1393 let c = unsafe {
1394 std::slice::from_raw_parts_mut(
1395 sp_o.at(start * hd),
1396 (end - start) * hd,
1397 )
1398 };
1399 crate::qtensor::neon_gemm_rm(
1400 end - start,
1401 hd,
1402 n,
1403 1.0,
1404 a,
1405 n,
1406 vmat,
1407 hd,
1408 false,
1409 c,
1410 hd,
1411 );
1412 }
1413 };
1414 match pool {
1415 Some(p) if m >= 64 => p.run_rows(m, &run),
1416 _ => run(0, m),
1417 }
1418 } else {
1419 crate::qtensor::sgemm_rm(
1420 m, hd, n, 1.0, scores, n, vmat, hd, false, aopanel, hd,
1421 );
1422 }
1423 for hl in 0..heads_per_kv {
1424 let hh = g * heads_per_kv + hl;
1425 for bi in 0..b {
1426 out[bi * nh * hd + hh * hd..][..hd]
1427 .copy_from_slice(&aopanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]);
1428 }
1429 }
1430 }
1431 });
1432 }
1433
1434 pub fn truncate_last(&mut self, n_drop: usize) {
1436 self.discard_linear_scratch();
1437 let d = n_drop.min(self.seq_len);
1438 if let Some(b) = self.bounded.as_mut() {
1439 let rolled = b.rollback(d);
1442 if rolled < d {
1443 tracing::warn!(
1444 "bounded anchor: rollback of {d} exceeds the undo depth ({rolled} restored)"
1445 );
1446 }
1447 self.seq_len = b.seen;
1448 return;
1449 }
1450 if let Some(w) = self.tail
1454 && self.base > 0
1455 && self.seq_len - d < w.saturating_sub(1)
1456 {
1457 tracing::warn!(
1458 "sliding tail: rollback of {d} leaves {} rows under the {w}-row window",
1459 self.seq_len - d
1460 );
1461 }
1462 let mut first_changed = usize::MAX;
1463 for h in 0..self.num_kv_heads {
1464 let keep = self.k[h].len().saturating_sub(d * self.head_dim);
1465 if !self.k[h].is_empty() || !self.v[h].is_empty() {
1466 first_changed = first_changed.min(keep / self.head_dim.max(1));
1467 }
1468 self.k[h].truncate(keep);
1469 self.v[h].truncate(keep);
1470 let ngk = self.head_dim.div_ceil(KV_K_GROUP);
1471 let keep_q = self.kq[h].len().saturating_sub(d * self.head_dim);
1472 self.kq[h].truncate(keep_q);
1473 let keep_vq = self.vq[h].len().saturating_sub(d * self.head_dim);
1474 self.vq[h].truncate(keep_vq);
1475 let keep_ks = self.ks[h].len().saturating_sub(d * ngk);
1476 self.ks[h].truncate(keep_ks);
1477 let keep_vs = self.vs[h].len().saturating_sub(d);
1478 self.vs[h].truncate(keep_vs);
1479 }
1480 self.imp.truncate(self.imp.len().saturating_sub(d));
1481 self.seq_len -= d;
1482 self.kv_rows_changed(first_changed);
1483 }
1484
1485 pub fn accumulate_imp(&mut self, probs: &[f32]) {
1487 for (dst, &p) in self.imp.iter_mut().zip(probs) {
1488 *dst += p;
1489 }
1490 }
1491
1492 pub fn head_keys(&self, kv_head: usize) -> &[f32] {
1494 &self.k[kv_head]
1495 }
1496
1497 pub fn head_values(&self, kv_head: usize) -> &[f32] {
1498 &self.v[kv_head]
1499 }
1500
1501 pub fn head_len(&self, kv_head: usize) -> usize {
1503 let ng = self.head_dim.div_ceil(KV_K_GROUP);
1504 (self.k[kv_head].len() / self.head_dim)
1505 .max(self.ks[kv_head].len() / ng)
1506 .max(self.vs[kv_head].len())
1507 }
1508
1509 pub fn clear(&mut self) {
1511 self.kv_rows_changed(0);
1512 for h in 0..self.num_kv_heads {
1513 self.k[h].clear();
1514 self.v[h].clear();
1515 self.kq[h].clear();
1516 self.ks[h].clear();
1517 self.vq[h].clear();
1518 self.vs[h].clear();
1519 self.kcol[h].clear();
1520 self.vcol[h].clear();
1521 }
1522 self.imp.clear();
1523 self.linear_state.clear();
1524 self.discard_linear_scratch();
1525 self.o1 = None;
1528 self.o1_error = None;
1529 self.o1_transitioned = false;
1530 if let Some(b) = self.bounded.as_mut() {
1533 b.clear();
1534 }
1535 self.seq_len = 0;
1536 self.base = 0;
1538 self.generation = next_gen();
1539 }
1540
1541 pub fn export_wire(&self, f16: bool) -> Result<Vec<u8>, String> {
1558 if !matches!(self.mode, KvMode::F32) {
1559 return Err("kv export: only the F32 cache is described by this format (CMF_KV=q8 stores int8 rows and per-row scales)"
1560 .into());
1561 }
1562 if self.o1.is_some() {
1563 return Err("kv export: an O(1) Nyström overlay is not part of this format — the skeletons are irreversible and would have to travel with it"
1564 .into());
1565 }
1566 if self.kcol.iter().any(|c| !c.is_empty()) || self.vcol.iter().any(|c| !c.is_empty()) {
1570 return Err(
1571 "kv export: frozen columns under an F32 cache — refusing to ship \
1572 a state this format does not describe"
1573 .into(),
1574 );
1575 }
1576 let mut out = Vec::with_capacity(self.memory_bytes() / if f16 { 2 } else { 1 } + 64);
1577 let u = |v: u32, o: &mut Vec<u8>| o.extend_from_slice(&v.to_le_bytes());
1578 out.extend_from_slice(WIRE_MAGIC);
1579 u(WIRE_VERSION, &mut out);
1580 out.extend_from_slice(&self.wire_identity.to_le_bytes());
1581 u(self.wire_layer, &mut out);
1582 let kind = match self.wire_kind {
1584 WireKind::Full if self.base > 0 => WireKind::FullTail,
1585 k => k,
1586 };
1587 out.push(kind as u8);
1588 out.push(u8::from(f16));
1589 out.extend_from_slice(&0u16.to_le_bytes());
1590 out.extend_from_slice(&(self.pos_len() as u64).to_le_bytes());
1592 let push = |xs: &[f32], o: &mut Vec<u8>| {
1593 if f16 {
1594 for &x in xs {
1595 o.extend_from_slice(&cortiq_core::quant::f32_to_f16(x).to_le_bytes());
1596 }
1597 } else {
1598 for &x in xs {
1599 o.extend_from_slice(&x.to_le_bytes());
1600 }
1601 }
1602 };
1603 match kind {
1604 WireKind::Full => self.export_full_body(f16, &mut out),
1605 WireKind::FullTail => {
1606 out.extend_from_slice(&(self.base as u64).to_le_bytes());
1607 self.export_full_body(f16, &mut out);
1608 }
1609 WireKind::Linear => {
1610 u(self.num_kv_heads as u32, &mut out);
1611 u(self.head_dim as u32, &mut out);
1612 u(self.linear_state.len() as u32, &mut out);
1613 for &x in &self.linear_state {
1614 out.extend_from_slice(&x.to_le_bytes());
1615 }
1616 }
1617 WireKind::Bounded => {
1618 let b = self
1619 .bounded
1620 .as_ref()
1621 .ok_or("kv export: bounded wire kind without an installed ring")?;
1622 u(self.num_kv_heads as u32, &mut out);
1623 u(self.head_dim as u32, &mut out);
1624 u(b.window as u32, &mut out);
1625 u(b.len() as u32, &mut out);
1626 u(b.head() as u32, &mut out);
1627 push(&b.ring_k, &mut out);
1628 push(&b.ring_v, &mut out);
1629 }
1630 }
1631 Ok(out)
1632 }
1633
1634 fn export_full_body(&self, f16: bool, out: &mut Vec<u8>) {
1637 let u = |v: u32, o: &mut Vec<u8>| o.extend_from_slice(&v.to_le_bytes());
1638 u(u8::from(f16) as u32, out);
1639 u(self.seq_len as u32, out);
1640 u(self.num_kv_heads as u32, out);
1641 u(self.head_dim as u32, out);
1642 u(self.linear_state.len() as u32, out);
1643 u(self.imp.len() as u32, out);
1648 let push = |xs: &[f32], o: &mut Vec<u8>| {
1649 if f16 {
1650 for &x in xs {
1651 o.extend_from_slice(&cortiq_core::quant::f32_to_f16(x).to_le_bytes());
1652 }
1653 } else {
1654 for &x in xs {
1655 o.extend_from_slice(&x.to_le_bytes());
1656 }
1657 }
1658 };
1659 for &x in &self.linear_state {
1660 out.extend_from_slice(&x.to_le_bytes());
1661 }
1662 for &x in &self.imp {
1663 out.extend_from_slice(&x.to_le_bytes());
1664 }
1665 for h in 0..self.num_kv_heads {
1666 u(self.k[h].len() as u32, out);
1667 push(&self.k[h], out);
1668 u(self.v[h].len() as u32, out);
1669 push(&self.v[h], out);
1670 }
1671 }
1672
1673 pub fn import_wire(&mut self, buf: &[u8]) -> Result<(), String> {
1678 if buf.len() >= 4 && &buf[..4] == WIRE_MAGIC {
1679 return self.import_wire_v2(buf);
1680 }
1681 if !self.linear_wire_allowed {
1683 return Err(
1684 "kv import: Delta linear state cannot use the unversioned cache wire; refusing until the wire carries operator identity".into(),
1685 );
1686 }
1687 if self.bounded.is_some() {
1688 return Err(
1689 "kv import: a bounded anchor takes only the versioned wire (v2) — the \
1690 unversioned body has no ring record"
1691 .into(),
1692 );
1693 }
1694 let n = self.import_full_body(buf)?;
1695 if n != buf.len() {
1696 return Err(format!(
1697 "kv import: {} trailing byte(s) after the record",
1698 buf.len() - n
1699 ));
1700 }
1701 Ok(())
1702 }
1703
1704 fn import_wire_v2(&mut self, buf: &[u8]) -> Result<(), String> {
1705 let need = |n: usize, o: usize| -> Result<(), String> {
1706 if o + n > buf.len() {
1707 Err("kv import: truncated header".into())
1708 } else {
1709 Ok(())
1710 }
1711 };
1712 need(28, 0)?;
1713 let version = u32::from_le_bytes(buf[4..8].try_into().unwrap());
1714 if version != WIRE_VERSION {
1715 return Err(format!(
1716 "kv import: wire version {version}, this runtime speaks {WIRE_VERSION}"
1717 ));
1718 }
1719 let identity = u64::from_le_bytes(buf[8..16].try_into().unwrap());
1720 let layer = u32::from_le_bytes(buf[16..20].try_into().unwrap());
1721 let kind = WireKind::from_u8(buf[20])
1722 .ok_or_else(|| format!("kv import: unknown state kind {}", buf[20]))?;
1723 let f16 = buf[21] != 0;
1724 let position = u64::from_le_bytes(buf[24..32].try_into().unwrap()) as usize;
1725 if identity != self.wire_identity {
1726 return Err(format!(
1727 "kv import: peer operator identity {identity:016x} != mine {:016x} — \
1728 the two sides do not hold the same operator",
1729 self.wire_identity
1730 ));
1731 }
1732 if layer != self.wire_layer {
1733 return Err(format!(
1734 "kv import: record is for layer {layer}, this is layer {}",
1735 self.wire_layer
1736 ));
1737 }
1738 let fits = kind == self.wire_kind
1740 || (kind == WireKind::FullTail && self.wire_kind == WireKind::Full);
1741 if !fits {
1742 return Err(format!(
1743 "kv import: record kind {kind:?} does not match this layer's {:?}",
1744 self.wire_kind
1745 ));
1746 }
1747 let mut o = 32usize;
1748 let u32_at = |o: &mut usize| -> Result<u32, String> {
1749 if *o + 4 > buf.len() {
1750 return Err("kv import: truncated record".into());
1751 }
1752 let v = u32::from_le_bytes(buf[*o..*o + 4].try_into().unwrap());
1753 *o += 4;
1754 Ok(v)
1755 };
1756 let need_payload = |n: usize, o: usize| -> Result<(), String> {
1757 if o + n > buf.len() {
1758 Err("kv import: truncated payload".into())
1759 } else {
1760 Ok(())
1761 }
1762 };
1763 match kind {
1764 WireKind::Full => {
1765 let n = self.import_full_body(&buf[o..])?;
1766 o += n;
1767 if self.seq_len != position {
1768 return Err(format!(
1769 "kv import: header position {position} != record seq_len {}",
1770 self.seq_len
1771 ));
1772 }
1773 }
1774 WireKind::FullTail => {
1775 need_payload(8, o)?;
1776 let base = u64::from_le_bytes(buf[o..o + 8].try_into().unwrap()) as usize;
1777 o += 8;
1778 let n = self.import_full_body(&buf[o..])?;
1779 o += n;
1780 if base.checked_add(self.seq_len) != Some(position) {
1781 let rows = self.seq_len;
1782 self.reset_per_position_storage();
1783 self.seq_len = 0;
1784 return Err(format!(
1785 "kv import: tail base {base} + {rows} rows != header position {position}"
1786 ));
1787 }
1788 self.base = base;
1789 }
1790 WireKind::Linear => {
1791 let heads = u32_at(&mut o)? as usize;
1792 let hd = u32_at(&mut o)? as usize;
1793 if heads != self.num_kv_heads || hd != self.head_dim {
1794 return Err(format!(
1795 "kv import: peer sent {heads}×{hd} per position, this layer is {}×{}",
1796 self.num_kv_heads, self.head_dim
1797 ));
1798 }
1799 let lin = u32_at(&mut o)? as usize;
1800 need_payload(lin * 4, o)?;
1801 self.linear_state = (0..lin)
1802 .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1803 .collect();
1804 o += lin * 4;
1805 self.reset_per_position_storage();
1806 self.seq_len = position;
1807 }
1808 WireKind::Bounded => {
1809 let heads = u32_at(&mut o)? as usize;
1810 let hd = u32_at(&mut o)? as usize;
1811 let window = u32_at(&mut o)? as usize;
1812 let len = u32_at(&mut o)? as usize;
1813 let head = u32_at(&mut o)? as usize;
1814 let b = self
1815 .bounded
1816 .as_mut()
1817 .ok_or("kv import: bounded record for a layer without a ring")?;
1818 if heads != b.num_kv_heads || hd != b.head_dim || window != b.window {
1819 return Err(format!(
1820 "kv import: bounded record {heads}×{window}×{hd} does not fit this \
1821 layer's ring {}×{}×{}",
1822 b.num_kv_heads, b.window, b.head_dim
1823 ));
1824 }
1825 if len != position.min(window) || head != position % window {
1826 return Err(format!(
1827 "kv import: bounded record len/head {len}/{head} inconsistent with \
1828 position {position} (window {window})"
1829 ));
1830 }
1831 let n = b.ring_k.len();
1832 let w = if f16 { 2 } else { 4 };
1833 need_payload(2 * n * w, o)?;
1834 let read = |o: usize, dst: &mut [f32]| {
1835 for (i, d) in dst.iter_mut().enumerate() {
1836 let at = o + i * w;
1837 *d = if f16 {
1838 cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
1839 buf[at..at + 2].try_into().unwrap(),
1840 ))
1841 } else {
1842 f32::from_le_bytes(buf[at..at + 4].try_into().unwrap())
1843 };
1844 }
1845 };
1846 read(o, &mut b.ring_k);
1847 o += n * w;
1848 read(o, &mut b.ring_v);
1849 o += n * w;
1850 b.seen = position;
1851 let snap = b.snapshot();
1853 b.restore(&snap);
1854 self.linear_state = Vec::new();
1855 self.reset_per_position_storage();
1856 self.seq_len = position;
1857 }
1858 }
1859 if o != buf.len() {
1860 return Err(format!(
1861 "kv import: {} trailing byte(s) after the record",
1862 buf.len() - o
1863 ));
1864 }
1865 Ok(())
1866 }
1867
1868 fn reset_per_position_storage(&mut self) {
1871 let heads = self.num_kv_heads;
1872 self.mode = KvMode::F32;
1873 self.k = vec![Vec::new(); heads];
1874 self.v = vec![Vec::new(); heads];
1875 self.kv_rows_changed(0);
1876 self.kq = vec![Vec::new(); heads];
1877 self.ks = vec![Vec::new(); heads];
1878 self.vq = vec![Vec::new(); heads];
1879 self.vs = vec![Vec::new(); heads];
1880 self.kcol = vec![Vec::new(); heads];
1881 self.vcol = vec![Vec::new(); heads];
1882 self.imp = Vec::new();
1883 self.discard_linear_scratch();
1884 self.o1 = None;
1885 self.o1_error = None;
1886 self.o1_transitioned = false;
1887 self.base = 0;
1888 self.generation = next_gen();
1889 }
1890
1891 fn import_full_body(&mut self, buf: &[u8]) -> Result<usize, String> {
1894 let mut o = 0usize;
1895 let u32_at = |o: &mut usize| -> Result<u32, String> {
1896 if *o + 4 > buf.len() {
1897 return Err("kv import: truncated header".into());
1898 }
1899 let v = u32::from_le_bytes(buf[*o..*o + 4].try_into().unwrap());
1900 *o += 4;
1901 Ok(v)
1902 };
1903 let f16 = u32_at(&mut o)? != 0;
1904 let seq_len = u32_at(&mut o)? as usize;
1905 let heads = u32_at(&mut o)? as usize;
1906 let hd = u32_at(&mut o)? as usize;
1907 let lin = u32_at(&mut o)? as usize;
1908 let nimp = u32_at(&mut o)? as usize;
1909 if heads != self.num_kv_heads || hd != self.head_dim {
1910 return Err(format!(
1911 "kv import: peer sent {heads}×{hd} per position, this layer is {}×{}",
1912 self.num_kv_heads, self.head_dim
1913 ));
1914 }
1915 let w = if f16 { 2 } else { 4 };
1916 let need = |n: usize, o: usize| -> Result<(), String> {
1917 if o + n > buf.len() {
1918 Err("kv import: truncated payload".into())
1919 } else {
1920 Ok(())
1921 }
1922 };
1923 need(lin * 4, o)?;
1924 self.linear_state = (0..lin)
1925 .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1926 .collect();
1927 o += lin * 4;
1928 need(nimp * 4, o)?;
1929 let imp: Vec<f32> = (0..nimp)
1930 .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1931 .collect();
1932 o += nimp * 4;
1933 let mut k: Vec<Vec<f32>> = Vec::with_capacity(heads);
1934 let mut v: Vec<Vec<f32>> = Vec::with_capacity(heads);
1935 for _ in 0..heads {
1936 for which in 0..2 {
1937 let n = u32_at(&mut o)? as usize;
1938 need(n * w, o)?;
1939 let xs: Vec<f32> = (0..n)
1940 .map(|i| {
1941 let at = o + i * w;
1942 if f16 {
1943 cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
1944 buf[at..at + 2].try_into().unwrap(),
1945 ))
1946 } else {
1947 f32::from_le_bytes(buf[at..at + 4].try_into().unwrap())
1948 }
1949 })
1950 .collect();
1951 o += n * w;
1952 if which == 0 { k.push(xs) } else { v.push(xs) }
1953 }
1954 }
1955 self.reset_per_position_storage();
1956 self.k = k;
1957 self.v = v;
1958 self.imp = imp;
1959 self.seq_len = seq_len;
1960 Ok(o)
1961 }
1962
1963 pub fn memory_bytes(&self) -> usize {
1964 let floats: usize = self.k.iter().map(Vec::len).sum::<usize>()
1965 + self.v.iter().map(Vec::len).sum::<usize>()
1966 + self.ks.iter().map(Vec::len).sum::<usize>()
1967 + self.vs.iter().map(Vec::len).sum::<usize>()
1968 + self.kcol.iter().map(Vec::len).sum::<usize>()
1969 + self.vcol.iter().map(Vec::len).sum::<usize>();
1970 let bytes: usize = self.kq.iter().map(Vec::len).sum::<usize>()
1971 + self.vq.iter().map(Vec::len).sum::<usize>();
1972 floats * std::mem::size_of::<f32>()
1973 + bytes
1974 + self.linear_state.len() * std::mem::size_of::<f32>()
1978 + self.o1_memory_bytes()
1981 + self.bounded_state_bytes()
1984 }
1985
1986 fn evict(&mut self, keep_last: usize) {
1988 if self.o1.is_some()
1999 || self.bounded.is_some()
2000 || self.tail.is_some()
2001 || self.seq_len <= keep_last
2002 {
2003 return;
2004 }
2005 self.generation = next_gen();
2006 let drop = self.seq_len - keep_last;
2007 for h in 0..self.num_kv_heads {
2008 let stored = self.head_len(h);
2010 let d = drop.min(stored);
2011 let hd = self.head_dim;
2012 fn drop_front<T>(v: &mut Vec<T>, n: usize) {
2013 let n = n.min(v.len());
2014 v.drain(..n);
2015 }
2016 drop_front(&mut self.k[h], d * hd);
2017 drop_front(&mut self.v[h], d * hd);
2018 self.kv_rows_changed(0);
2019 drop_front(&mut self.kq[h], d * hd);
2020 drop_front(&mut self.vq[h], d * hd);
2021 drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
2022 drop_front(&mut self.vs[h], d);
2023 }
2024 let d = drop.min(self.imp.len());
2025 self.imp.drain(..d);
2026 self.seq_len = keep_last;
2027 }
2028
2029 fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
2034 if self.o1.is_some() || self.bounded.is_some() || self.tail.is_some() {
2035 return;
2040 }
2041 let stored = self.imp.len();
2042 if stored <= keep_last {
2043 return;
2044 }
2045 self.generation = next_gen();
2046 let sink_n = sink.min(keep_last);
2049 let recent_n = recent.min(keep_last - sink_n);
2050 let mut keep = vec![false; stored];
2051 for k in keep.iter_mut().take(sink_n) {
2052 *k = true;
2053 }
2054 for k in keep.iter_mut().skip(stored.saturating_sub(recent_n)) {
2055 *k = true;
2056 }
2057 let mut budget = keep_last.saturating_sub(keep.iter().filter(|&&x| x).count());
2058 let mut order: Vec<usize> = (0..stored).filter(|&i| !keep[i]).collect();
2060 order.sort_by(|&a, &b| {
2061 self.imp[b]
2062 .partial_cmp(&self.imp[a])
2063 .unwrap_or(std::cmp::Ordering::Equal)
2064 });
2065 for i in order {
2066 if budget == 0 {
2067 break;
2068 }
2069 keep[i] = true;
2070 budget -= 1;
2071 }
2072
2073 let kept: Vec<usize> = (0..stored).filter(|&i| keep[i]).collect();
2074 let hd = self.head_dim;
2075 fn gather<T: Copy>(src: &[T], kept: &[usize], step: usize) -> Vec<T> {
2076 let mut out = Vec::with_capacity(kept.len() * step);
2077 for &i in kept {
2078 out.extend_from_slice(&src[i * step..(i + 1) * step]);
2079 }
2080 out
2081 }
2082 for h in 0..self.num_kv_heads {
2087 if !self.k[h].is_empty() {
2088 self.k[h] = gather(&self.k[h], &kept, hd);
2089 self.kv_rows_changed(0);
2090 }
2091 if !self.v[h].is_empty() {
2092 self.v[h] = gather(&self.v[h], &kept, hd);
2093 self.kv_rows_changed(0);
2094 }
2095 if !self.kq[h].is_empty() {
2096 self.kq[h] = gather(&self.kq[h], &kept, hd);
2097 self.ks[h] = gather(&self.ks[h], &kept, hd.div_ceil(KV_K_GROUP));
2098 }
2099 if !self.vq[h].is_empty() {
2100 self.vq[h] = gather(&self.vq[h], &kept, hd);
2101 self.vs[h] = gather(&self.vs[h], &kept, 1);
2102 }
2103 }
2104 self.imp = kept.iter().map(|&i| self.imp[i]).collect();
2105 self.seq_len = kept.len();
2106 }
2107}
2108
2109#[derive(Debug, Clone, Copy, PartialEq, Eq)]
2111pub enum EvictionPolicy {
2112 Recent,
2114 Born { sink: usize },
2116}
2117
2118#[derive(Debug)]
2120pub struct KvCache {
2121 pub layers: Vec<LayerKvCache>,
2122 pub max_seq_len: usize,
2123 pub policy: EvictionPolicy,
2124}
2125
2126impl KvCache {
2127 pub fn new(
2128 num_layers: usize,
2129 num_kv_heads: usize,
2130 head_dim: usize,
2131 max_seq_len: usize,
2132 ) -> Self {
2133 let layers = (0..num_layers)
2134 .map(|li| {
2135 let mut l = LayerKvCache::new(num_kv_heads, head_dim);
2136 l.wire_layer = li as u32;
2137 l
2138 })
2139 .collect();
2140 Self {
2141 layers,
2142 max_seq_len,
2143 policy: EvictionPolicy::Born { sink: 4 },
2144 }
2145 }
2146
2147 pub fn clear(&mut self) {
2148 for layer in &mut self.layers {
2149 layer.clear();
2150 }
2151 }
2152
2153 pub fn total_memory_bytes(&self) -> usize {
2154 self.layers.iter().map(|l| l.memory_bytes()).sum()
2155 }
2156
2157 pub fn recurrent_state_bytes(&self) -> usize {
2162 let floats: usize = self
2163 .layers
2164 .iter()
2165 .map(|l| l.linear_state.len() + l.linear_scratch.len())
2166 .sum();
2167 floats * std::mem::size_of::<f32>()
2168 }
2169
2170 pub fn attention_state_bytes(&self) -> usize {
2173 self.total_memory_bytes()
2174 .saturating_sub(self.recurrent_state_bytes())
2175 }
2176
2177 pub fn seq_len(&self) -> usize {
2181 self.layers.iter().map(|l| l.pos_len()).max().unwrap_or(0)
2182 }
2183
2184 pub fn bounded_state_bytes(&self) -> usize {
2186 self.layers.iter().map(|l| l.bounded_state_bytes()).sum()
2187 }
2188
2189 pub fn needs_eviction(&self) -> bool {
2194 self.layers
2195 .iter()
2196 .filter(|l| l.bounded.is_none() && l.tail.is_none())
2197 .map(|l| l.seq_len)
2198 .max()
2199 .unwrap_or(0)
2200 >= self.max_seq_len
2201 }
2202
2203 pub fn evict(&mut self, keep_last: usize) {
2205 match self.policy {
2206 EvictionPolicy::Recent => {
2207 for layer in &mut self.layers {
2208 layer.evict(keep_last);
2209 }
2210 }
2211 EvictionPolicy::Born { sink } => {
2212 let recent = (keep_last / 2).max(1);
2213 for layer in &mut self.layers {
2214 layer.evict_born(keep_last, sink, recent);
2215 }
2216 }
2217 }
2218 }
2219}
2220
2221#[cfg(test)]
2222mod tests {
2223 use super::*;
2224
2225 #[test]
2231 fn kv_abs_max_follows_appends_truncation_and_clear() {
2232 let (heads, hd) = (2usize, 4usize);
2233 let mut c = LayerKvCache::new(heads, hd);
2234 c.mode = KvMode::F32;
2235 let row = |x: f32| vec![x; heads * hd];
2236 for _ in 0..4 {
2237 c.append(&row(0.5), &row(-0.25), &[]);
2238 }
2239 assert_eq!(c.kv_abs_max(), (0.5, 0.25));
2240 c.truncate_last(2);
2241 c.append(&row(-9.0), &row(3.0), &[]);
2242 c.append(&row(0.5), &row(0.25), &[]);
2243 assert_eq!(c.kv_abs_max(), (9.0, 3.0), "rewritten rows were not rescanned");
2244 c.append(&row(f32::NAN), &row(0.0), &[]);
2245 assert_eq!(c.kv_abs_max().0, f32::INFINITY);
2246 c.clear();
2247 c.append(&row(1.0), &row(2.0), &[]);
2248 assert_eq!(c.kv_abs_max(), (1.0, 2.0));
2249 c.append(&row(4.0), &row(1.0), &[]);
2250 assert_eq!(c.kv_abs_max_after((1, 2, 4.0, 1.0)), (4.0, 2.0));
2252 c.append(&row(8.0), &row(1.0), &[]);
2253 assert_eq!(c.kv_abs_max_after((0, 3, 0.0, 0.0)), (8.0, 2.0));
2255 }
2256
2257 #[test]
2261 fn kv_abs_max_scans_rows_appended_after_a_trim() {
2262 let (heads, hd, w) = (1usize, 4usize, 8usize);
2263 let mut c = LayerKvCache::new(heads, hd);
2264 c.mode = KvMode::F32;
2265 let row = |x: f32| vec![x; heads * hd];
2266 for _ in 0..(2 * w + 4) {
2267 c.append(&row(0.5), &row(0.5), &[]);
2268 }
2269 assert_eq!(c.kv_abs_max(), (0.5, 0.5));
2270 let d = c.trim_window(w, 2, 2);
2271 assert!(d > 0, "the trim must drop rows for this test");
2272 let first_new = c.kv_rows();
2276 c.append(&row(-7.0), &row(6.0), &[]);
2277 for _ in 0..(d + 3) {
2278 c.append(&row(0.5), &row(0.5), &[]);
2279 }
2280 assert!(c.kv_rows() > 2 * w + 4 && first_new < 2 * w + 4);
2281 assert_eq!(c.kv_abs_max(), (7.0, 6.0));
2282 }
2283
2284 #[test]
2285 fn memory_breakdown_separates_recurrent_and_attention_state() {
2286 let mut cache = KvCache::new(1, 1, 4, 16);
2287 cache.layers[0].linear_state = vec![0.0; 8];
2288 cache.layers[0].linear_scratch = vec![0.0; 4];
2289 cache.layers[0].append(&[0.0; 4], &[1.0; 4], &[true]);
2290 let recurrent = cache.recurrent_state_bytes();
2291 assert_eq!(recurrent, 12 * std::mem::size_of::<f32>());
2292 assert_eq!(
2293 cache.attention_state_bytes() + recurrent,
2294 cache.total_memory_bytes()
2295 );
2296 assert!(cache.attention_state_bytes() > 0);
2297 }
2298
2299 #[test]
2300 fn wire_round_trip_reproduces_attention() {
2301 let (heads, hd) = (2usize, 4usize);
2305 let mut a = LayerKvCache::new(heads, hd);
2306 for p in 0..5 {
2307 let k: Vec<f32> = (0..heads * hd)
2308 .map(|i| (p * 10 + i) as f32 * 0.031)
2309 .collect();
2310 let v: Vec<f32> = (0..heads * hd)
2311 .map(|i| (p * 7 + i) as f32 * -0.017)
2312 .collect();
2313 a.append(&k, &v, &[true, true]);
2314 }
2315 a.linear_state = vec![0.5, -0.25, 1.0];
2316 let q: Vec<f32> = (0..hd).map(|i| 0.1 * (i as f32 + 1.0)).collect();
2317
2318 let bytes = a.export_wire(false).expect("f32 cache exports");
2319 let mut b = LayerKvCache::new(heads, hd);
2320 b.linear_scratch = vec![9.0; 3];
2321 b.import_wire(&bytes).expect("import");
2322 assert!(
2323 b.linear_scratch.is_empty(),
2324 "import must discard tentative state"
2325 );
2326
2327 assert_eq!(b.seq_len, a.seq_len);
2328 assert_eq!(b.linear_state, a.linear_state);
2329 for h in 0..heads {
2330 let (oa, sa) = a.attend(&q, h);
2331 let (ob, sb) = b.attend(&q, h);
2332 assert_eq!(oa, ob, "head {h} attention output diverged");
2333 assert_eq!(sa, sb, "head {h} attention scores diverged");
2334 }
2335 }
2336
2337 #[test]
2338 fn wire_refuses_what_it_cannot_describe() {
2339 let mut c = LayerKvCache::new(1, 4);
2342 c.mode = KvMode::Q8 { k: true, v: true };
2343 let err = c.export_wire(false).unwrap_err();
2344 assert!(err.contains("F32"), "{err}");
2345 }
2346
2347 #[test]
2348 fn wire_refuses_unversioned_delta_state() {
2349 let mut c = LayerKvCache::new(1, 4);
2352 c.set_linear_wire_allowed(false);
2353 let bytes = c.export_wire(false).expect("v2 export carries identity");
2354 assert_eq!(&bytes[..4], WIRE_MAGIC);
2355 let err = c.import_wire(&[0, 0, 0, 0]).unwrap_err();
2356 assert!(err.contains("operator identity"), "{err}");
2357 }
2358
2359 #[test]
2360 fn wire_v2_round_trips_linear_and_bounded_records() {
2361 let mut a = LayerKvCache::new(1, 4);
2364 a.wire_kind = WireKind::Linear;
2365 a.wire_identity = 0xC0FFEE;
2366 a.linear_state = vec![0.5, -0.25, 1.0, 3.5];
2367 a.seq_len = 9;
2368 let bytes = a.export_wire(true).unwrap();
2369 let mut b = LayerKvCache::new(1, 4);
2370 b.wire_kind = WireKind::Linear;
2371 b.wire_identity = 0xC0FFEE;
2372 b.import_wire(&bytes).unwrap();
2373 assert_eq!(b.linear_state, a.linear_state);
2374 assert_eq!(b.seq_len, 9);
2375 let mut c = LayerKvCache::new(1, 4);
2377 c.wire_kind = WireKind::Linear;
2378 let err = c.import_wire(&bytes).unwrap_err();
2379 assert!(err.contains("operator identity"), "{err}");
2380
2381 let (kvh, hd, w) = (2, 4, 8);
2383 let mut a = LayerKvCache::new(kvh, hd);
2384 a.install_bounded(w);
2385 let k: Vec<f32> = (0..kvh * hd).map(|i| i as f32 * 0.125).collect();
2386 for p in 0..11 {
2387 a.bounded.as_mut().unwrap().insert(&k, &k);
2388 a.seq_len = p + 1;
2389 }
2390 for f16 in [false, true] {
2391 let bytes = a.export_wire(f16).unwrap();
2392 let mut b = LayerKvCache::new(kvh, hd);
2393 b.install_bounded(w);
2394 b.import_wire(&bytes).unwrap();
2395 let (ra, rb) = (a.bounded.as_ref().unwrap(), b.bounded.as_ref().unwrap());
2396 assert_eq!(rb.seen, 11);
2397 assert_eq!(b.seq_len, 11);
2398 assert!(ra.same_state(rb), "f16={f16}");
2400 let mut c = LayerKvCache::new(kvh, hd);
2402 c.install_bounded(w * 2);
2403 assert!(c.import_wire(&bytes).is_err());
2404 }
2405 }
2406
2407 #[test]
2408 fn wire_import_checks_geometry() {
2409 let a = LayerKvCache::new(2, 4);
2410 let bytes = a.export_wire(false).unwrap();
2411 let mut wrong = LayerKvCache::new(2, 8);
2412 let err = wrong.import_wire(&bytes).unwrap_err();
2413 assert!(err.contains("2×4"), "{err}");
2414 }
2415
2416 #[test]
2417 fn append_tracks_seq_len_and_layout() {
2418 let mut cache = LayerKvCache::new(4, 8);
2419 cache.mode = KvMode::F32;
2420 assert_eq!(cache.seq_len, 0);
2421
2422 let k: Vec<f32> = (0..32).map(|i| i as f32).collect();
2423 let v = vec![2.0f32; 32];
2424 cache.append(&k, &v, &[true; 4]);
2425
2426 assert_eq!(cache.seq_len, 1);
2427 assert_eq!(cache.head_len(0), 1);
2428 assert_eq!(cache.head_keys(1), &k[8..16]);
2430 assert_eq!(cache.memory_bytes(), 256);
2431 }
2432
2433 #[test]
2434 fn dead_head_stores_nothing() {
2435 let mut cache = LayerKvCache::new(2, 4);
2436 cache.mode = KvMode::F32;
2437 let k = vec![1.0f32; 8];
2438 let v = vec![2.0f32; 8];
2439 cache.append(&k, &v, &[true, false]);
2440 cache.append(&k, &v, &[true, false]);
2441
2442 assert_eq!(cache.seq_len, 2);
2443 assert_eq!(cache.head_len(0), 2);
2444 assert_eq!(cache.head_len(1), 0, "dead head must not store KV");
2445 assert_eq!(cache.memory_bytes(), 2 * 2 * 4 * 4);
2446 }
2447
2448 #[test]
2449 fn eviction_keeps_recent() {
2450 let mut cache = KvCache::new(2, 4, 8, 10);
2451 cache.policy = EvictionPolicy::Recent;
2452 for l in &mut cache.layers {
2453 l.mode = KvMode::F32;
2454 }
2455 let k = vec![1.0f32; 32];
2456 let v = vec![2.0f32; 32];
2457 for _ in 0..8 {
2458 for layer in &mut cache.layers {
2459 layer.append(&k, &v, &[true; 4]);
2460 }
2461 }
2462 assert_eq!(cache.seq_len(), 8);
2463 assert!(!cache.needs_eviction());
2464
2465 cache.evict(4);
2466 assert_eq!(cache.seq_len(), 4);
2467 assert_eq!(cache.layers[0].head_len(0), 4);
2468 }
2469
2470 #[test]
2471 fn collecting_o1_eviction_retains_exact_storage_until_boundary() {
2472 const B: usize = 19;
2473 let q = vec![0.1f32; 8];
2474 let k = vec![0.2f32; 4];
2475 let v = vec![0.3f32; 4];
2476
2477 for policy in [EvictionPolicy::Recent, EvictionPolicy::Born { sink: 2 }] {
2478 let mut cache = KvCache::new(1, 1, 4, 6);
2479 cache.policy = policy;
2480 cache.layers[0].mode = KvMode::F32;
2481 cache.layers[0].o1_begin_with_boundary(
2482 4,
2483 8,
2484 2,
2485 crate::nystrom::O1Rect::Aggregate,
2486 Some(B),
2487 );
2488
2489 for pos in 0..B {
2490 {
2491 let layer = &mut cache.layers[0];
2492 layer.o1_push_q(&q);
2493 layer.append(&k, &v, &[]);
2494 }
2495 if pos + 1 < B {
2496 cache.evict(3);
2497 }
2498 }
2499
2500 let layer = &cache.layers[0];
2501 let rows = B * layer.head_dim;
2502 assert_eq!(layer.seq_len, B, "policy {policy:?} retained depth");
2503 assert_eq!(layer.k[0].len(), rows, "policy {policy:?} K rows");
2504 assert_eq!(layer.v[0].len(), rows, "policy {policy:?} V rows");
2505 assert!(
2506 layer.k[0].capacity() >= rows,
2507 "policy {policy:?} K capacity"
2508 );
2509 assert!(
2510 layer.v[0].capacity() >= rows,
2511 "policy {policy:?} V capacity"
2512 );
2513 let q_capacity = match layer.o1.as_ref() {
2514 Some(O1State::Collecting { q_buf, .. }) => q_buf.capacity(),
2515 other => panic!("policy {policy:?} changed state early: {other:?}"),
2516 };
2517 assert!(
2518 q_capacity >= B * 8,
2519 "policy {policy:?} Q capacity must cover the exact prefix"
2520 );
2521
2522 assert!(cache.layers[0].o1_seal_checked(2).unwrap());
2523 assert_eq!(cache.layers[0].k[0].capacity(), 0, "K released after seal");
2524 assert_eq!(cache.layers[0].v[0].capacity(), 0, "V released after seal");
2525 }
2526 }
2527
2528 #[test]
2529 fn truncate_rolls_back_speculative_positions() {
2530 let mut cache = LayerKvCache::new(2, 4);
2531 cache.mode = KvMode::F32;
2532 for pos in 0..5 {
2533 let k = vec![pos as f32; 8];
2534 let v = vec![pos as f32; 8];
2535 cache.append(&k, &v, &[true; 2]);
2536 }
2537 cache.truncate_last(2);
2538 assert_eq!(cache.seq_len, 3);
2539 assert_eq!(cache.head_len(0), 3);
2540 assert_eq!(cache.head_keys(0)[2 * 4], 2.0, "position 2 survives");
2541 }
2542
2543 #[test]
2547 fn q8_attend_matches_f32_within_grid() {
2548 let (heads, hd) = (2, 32);
2549 let mut f = LayerKvCache::new(heads, hd);
2550 f.mode = KvMode::F32;
2551 let mut q8 = LayerKvCache::new(heads, hd);
2552 q8.mode = KvMode::Q8 { k: true, v: true };
2553
2554 let synth = |p: usize, salt: usize| -> Vec<f32> {
2555 (0..heads * hd)
2556 .map(|i| {
2557 let x = ((i * 31 + p * 17 + salt * 7 + 3) % 97) as f32 / 97.0 - 0.5;
2558 if i % 2 == 0 { x * 4.0 } else { x * 0.25 }
2560 })
2561 .collect()
2562 };
2563 for p in 0..100 {
2564 let k = synth(p, 1);
2565 let v = synth(p, 2);
2566 f.append(&k, &v, &[true; 2]);
2567 q8.append(&k, &v, &[true; 2]);
2568 }
2569 let q: Vec<f32> = (0..hd)
2570 .map(|i| ((i * 13 + 5) % 89) as f32 / 89.0 - 0.5)
2571 .collect();
2572 for g in 0..heads {
2573 let (of, pf) = f.attend(&q, g);
2574 let (o8, p8) = q8.attend(&q, g);
2575 let scale = of.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-6);
2576 for d in 0..hd {
2577 assert!(
2578 (of[d] - o8[d]).abs() <= scale * 0.03 + 1e-3,
2579 "g{g} d{d}: f32 {} vs q8 {}",
2580 of[d],
2581 o8[d]
2582 );
2583 }
2584 for p in 0..100 {
2585 assert!((pf[p] - p8[p]).abs() < 0.02, "prob p{p}");
2586 }
2587 }
2588 q8.truncate_last(30);
2590 assert_eq!(q8.head_len(0), 70);
2591 let imp: Vec<f32> = (0..70).map(|i| i as f32).collect();
2592 q8.accumulate_imp(&imp);
2593 q8.evict_born(20, 2, 8);
2594 assert_eq!(q8.head_len(0), 20);
2595 let (o, _) = q8.attend(&q, 0);
2596 assert!(o.iter().all(|x| x.is_finite()));
2597 assert!(q8.memory_bytes() * 3 < f.memory_bytes());
2599 }
2600
2601 #[test]
2604 fn attend_group_equals_per_head_attend_bitexact() {
2605 let (kv_heads, hd, hpk) = (2usize, 32usize, 3usize); for mode in [KvMode::F32, KvMode::Q8 { k: true, v: true }] {
2607 let mut c = LayerKvCache::new(kv_heads, hd);
2608 c.mode = mode;
2609 for p in 0..70 {
2610 let k: Vec<f32> = (0..kv_heads * hd)
2611 .map(|i| ((i * 31 + p * 17 + 3) % 97) as f32 / 97.0 - 0.5)
2612 .collect();
2613 let v: Vec<f32> = (0..kv_heads * hd)
2614 .map(|i| ((i * 13 + p * 29 + 7) % 89) as f32 / 89.0 - 0.5)
2615 .collect();
2616 c.append(&k, &v, &[true; 2]);
2617 }
2618 let q: Vec<f32> = (0..kv_heads * hpk * hd)
2619 .map(|i| ((i * 11 + 5) % 83) as f32 / 83.0 - 0.5)
2620 .collect();
2621 for g in 0..kv_heads {
2622 let span = g * hpk * hd..(g + 1) * hpk * hd;
2623 let mut out = vec![0f32; hpk * hd];
2624 let mut imp = vec![0f32; 70];
2625 c.attend_group(
2626 &q[span.clone()],
2627 g,
2628 &mut out,
2629 &mut imp,
2630 1.0 / (hd as f32).sqrt(),
2631 0,
2632 0.0,
2633 &[],
2634 );
2635 let mut imp_ref = vec![0f32; 70];
2636 for h in 0..hpk {
2637 let qh = &q[span.start + h * hd..span.start + (h + 1) * hd];
2638 let (o, probs) = c.attend(qh, g);
2639 assert_eq!(
2640 &out[h * hd..(h + 1) * hd],
2641 &o[..],
2642 "mode {mode:?} g{g} h{h}: grouped attend must be bit-identical"
2643 );
2644 for (dst, &p) in imp_ref.iter_mut().zip(&probs) {
2645 *dst += p;
2646 }
2647 }
2648 assert_eq!(
2649 imp, imp_ref,
2650 "mode {mode:?} g{g}: attention mass must match"
2651 );
2652 }
2653 }
2654 }
2655
2656 #[test]
2663 fn sink_attend_matches_explicit_sink_column() {
2664 let (nkv, hd, hpk) = (2usize, 8usize, 3usize);
2665 let rows = 9usize;
2666 let mut c = LayerKvCache::new(nkv, hd);
2667 c.mode = KvMode::F32;
2668 let kv = |r: usize, i: usize, a: usize, m: usize| {
2669 (((r * a + i * 7 + 3) % m) as f32 / m as f32 - 0.5) * 2.0
2670 };
2671 let mut ks = Vec::new();
2672 let mut vs = Vec::new();
2673 for r in 0..rows {
2674 let k: Vec<f32> = (0..nkv * hd).map(|i| kv(r, i, 31, 97)).collect();
2675 let v: Vec<f32> = (0..nkv * hd).map(|i| kv(r, i, 17, 89)).collect();
2676 c.append(&k, &v, &[]);
2677 ks.push(k);
2678 vs.push(v);
2679 }
2680 let q: Vec<f32> = (0..nkv * hpk * hd)
2681 .map(|i| (((i * 11 + 5) % 83) as f32 / 83.0 - 0.5) * 3.0)
2682 .collect();
2683 let sinks = [0.7f32, -1.3, 2.5, 0.0, -4.0, 6.0];
2685 let scale = 1.0 / (hd as f32).sqrt();
2686 let mut checked = 0usize;
2687 for upto in [1usize, 2, 5, 9] {
2688 for window in [None, Some(3usize), Some(1)] {
2689 let first = window.map(|w| upto.saturating_sub(w)).unwrap_or(0);
2690 for g in 0..nkv {
2691 let qg = &q[g * hpk * hd..(g + 1) * hpk * hd];
2692 let sg = &sinks[g * hpk..(g + 1) * hpk];
2693 let mut out = vec![0f32; hpk * hd];
2694 let mut imp = vec![0f32; upto];
2695 c.attend_group_upto(qg, g, &mut out, &mut imp, scale, first, 0.0, upto, sg);
2696 let mut imp_ref = vec![0f64; upto];
2697 for h in 0..hpk {
2698 let qh = &qg[h * hd..(h + 1) * hd];
2699 let mut z: Vec<f64> = (first..upto)
2701 .map(|p| {
2702 let k = &ks[p][g * hd..(g + 1) * hd];
2703 qh.iter()
2704 .zip(k)
2705 .map(|(&a, &b)| a as f64 * b as f64)
2706 .sum::<f64>()
2707 * scale as f64
2708 })
2709 .collect();
2710 z.push(sg[h] as f64);
2711 let m = z.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
2712 let e: Vec<f64> = z.iter().map(|&x| (x - m).exp()).collect();
2713 let s: f64 = e.iter().sum();
2714 let p: Vec<f64> = e.iter().map(|&x| x / s).collect();
2715 for d in 0..hd {
2717 let want: f64 = (first..upto)
2718 .map(|r| p[r - first] * vs[r][g * hd + d] as f64)
2719 .sum();
2720 let got = out[h * hd + d] as f64;
2721 assert!(
2722 (got - want).abs() < 1e-6,
2723 "upto {upto} window {window:?} g{g} h{h} d{d}: {got} vs {want}"
2724 );
2725 }
2726 for r in first..upto {
2727 imp_ref[r] += p[r - first];
2728 }
2729 checked += 1;
2730 }
2731 for r in 0..upto {
2732 assert!(
2733 (imp[r] as f64 - imp_ref[r]).abs() < 1e-6,
2734 "imp upto {upto} window {window:?} g{g} row {r}: {} vs {}",
2735 imp[r],
2736 imp_ref[r]
2737 );
2738 }
2739 let row_mass: f32 = imp.iter().sum();
2741 assert!(row_mass < hpk as f32, "sinks must absorb some mass");
2742 }
2743 }
2744 }
2745 assert_eq!(checked, 4 * 3 * nkv * hpk);
2746 }
2747
2748 #[test]
2753 fn windowed_attend_equals_masked_full_row() {
2754 let (nkv, hd, hpk) = (1usize, 16usize, 2usize);
2755 let rows = 40usize;
2756 let mut c = LayerKvCache::new(nkv, hd);
2757 c.mode = KvMode::F32;
2758 for r in 0..rows {
2759 let k: Vec<f32> = (0..hd)
2760 .map(|i| ((r * 13 + i * 5) % 29) as f32 / 29.0 - 0.5)
2761 .collect();
2762 let v: Vec<f32> = (0..hd)
2763 .map(|i| ((r * 7 + i * 3) % 31) as f32 / 31.0 - 0.5)
2764 .collect();
2765 c.append(&k, &v, &[]);
2766 }
2767 let q: Vec<f32> = (0..hpk * hd)
2768 .map(|i| ((i * 19) % 23) as f32 / 23.0 - 0.5)
2769 .collect();
2770 let scale = 0.25f32;
2771 for w in [1usize, 7, 39, 40, 100] {
2772 let first = rows.saturating_sub(w);
2773 let mut out = vec![0f32; hpk * hd];
2774 let mut imp = vec![0f32; rows];
2775 c.attend_group(&q, 0, &mut out, &mut imp, scale, first, 0.0, &[]);
2776 let mut out_ref = vec![0f32; hpk * hd];
2778 let mut imp_ref = vec![0f32; rows];
2779 for h in 0..hpk {
2780 let mut s = vec![f32::NEG_INFINITY; rows];
2781 for p in first..rows {
2782 s[p] = crate::attention::dot_f32(
2783 &q[h * hd..(h + 1) * hd],
2784 &c.head_keys(0)[p * hd..(p + 1) * hd],
2785 ) * scale;
2786 }
2787 let m = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2788 let mut sum = 0f32;
2789 for v in s.iter_mut() {
2790 *v = (*v - m).exp();
2791 sum += *v;
2792 }
2793 for v in s.iter_mut() {
2794 *v /= sum;
2795 }
2796 for p in first..rows {
2797 if s[p].abs() < 1e-12 {
2798 continue;
2799 }
2800 crate::attention::axpy_f32(
2801 &mut out_ref[h * hd..(h + 1) * hd],
2802 &c.head_values(0)[p * hd..(p + 1) * hd],
2803 s[p],
2804 );
2805 }
2806 for (d, &p) in imp_ref.iter_mut().zip(&s) {
2807 *d += p;
2808 }
2809 }
2810 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
2811 assert_eq!(bits(&out), bits(&out_ref), "out window {w}");
2812 assert_eq!(bits(&imp), bits(&imp_ref), "imp window {w}");
2813 }
2814 }
2815
2816 #[test]
2819 fn sinks_survive_clear_and_wire_import() {
2820 let mut c = LayerKvCache::new(1, 4);
2821 c.mode = KvMode::F32;
2822 c.sinks = Some(vec![0.5, -0.5]);
2823 c.append(&[1.0; 4], &[2.0; 4], &[]);
2824 let wire = c.export_wire(false).unwrap();
2825 c.clear();
2826 assert_eq!(c.sinks.as_deref(), Some(&[0.5f32, -0.5][..]));
2827 c.import_wire(&wire).unwrap();
2828 assert_eq!(c.sinks.as_deref(), Some(&[0.5f32, -0.5][..]));
2829 assert_eq!(c.seq_len, 1);
2830 }
2831
2832 #[test]
2836 fn born_eviction_mixed_modes_stay_consistent() {
2837 for (mk, mv) in [(false, true), (true, false)] {
2838 let mut c = LayerKvCache::new(1, 4);
2839 c.mode = KvMode::Q8 { k: mk, v: mv };
2840 for p in 0..80 {
2841 let k = vec![p as f32 * 0.01; 4];
2842 let v = vec![p as f32; 4];
2843 c.append(&k, &v, &[true]);
2844 }
2845 let imp: Vec<f32> = (0..80).map(|i| i as f32).collect();
2846 c.accumulate_imp(&imp);
2847 let before = c.memory_bytes();
2848 c.evict_born(20, 4, 8); assert_eq!(c.head_len(0), 20, "k={mk} v={mv}");
2850 assert!(
2851 c.memory_bytes() < before / 2,
2852 "memory must shrink (k={mk} v={mv})"
2853 );
2854 let (out, _) = c.attend(&[1.0, 1.0, 1.0, 1.0], 0);
2857 assert!(
2858 out[0] > 30.0,
2859 "V from the kept tail, not the stale head (k={mk} v={mv}, out {})",
2860 out[0]
2861 );
2862 }
2863 }
2864
2865 #[test]
2866 fn born_eviction_keeps_high_mass_position() {
2867 let mut cache = KvCache::new(1, 1, 2, 16);
2868 cache.policy = EvictionPolicy::Born { sink: 1 };
2869 for l in &mut cache.layers {
2870 l.mode = KvMode::F32;
2871 }
2872 let layer = &mut cache.layers[0];
2873 for pos in 0..8 {
2876 let k = vec![pos as f32; 2];
2877 let v = vec![pos as f32 + 100.0; 2];
2878 layer.append(&k, &v, &[true]);
2879 }
2880 let mut imp = vec![0.05f32; 8];
2882 imp[3] = 5.0;
2883 layer.accumulate_imp(&imp);
2884
2885 cache.evict(4); let layer = &cache.layers[0];
2887 assert_eq!(layer.seq_len, 4);
2888 let kept_keys: Vec<f32> = (0..4).map(|i| layer.head_keys(0)[i * 2]).collect();
2889 assert_eq!(
2890 kept_keys,
2891 vec![0.0, 3.0, 6.0, 7.0],
2892 "kept = sink(0) + mass-top(3) + recent(6,7)"
2893 );
2894 assert_eq!(layer.head_len(0), 4);
2896 }
2897
2898 fn swa_row(p: usize, n: usize, salt: usize) -> Vec<f32> {
2902 (0..n)
2903 .map(|i| ((p * 7919 + i * 104_729 + salt * 1_299_709) % 2003) as f32 / 1001.5 - 1.0)
2904 .collect()
2905 }
2906
2907 fn bits(v: &[f32]) -> Vec<u32> {
2908 v.iter().map(|x| x.to_bits()).collect()
2909 }
2910
2911 #[test]
2914 fn swa_trim_keeps_contiguous_tail() {
2915 let (nkv, hd, w, slack, align) = (2usize, 8usize, 6usize, 2usize, 4usize);
2916 let alive = [true, false];
2917 let mut full = LayerKvCache::new(nkv, hd);
2918 full.mode = KvMode::F32;
2919 let mut t = full.clone();
2920 let mut trims = 0;
2921 for p in 0..100 {
2922 let (k, v) = (swa_row(p, nkv * hd, 1), swa_row(p, nkv * hd, 2));
2923 full.append(&k, &v, &alive);
2924 t.append(&k, &v, &alive);
2925 if t.trim_window(w, slack, align) > 0 {
2926 trims += 1;
2927 assert!(
2928 (w + slack..w + slack + align).contains(&t.seq_len),
2929 "rows after a trim: {}",
2930 t.seq_len
2931 );
2932 }
2933 assert!(t.seq_len <= 2 * w, "rows {} past the trigger", t.seq_len);
2934 assert_eq!(t.pos_len(), p + 1);
2935 assert_eq!(t.base() % align, 0);
2936 assert_eq!(t.base() + t.seq_len, full.seq_len);
2937 let b = t.base();
2938 assert_eq!(t.head_keys(0), &full.head_keys(0)[b * hd..]);
2939 assert_eq!(t.head_values(0), &full.head_values(0)[b * hd..]);
2940 assert!(t.head_keys(1).is_empty() && t.head_values(1).is_empty());
2941 assert_eq!(t.imp.len(), t.seq_len);
2942 }
2943 assert!(trims > 5, "only {trims} trims in 100 positions");
2944 assert_eq!(t.tail_window(), Some(w));
2945 assert_eq!(full.tail_window(), None);
2946 }
2947
2948 #[test]
2953 fn swa_trimmed_attend_equals_untrimmed() {
2954 let (nkv, hd, hpk) = (2usize, 16usize, 2usize);
2955 for mode in [KvMode::F32, KvMode::Q8 { k: true, v: true }] {
2956 for (w, slack, align) in [(40usize, 2usize, 4usize), (37, 5, 16), (64, 64, 64)] {
2957 let mut full = LayerKvCache::new(nkv, hd);
2958 full.mode = mode;
2959 let mut t = full.clone();
2960 let scale = 0.3f32;
2961 let mut trims = 0;
2962 for p in 0..400 {
2963 let (k, v) = (swa_row(p, nkv * hd, 3), swa_row(p, nkv * hd, 4));
2964 full.append(&k, &v, &[]);
2965 t.append(&k, &v, &[]);
2966 let q = swa_row(p, nkv * hpk * hd, 5);
2967 for g in 0..nkv {
2968 let qg = &q[g * hpk * hd..(g + 1) * hpk * hd];
2969 let (nf, nt) = (full.head_len(g), t.head_len(g));
2970 let (mut of, mut ot) = (vec![0f32; hpk * hd], vec![0f32; hpk * hd]);
2971 let (mut imf, mut imt) = (vec![0f32; nf], vec![0f32; nt]);
2972 full.attend_group(
2973 qg,
2974 g,
2975 &mut of,
2976 &mut imf,
2977 scale,
2978 nf.saturating_sub(w),
2979 0.0,
2980 &[],
2981 );
2982 t.attend_group(
2983 qg,
2984 g,
2985 &mut ot,
2986 &mut imt,
2987 scale,
2988 nt.saturating_sub(w),
2989 0.0,
2990 &[],
2991 );
2992 assert_eq!(bits(&of), bits(&ot), "{mode:?} w {w} pos {p} group {g}");
2993 assert_eq!(bits(&imf[t.base()..]), bits(&imt), "{mode:?} w {w} pos {p}");
2994 full.accumulate_imp(&imf);
2995 t.accumulate_imp(&imt);
2996 }
2997 if t.trim_window(w, slack, align) > 0 {
2998 trims += 1;
2999 }
3000 assert_eq!(bits(&full.imp[t.base()..]), bits(&t.imp));
3001 }
3002 assert!(trims > 0, "{mode:?} w {w}: never trimmed");
3003 }
3004 }
3005 }
3006
3007 #[test]
3010 fn swa_truncate_after_trim() {
3011 let (nkv, hd, w) = (1usize, 8usize, 10usize);
3012 let mut full = LayerKvCache::new(nkv, hd);
3013 full.mode = KvMode::F32;
3014 let mut t = full.clone();
3015 let mut p = 0usize;
3016 for round in 0..60 {
3017 for _ in 0..3 {
3018 let (k, v) = (swa_row(p, hd, 6), swa_row(p, hd, 7));
3019 full.append(&k, &v, &[]);
3020 t.append(&k, &v, &[]);
3021 p += 1;
3022 }
3023 t.trim_window(w, 2, 4);
3024 full.truncate_last(2);
3026 t.truncate_last(2);
3027 p -= 2;
3028 assert_eq!(t.pos_len(), full.seq_len);
3029 for _ in 0..2 {
3030 let (k, v) = (swa_row(p, hd, 8 + round), swa_row(p, hd, 9 + round));
3031 full.append(&k, &v, &[]);
3032 t.append(&k, &v, &[]);
3033 p += 1;
3034 }
3035 let q = swa_row(p, hd, 10);
3036 let (nf, nt) = (full.head_len(0), t.head_len(0));
3037 let (mut of, mut ot) = (vec![0f32; hd], vec![0f32; hd]);
3038 let (mut imf, mut imt) = (vec![0f32; nf], vec![0f32; nt]);
3039 full.attend_group(&q, 0, &mut of, &mut imf, 0.5, nf - w.min(nf), 0.0, &[]);
3040 t.attend_group(&q, 0, &mut ot, &mut imt, 0.5, nt - w.min(nt), 0.0, &[]);
3041 assert_eq!(bits(&of), bits(&ot), "round {round}");
3042 }
3043 assert!(t.base() > 0);
3044 }
3045
3046 #[test]
3049 fn swa_evict_skips_tail() {
3050 for policy in [EvictionPolicy::Recent, EvictionPolicy::Born { sink: 1 }] {
3051 let mut c = KvCache::new(2, 1, 4, 30);
3052 c.policy = policy;
3053 for p in 0..25 {
3054 for l in &mut c.layers {
3055 l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3056 }
3057 c.layers[0].trim_window(8, 2, 4);
3058 }
3059 let t_rows = c.layers[0].seq_len;
3060 let t_base = c.layers[0].base();
3061 assert!(t_base > 0);
3062 assert!(!c.needs_eviction(), "25 rows < cap 30");
3063 for p in 25..30 {
3064 for l in &mut c.layers {
3065 l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3066 }
3067 }
3068 assert!(c.needs_eviction());
3070 let gen0 = c.layers[0].generation();
3071 c.evict(10);
3072 assert_eq!(c.layers[0].seq_len, t_rows + 5, "{policy:?}: tail evicted");
3073 assert_eq!(c.layers[0].base(), t_base);
3074 assert_eq!(c.layers[0].generation(), gen0);
3075 assert_eq!(
3076 c.layers[1].seq_len, 10,
3077 "{policy:?}: full layer not evicted"
3078 );
3079 assert_eq!(c.seq_len(), 30, "absolute depth from the tail");
3080 }
3081 }
3082
3083 #[test]
3086 fn swa_gen_bumps() {
3087 let mut l = LayerKvCache::new(1, 4);
3088 l.mode = KvMode::F32;
3089 let g0 = l.generation();
3090 assert_ne!(
3091 LayerKvCache::new(1, 4).generation(),
3092 g0,
3093 "fresh caches differ"
3094 );
3095 for p in 0..20 {
3096 l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3097 }
3098 assert_eq!(l.generation(), g0, "append");
3099 l.truncate_last(2);
3100 assert_eq!(l.generation(), g0, "truncate_last");
3101 assert_eq!(l.trim_window(10, 2, 4), 0, "18 rows ≤ 2w: no trim");
3102 assert_eq!(l.generation(), g0, "a no-op trim");
3103 for p in 18..21 {
3104 l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3105 }
3106 assert!(l.trim_window(10, 2, 4) > 0);
3107 let g1 = l.generation();
3108 assert_ne!(g1, g0, "trim");
3109 let snap = l.clone();
3110 l.clear();
3111 let g2 = l.generation();
3112 assert_ne!(g2, g1, "clear");
3113 assert_eq!((l.base(), l.pos_len(), l.tail_window()), (0, 0, Some(10)));
3114 let mut e = LayerKvCache::new(1, 4);
3115 e.mode = KvMode::F32;
3116 for p in 0..12 {
3117 e.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3118 }
3119 let ge = e.generation();
3120 e.evict(20);
3121 assert_eq!(e.generation(), ge, "an eviction that does nothing");
3122 e.evict(6);
3123 assert_ne!(e.generation(), ge, "evict");
3124 let ge = e.generation();
3125 e.evict_born(3, 1, 1);
3126 assert_ne!(e.generation(), ge, "evict_born");
3127 let mut i = LayerKvCache::new(1, 4);
3128 let gi = i.generation();
3129 i.import_wire(&snap.export_wire(false).unwrap()).unwrap();
3130 assert_ne!(i.generation(), gi, "import");
3131 assert_ne!(
3132 i.generation(),
3133 snap.generation(),
3134 "an import is new storage"
3135 );
3136 }
3137
3138 #[test]
3142 fn wire_full_tail_roundtrip() {
3143 let (nkv, hd, w) = (2usize, 4usize, 6usize);
3144 let mut a = LayerKvCache::new(nkv, hd);
3145 a.mode = KvMode::F32;
3146 for p in 0..5 {
3147 a.append(&swa_row(p, nkv * hd, 1), &swa_row(p, nkv * hd, 2), &[]);
3148 }
3149 let plain = a.export_wire(false).unwrap();
3150 assert_eq!(plain[20], WireKind::Full as u8, "untrimmed stays Full");
3151 assert_eq!(u64::from_le_bytes(plain[24..32].try_into().unwrap()), 5);
3152 for p in 5..40 {
3153 a.append(&swa_row(p, nkv * hd, 1), &swa_row(p, nkv * hd, 2), &[]);
3154 a.trim_window(w, 2, 4);
3155 }
3156 assert!(a.base() > 0);
3157 for f16 in [false, true] {
3158 let bytes = a.export_wire(f16).unwrap();
3159 assert_eq!(bytes[20], WireKind::FullTail as u8);
3160 assert_eq!(u64::from_le_bytes(bytes[24..32].try_into().unwrap()), 40);
3161 let mut b = LayerKvCache::new(nkv, hd);
3162 b.import_wire(&bytes).expect("tail import");
3163 assert_eq!(
3164 (b.base(), b.seq_len, b.pos_len()),
3165 (a.base(), a.seq_len, 40)
3166 );
3167 if !f16 {
3168 let q = swa_row(99, 2 * hd, 3);
3169 for g in 0..nkv {
3170 let n = a.head_len(g);
3171 let (mut oa, mut ob) = (vec![0f32; 2 * hd], vec![0f32; 2 * hd]);
3172 let (mut ia, mut ib) = (vec![0f32; n], vec![0f32; n]);
3173 a.attend_group(&q, g, &mut oa, &mut ia, 0.5, n - w, 0.0, &[]);
3174 b.attend_group(&q, g, &mut ob, &mut ib, 0.5, n - w, 0.0, &[]);
3175 assert_eq!(bits(&oa), bits(&ob));
3176 }
3177 b.import_wire(&plain).unwrap();
3179 assert_eq!((b.base(), b.pos_len()), (0, 5));
3180 }
3181 }
3182 let mut bad = a.export_wire(false).unwrap();
3184 bad[24..32].copy_from_slice(&41u64.to_le_bytes());
3185 assert!(LayerKvCache::new(nkv, hd).import_wire(&bad).is_err());
3186 let mut r = LayerKvCache::new(nkv, hd);
3188 r.install_bounded(8);
3189 let err = r.import_wire(&a.export_wire(false).unwrap()).unwrap_err();
3190 assert!(err.contains("does not match"), "{err}");
3191 }
3192}