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
51#[derive(Debug, Clone)]
62pub enum O1State {
63 Collecting {
64 m: usize,
65 w: usize,
66 sink: usize,
67 rect: crate::nystrom::O1Rect,
68 seal_at: Option<usize>,
72 q_buf: Vec<f32>,
74 },
75 Sealed {
81 groups: Vec<crate::nystrom::NystromState>,
82 },
83}
84
85#[derive(Debug, Clone)]
87pub struct LayerKvCache {
88 pub mode: KvMode,
89 k: Vec<Vec<f32>>,
91 v: Vec<Vec<f32>>,
93 kq: Vec<Vec<i8>>,
95 ks: Vec<Vec<f32>>,
96 vq: Vec<Vec<i8>>,
97 vs: Vec<Vec<f32>>,
98 kcol: Vec<Vec<f32>>,
100 vcol: Vec<Vec<f32>>,
101 imp: Vec<f32>,
104 pub seq_len: usize,
106 pub num_kv_heads: usize,
107 pub head_dim: usize,
108 pub linear_state: Vec<f32>,
110 pub linear_scratch: Vec<f32>,
112 linear_wire_allowed: bool,
116 pub o1: Option<O1State>,
118 o1_error: Option<String>,
122 o1_transitioned: bool,
125 pub sinks: Option<Vec<f32>>,
132 pub bounded: Option<crate::bounded::BoundedState>,
136 pub wire_kind: WireKind,
138 pub wire_layer: u32,
140 pub wire_identity: u64,
143}
144
145#[derive(Debug, Clone, Copy, PartialEq, Eq)]
147#[repr(u8)]
148pub enum WireKind {
149 Full = 0,
151 Linear = 1,
153 Bounded = 2,
155}
156
157impl WireKind {
158 fn from_u8(v: u8) -> Option<Self> {
159 match v {
160 0 => Some(WireKind::Full),
161 1 => Some(WireKind::Linear),
162 2 => Some(WireKind::Bounded),
163 _ => None,
164 }
165 }
166}
167
168pub const WIRE_MAGIC: &[u8; 4] = b"CMFS";
170pub const WIRE_VERSION: u32 = 2;
172
173impl LayerKvCache {
174 pub fn new(num_kv_heads: usize, head_dim: usize) -> Self {
175 Self {
176 sinks: None,
177 mode: KvMode::from_env(),
178 k: vec![Vec::new(); num_kv_heads],
179 v: vec![Vec::new(); num_kv_heads],
180 kq: vec![Vec::new(); num_kv_heads],
181 ks: vec![Vec::new(); num_kv_heads],
182 vq: vec![Vec::new(); num_kv_heads],
183 vs: vec![Vec::new(); num_kv_heads],
184 kcol: vec![Vec::new(); num_kv_heads],
185 vcol: vec![Vec::new(); num_kv_heads],
186 imp: Vec::new(),
187 seq_len: 0,
188 num_kv_heads,
189 head_dim,
190 linear_state: Vec::new(),
191 linear_scratch: Vec::new(),
192 linear_wire_allowed: true,
193 o1: None,
194 o1_error: None,
195 o1_transitioned: false,
196 bounded: None,
197 wire_kind: WireKind::Full,
198 wire_layer: 0,
199 wire_identity: 0,
200 }
201 }
202
203 pub fn install_bounded(&mut self, window: usize) {
209 self.bounded = Some(crate::bounded::BoundedState::new(
210 self.num_kv_heads,
211 self.head_dim,
212 window,
213 ));
214 self.wire_kind = WireKind::Bounded;
215 }
216
217 #[allow(clippy::too_many_arguments)]
222 pub fn bounded_step(
223 &mut self,
224 q: &[f32],
225 k: &[f32],
226 v: &[f32],
227 w: &crate::bounded::BoundedWeights,
228 rope: &crate::bounded::BoundedRope,
229 scale: f32,
230 num_heads: usize,
231 out: &mut [f32],
232 ) {
233 let st = self
234 .bounded
235 .as_mut()
236 .expect("bounded_step on a layer without an installed ring");
237 st.insert(k, v);
238 st.attend(q, num_heads, &w.sink_k, &w.sink_v, w.sink, rope, scale, out);
239 self.seq_len += 1;
242 }
243
244 pub fn bounded_state_bytes(&self) -> usize {
246 self.bounded.as_ref().map(|b| b.state_bytes()).unwrap_or(0)
247 }
248
249 pub fn bounded_snapshot(&self) -> Option<crate::bounded::BoundedSnapshot> {
251 self.bounded.as_ref().map(|b| b.snapshot())
252 }
253
254 pub fn bounded_restore(&mut self, s: &crate::bounded::BoundedSnapshot) {
257 if let Some(b) = self.bounded.as_mut() {
258 b.restore(s);
259 self.seq_len = b.seen;
260 }
261 }
262
263 pub fn set_linear_wire_allowed(&mut self, allowed: bool) {
267 self.linear_wire_allowed = allowed;
268 }
269
270 pub fn discard_linear_scratch(&mut self) {
273 self.linear_scratch.clear();
274 }
275
276 pub fn k_heads(&self) -> &[Vec<f32>] {
278 &self.k
279 }
280 pub fn v_heads(&self) -> &[Vec<f32>] {
282 &self.v
283 }
284
285 pub fn o1_begin(&mut self, m: usize, w: usize, sink: usize, rect: crate::nystrom::O1Rect) {
289 self.o1_begin_with_boundary(m, w, sink, rect, None);
290 }
291
292 pub(crate) fn o1_begin_with_boundary(
296 &mut self,
297 m: usize,
298 w: usize,
299 sink: usize,
300 rect: crate::nystrom::O1Rect,
301 seal_at: Option<usize>,
302 ) {
303 self.o1 = Some(O1State::Collecting {
304 m,
305 w,
306 sink,
307 rect,
308 seal_at,
309 q_buf: Vec::new(),
310 });
311 self.o1_error = None;
312 self.o1_transitioned = false;
313 }
314
315 pub fn o1_push_q(&mut self, q_all: &[f32]) {
320 if let Some(O1State::Collecting { q_buf, .. }) = &mut self.o1 {
321 q_buf.extend_from_slice(q_all);
322 }
323 }
324
325 pub fn o1_sealed(&self) -> bool {
326 matches!(self.o1, Some(O1State::Sealed { .. }))
327 }
328
329 pub(crate) fn o1_pending_boundary(&self) -> Option<usize> {
332 match &self.o1 {
333 Some(O1State::Collecting { seal_at, .. }) => *seal_at,
334 _ => None,
335 }
336 }
337
338 pub(crate) fn o1_boundary_crossed_by(&self, count: usize) -> bool {
342 let Some(target) = self.o1_pending_boundary() else {
343 return false;
344 };
345 target <= self.seq_len
346 || self
347 .seq_len
348 .checked_add(count)
349 .map_or(true, |next| next >= target)
350 }
351
352 pub(crate) fn take_o1_transition(&mut self) -> bool {
353 std::mem::take(&mut self.o1_transitioned)
354 }
355
356 pub(crate) fn take_o1_error(&self) -> Option<String> {
357 self.o1_error.clone()
363 }
364
365 pub(crate) fn o1_abort(&mut self, err: String) {
370 self.k.iter_mut().for_each(Vec::clear);
371 self.v.iter_mut().for_each(Vec::clear);
372 self.kq.iter_mut().for_each(Vec::clear);
373 self.ks.iter_mut().for_each(Vec::clear);
374 self.vq.iter_mut().for_each(Vec::clear);
375 self.vs.iter_mut().for_each(Vec::clear);
376 self.kcol.iter_mut().for_each(Vec::clear);
377 self.vcol.iter_mut().for_each(Vec::clear);
378 self.imp.clear();
379 self.o1 = None;
380 self.seq_len = 0;
381 self.o1_transitioned = false;
382 self.o1_error = Some(err);
383 }
384
385 pub fn o1_seal(&mut self, num_heads: usize) -> bool {
392 match self.o1_seal_checked(num_heads) {
393 Ok(sealed) => sealed,
394 Err(err) => {
395 tracing::error!("o1: seal aborted: {err}");
396 self.o1_abort(err);
397 false
398 }
399 }
400 }
401
402 pub(crate) fn o1_seal_checked(&mut self, num_heads: usize) -> Result<bool, String> {
406 if let Some(err) = self.o1_error.clone() {
407 return Err(err);
408 }
409 if !matches!(self.o1, Some(O1State::Collecting { .. })) {
412 return Ok(self.o1_sealed());
413 }
414 let (m, w, sink, requested_boundary, q_len) = match &self.o1 {
415 Some(O1State::Collecting {
416 m,
417 w,
418 sink,
419 rect: _,
420 seal_at,
421 q_buf,
422 }) => (*m, *w, *sink, *seal_at, q_buf.len()),
423 _ => unreachable!("checked above"),
424 };
425 let floor = crate::nystrom::o1_deferred_boundary(w, sink)
426 .ok_or_else(|| "o1 seal: w + sink + slack + 1 overflow".to_string())?;
427 let target = requested_boundary.unwrap_or(floor).max(floor);
428 let t = self.seq_len;
429 if t < target {
430 if let Some(O1State::Collecting { seal_at, .. }) = &mut self.o1 {
431 if *seal_at != Some(target) {
432 *seal_at = Some(target);
433 tracing::info!(
434 "o1 deferred seal: current rows={t}, boundary={target} (floor={floor})"
435 );
436 }
437 }
438 return Ok(false);
439 }
440
441 let hd = self.head_dim;
442 if t == 0 {
443 return Err("o1 seal: cannot seal an empty layer".into());
444 }
445 if self.mode != KvMode::F32 {
446 return Err("o1 seal: requires dense F32 KV storage".into());
447 }
448 if self.num_kv_heads == 0 || num_heads == 0 || num_heads % self.num_kv_heads != 0 {
449 return Err(format!(
450 "o1 seal: invalid GQA geometry num_heads={num_heads} num_kv_heads={}",
451 self.num_kv_heads
452 ));
453 }
454 let hpk = num_heads / self.num_kv_heads;
455 let expected_k = t
456 .checked_mul(hd)
457 .ok_or_else(|| "o1 seal: KV row length overflow".to_string())?;
458 let expected_q = expected_k
459 .checked_mul(num_heads)
460 .ok_or_else(|| "o1 seal: query trace length overflow".to_string())?;
461 if q_len != expected_q {
462 return Err(format!(
463 "o1 seal: query trace has {q_len} values, expected {expected_q}"
464 ));
465 }
466 if (0..self.num_kv_heads)
467 .any(|g| self.k[g].len() != expected_k || self.v[g].len() != expected_k)
468 {
469 return Err("o1 seal: KV heads are not densely populated".into());
470 }
471 if m < 4 || w == 0 {
472 return Err(format!("o1 seal: invalid geometry m={m} w={w}"));
473 }
474
475 let Some(O1State::Collecting {
476 m,
477 w,
478 sink,
479 rect,
480 q_buf,
481 ..
482 }) = self.o1.take()
483 else {
484 unreachable!("collecting state disappeared after validation");
485 };
486 let mut groups = Vec::with_capacity(self.num_kv_heads);
487 let mut qh = vec![0.0f32; hpk * t * hd];
490 for g in 0..self.num_kv_heads {
491 for hh in 0..hpk {
492 let h = g * hpk + hh;
493 for p in 0..t {
494 let src = (p * num_heads + h) * hd;
495 let dst = (hh * t + p) * hd;
496 qh[dst..dst + hd].copy_from_slice(&q_buf[src..src + hd]);
497 }
498 }
499 let qs: Vec<&[f32]> = (0..hpk)
500 .map(|hh| &qh[hh * t * hd..(hh + 1) * t * hd])
501 .collect();
502 let mut st = crate::nystrom::NystromState::new_group(m, w, sink, hpk).with_rect(rect);
503 st.prefill_group(&qs, &self.k[g], &self.v[g], t, hd, hd);
504 groups.push(st);
505 }
506 for h in 0..self.num_kv_heads {
509 self.k[h] = Vec::new();
510 self.v[h] = Vec::new();
511 }
512 self.imp = Vec::new();
513 self.o1 = Some(O1State::Sealed { groups });
514 self.o1_transitioned = true;
515 Ok(true)
516 }
517
518 pub fn o1_views(&self) -> Option<Vec<crate::nystrom::O1DeviceView<'_>>> {
527 let Some(O1State::Sealed { groups }) = &self.o1 else {
528 return None;
529 };
530 let views: Vec<_> = groups.iter().map(|g| g.device_view()).collect();
531 if views.iter().any(|v| v.exact_only) {
532 return None;
533 }
534 Some(views)
535 }
536
537 pub fn o1_step(
538 &mut self,
539 q_all: &[f32],
540 k_new: &[f32],
541 v_new: &[f32],
542 num_heads: usize,
543 ) -> Vec<f32> {
544 let hd = self.head_dim;
545 let hpk = num_heads / self.num_kv_heads.max(1);
546 let mut out = vec![0.0f32; num_heads * hd];
547 let Some(O1State::Sealed { groups }) = &mut self.o1 else {
548 debug_assert!(false, "o1_step on an unsealed layer");
549 return out;
550 };
551 for (g, st) in groups.iter_mut().enumerate() {
552 let (lo, hi) = (g * hpk * hd, (g + 1) * hpk * hd);
553 st.step_group(
554 &q_all[lo..hi],
555 &k_new[g * hd..(g + 1) * hd],
556 &v_new[g * hd..(g + 1) * hd],
557 &mut out[lo..hi],
558 );
559 }
560 self.seq_len += 1;
563 out
564 }
565
566 pub fn o1_memory_bytes(&self) -> usize {
569 match &self.o1 {
570 Some(O1State::Collecting { q_buf, .. }) => q_buf.len() * std::mem::size_of::<f32>(),
571 Some(O1State::Sealed { groups }) => groups.iter().map(|s| s.memory_bytes()).sum(),
572 None => 0,
573 }
574 }
575
576 fn quant_row(row: &[f32], col: &[f32], q: &mut Vec<i8>, sc: &mut Vec<f32>, group: usize) {
579 let mut resid = vec![0.0f32; row.len()];
580 for (d, &x) in row.iter().enumerate() {
581 resid[d] = if col.is_empty() { x } else { x / col[d] };
582 }
583 for g0 in (0..row.len()).step_by(group) {
584 let g1 = (g0 + group).min(row.len());
585 let mut absmax = 0.0f32;
586 for &r in &resid[g0..g1] {
587 absmax = absmax.max(r.abs());
588 }
589 let s = (absmax / 127.0).max(1e-12);
590 sc.push(s);
591 for &r in &resid[g0..g1] {
592 q.push((r / s).round().clamp(-127.0, 127.0) as i8);
593 }
594 }
595 }
596
597 fn freeze_cols(&mut self) {
600 let hd = self.head_dim;
601 let ngk = hd.div_ceil(KV_K_GROUP);
602 for h in 0..self.num_kv_heads {
603 for (qv, sv, colv, group) in [
604 (
605 &mut self.kq[h],
606 &mut self.ks[h],
607 &mut self.kcol[h],
608 KV_K_GROUP,
609 ),
610 (&mut self.vq[h], &mut self.vs[h], &mut self.vcol[h], hd),
611 ] {
612 let spp = if group == hd { 1 } else { ngk }; let n = sv.len() / spp;
614 if n == 0 {
615 continue;
616 }
617 let mut rows = vec![0.0f32; n * hd];
619 for p in 0..n {
620 for d in 0..hd {
621 rows[p * hd + d] = qv[p * hd + d] as f32 * sv[p * spp + d / group];
622 }
623 }
624 let mut col = vec![0.0f32; hd];
625 for p in 0..n {
626 for d in 0..hd {
627 col[d] += rows[p * hd + d] * rows[p * hd + d];
628 }
629 }
630 for c in col.iter_mut() {
631 *c = (*c / n as f32).sqrt().max(1e-6);
632 }
633 qv.clear();
634 sv.clear();
635 for p in 0..n {
636 Self::quant_row(&rows[p * hd..(p + 1) * hd], &col, qv, sv, group);
637 }
638 *colv = col;
639 }
640 }
641 }
642
643 pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
647 if self.o1_error.is_some() {
651 return;
652 }
653 debug_assert_eq!(k_new.len(), self.num_kv_heads * self.head_dim);
654 debug_assert_eq!(v_new.len(), self.num_kv_heads * self.head_dim);
655 if matches!(self.mode, KvMode::Q8 { .. })
661 && self.seq_len >= KV_COL_WARMUP
662 && self.kcol.iter().all(Vec::is_empty)
663 && self.vcol.iter().all(Vec::is_empty)
664 {
665 self.freeze_cols();
666 }
667 for h in 0..self.num_kv_heads {
668 if !alive.get(h).copied().unwrap_or(true) {
669 continue;
670 }
671 let s = h * self.head_dim;
672 if self.mode.quant_k() {
673 Self::quant_row(
674 &k_new[s..s + self.head_dim],
675 &self.kcol[h],
676 &mut self.kq[h],
677 &mut self.ks[h],
678 KV_K_GROUP,
679 );
680 } else {
681 self.k[h].extend_from_slice(&k_new[s..s + self.head_dim]);
682 }
683 if self.mode.quant_v() {
684 Self::quant_row(
685 &v_new[s..s + self.head_dim],
686 &self.vcol[h],
687 &mut self.vq[h],
688 &mut self.vs[h],
689 self.head_dim,
690 );
691 } else {
692 self.v[h].extend_from_slice(&v_new[s..s + self.head_dim]);
693 }
694 }
695 self.imp.push(0.0);
696 self.seq_len += 1;
697 }
698
699 pub fn attend(&self, q: &[f32], kv_head: usize) -> (Vec<f32>, Vec<f32>) {
704 let hd = self.head_dim;
705 if self.mode == KvMode::F32 {
706 let stored = self.k[kv_head].len() / hd;
707 return crate::attention::attention_head(
708 q,
709 &self.k[kv_head],
710 &self.v[kv_head],
711 hd,
712 stored,
713 );
714 }
715 let stored = self.head_len(kv_head);
716 let scale = 1.0 / (hd as f32).sqrt();
717 let mut scores = vec![0.0f32; stored];
718 if self.mode.quant_k() {
719 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
720 let kcol = &self.kcol[kv_head];
722 let mut qc = vec![0.0f32; hd];
723 for d in 0..hd {
724 qc[d] = if kcol.is_empty() {
725 q[d]
726 } else {
727 q[d] * kcol[d]
728 };
729 }
730 let ng = hd.div_ceil(KV_K_GROUP);
731 for p in 0..stored {
732 let row = &kq[p * hd..(p + 1) * hd];
733 let row_u8 =
736 unsafe { std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len()) };
737 let mut dot = 0.0f32;
738 for g in 0..ng {
739 let g0 = g * KV_K_GROUP;
740 let g1 = (g0 + KV_K_GROUP).min(hd);
741 dot +=
742 crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qc[g0..g1]) * ks[p * ng + g];
743 }
744 scores[p] = dot * scale;
745 }
746 } else {
747 let k = &self.k[kv_head];
748 for p in 0..stored {
749 let row = &k[p * hd..(p + 1) * hd];
750 scores[p] = crate::attention::dot_f32(q, row) * scale;
751 }
752 }
753 let max_score = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
754 let mut sum = 0.0f32;
755 for s in scores.iter_mut() {
756 *s = (*s - max_score).exp();
757 sum += *s;
758 }
759 if sum > 0.0 {
760 for s in scores.iter_mut() {
761 *s /= sum;
762 }
763 }
764 let mut acc = vec![0.0f32; hd];
765 if self.mode.quant_v() {
766 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
767 for p in 0..stored {
768 let w = scores[p] * vs[p];
769 if w.abs() < 1e-12 {
770 continue;
771 }
772 crate::qtensor::axpy_i8_f32(&mut acc, &vq[p * hd..(p + 1) * hd], w);
773 }
774 let vcol = &self.vcol[kv_head];
775 if !vcol.is_empty() {
776 for d in 0..hd {
777 acc[d] *= vcol[d];
778 }
779 }
780 } else {
781 let v = &self.v[kv_head];
782 for p in 0..stored {
783 let w = scores[p];
784 if w.abs() < 1e-12 {
785 continue;
786 }
787 crate::attention::axpy_f32(&mut acc, &v[p * hd..(p + 1) * hd], w);
788 }
789 }
790 (acc, scores)
791 }
792
793 #[allow(clippy::too_many_arguments)]
810 pub fn attend_group(
811 &self,
812 q_group: &[f32],
813 kv_head: usize,
814 out: &mut [f32],
815 imp_acc: &mut [f32],
816 scale: f32,
817 first: usize,
818 softcap: f32,
819 sinks: &[f32],
820 ) {
821 self.attend_group_upto(
822 q_group,
823 kv_head,
824 out,
825 imp_acc,
826 scale,
827 first,
828 softcap,
829 usize::MAX,
830 sinks,
831 )
832 }
833
834 #[allow(clippy::too_many_arguments)]
852 pub fn attend_group_upto(
853 &self,
854 q_group: &[f32],
855 kv_head: usize,
856 out: &mut [f32],
857 imp_acc: &mut [f32],
858 scale: f32,
859 first: usize,
860 softcap: f32,
861 upto: usize,
862 sinks: &[f32],
863 ) {
864 let hd = self.head_dim;
865 let nheads = q_group.len() / hd;
866 debug_assert_eq!(out.len(), nheads * hd);
867 assert!(
868 sinks.is_empty() || sinks.len() == nheads,
869 "attend_group: {} sink logits for {nheads} heads",
870 sinks.len()
871 );
872 let stored = if self.mode == KvMode::F32 {
873 self.k[kv_head].len() / hd
874 } else {
875 self.head_len(kv_head)
876 }
877 .min(upto);
878 if stored == 0 {
879 out.fill(0.0);
880 return;
881 }
882 let first = first.min(stored.saturating_sub(1));
883 let span = stored - first;
885
886 thread_local! {
887 static GQA_SCORES: std::cell::RefCell<Vec<f32>> =
889 const { std::cell::RefCell::new(Vec::new()) };
890 static GQA_QC: std::cell::RefCell<Vec<f32>> =
892 const { std::cell::RefCell::new(Vec::new()) };
893 }
894
895 GQA_SCORES.with(|sc| {
896 let mut scores = sc.borrow_mut();
897 scores.resize(nheads * span, 0.0);
898
899 if self.mode.quant_k() {
901 let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
902 let kcol = &self.kcol[kv_head];
903 let ng = hd.div_ceil(KV_K_GROUP);
904 GQA_QC.with(|qc| {
905 let mut qcb = qc.borrow_mut();
906 qcb.resize(nheads * hd, 0.0);
907 for h in 0..nheads {
908 for d in 0..hd {
909 let qv = q_group[h * hd + d];
910 qcb[h * hd + d] = if kcol.is_empty() { qv } else { qv * kcol[d] };
911 }
912 }
913 for p in first..stored {
914 let row = &kq[p * hd..(p + 1) * hd];
915 let row_u8 = unsafe {
918 std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len())
919 };
920 for h in 0..nheads {
921 let qch = &qcb[h * hd..(h + 1) * hd];
922 let mut dot = 0.0f32;
923 for g in 0..ng {
924 let g0 = g * KV_K_GROUP;
925 let g1 = (g0 + KV_K_GROUP).min(hd);
926 dot += crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qch[g0..g1])
927 * ks[p * ng + g];
928 }
929 scores[h * span + (p - first)] = dot * scale;
930 }
931 }
932 });
933 } else {
934 let k = &self.k[kv_head];
935 for p in first..stored {
936 let row = &k[p * hd..(p + 1) * hd];
937 for h in 0..nheads {
938 scores[h * span + (p - first)] =
939 crate::attention::dot_f32(&q_group[h * hd..(h + 1) * hd], row) * scale;
940 }
941 }
942 }
943
944 if softcap > 0.0 {
948 for v in scores.iter_mut() {
949 *v = softcap * (*v / softcap).tanh();
950 }
951 }
952
953 for h in 0..nheads {
956 let s = &mut scores[h * span..(h + 1) * span];
957 let row_max = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
958 let sink = sinks.get(h).copied();
959 let max_score = match sink {
960 Some(z) => row_max.max(z),
961 None => row_max,
962 };
963 let mut sum = 0.0f32;
964 for v in s.iter_mut() {
965 *v = (*v - max_score).exp();
966 sum += *v;
967 }
968 if let Some(z) = sink {
969 sum += (z - max_score).exp();
970 }
971 if sum > 0.0 {
972 for v in s.iter_mut() {
973 *v /= sum;
974 }
975 }
976 }
977
978 out.fill(0.0);
980 if self.mode.quant_v() {
981 let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
982 for p in first..stored {
983 let row = &vq[p * hd..(p + 1) * hd];
984 for h in 0..nheads {
985 let w = scores[h * span + (p - first)] * vs[p];
986 if w.abs() < 1e-12 {
987 continue;
988 }
989 crate::qtensor::axpy_i8_f32(&mut out[h * hd..(h + 1) * hd], row, w);
990 }
991 }
992 let vcol = &self.vcol[kv_head];
993 if !vcol.is_empty() {
994 for h in 0..nheads {
995 for d in 0..hd {
996 out[h * hd + d] *= vcol[d];
997 }
998 }
999 }
1000 } else {
1001 let v = &self.v[kv_head];
1002 for p in first..stored {
1003 let row = &v[p * hd..(p + 1) * hd];
1004 for h in 0..nheads {
1005 let w = scores[h * span + (p - first)];
1006 if w.abs() < 1e-12 {
1007 continue;
1008 }
1009 crate::attention::axpy_f32(&mut out[h * hd..(h + 1) * hd], row, w);
1010 }
1011 }
1012 }
1013
1014 let n = imp_acc.len().min(stored);
1018 if n > first {
1019 for h in 0..nheads {
1020 let s = &scores[h * span..(h + 1) * span];
1021 for (dst, &p) in imp_acc[first..n].iter_mut().zip(s) {
1022 *dst += p;
1023 }
1024 }
1025 }
1026 });
1027 }
1028
1029 #[cfg(target_arch = "aarch64")]
1037 #[allow(clippy::too_many_arguments)]
1038 pub fn attend_chunk(
1039 &mut self,
1040 q_all: &[f32],
1041 b: usize,
1042 s0: usize,
1043 nh: usize,
1044 heads_per_kv: usize,
1045 hd: usize,
1046 out: &mut [f32],
1047 pool: Option<&crate::pool::Pool>,
1048 scale: f32,
1049 window: Option<usize>,
1050 ) {
1051 let n = s0 + b;
1052 struct SendPtr(*mut f32);
1053 unsafe impl Send for SendPtr {}
1054 unsafe impl Sync for SendPtr {}
1055 impl SendPtr {
1056 fn at(&self, i: usize) -> *mut f32 {
1057 unsafe { self.0.add(i) }
1060 }
1061 }
1062 thread_local! {
1063 static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>)> =
1064 const { std::cell::RefCell::new((Vec::new(), Vec::new(), Vec::new(), Vec::new())) };
1065 }
1066 let neon_gemm = cfg!(not(target_os = "macos"))
1071 || std::env::var("CMF_FORCE_NEON_GEMM")
1072 .map(|v| v == "1")
1073 .unwrap_or(false);
1074 SCRATCH.with(|s| {
1075 let mut s = s.borrow_mut();
1076 let (qpanel, scores, aopanel, ktpack) = &mut *s;
1077 let m = heads_per_kv * b;
1082 qpanel.resize(m * hd, 0.0);
1083 scores.resize(m * n, 0.0);
1084 aopanel.resize(m * hd, 0.0);
1085 for g in 0..self.num_kv_heads {
1086 let kmat = &self.k[g];
1087 let vmat = &self.v[g];
1088 debug_assert_eq!(kmat.len(), n * hd);
1089 for hl in 0..heads_per_kv {
1090 let hh = g * heads_per_kv + hl;
1091 for bi in 0..b {
1092 qpanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]
1093 .copy_from_slice(&q_all[bi * nh * hd + hh * hd..][..hd]);
1094 }
1095 }
1096 if neon_gemm {
1097 ktpack.resize(hd * n, 0.0);
1098 for p in 0..n {
1099 let row = &kmat[p * hd..(p + 1) * hd];
1100 for (d, &v) in row.iter().enumerate() {
1101 ktpack[d * n + p] = v;
1102 }
1103 }
1104 let sp_q = SendPtr(qpanel.as_ptr() as *mut f32);
1107 let sp_s = SendPtr(scores.as_mut_ptr());
1108 let kt = &*ktpack;
1109 let run = |start: usize, end: usize| {
1110 if end > start {
1111 let a = unsafe {
1113 std::slice::from_raw_parts(sp_q.at(start * hd), (end - start) * hd)
1114 };
1115 let c = unsafe {
1116 std::slice::from_raw_parts_mut(
1117 sp_s.at(start * n),
1118 (end - start) * n,
1119 )
1120 };
1121 crate::qtensor::neon_gemm_rm(
1122 end - start,
1123 n,
1124 hd,
1125 scale,
1126 a,
1127 hd,
1128 kt,
1129 n,
1130 false,
1131 c,
1132 n,
1133 );
1134 }
1135 };
1136 match pool {
1137 Some(p) if m >= 64 => p.run_rows(m, &run),
1138 _ => run(0, m),
1139 }
1140 } else {
1141 crate::qtensor::sgemm_rm(
1142 m, n, hd, scale, qpanel, hd, kmat, hd, true, scores, n,
1143 );
1144 }
1145 let sp = SendPtr(scores.as_mut_ptr());
1147 let run = |start: usize, end: usize| {
1148 for r in start..end {
1149 let allowed = s0 + (r % b) + 1;
1150 let lo = window.map(|w| allowed.saturating_sub(w)).unwrap_or(0);
1154 let row = unsafe { std::slice::from_raw_parts_mut(sp.at(r * n), n) };
1156 crate::attention::softmax_row(&mut row[lo..allowed]);
1157 row[..lo].fill(0.0);
1158 row[allowed..].fill(0.0);
1159 }
1160 };
1161 match pool {
1162 Some(p) if m >= 64 => p.run_rows(m, &run),
1163 _ => run(0, m),
1164 }
1165 let ni = self.imp.len().min(n);
1169 for r in 0..m {
1170 let al = (s0 + (r % b) + 1).min(ni);
1171 for (dst, &p) in self.imp[..al].iter_mut().zip(&scores[r * n..r * n + al]) {
1172 *dst += p;
1173 }
1174 }
1175 if neon_gemm {
1176 let sp_s = SendPtr(scores.as_mut_ptr());
1177 let sp_o = SendPtr(aopanel.as_mut_ptr());
1178 let run = |start: usize, end: usize| {
1179 if end > start {
1180 let a = unsafe {
1182 std::slice::from_raw_parts(sp_s.at(start * n), (end - start) * n)
1183 };
1184 let c = unsafe {
1185 std::slice::from_raw_parts_mut(
1186 sp_o.at(start * hd),
1187 (end - start) * hd,
1188 )
1189 };
1190 crate::qtensor::neon_gemm_rm(
1191 end - start,
1192 hd,
1193 n,
1194 1.0,
1195 a,
1196 n,
1197 vmat,
1198 hd,
1199 false,
1200 c,
1201 hd,
1202 );
1203 }
1204 };
1205 match pool {
1206 Some(p) if m >= 64 => p.run_rows(m, &run),
1207 _ => run(0, m),
1208 }
1209 } else {
1210 crate::qtensor::sgemm_rm(
1211 m, hd, n, 1.0, scores, n, vmat, hd, false, aopanel, hd,
1212 );
1213 }
1214 for hl in 0..heads_per_kv {
1215 let hh = g * heads_per_kv + hl;
1216 for bi in 0..b {
1217 out[bi * nh * hd + hh * hd..][..hd]
1218 .copy_from_slice(&aopanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]);
1219 }
1220 }
1221 }
1222 });
1223 }
1224
1225 pub fn truncate_last(&mut self, n_drop: usize) {
1227 self.discard_linear_scratch();
1228 let d = n_drop.min(self.seq_len);
1229 if let Some(b) = self.bounded.as_mut() {
1230 let rolled = b.rollback(d);
1233 if rolled < d {
1234 tracing::warn!(
1235 "bounded anchor: rollback of {d} exceeds the undo depth ({rolled} restored)"
1236 );
1237 }
1238 self.seq_len = b.seen;
1239 return;
1240 }
1241 for h in 0..self.num_kv_heads {
1242 let keep = self.k[h].len().saturating_sub(d * self.head_dim);
1243 self.k[h].truncate(keep);
1244 self.v[h].truncate(keep);
1245 let ngk = self.head_dim.div_ceil(KV_K_GROUP);
1246 let keep_q = self.kq[h].len().saturating_sub(d * self.head_dim);
1247 self.kq[h].truncate(keep_q);
1248 let keep_vq = self.vq[h].len().saturating_sub(d * self.head_dim);
1249 self.vq[h].truncate(keep_vq);
1250 let keep_ks = self.ks[h].len().saturating_sub(d * ngk);
1251 self.ks[h].truncate(keep_ks);
1252 let keep_vs = self.vs[h].len().saturating_sub(d);
1253 self.vs[h].truncate(keep_vs);
1254 }
1255 self.imp.truncate(self.imp.len().saturating_sub(d));
1256 self.seq_len -= d;
1257 }
1258
1259 pub fn accumulate_imp(&mut self, probs: &[f32]) {
1261 for (dst, &p) in self.imp.iter_mut().zip(probs) {
1262 *dst += p;
1263 }
1264 }
1265
1266 pub fn head_keys(&self, kv_head: usize) -> &[f32] {
1268 &self.k[kv_head]
1269 }
1270
1271 pub fn head_values(&self, kv_head: usize) -> &[f32] {
1272 &self.v[kv_head]
1273 }
1274
1275 pub fn head_len(&self, kv_head: usize) -> usize {
1277 let ng = self.head_dim.div_ceil(KV_K_GROUP);
1278 (self.k[kv_head].len() / self.head_dim)
1279 .max(self.ks[kv_head].len() / ng)
1280 .max(self.vs[kv_head].len())
1281 }
1282
1283 pub fn clear(&mut self) {
1285 for h in 0..self.num_kv_heads {
1286 self.k[h].clear();
1287 self.v[h].clear();
1288 self.kq[h].clear();
1289 self.ks[h].clear();
1290 self.vq[h].clear();
1291 self.vs[h].clear();
1292 self.kcol[h].clear();
1293 self.vcol[h].clear();
1294 }
1295 self.imp.clear();
1296 self.linear_state.clear();
1297 self.discard_linear_scratch();
1298 self.o1 = None;
1301 self.o1_error = None;
1302 self.o1_transitioned = false;
1303 if let Some(b) = self.bounded.as_mut() {
1306 b.clear();
1307 }
1308 self.seq_len = 0;
1309 }
1310
1311 pub fn export_wire(&self, f16: bool) -> Result<Vec<u8>, String> {
1328 if !matches!(self.mode, KvMode::F32) {
1329 return Err("kv export: only the F32 cache is described by this format (CMF_KV=q8 stores int8 rows and per-row scales)"
1330 .into());
1331 }
1332 if self.o1.is_some() {
1333 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"
1334 .into());
1335 }
1336 if self.kcol.iter().any(|c| !c.is_empty()) || self.vcol.iter().any(|c| !c.is_empty()) {
1340 return Err(
1341 "kv export: frozen columns under an F32 cache — refusing to ship \
1342 a state this format does not describe"
1343 .into(),
1344 );
1345 }
1346 let mut out = Vec::with_capacity(self.memory_bytes() / if f16 { 2 } else { 1 } + 64);
1347 let u = |v: u32, o: &mut Vec<u8>| o.extend_from_slice(&v.to_le_bytes());
1348 out.extend_from_slice(WIRE_MAGIC);
1349 u(WIRE_VERSION, &mut out);
1350 out.extend_from_slice(&self.wire_identity.to_le_bytes());
1351 u(self.wire_layer, &mut out);
1352 out.push(self.wire_kind as u8);
1353 out.push(u8::from(f16));
1354 out.extend_from_slice(&0u16.to_le_bytes());
1355 out.extend_from_slice(&(self.seq_len as u64).to_le_bytes());
1356 let push = |xs: &[f32], o: &mut Vec<u8>| {
1357 if f16 {
1358 for &x in xs {
1359 o.extend_from_slice(&cortiq_core::quant::f32_to_f16(x).to_le_bytes());
1360 }
1361 } else {
1362 for &x in xs {
1363 o.extend_from_slice(&x.to_le_bytes());
1364 }
1365 }
1366 };
1367 match self.wire_kind {
1368 WireKind::Full => self.export_full_body(f16, &mut out),
1369 WireKind::Linear => {
1370 u(self.num_kv_heads as u32, &mut out);
1371 u(self.head_dim as u32, &mut out);
1372 u(self.linear_state.len() as u32, &mut out);
1373 for &x in &self.linear_state {
1374 out.extend_from_slice(&x.to_le_bytes());
1375 }
1376 }
1377 WireKind::Bounded => {
1378 let b = self
1379 .bounded
1380 .as_ref()
1381 .ok_or("kv export: bounded wire kind without an installed ring")?;
1382 u(self.num_kv_heads as u32, &mut out);
1383 u(self.head_dim as u32, &mut out);
1384 u(b.window as u32, &mut out);
1385 u(b.len() as u32, &mut out);
1386 u(b.head() as u32, &mut out);
1387 push(&b.ring_k, &mut out);
1388 push(&b.ring_v, &mut out);
1389 }
1390 }
1391 Ok(out)
1392 }
1393
1394 fn export_full_body(&self, f16: bool, out: &mut Vec<u8>) {
1397 let u = |v: u32, o: &mut Vec<u8>| o.extend_from_slice(&v.to_le_bytes());
1398 u(u8::from(f16) as u32, out);
1399 u(self.seq_len as u32, out);
1400 u(self.num_kv_heads as u32, out);
1401 u(self.head_dim as u32, out);
1402 u(self.linear_state.len() as u32, out);
1403 u(self.imp.len() as u32, out);
1408 let push = |xs: &[f32], o: &mut Vec<u8>| {
1409 if f16 {
1410 for &x in xs {
1411 o.extend_from_slice(&cortiq_core::quant::f32_to_f16(x).to_le_bytes());
1412 }
1413 } else {
1414 for &x in xs {
1415 o.extend_from_slice(&x.to_le_bytes());
1416 }
1417 }
1418 };
1419 for &x in &self.linear_state {
1420 out.extend_from_slice(&x.to_le_bytes());
1421 }
1422 for &x in &self.imp {
1423 out.extend_from_slice(&x.to_le_bytes());
1424 }
1425 for h in 0..self.num_kv_heads {
1426 u(self.k[h].len() as u32, out);
1427 push(&self.k[h], out);
1428 u(self.v[h].len() as u32, out);
1429 push(&self.v[h], out);
1430 }
1431 }
1432
1433 pub fn import_wire(&mut self, buf: &[u8]) -> Result<(), String> {
1438 if buf.len() >= 4 && &buf[..4] == WIRE_MAGIC {
1439 return self.import_wire_v2(buf);
1440 }
1441 if !self.linear_wire_allowed {
1443 return Err(
1444 "kv import: Delta linear state cannot use the unversioned cache wire; refusing until the wire carries operator identity".into(),
1445 );
1446 }
1447 if self.bounded.is_some() {
1448 return Err(
1449 "kv import: a bounded anchor takes only the versioned wire (v2) — the \
1450 unversioned body has no ring record"
1451 .into(),
1452 );
1453 }
1454 let n = self.import_full_body(buf)?;
1455 if n != buf.len() {
1456 return Err(format!(
1457 "kv import: {} trailing byte(s) after the record",
1458 buf.len() - n
1459 ));
1460 }
1461 Ok(())
1462 }
1463
1464 fn import_wire_v2(&mut self, buf: &[u8]) -> Result<(), String> {
1465 let need = |n: usize, o: usize| -> Result<(), String> {
1466 if o + n > buf.len() {
1467 Err("kv import: truncated header".into())
1468 } else {
1469 Ok(())
1470 }
1471 };
1472 need(28, 0)?;
1473 let version = u32::from_le_bytes(buf[4..8].try_into().unwrap());
1474 if version != WIRE_VERSION {
1475 return Err(format!(
1476 "kv import: wire version {version}, this runtime speaks {WIRE_VERSION}"
1477 ));
1478 }
1479 let identity = u64::from_le_bytes(buf[8..16].try_into().unwrap());
1480 let layer = u32::from_le_bytes(buf[16..20].try_into().unwrap());
1481 let kind = WireKind::from_u8(buf[20])
1482 .ok_or_else(|| format!("kv import: unknown state kind {}", buf[20]))?;
1483 let f16 = buf[21] != 0;
1484 let position = u64::from_le_bytes(buf[24..32].try_into().unwrap()) as usize;
1485 if identity != self.wire_identity {
1486 return Err(format!(
1487 "kv import: peer operator identity {identity:016x} != mine {:016x} — \
1488 the two sides do not hold the same operator",
1489 self.wire_identity
1490 ));
1491 }
1492 if layer != self.wire_layer {
1493 return Err(format!(
1494 "kv import: record is for layer {layer}, this is layer {}",
1495 self.wire_layer
1496 ));
1497 }
1498 if kind != self.wire_kind {
1499 return Err(format!(
1500 "kv import: record kind {kind:?} does not match this layer's {:?}",
1501 self.wire_kind
1502 ));
1503 }
1504 let mut o = 32usize;
1505 let u32_at = |o: &mut usize| -> Result<u32, String> {
1506 if *o + 4 > buf.len() {
1507 return Err("kv import: truncated record".into());
1508 }
1509 let v = u32::from_le_bytes(buf[*o..*o + 4].try_into().unwrap());
1510 *o += 4;
1511 Ok(v)
1512 };
1513 let need_payload = |n: usize, o: usize| -> Result<(), String> {
1514 if o + n > buf.len() {
1515 Err("kv import: truncated payload".into())
1516 } else {
1517 Ok(())
1518 }
1519 };
1520 match kind {
1521 WireKind::Full => {
1522 let n = self.import_full_body(&buf[o..])?;
1523 o += n;
1524 if self.seq_len != position {
1525 return Err(format!(
1526 "kv import: header position {position} != record seq_len {}",
1527 self.seq_len
1528 ));
1529 }
1530 }
1531 WireKind::Linear => {
1532 let heads = u32_at(&mut o)? as usize;
1533 let hd = u32_at(&mut o)? as usize;
1534 if heads != self.num_kv_heads || hd != self.head_dim {
1535 return Err(format!(
1536 "kv import: peer sent {heads}×{hd} per position, this layer is {}×{}",
1537 self.num_kv_heads, self.head_dim
1538 ));
1539 }
1540 let lin = u32_at(&mut o)? as usize;
1541 need_payload(lin * 4, o)?;
1542 self.linear_state = (0..lin)
1543 .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1544 .collect();
1545 o += lin * 4;
1546 self.reset_per_position_storage();
1547 self.seq_len = position;
1548 }
1549 WireKind::Bounded => {
1550 let heads = u32_at(&mut o)? as usize;
1551 let hd = u32_at(&mut o)? as usize;
1552 let window = u32_at(&mut o)? as usize;
1553 let len = u32_at(&mut o)? as usize;
1554 let head = u32_at(&mut o)? as usize;
1555 let b = self
1556 .bounded
1557 .as_mut()
1558 .ok_or("kv import: bounded record for a layer without a ring")?;
1559 if heads != b.num_kv_heads || hd != b.head_dim || window != b.window {
1560 return Err(format!(
1561 "kv import: bounded record {heads}×{window}×{hd} does not fit this \
1562 layer's ring {}×{}×{}",
1563 b.num_kv_heads, b.window, b.head_dim
1564 ));
1565 }
1566 if len != position.min(window) || head != position % window {
1567 return Err(format!(
1568 "kv import: bounded record len/head {len}/{head} inconsistent with \
1569 position {position} (window {window})"
1570 ));
1571 }
1572 let n = b.ring_k.len();
1573 let w = if f16 { 2 } else { 4 };
1574 need_payload(2 * n * w, o)?;
1575 let read = |o: usize, dst: &mut [f32]| {
1576 for (i, d) in dst.iter_mut().enumerate() {
1577 let at = o + i * w;
1578 *d = if f16 {
1579 cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
1580 buf[at..at + 2].try_into().unwrap(),
1581 ))
1582 } else {
1583 f32::from_le_bytes(buf[at..at + 4].try_into().unwrap())
1584 };
1585 }
1586 };
1587 read(o, &mut b.ring_k);
1588 o += n * w;
1589 read(o, &mut b.ring_v);
1590 o += n * w;
1591 b.seen = position;
1592 let snap = b.snapshot();
1594 b.restore(&snap);
1595 self.linear_state = Vec::new();
1596 self.reset_per_position_storage();
1597 self.seq_len = position;
1598 }
1599 }
1600 if o != buf.len() {
1601 return Err(format!(
1602 "kv import: {} trailing byte(s) after the record",
1603 buf.len() - o
1604 ));
1605 }
1606 Ok(())
1607 }
1608
1609 fn reset_per_position_storage(&mut self) {
1612 let heads = self.num_kv_heads;
1613 self.mode = KvMode::F32;
1614 self.k = vec![Vec::new(); heads];
1615 self.v = vec![Vec::new(); heads];
1616 self.kq = vec![Vec::new(); heads];
1617 self.ks = vec![Vec::new(); heads];
1618 self.vq = vec![Vec::new(); heads];
1619 self.vs = vec![Vec::new(); heads];
1620 self.kcol = vec![Vec::new(); heads];
1621 self.vcol = vec![Vec::new(); heads];
1622 self.imp = Vec::new();
1623 self.discard_linear_scratch();
1624 self.o1 = None;
1625 self.o1_error = None;
1626 self.o1_transitioned = false;
1627 }
1628
1629 fn import_full_body(&mut self, buf: &[u8]) -> Result<usize, String> {
1632 let mut o = 0usize;
1633 let u32_at = |o: &mut usize| -> Result<u32, String> {
1634 if *o + 4 > buf.len() {
1635 return Err("kv import: truncated header".into());
1636 }
1637 let v = u32::from_le_bytes(buf[*o..*o + 4].try_into().unwrap());
1638 *o += 4;
1639 Ok(v)
1640 };
1641 let f16 = u32_at(&mut o)? != 0;
1642 let seq_len = u32_at(&mut o)? as usize;
1643 let heads = u32_at(&mut o)? as usize;
1644 let hd = u32_at(&mut o)? as usize;
1645 let lin = u32_at(&mut o)? as usize;
1646 let nimp = u32_at(&mut o)? as usize;
1647 if heads != self.num_kv_heads || hd != self.head_dim {
1648 return Err(format!(
1649 "kv import: peer sent {heads}×{hd} per position, this layer is {}×{}",
1650 self.num_kv_heads, self.head_dim
1651 ));
1652 }
1653 let w = if f16 { 2 } else { 4 };
1654 let need = |n: usize, o: usize| -> Result<(), String> {
1655 if o + n > buf.len() {
1656 Err("kv import: truncated payload".into())
1657 } else {
1658 Ok(())
1659 }
1660 };
1661 need(lin * 4, o)?;
1662 self.linear_state = (0..lin)
1663 .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1664 .collect();
1665 o += lin * 4;
1666 need(nimp * 4, o)?;
1667 let imp: Vec<f32> = (0..nimp)
1668 .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1669 .collect();
1670 o += nimp * 4;
1671 let mut k: Vec<Vec<f32>> = Vec::with_capacity(heads);
1672 let mut v: Vec<Vec<f32>> = Vec::with_capacity(heads);
1673 for _ in 0..heads {
1674 for which in 0..2 {
1675 let n = u32_at(&mut o)? as usize;
1676 need(n * w, o)?;
1677 let xs: Vec<f32> = (0..n)
1678 .map(|i| {
1679 let at = o + i * w;
1680 if f16 {
1681 cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
1682 buf[at..at + 2].try_into().unwrap(),
1683 ))
1684 } else {
1685 f32::from_le_bytes(buf[at..at + 4].try_into().unwrap())
1686 }
1687 })
1688 .collect();
1689 o += n * w;
1690 if which == 0 { k.push(xs) } else { v.push(xs) }
1691 }
1692 }
1693 self.reset_per_position_storage();
1694 self.k = k;
1695 self.v = v;
1696 self.imp = imp;
1697 self.seq_len = seq_len;
1698 Ok(o)
1699 }
1700
1701 pub fn memory_bytes(&self) -> usize {
1702 let floats: usize = self.k.iter().map(Vec::len).sum::<usize>()
1703 + self.v.iter().map(Vec::len).sum::<usize>()
1704 + self.ks.iter().map(Vec::len).sum::<usize>()
1705 + self.vs.iter().map(Vec::len).sum::<usize>()
1706 + self.kcol.iter().map(Vec::len).sum::<usize>()
1707 + self.vcol.iter().map(Vec::len).sum::<usize>();
1708 let bytes: usize = self.kq.iter().map(Vec::len).sum::<usize>()
1709 + self.vq.iter().map(Vec::len).sum::<usize>();
1710 floats * std::mem::size_of::<f32>()
1711 + bytes
1712 + self.linear_state.len() * std::mem::size_of::<f32>()
1716 + self.o1_memory_bytes()
1719 + self.bounded_state_bytes()
1722 }
1723
1724 fn evict(&mut self, keep_last: usize) {
1726 if self.o1.is_some() || self.bounded.is_some() || self.seq_len <= keep_last {
1735 return;
1736 }
1737 let drop = self.seq_len - keep_last;
1738 for h in 0..self.num_kv_heads {
1739 let stored = self.head_len(h);
1741 let d = drop.min(stored);
1742 let hd = self.head_dim;
1743 fn drop_front<T>(v: &mut Vec<T>, n: usize) {
1744 let n = n.min(v.len());
1745 v.drain(..n);
1746 }
1747 drop_front(&mut self.k[h], d * hd);
1748 drop_front(&mut self.v[h], d * hd);
1749 drop_front(&mut self.kq[h], d * hd);
1750 drop_front(&mut self.vq[h], d * hd);
1751 drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
1752 drop_front(&mut self.vs[h], d);
1753 }
1754 let d = drop.min(self.imp.len());
1755 self.imp.drain(..d);
1756 self.seq_len = keep_last;
1757 }
1758
1759 fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
1764 if self.o1.is_some() || self.bounded.is_some() {
1765 return;
1769 }
1770 let stored = self.imp.len();
1771 if stored <= keep_last {
1772 return;
1773 }
1774 let sink_n = sink.min(keep_last);
1777 let recent_n = recent.min(keep_last - sink_n);
1778 let mut keep = vec![false; stored];
1779 for k in keep.iter_mut().take(sink_n) {
1780 *k = true;
1781 }
1782 for k in keep.iter_mut().skip(stored.saturating_sub(recent_n)) {
1783 *k = true;
1784 }
1785 let mut budget = keep_last.saturating_sub(keep.iter().filter(|&&x| x).count());
1786 let mut order: Vec<usize> = (0..stored).filter(|&i| !keep[i]).collect();
1788 order.sort_by(|&a, &b| {
1789 self.imp[b]
1790 .partial_cmp(&self.imp[a])
1791 .unwrap_or(std::cmp::Ordering::Equal)
1792 });
1793 for i in order {
1794 if budget == 0 {
1795 break;
1796 }
1797 keep[i] = true;
1798 budget -= 1;
1799 }
1800
1801 let kept: Vec<usize> = (0..stored).filter(|&i| keep[i]).collect();
1802 let hd = self.head_dim;
1803 fn gather<T: Copy>(src: &[T], kept: &[usize], step: usize) -> Vec<T> {
1804 let mut out = Vec::with_capacity(kept.len() * step);
1805 for &i in kept {
1806 out.extend_from_slice(&src[i * step..(i + 1) * step]);
1807 }
1808 out
1809 }
1810 for h in 0..self.num_kv_heads {
1815 if !self.k[h].is_empty() {
1816 self.k[h] = gather(&self.k[h], &kept, hd);
1817 }
1818 if !self.v[h].is_empty() {
1819 self.v[h] = gather(&self.v[h], &kept, hd);
1820 }
1821 if !self.kq[h].is_empty() {
1822 self.kq[h] = gather(&self.kq[h], &kept, hd);
1823 self.ks[h] = gather(&self.ks[h], &kept, hd.div_ceil(KV_K_GROUP));
1824 }
1825 if !self.vq[h].is_empty() {
1826 self.vq[h] = gather(&self.vq[h], &kept, hd);
1827 self.vs[h] = gather(&self.vs[h], &kept, 1);
1828 }
1829 }
1830 self.imp = kept.iter().map(|&i| self.imp[i]).collect();
1831 self.seq_len = kept.len();
1832 }
1833}
1834
1835#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1837pub enum EvictionPolicy {
1838 Recent,
1840 Born { sink: usize },
1842}
1843
1844#[derive(Debug)]
1846pub struct KvCache {
1847 pub layers: Vec<LayerKvCache>,
1848 pub max_seq_len: usize,
1849 pub policy: EvictionPolicy,
1850}
1851
1852impl KvCache {
1853 pub fn new(
1854 num_layers: usize,
1855 num_kv_heads: usize,
1856 head_dim: usize,
1857 max_seq_len: usize,
1858 ) -> Self {
1859 let layers = (0..num_layers)
1860 .map(|li| {
1861 let mut l = LayerKvCache::new(num_kv_heads, head_dim);
1862 l.wire_layer = li as u32;
1863 l
1864 })
1865 .collect();
1866 Self {
1867 layers,
1868 max_seq_len,
1869 policy: EvictionPolicy::Born { sink: 4 },
1870 }
1871 }
1872
1873 pub fn clear(&mut self) {
1874 for layer in &mut self.layers {
1875 layer.clear();
1876 }
1877 }
1878
1879 pub fn total_memory_bytes(&self) -> usize {
1880 self.layers.iter().map(|l| l.memory_bytes()).sum()
1881 }
1882
1883 pub fn recurrent_state_bytes(&self) -> usize {
1888 let floats: usize = self
1889 .layers
1890 .iter()
1891 .map(|l| l.linear_state.len() + l.linear_scratch.len())
1892 .sum();
1893 floats * std::mem::size_of::<f32>()
1894 }
1895
1896 pub fn attention_state_bytes(&self) -> usize {
1899 self.total_memory_bytes()
1900 .saturating_sub(self.recurrent_state_bytes())
1901 }
1902
1903 pub fn seq_len(&self) -> usize {
1905 self.layers.iter().map(|l| l.seq_len).max().unwrap_or(0)
1906 }
1907
1908 pub fn bounded_state_bytes(&self) -> usize {
1910 self.layers.iter().map(|l| l.bounded_state_bytes()).sum()
1911 }
1912
1913 pub fn needs_eviction(&self) -> bool {
1917 self.layers
1918 .iter()
1919 .filter(|l| l.bounded.is_none())
1920 .map(|l| l.seq_len)
1921 .max()
1922 .unwrap_or(0)
1923 >= self.max_seq_len
1924 }
1925
1926 pub fn evict(&mut self, keep_last: usize) {
1928 match self.policy {
1929 EvictionPolicy::Recent => {
1930 for layer in &mut self.layers {
1931 layer.evict(keep_last);
1932 }
1933 }
1934 EvictionPolicy::Born { sink } => {
1935 let recent = (keep_last / 2).max(1);
1936 for layer in &mut self.layers {
1937 layer.evict_born(keep_last, sink, recent);
1938 }
1939 }
1940 }
1941 }
1942}
1943
1944#[cfg(test)]
1945mod tests {
1946 use super::*;
1947
1948 #[test]
1949 fn memory_breakdown_separates_recurrent_and_attention_state() {
1950 let mut cache = KvCache::new(1, 1, 4, 16);
1951 cache.layers[0].linear_state = vec![0.0; 8];
1952 cache.layers[0].linear_scratch = vec![0.0; 4];
1953 cache.layers[0].append(&[0.0; 4], &[1.0; 4], &[true]);
1954 let recurrent = cache.recurrent_state_bytes();
1955 assert_eq!(recurrent, 12 * std::mem::size_of::<f32>());
1956 assert_eq!(
1957 cache.attention_state_bytes() + recurrent,
1958 cache.total_memory_bytes()
1959 );
1960 assert!(cache.attention_state_bytes() > 0);
1961 }
1962
1963 #[test]
1964 fn wire_round_trip_reproduces_attention() {
1965 let (heads, hd) = (2usize, 4usize);
1969 let mut a = LayerKvCache::new(heads, hd);
1970 for p in 0..5 {
1971 let k: Vec<f32> = (0..heads * hd)
1972 .map(|i| (p * 10 + i) as f32 * 0.031)
1973 .collect();
1974 let v: Vec<f32> = (0..heads * hd)
1975 .map(|i| (p * 7 + i) as f32 * -0.017)
1976 .collect();
1977 a.append(&k, &v, &[true, true]);
1978 }
1979 a.linear_state = vec![0.5, -0.25, 1.0];
1980 let q: Vec<f32> = (0..hd).map(|i| 0.1 * (i as f32 + 1.0)).collect();
1981
1982 let bytes = a.export_wire(false).expect("f32 cache exports");
1983 let mut b = LayerKvCache::new(heads, hd);
1984 b.linear_scratch = vec![9.0; 3];
1985 b.import_wire(&bytes).expect("import");
1986 assert!(
1987 b.linear_scratch.is_empty(),
1988 "import must discard tentative state"
1989 );
1990
1991 assert_eq!(b.seq_len, a.seq_len);
1992 assert_eq!(b.linear_state, a.linear_state);
1993 for h in 0..heads {
1994 let (oa, sa) = a.attend(&q, h);
1995 let (ob, sb) = b.attend(&q, h);
1996 assert_eq!(oa, ob, "head {h} attention output diverged");
1997 assert_eq!(sa, sb, "head {h} attention scores diverged");
1998 }
1999 }
2000
2001 #[test]
2002 fn wire_refuses_what_it_cannot_describe() {
2003 let mut c = LayerKvCache::new(1, 4);
2006 c.mode = KvMode::Q8 { k: true, v: true };
2007 let err = c.export_wire(false).unwrap_err();
2008 assert!(err.contains("F32"), "{err}");
2009 }
2010
2011 #[test]
2012 fn wire_refuses_unversioned_delta_state() {
2013 let mut c = LayerKvCache::new(1, 4);
2016 c.set_linear_wire_allowed(false);
2017 let bytes = c.export_wire(false).expect("v2 export carries identity");
2018 assert_eq!(&bytes[..4], WIRE_MAGIC);
2019 let err = c.import_wire(&[0, 0, 0, 0]).unwrap_err();
2020 assert!(err.contains("operator identity"), "{err}");
2021 }
2022
2023 #[test]
2024 fn wire_v2_round_trips_linear_and_bounded_records() {
2025 let mut a = LayerKvCache::new(1, 4);
2028 a.wire_kind = WireKind::Linear;
2029 a.wire_identity = 0xC0FFEE;
2030 a.linear_state = vec![0.5, -0.25, 1.0, 3.5];
2031 a.seq_len = 9;
2032 let bytes = a.export_wire(true).unwrap();
2033 let mut b = LayerKvCache::new(1, 4);
2034 b.wire_kind = WireKind::Linear;
2035 b.wire_identity = 0xC0FFEE;
2036 b.import_wire(&bytes).unwrap();
2037 assert_eq!(b.linear_state, a.linear_state);
2038 assert_eq!(b.seq_len, 9);
2039 let mut c = LayerKvCache::new(1, 4);
2041 c.wire_kind = WireKind::Linear;
2042 let err = c.import_wire(&bytes).unwrap_err();
2043 assert!(err.contains("operator identity"), "{err}");
2044
2045 let (kvh, hd, w) = (2, 4, 8);
2047 let mut a = LayerKvCache::new(kvh, hd);
2048 a.install_bounded(w);
2049 let k: Vec<f32> = (0..kvh * hd).map(|i| i as f32 * 0.125).collect();
2050 for p in 0..11 {
2051 a.bounded.as_mut().unwrap().insert(&k, &k);
2052 a.seq_len = p + 1;
2053 }
2054 for f16 in [false, true] {
2055 let bytes = a.export_wire(f16).unwrap();
2056 let mut b = LayerKvCache::new(kvh, hd);
2057 b.install_bounded(w);
2058 b.import_wire(&bytes).unwrap();
2059 let (ra, rb) = (a.bounded.as_ref().unwrap(), b.bounded.as_ref().unwrap());
2060 assert_eq!(rb.seen, 11);
2061 assert_eq!(b.seq_len, 11);
2062 assert!(ra.same_state(rb), "f16={f16}");
2064 let mut c = LayerKvCache::new(kvh, hd);
2066 c.install_bounded(w * 2);
2067 assert!(c.import_wire(&bytes).is_err());
2068 }
2069 }
2070
2071 #[test]
2072 fn wire_import_checks_geometry() {
2073 let a = LayerKvCache::new(2, 4);
2074 let bytes = a.export_wire(false).unwrap();
2075 let mut wrong = LayerKvCache::new(2, 8);
2076 let err = wrong.import_wire(&bytes).unwrap_err();
2077 assert!(err.contains("2×4"), "{err}");
2078 }
2079
2080 #[test]
2081 fn append_tracks_seq_len_and_layout() {
2082 let mut cache = LayerKvCache::new(4, 8);
2083 cache.mode = KvMode::F32;
2084 assert_eq!(cache.seq_len, 0);
2085
2086 let k: Vec<f32> = (0..32).map(|i| i as f32).collect();
2087 let v = vec![2.0f32; 32];
2088 cache.append(&k, &v, &[true; 4]);
2089
2090 assert_eq!(cache.seq_len, 1);
2091 assert_eq!(cache.head_len(0), 1);
2092 assert_eq!(cache.head_keys(1), &k[8..16]);
2094 assert_eq!(cache.memory_bytes(), 256);
2095 }
2096
2097 #[test]
2098 fn dead_head_stores_nothing() {
2099 let mut cache = LayerKvCache::new(2, 4);
2100 cache.mode = KvMode::F32;
2101 let k = vec![1.0f32; 8];
2102 let v = vec![2.0f32; 8];
2103 cache.append(&k, &v, &[true, false]);
2104 cache.append(&k, &v, &[true, false]);
2105
2106 assert_eq!(cache.seq_len, 2);
2107 assert_eq!(cache.head_len(0), 2);
2108 assert_eq!(cache.head_len(1), 0, "dead head must not store KV");
2109 assert_eq!(cache.memory_bytes(), 2 * 2 * 4 * 4);
2110 }
2111
2112 #[test]
2113 fn eviction_keeps_recent() {
2114 let mut cache = KvCache::new(2, 4, 8, 10);
2115 cache.policy = EvictionPolicy::Recent;
2116 for l in &mut cache.layers {
2117 l.mode = KvMode::F32;
2118 }
2119 let k = vec![1.0f32; 32];
2120 let v = vec![2.0f32; 32];
2121 for _ in 0..8 {
2122 for layer in &mut cache.layers {
2123 layer.append(&k, &v, &[true; 4]);
2124 }
2125 }
2126 assert_eq!(cache.seq_len(), 8);
2127 assert!(!cache.needs_eviction());
2128
2129 cache.evict(4);
2130 assert_eq!(cache.seq_len(), 4);
2131 assert_eq!(cache.layers[0].head_len(0), 4);
2132 }
2133
2134 #[test]
2135 fn collecting_o1_eviction_retains_exact_storage_until_boundary() {
2136 const B: usize = 19;
2137 let q = vec![0.1f32; 8];
2138 let k = vec![0.2f32; 4];
2139 let v = vec![0.3f32; 4];
2140
2141 for policy in [EvictionPolicy::Recent, EvictionPolicy::Born { sink: 2 }] {
2142 let mut cache = KvCache::new(1, 1, 4, 6);
2143 cache.policy = policy;
2144 cache.layers[0].mode = KvMode::F32;
2145 cache.layers[0].o1_begin_with_boundary(
2146 4,
2147 8,
2148 2,
2149 crate::nystrom::O1Rect::Aggregate,
2150 Some(B),
2151 );
2152
2153 for pos in 0..B {
2154 {
2155 let layer = &mut cache.layers[0];
2156 layer.o1_push_q(&q);
2157 layer.append(&k, &v, &[]);
2158 }
2159 if pos + 1 < B {
2160 cache.evict(3);
2161 }
2162 }
2163
2164 let layer = &cache.layers[0];
2165 let rows = B * layer.head_dim;
2166 assert_eq!(layer.seq_len, B, "policy {policy:?} retained depth");
2167 assert_eq!(layer.k[0].len(), rows, "policy {policy:?} K rows");
2168 assert_eq!(layer.v[0].len(), rows, "policy {policy:?} V rows");
2169 assert!(
2170 layer.k[0].capacity() >= rows,
2171 "policy {policy:?} K capacity"
2172 );
2173 assert!(
2174 layer.v[0].capacity() >= rows,
2175 "policy {policy:?} V capacity"
2176 );
2177 let q_capacity = match layer.o1.as_ref() {
2178 Some(O1State::Collecting { q_buf, .. }) => q_buf.capacity(),
2179 other => panic!("policy {policy:?} changed state early: {other:?}"),
2180 };
2181 assert!(
2182 q_capacity >= B * 8,
2183 "policy {policy:?} Q capacity must cover the exact prefix"
2184 );
2185
2186 assert!(cache.layers[0].o1_seal_checked(2).unwrap());
2187 assert_eq!(cache.layers[0].k[0].capacity(), 0, "K released after seal");
2188 assert_eq!(cache.layers[0].v[0].capacity(), 0, "V released after seal");
2189 }
2190 }
2191
2192 #[test]
2193 fn truncate_rolls_back_speculative_positions() {
2194 let mut cache = LayerKvCache::new(2, 4);
2195 cache.mode = KvMode::F32;
2196 for pos in 0..5 {
2197 let k = vec![pos as f32; 8];
2198 let v = vec![pos as f32; 8];
2199 cache.append(&k, &v, &[true; 2]);
2200 }
2201 cache.truncate_last(2);
2202 assert_eq!(cache.seq_len, 3);
2203 assert_eq!(cache.head_len(0), 3);
2204 assert_eq!(cache.head_keys(0)[2 * 4], 2.0, "position 2 survives");
2205 }
2206
2207 #[test]
2211 fn q8_attend_matches_f32_within_grid() {
2212 let (heads, hd) = (2, 32);
2213 let mut f = LayerKvCache::new(heads, hd);
2214 f.mode = KvMode::F32;
2215 let mut q8 = LayerKvCache::new(heads, hd);
2216 q8.mode = KvMode::Q8 { k: true, v: true };
2217
2218 let synth = |p: usize, salt: usize| -> Vec<f32> {
2219 (0..heads * hd)
2220 .map(|i| {
2221 let x = ((i * 31 + p * 17 + salt * 7 + 3) % 97) as f32 / 97.0 - 0.5;
2222 if i % 2 == 0 { x * 4.0 } else { x * 0.25 }
2224 })
2225 .collect()
2226 };
2227 for p in 0..100 {
2228 let k = synth(p, 1);
2229 let v = synth(p, 2);
2230 f.append(&k, &v, &[true; 2]);
2231 q8.append(&k, &v, &[true; 2]);
2232 }
2233 let q: Vec<f32> = (0..hd)
2234 .map(|i| ((i * 13 + 5) % 89) as f32 / 89.0 - 0.5)
2235 .collect();
2236 for g in 0..heads {
2237 let (of, pf) = f.attend(&q, g);
2238 let (o8, p8) = q8.attend(&q, g);
2239 let scale = of.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-6);
2240 for d in 0..hd {
2241 assert!(
2242 (of[d] - o8[d]).abs() <= scale * 0.03 + 1e-3,
2243 "g{g} d{d}: f32 {} vs q8 {}",
2244 of[d],
2245 o8[d]
2246 );
2247 }
2248 for p in 0..100 {
2249 assert!((pf[p] - p8[p]).abs() < 0.02, "prob p{p}");
2250 }
2251 }
2252 q8.truncate_last(30);
2254 assert_eq!(q8.head_len(0), 70);
2255 let imp: Vec<f32> = (0..70).map(|i| i as f32).collect();
2256 q8.accumulate_imp(&imp);
2257 q8.evict_born(20, 2, 8);
2258 assert_eq!(q8.head_len(0), 20);
2259 let (o, _) = q8.attend(&q, 0);
2260 assert!(o.iter().all(|x| x.is_finite()));
2261 assert!(q8.memory_bytes() * 3 < f.memory_bytes());
2263 }
2264
2265 #[test]
2268 fn attend_group_equals_per_head_attend_bitexact() {
2269 let (kv_heads, hd, hpk) = (2usize, 32usize, 3usize); for mode in [KvMode::F32, KvMode::Q8 { k: true, v: true }] {
2271 let mut c = LayerKvCache::new(kv_heads, hd);
2272 c.mode = mode;
2273 for p in 0..70 {
2274 let k: Vec<f32> = (0..kv_heads * hd)
2275 .map(|i| ((i * 31 + p * 17 + 3) % 97) as f32 / 97.0 - 0.5)
2276 .collect();
2277 let v: Vec<f32> = (0..kv_heads * hd)
2278 .map(|i| ((i * 13 + p * 29 + 7) % 89) as f32 / 89.0 - 0.5)
2279 .collect();
2280 c.append(&k, &v, &[true; 2]);
2281 }
2282 let q: Vec<f32> = (0..kv_heads * hpk * hd)
2283 .map(|i| ((i * 11 + 5) % 83) as f32 / 83.0 - 0.5)
2284 .collect();
2285 for g in 0..kv_heads {
2286 let span = g * hpk * hd..(g + 1) * hpk * hd;
2287 let mut out = vec![0f32; hpk * hd];
2288 let mut imp = vec![0f32; 70];
2289 c.attend_group(
2290 &q[span.clone()],
2291 g,
2292 &mut out,
2293 &mut imp,
2294 1.0 / (hd as f32).sqrt(),
2295 0,
2296 0.0,
2297 &[],
2298 );
2299 let mut imp_ref = vec![0f32; 70];
2300 for h in 0..hpk {
2301 let qh = &q[span.start + h * hd..span.start + (h + 1) * hd];
2302 let (o, probs) = c.attend(qh, g);
2303 assert_eq!(
2304 &out[h * hd..(h + 1) * hd],
2305 &o[..],
2306 "mode {mode:?} g{g} h{h}: grouped attend must be bit-identical"
2307 );
2308 for (dst, &p) in imp_ref.iter_mut().zip(&probs) {
2309 *dst += p;
2310 }
2311 }
2312 assert_eq!(
2313 imp, imp_ref,
2314 "mode {mode:?} g{g}: attention mass must match"
2315 );
2316 }
2317 }
2318 }
2319
2320 #[test]
2327 fn sink_attend_matches_explicit_sink_column() {
2328 let (nkv, hd, hpk) = (2usize, 8usize, 3usize);
2329 let rows = 9usize;
2330 let mut c = LayerKvCache::new(nkv, hd);
2331 c.mode = KvMode::F32;
2332 let kv = |r: usize, i: usize, a: usize, m: usize| {
2333 (((r * a + i * 7 + 3) % m) as f32 / m as f32 - 0.5) * 2.0
2334 };
2335 let mut ks = Vec::new();
2336 let mut vs = Vec::new();
2337 for r in 0..rows {
2338 let k: Vec<f32> = (0..nkv * hd).map(|i| kv(r, i, 31, 97)).collect();
2339 let v: Vec<f32> = (0..nkv * hd).map(|i| kv(r, i, 17, 89)).collect();
2340 c.append(&k, &v, &[]);
2341 ks.push(k);
2342 vs.push(v);
2343 }
2344 let q: Vec<f32> = (0..nkv * hpk * hd)
2345 .map(|i| (((i * 11 + 5) % 83) as f32 / 83.0 - 0.5) * 3.0)
2346 .collect();
2347 let sinks = [0.7f32, -1.3, 2.5, 0.0, -4.0, 6.0];
2349 let scale = 1.0 / (hd as f32).sqrt();
2350 let mut checked = 0usize;
2351 for upto in [1usize, 2, 5, 9] {
2352 for window in [None, Some(3usize), Some(1)] {
2353 let first = window.map(|w| upto.saturating_sub(w)).unwrap_or(0);
2354 for g in 0..nkv {
2355 let qg = &q[g * hpk * hd..(g + 1) * hpk * hd];
2356 let sg = &sinks[g * hpk..(g + 1) * hpk];
2357 let mut out = vec![0f32; hpk * hd];
2358 let mut imp = vec![0f32; upto];
2359 c.attend_group_upto(qg, g, &mut out, &mut imp, scale, first, 0.0, upto, sg);
2360 let mut imp_ref = vec![0f64; upto];
2361 for h in 0..hpk {
2362 let qh = &qg[h * hd..(h + 1) * hd];
2363 let mut z: Vec<f64> = (first..upto)
2365 .map(|p| {
2366 let k = &ks[p][g * hd..(g + 1) * hd];
2367 qh.iter()
2368 .zip(k)
2369 .map(|(&a, &b)| a as f64 * b as f64)
2370 .sum::<f64>()
2371 * scale as f64
2372 })
2373 .collect();
2374 z.push(sg[h] as f64);
2375 let m = z.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
2376 let e: Vec<f64> = z.iter().map(|&x| (x - m).exp()).collect();
2377 let s: f64 = e.iter().sum();
2378 let p: Vec<f64> = e.iter().map(|&x| x / s).collect();
2379 for d in 0..hd {
2381 let want: f64 = (first..upto)
2382 .map(|r| p[r - first] * vs[r][g * hd + d] as f64)
2383 .sum();
2384 let got = out[h * hd + d] as f64;
2385 assert!(
2386 (got - want).abs() < 1e-6,
2387 "upto {upto} window {window:?} g{g} h{h} d{d}: {got} vs {want}"
2388 );
2389 }
2390 for r in first..upto {
2391 imp_ref[r] += p[r - first];
2392 }
2393 checked += 1;
2394 }
2395 for r in 0..upto {
2396 assert!(
2397 (imp[r] as f64 - imp_ref[r]).abs() < 1e-6,
2398 "imp upto {upto} window {window:?} g{g} row {r}: {} vs {}",
2399 imp[r],
2400 imp_ref[r]
2401 );
2402 }
2403 let row_mass: f32 = imp.iter().sum();
2405 assert!(row_mass < hpk as f32, "sinks must absorb some mass");
2406 }
2407 }
2408 }
2409 assert_eq!(checked, 4 * 3 * nkv * hpk);
2410 }
2411
2412 #[test]
2417 fn windowed_attend_equals_masked_full_row() {
2418 let (nkv, hd, hpk) = (1usize, 16usize, 2usize);
2419 let rows = 40usize;
2420 let mut c = LayerKvCache::new(nkv, hd);
2421 c.mode = KvMode::F32;
2422 for r in 0..rows {
2423 let k: Vec<f32> = (0..hd)
2424 .map(|i| ((r * 13 + i * 5) % 29) as f32 / 29.0 - 0.5)
2425 .collect();
2426 let v: Vec<f32> = (0..hd)
2427 .map(|i| ((r * 7 + i * 3) % 31) as f32 / 31.0 - 0.5)
2428 .collect();
2429 c.append(&k, &v, &[]);
2430 }
2431 let q: Vec<f32> = (0..hpk * hd)
2432 .map(|i| ((i * 19) % 23) as f32 / 23.0 - 0.5)
2433 .collect();
2434 let scale = 0.25f32;
2435 for w in [1usize, 7, 39, 40, 100] {
2436 let first = rows.saturating_sub(w);
2437 let mut out = vec![0f32; hpk * hd];
2438 let mut imp = vec![0f32; rows];
2439 c.attend_group(&q, 0, &mut out, &mut imp, scale, first, 0.0, &[]);
2440 let mut out_ref = vec![0f32; hpk * hd];
2442 let mut imp_ref = vec![0f32; rows];
2443 for h in 0..hpk {
2444 let mut s = vec![f32::NEG_INFINITY; rows];
2445 for p in first..rows {
2446 s[p] = crate::attention::dot_f32(
2447 &q[h * hd..(h + 1) * hd],
2448 &c.head_keys(0)[p * hd..(p + 1) * hd],
2449 ) * scale;
2450 }
2451 let m = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2452 let mut sum = 0f32;
2453 for v in s.iter_mut() {
2454 *v = (*v - m).exp();
2455 sum += *v;
2456 }
2457 for v in s.iter_mut() {
2458 *v /= sum;
2459 }
2460 for p in first..rows {
2461 if s[p].abs() < 1e-12 {
2462 continue;
2463 }
2464 crate::attention::axpy_f32(
2465 &mut out_ref[h * hd..(h + 1) * hd],
2466 &c.head_values(0)[p * hd..(p + 1) * hd],
2467 s[p],
2468 );
2469 }
2470 for (d, &p) in imp_ref.iter_mut().zip(&s) {
2471 *d += p;
2472 }
2473 }
2474 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
2475 assert_eq!(bits(&out), bits(&out_ref), "out window {w}");
2476 assert_eq!(bits(&imp), bits(&imp_ref), "imp window {w}");
2477 }
2478 }
2479
2480 #[test]
2483 fn sinks_survive_clear_and_wire_import() {
2484 let mut c = LayerKvCache::new(1, 4);
2485 c.mode = KvMode::F32;
2486 c.sinks = Some(vec![0.5, -0.5]);
2487 c.append(&[1.0; 4], &[2.0; 4], &[]);
2488 let wire = c.export_wire(false).unwrap();
2489 c.clear();
2490 assert_eq!(c.sinks.as_deref(), Some(&[0.5f32, -0.5][..]));
2491 c.import_wire(&wire).unwrap();
2492 assert_eq!(c.sinks.as_deref(), Some(&[0.5f32, -0.5][..]));
2493 assert_eq!(c.seq_len, 1);
2494 }
2495
2496 #[test]
2500 fn born_eviction_mixed_modes_stay_consistent() {
2501 for (mk, mv) in [(false, true), (true, false)] {
2502 let mut c = LayerKvCache::new(1, 4);
2503 c.mode = KvMode::Q8 { k: mk, v: mv };
2504 for p in 0..80 {
2505 let k = vec![p as f32 * 0.01; 4];
2506 let v = vec![p as f32; 4];
2507 c.append(&k, &v, &[true]);
2508 }
2509 let imp: Vec<f32> = (0..80).map(|i| i as f32).collect();
2510 c.accumulate_imp(&imp);
2511 let before = c.memory_bytes();
2512 c.evict_born(20, 4, 8); assert_eq!(c.head_len(0), 20, "k={mk} v={mv}");
2514 assert!(
2515 c.memory_bytes() < before / 2,
2516 "memory must shrink (k={mk} v={mv})"
2517 );
2518 let (out, _) = c.attend(&[1.0, 1.0, 1.0, 1.0], 0);
2521 assert!(
2522 out[0] > 30.0,
2523 "V from the kept tail, not the stale head (k={mk} v={mv}, out {})",
2524 out[0]
2525 );
2526 }
2527 }
2528
2529 #[test]
2530 fn born_eviction_keeps_high_mass_position() {
2531 let mut cache = KvCache::new(1, 1, 2, 16);
2532 cache.policy = EvictionPolicy::Born { sink: 1 };
2533 for l in &mut cache.layers {
2534 l.mode = KvMode::F32;
2535 }
2536 let layer = &mut cache.layers[0];
2537 for pos in 0..8 {
2540 let k = vec![pos as f32; 2];
2541 let v = vec![pos as f32 + 100.0; 2];
2542 layer.append(&k, &v, &[true]);
2543 }
2544 let mut imp = vec![0.05f32; 8];
2546 imp[3] = 5.0;
2547 layer.accumulate_imp(&imp);
2548
2549 cache.evict(4); let layer = &cache.layers[0];
2551 assert_eq!(layer.seq_len, 4);
2552 let kept_keys: Vec<f32> = (0..4).map(|i| layer.head_keys(0)[i * 2]).collect();
2553 assert_eq!(
2554 kept_keys,
2555 vec![0.0, 3.0, 6.0, 7.0],
2556 "kept = sink(0) + mass-top(3) + recent(6,7)"
2557 );
2558 assert_eq!(layer.head_len(0), 4);
2560 }
2561}