Skip to main content

cortiq_engine/
kv_cache.rs

1//! KV cache โ€” per-layer, head-major storage.
2//!
3//! Layout: one contiguous `Vec<f32>` per KV head (`[pos ร— head_dim]`),
4//! so per-head attention reads a straight slice โ€” no per-head gather
5//! copies per token. Dead GQA groups (all Q heads masked) store
6//! nothing at all: masked heads cost neither FLOPs nor memory.
7
8/// KV storage mode. `CMF_KV=q8` enables the q8_2f cache: an int8 row per
9/// (position, head) + an f32 scale per row + a per-channel field ๐’ฒร—ฮธ,
10/// frozen after WARMUP positions with retroactive requantization
11/// (D4: "KV-quant 2f"). Memory ร—~3.7 smaller than f32.
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum KvMode {
14    F32,
15    /// Quantized components: (K, V) โ€” sensitivity diagnostics.
16    Q8 { k: bool, v: bool },
17}
18
19impl KvMode {
20    pub fn from_env() -> Self {
21        match std::env::var("CMF_KV").as_deref() {
22            Ok("q8") | Ok("q8_2f") => KvMode::Q8 { k: true, v: true },
23            Ok("q8k") => KvMode::Q8 { k: true, v: false },
24            Ok("q8v") => KvMode::Q8 { k: false, v: true },
25            _ => KvMode::F32,
26        }
27    }
28
29    fn quant_k(self) -> bool {
30        matches!(self, KvMode::Q8 { k: true, .. })
31    }
32
33    fn quant_v(self) -> bool {
34        matches!(self, KvMode::Q8 { v: true, .. })
35    }
36}
37
38/// Positions before freezing the per-channel field (2f): before โ€” col โ‰ก 1,
39/// after โ€” col = RMS over channels of the stored rows, old rows are requantized.
40const KV_COL_WARMUP: usize = 64;
41
42/// K-rows are quantized in groups of 32 channels (scale per group):
43/// attention logits are sensitive to the dot-product error, per-group scales
44/// localize it along RoPE bands (35B: +4.6% PPL with a per-row scale
45/// โ†’ target <1% with a per-group one). V โ€” per-row scale (measured +0.56%).
46const KV_K_GROUP: usize = 32;
47
48/// KV cache for a single layer, head-major.
49#[derive(Debug, Clone)]
50pub struct LayerKvCache {
51    pub mode: KvMode,
52    /// Per-KV-head keys: `k[h]` is `[seq_len ร— head_dim]` (empty if head is dead).
53    k: Vec<Vec<f32>>,
54    /// Per-KV-head values, same layout.
55    v: Vec<Vec<f32>>,
56    /// q8 storage (mode == Q8_2F): int8 rows + f32 scale per row.
57    kq: Vec<Vec<i8>>,
58    ks: Vec<Vec<f32>>,
59    vq: Vec<Vec<i8>>,
60    vs: Vec<Vec<f32>>,
61    /// Per-channel fields ๐’ฒร—ฮธ per head [head_dim]; empty until frozen.
62    kcol: Vec<Vec<f32>>,
63    vcol: Vec<Vec<f32>>,
64    /// Accumulated attention mass per stored position (Born rule:
65    /// importance of a position = how much probability mass reads it).
66    imp: Vec<f32>,
67    /// Positions appended so far (grows once per token, dead heads included).
68    pub seq_len: usize,
69    pub num_kv_heads: usize,
70    pub head_dim: usize,
71    /// Linear-core condensate S (vmf_phase), f64; empty on full layers.
72    pub linear_state: Vec<f64>,
73    /// Tentative lane-2 state during speculative verify.
74    pub linear_scratch: Vec<f64>,
75}
76
77impl LayerKvCache {
78    pub fn new(num_kv_heads: usize, head_dim: usize) -> Self {
79        Self {
80            mode: KvMode::from_env(),
81            k: vec![Vec::new(); num_kv_heads],
82            v: vec![Vec::new(); num_kv_heads],
83            kq: vec![Vec::new(); num_kv_heads],
84            ks: vec![Vec::new(); num_kv_heads],
85            vq: vec![Vec::new(); num_kv_heads],
86            vs: vec![Vec::new(); num_kv_heads],
87            kcol: vec![Vec::new(); num_kv_heads],
88            vcol: vec![Vec::new(); num_kv_heads],
89            imp: Vec::new(),
90            seq_len: 0,
91            num_kv_heads,
92            head_dim,
93            linear_state: Vec::new(),
94            linear_scratch: Vec::new(),
95        }
96    }
97
98    /// Quantize one row against the per-channel field (empty col = 1);
99    /// `group` โ€” elements per scale (the whole row or KV_K_GROUP).
100    fn quant_row(row: &[f32], col: &[f32], q: &mut Vec<i8>, sc: &mut Vec<f32>,
101                 group: usize) {
102        let mut resid = vec![0.0f32; row.len()];
103        for (d, &x) in row.iter().enumerate() {
104            resid[d] = if col.is_empty() { x } else { x / col[d] };
105        }
106        for g0 in (0..row.len()).step_by(group) {
107            let g1 = (g0 + group).min(row.len());
108            let mut absmax = 0.0f32;
109            for &r in &resid[g0..g1] {
110                absmax = absmax.max(r.abs());
111            }
112            let s = (absmax / 127.0).max(1e-12);
113            sc.push(s);
114            for &r in &resid[g0..g1] {
115                q.push((r / s).round().clamp(-127.0, 127.0) as i8);
116            }
117        }
118    }
119
120    /// Freeze the 2f field: col = RMS of channels over stored rows, old
121    /// rows are requantized against the new field (once per conversation).
122    fn freeze_cols(&mut self) {
123        let hd = self.head_dim;
124        let ngk = hd.div_ceil(KV_K_GROUP);
125        for h in 0..self.num_kv_heads {
126            for (qv, sv, colv, group) in [
127                (&mut self.kq[h], &mut self.ks[h], &mut self.kcol[h], KV_K_GROUP),
128                (&mut self.vq[h], &mut self.vs[h], &mut self.vcol[h], hd),
129            ] {
130                let spp = if group == hd { 1 } else { ngk }; // scales per position
131                let n = sv.len() / spp;
132                if n == 0 {
133                    continue;
134                }
135                // Dequantize to f32, RMS over channels, requantize.
136                let mut rows = vec![0.0f32; n * hd];
137                for p in 0..n {
138                    for d in 0..hd {
139                        rows[p * hd + d] =
140                            qv[p * hd + d] as f32 * sv[p * spp + d / group];
141                    }
142                }
143                let mut col = vec![0.0f32; hd];
144                for p in 0..n {
145                    for d in 0..hd {
146                        col[d] += rows[p * hd + d] * rows[p * hd + d];
147                    }
148                }
149                for c in col.iter_mut() {
150                    *c = (*c / n as f32).sqrt().max(1e-6);
151                }
152                qv.clear();
153                sv.clear();
154                for p in 0..n {
155                    Self::quant_row(&rows[p * hd..(p + 1) * hd], &col, qv, sv, group);
156                }
157                *colv = col;
158            }
159        }
160    }
161
162    /// Append K/V for one position. `k_new`/`v_new` are
163    /// `[num_kv_heads ร— head_dim]`; heads with `alive[h] == false` are
164    /// skipped (their slices stay empty).
165    pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
166        debug_assert_eq!(k_new.len(), self.num_kv_heads * self.head_dim);
167        debug_assert_eq!(v_new.len(), self.num_kv_heads * self.head_dim);
168        // Freeze the 2f field AT THE START of append: only rows that
169        // survived verify are visible (a rejected lane-2 draft does not
170        // pollute the field โ€” found in review), and the threshold uses >=
171        // rather than strict equality (in small windows eviction may
172        // oscillate across 64).
173        if matches!(self.mode, KvMode::Q8 { .. })
174            && self.seq_len >= KV_COL_WARMUP
175            && self.kcol.iter().all(Vec::is_empty)
176            && self.vcol.iter().all(Vec::is_empty)
177        {
178            self.freeze_cols();
179        }
180        for h in 0..self.num_kv_heads {
181            if !alive.get(h).copied().unwrap_or(true) {
182                continue;
183            }
184            let s = h * self.head_dim;
185            if self.mode.quant_k() {
186                Self::quant_row(&k_new[s..s + self.head_dim],
187                                &self.kcol[h], &mut self.kq[h], &mut self.ks[h],
188                                KV_K_GROUP);
189            } else {
190                self.k[h].extend_from_slice(&k_new[s..s + self.head_dim]);
191            }
192            if self.mode.quant_v() {
193                Self::quant_row(&v_new[s..s + self.head_dim],
194                                &self.vcol[h], &mut self.vq[h], &mut self.vs[h],
195                                self.head_dim);
196            } else {
197                self.v[h].extend_from_slice(&v_new[s..s + self.head_dim]);
198            }
199        }
200        self.imp.push(0.0);
201        self.seq_len += 1;
202    }
203
204    /// Per-head attention over its own storage: the f32 branch is
205    /// bit-for-bit equal to attention_head() over slices; the q8 branch
206    /// computes score = s_kยทโŸจqโŠ™col_k, k_qโŸฉ and the weighted sum of V in i8
207    /// with f32 accumulation. Returns (output [head_dim], probs [stored]).
208    pub fn attend(&self, q: &[f32], kv_head: usize) -> (Vec<f32>, Vec<f32>) {
209        let hd = self.head_dim;
210        if self.mode == KvMode::F32 {
211            let stored = self.k[kv_head].len() / hd;
212            return crate::attention::attention_head(
213                q, &self.k[kv_head], &self.v[kv_head], hd, stored);
214        }
215        let stored = self.head_len(kv_head);
216        let scale = 1.0 / (hd as f32).sqrt();
217        let mut scores = vec![0.0f32; stored];
218        if self.mode.quant_k() {
219            let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
220            // q โŠ™ col_k โ€” once per call.
221            let kcol = &self.kcol[kv_head];
222            let mut qc = vec![0.0f32; hd];
223            for d in 0..hd {
224                qc[d] = if kcol.is_empty() { q[d] } else { q[d] * kcol[d] };
225            }
226            let ng = hd.div_ceil(KV_K_GROUP);
227            for p in 0..stored {
228                let row = &kq[p * hd..(p + 1) * hd];
229                let mut dot = 0.0f32;
230                for g in 0..ng {
231                    let g0 = g * KV_K_GROUP;
232                    let g1 = (g0 + KV_K_GROUP).min(hd);
233                    let mut gd = 0.0f32;
234                    for d in g0..g1 {
235                        gd += qc[d] * row[d] as f32;
236                    }
237                    dot += gd * ks[p * ng + g];
238                }
239                scores[p] = dot * scale;
240            }
241        } else {
242            let k = &self.k[kv_head];
243            for p in 0..stored {
244                let row = &k[p * hd..(p + 1) * hd];
245                let mut dot = 0.0f32;
246                for d in 0..hd {
247                    dot += q[d] * row[d];
248                }
249                scores[p] = dot * scale;
250            }
251        }
252        let max_score = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
253        let mut sum = 0.0f32;
254        for s in scores.iter_mut() {
255            *s = (*s - max_score).exp();
256            sum += *s;
257        }
258        if sum > 0.0 {
259            for s in scores.iter_mut() {
260                *s /= sum;
261            }
262        }
263        let mut acc = vec![0.0f32; hd];
264        if self.mode.quant_v() {
265            let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
266            for p in 0..stored {
267                let w = scores[p] * vs[p];
268                if w.abs() < 1e-12 {
269                    continue;
270                }
271                let row = &vq[p * hd..(p + 1) * hd];
272                for d in 0..hd {
273                    acc[d] += w * row[d] as f32;
274                }
275            }
276            let vcol = &self.vcol[kv_head];
277            if !vcol.is_empty() {
278                for d in 0..hd {
279                    acc[d] *= vcol[d];
280                }
281            }
282        } else {
283            let v = &self.v[kv_head];
284            for p in 0..stored {
285                let w = scores[p];
286                if w.abs() < 1e-12 {
287                    continue;
288                }
289                let row = &v[p * hd..(p + 1) * hd];
290                for d in 0..hd {
291                    acc[d] += w * row[d];
292                }
293            }
294        }
295        (acc, scores)
296    }
297
298    /// Roll back the last `n_drop` positions (speculative-decode reject).
299    pub fn truncate_last(&mut self, n_drop: usize) {
300        let d = n_drop.min(self.seq_len);
301        for h in 0..self.num_kv_heads {
302            let keep = self.k[h].len().saturating_sub(d * self.head_dim);
303            self.k[h].truncate(keep);
304            self.v[h].truncate(keep);
305            let ngk = self.head_dim.div_ceil(KV_K_GROUP);
306            let keep_q = self.kq[h].len().saturating_sub(d * self.head_dim);
307            self.kq[h].truncate(keep_q);
308            let keep_vq = self.vq[h].len().saturating_sub(d * self.head_dim);
309            self.vq[h].truncate(keep_vq);
310            let keep_ks = self.ks[h].len().saturating_sub(d * ngk);
311            self.ks[h].truncate(keep_ks);
312            let keep_vs = self.vs[h].len().saturating_sub(d);
313            self.vs[h].truncate(keep_vs);
314        }
315        self.imp.truncate(self.imp.len().saturating_sub(d));
316        self.seq_len -= d;
317    }
318
319    /// Accumulate attention mass per stored position (summed over heads).
320    pub fn accumulate_imp(&mut self, probs: &[f32]) {
321        for (dst, &p) in self.imp.iter_mut().zip(probs) {
322            *dst += p;
323        }
324    }
325
326    /// Contiguous keys of one head: `[stored_len ร— head_dim]`.
327    pub fn head_keys(&self, kv_head: usize) -> &[f32] {
328        &self.k[kv_head]
329    }
330
331    pub fn head_values(&self, kv_head: usize) -> &[f32] {
332        &self.v[kv_head]
333    }
334
335    /// Number of positions actually stored for a head (0 for dead heads).
336    pub fn head_len(&self, kv_head: usize) -> usize {
337        let ng = self.head_dim.div_ceil(KV_K_GROUP);
338        (self.k[kv_head].len() / self.head_dim)
339            .max(self.ks[kv_head].len() / ng)
340            .max(self.vs[kv_head].len())
341    }
342
343    /// Clear cache (e.g. on new conversation or task switch).
344    pub fn clear(&mut self) {
345        for h in 0..self.num_kv_heads {
346            self.k[h].clear();
347            self.v[h].clear();
348            self.kq[h].clear();
349            self.ks[h].clear();
350            self.vq[h].clear();
351            self.vs[h].clear();
352            self.kcol[h].clear();
353            self.vcol[h].clear();
354        }
355        self.imp.clear();
356        self.linear_state.clear();
357        self.linear_scratch.clear();
358        self.seq_len = 0;
359    }
360
361    /// Memory usage in bytes.
362    pub fn memory_bytes(&self) -> usize {
363        let floats: usize = self.k.iter().map(Vec::len).sum::<usize>()
364            + self.v.iter().map(Vec::len).sum::<usize>()
365            + self.ks.iter().map(Vec::len).sum::<usize>()
366            + self.vs.iter().map(Vec::len).sum::<usize>()
367            + self.kcol.iter().map(Vec::len).sum::<usize>()
368            + self.vcol.iter().map(Vec::len).sum::<usize>();
369        let bytes: usize = self.kq.iter().map(Vec::len).sum::<usize>()
370            + self.vq.iter().map(Vec::len).sum::<usize>();
371        floats * std::mem::size_of::<f32>() + bytes
372    }
373
374    /// Drop oldest positions, keeping the last `keep_last`.
375    fn evict(&mut self, keep_last: usize) {
376        if self.seq_len <= keep_last {
377            return;
378        }
379        let drop = self.seq_len - keep_last;
380        for h in 0..self.num_kv_heads {
381            // Dead heads store fewer positions; drop proportionally.
382            let stored = self.head_len(h);
383            let d = drop.min(stored);
384            let hd = self.head_dim;
385            fn drop_front<T>(v: &mut Vec<T>, n: usize) {
386                let n = n.min(v.len());
387                v.drain(..n);
388            }
389            drop_front(&mut self.k[h], d * hd);
390            drop_front(&mut self.v[h], d * hd);
391            drop_front(&mut self.kq[h], d * hd);
392            drop_front(&mut self.vq[h], d * hd);
393            drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
394            drop_front(&mut self.vs[h], d);
395        }
396        let d = drop.min(self.imp.len());
397        self.imp.drain(..d);
398        self.seq_len = keep_last;
399    }
400
401    /// Born eviction: keep `sink` earliest positions (attention sinks),
402    /// the `recent` latest, and fill the rest of the `keep_last` budget
403    /// with the positions carrying the highest accumulated attention
404    /// mass (vmfcore: PPL 8.342 vs 8.687 for recency-only, full 8.295).
405    fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
406        let stored = self.imp.len();
407        if stored <= keep_last {
408            return;
409        }
410        // Budget discipline: sinks first, recents next, both clamped so
411        // the total never exceeds keep_last.
412        let sink_n = sink.min(keep_last);
413        let recent_n = recent.min(keep_last - sink_n);
414        let mut keep = vec![false; stored];
415        for k in keep.iter_mut().take(sink_n) {
416            *k = true;
417        }
418        for k in keep.iter_mut().skip(stored.saturating_sub(recent_n)) {
419            *k = true;
420        }
421        let mut budget = keep_last.saturating_sub(keep.iter().filter(|&&x| x).count());
422        // Highest accumulated mass first among the middle positions.
423        let mut order: Vec<usize> = (0..stored).filter(|&i| !keep[i]).collect();
424        order.sort_by(|&a, &b| {
425            self.imp[b].partial_cmp(&self.imp[a]).unwrap_or(std::cmp::Ordering::Equal)
426        });
427        for i in order {
428            if budget == 0 {
429                break;
430            }
431            keep[i] = true;
432            budget -= 1;
433        }
434
435        let kept: Vec<usize> = (0..stored).filter(|&i| keep[i]).collect();
436        let hd = self.head_dim;
437        fn gather<T: Copy>(src: &[T], kept: &[usize], step: usize) -> Vec<T> {
438            let mut out = Vec::with_capacity(kept.len() * step);
439            for &i in kept {
440                out.extend_from_slice(&src[i * step..(i + 1) * step]);
441            }
442            out
443        }
444        // Each storage is gathered INDEPENDENTLY: in mixed modes
445        // (q8k/q8v) K and V live in different storages โ€” the paired branch
446        // panicked (q8v) or silently left V uncompressed (q8k);
447        // found by adversarial review, closed by regression tests.
448        for h in 0..self.num_kv_heads {
449            if !self.k[h].is_empty() {
450                self.k[h] = gather(&self.k[h], &kept, hd);
451            }
452            if !self.v[h].is_empty() {
453                self.v[h] = gather(&self.v[h], &kept, hd);
454            }
455            if !self.kq[h].is_empty() {
456                self.kq[h] = gather(&self.kq[h], &kept, hd);
457                self.ks[h] = gather(&self.ks[h], &kept, hd.div_ceil(KV_K_GROUP));
458            }
459            if !self.vq[h].is_empty() {
460                self.vq[h] = gather(&self.vq[h], &kept, hd);
461                self.vs[h] = gather(&self.vs[h], &kept, 1);
462            }
463        }
464        self.imp = kept.iter().map(|&i| self.imp[i]).collect();
465        self.seq_len = kept.len();
466    }
467}
468
469/// Eviction policy for a bounded cache.
470#[derive(Debug, Clone, Copy, PartialEq, Eq)]
471pub enum EvictionPolicy {
472    /// Sliding window: keep only the most recent positions.
473    Recent,
474    /// Born rule: sinks + recents + top accumulated attention mass.
475    Born { sink: usize },
476}
477
478/// Full KV cache for all layers.
479#[derive(Debug)]
480pub struct KvCache {
481    pub layers: Vec<LayerKvCache>,
482    pub max_seq_len: usize,
483    pub policy: EvictionPolicy,
484}
485
486impl KvCache {
487    pub fn new(num_layers: usize, num_kv_heads: usize, head_dim: usize, max_seq_len: usize) -> Self {
488        let layers = (0..num_layers)
489            .map(|_| LayerKvCache::new(num_kv_heads, head_dim))
490            .collect();
491        Self {
492            layers,
493            max_seq_len,
494            policy: EvictionPolicy::Born { sink: 4 },
495        }
496    }
497
498    pub fn clear(&mut self) {
499        for layer in &mut self.layers {
500            layer.clear();
501        }
502    }
503
504    pub fn total_memory_bytes(&self) -> usize {
505        self.layers.iter().map(|l| l.memory_bytes()).sum()
506    }
507
508    /// Current sequence length (max across layers โ€” dead layers may lag).
509    pub fn seq_len(&self) -> usize {
510        self.layers.iter().map(|l| l.seq_len).max().unwrap_or(0)
511    }
512
513    pub fn needs_eviction(&self) -> bool {
514        self.seq_len() >= self.max_seq_len
515    }
516
517    /// Evict down to `keep_last` positions according to the policy.
518    pub fn evict(&mut self, keep_last: usize) {
519        match self.policy {
520            EvictionPolicy::Recent => {
521                for layer in &mut self.layers {
522                    layer.evict(keep_last);
523                }
524            }
525            EvictionPolicy::Born { sink } => {
526                let recent = (keep_last / 2).max(1);
527                for layer in &mut self.layers {
528                    layer.evict_born(keep_last, sink, recent);
529                }
530            }
531        }
532    }
533}
534
535#[cfg(test)]
536mod tests {
537    use super::*;
538
539    #[test]
540    fn append_tracks_seq_len_and_layout() {
541        let mut cache = LayerKvCache::new(4, 8);
542        cache.mode = KvMode::F32;
543        assert_eq!(cache.seq_len, 0);
544
545        let k: Vec<f32> = (0..32).map(|i| i as f32).collect();
546        let v = vec![2.0f32; 32];
547        cache.append(&k, &v, &[true; 4]);
548
549        assert_eq!(cache.seq_len, 1);
550        assert_eq!(cache.head_len(0), 1);
551        // head 1 slice is contiguous and equals its part of k_new
552        assert_eq!(cache.head_keys(1), &k[8..16]);
553        assert_eq!(cache.memory_bytes(), 256);
554    }
555
556    #[test]
557    fn dead_head_stores_nothing() {
558        let mut cache = LayerKvCache::new(2, 4);
559        cache.mode = KvMode::F32;
560        let k = vec![1.0f32; 8];
561        let v = vec![2.0f32; 8];
562        cache.append(&k, &v, &[true, false]);
563        cache.append(&k, &v, &[true, false]);
564
565        assert_eq!(cache.seq_len, 2);
566        assert_eq!(cache.head_len(0), 2);
567        assert_eq!(cache.head_len(1), 0, "dead head must not store KV");
568        assert_eq!(cache.memory_bytes(), 2 * 2 * 4 * 4);
569    }
570
571    #[test]
572    fn eviction_keeps_recent() {
573        let mut cache = KvCache::new(2, 4, 8, 10);
574        cache.policy = EvictionPolicy::Recent;
575        for l in &mut cache.layers { l.mode = KvMode::F32; }
576        let k = vec![1.0f32; 32];
577        let v = vec![2.0f32; 32];
578        for _ in 0..8 {
579            for layer in &mut cache.layers {
580                layer.append(&k, &v, &[true; 4]);
581            }
582        }
583        assert_eq!(cache.seq_len(), 8);
584        assert!(!cache.needs_eviction());
585
586        cache.evict(4);
587        assert_eq!(cache.seq_len(), 4);
588        assert_eq!(cache.layers[0].head_len(0), 4);
589    }
590
591    #[test]
592    fn truncate_rolls_back_speculative_positions() {
593        let mut cache = LayerKvCache::new(2, 4);
594        cache.mode = KvMode::F32;
595        for pos in 0..5 {
596            let k = vec![pos as f32; 8];
597            let v = vec![pos as f32; 8];
598            cache.append(&k, &v, &[true; 2]);
599        }
600        cache.truncate_last(2);
601        assert_eq!(cache.seq_len, 3);
602        assert_eq!(cache.head_len(0), 3);
603        assert_eq!(cache.head_keys(0)[2 * 4], 2.0, "position 2 survives");
604    }
605
606    /// q8_2f-attend โ‰ˆ f32-attend: 100 positions (crosses the field freeze
607    /// at the 64th), pseudo-random vectors, relative tolerance of the
608    /// int8 grid. Plus rollback and Born eviction on the q8 storage.
609    #[test]
610    fn q8_attend_matches_f32_within_grid() {
611        let (heads, hd) = (2, 32);
612        let mut f = LayerKvCache::new(heads, hd);
613        f.mode = KvMode::F32;
614        let mut q8 = LayerKvCache::new(heads, hd);
615        q8.mode = KvMode::Q8 { k: true, v: true };
616
617        let synth = |p: usize, salt: usize| -> Vec<f32> {
618            (0..heads * hd)
619                .map(|i| {
620                    let x = ((i * 31 + p * 17 + salt * 7 + 3) % 97) as f32 / 97.0 - 0.5;
621                    // channel structure: even channels ร—4 (checks the 2f field)
622                    if i % 2 == 0 { x * 4.0 } else { x * 0.25 }
623                })
624                .collect()
625        };
626        for p in 0..100 {
627            let k = synth(p, 1);
628            let v = synth(p, 2);
629            f.append(&k, &v, &[true; 2]);
630            q8.append(&k, &v, &[true; 2]);
631        }
632        let q: Vec<f32> = (0..hd).map(|i| ((i * 13 + 5) % 89) as f32 / 89.0 - 0.5).collect();
633        for g in 0..heads {
634            let (of, pf) = f.attend(&q, g);
635            let (o8, p8) = q8.attend(&q, g);
636            let scale = of.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-6);
637            for d in 0..hd {
638                assert!(
639                    (of[d] - o8[d]).abs() <= scale * 0.03 + 1e-3,
640                    "g{g} d{d}: f32 {} vs q8 {}", of[d], o8[d]
641                );
642            }
643            for p in 0..100 {
644                assert!((pf[p] - p8[p]).abs() < 0.02, "prob p{p}");
645            }
646        }
647        // rollback + eviction live on the q8 storage
648        q8.truncate_last(30);
649        assert_eq!(q8.head_len(0), 70);
650        let imp: Vec<f32> = (0..70).map(|i| i as f32).collect();
651        q8.accumulate_imp(&imp);
652        q8.evict_born(20, 2, 8);
653        assert_eq!(q8.head_len(0), 20);
654        let (o, _) = q8.attend(&q, 0);
655        assert!(o.iter().all(|x| x.is_finite()));
656        // memory: q8 โ‰ˆ 1 byte/element + scale per row (vs 4 for f32)
657        assert!(q8.memory_bytes() * 3 < f.memory_bytes());
658    }
659
660    /// Review regression: Born eviction in MIXED modes. q8v used to
661    /// panic (gather over an empty v[h]), q8k silently left raw V
662    /// uncompressed (stale rows under kept keys + memory leak).
663    #[test]
664    fn born_eviction_mixed_modes_stay_consistent() {
665        for (mk, mv) in [(false, true), (true, false)] {
666            let mut c = LayerKvCache::new(1, 4);
667            c.mode = KvMode::Q8 { k: mk, v: mv };
668            for p in 0..80 {
669                let k = vec![p as f32 * 0.01; 4];
670                let v = vec![p as f32; 4];
671                c.append(&k, &v, &[true]);
672            }
673            let imp: Vec<f32> = (0..80).map(|i| i as f32).collect();
674            c.accumulate_imp(&imp);
675            let before = c.memory_bytes();
676            c.evict_born(20, 4, 8); // q8v: used to panic here
677            assert_eq!(c.head_len(0), 20, "k={mk} v={mv}");
678            assert!(c.memory_bytes() < before / 2,
679                    "memory must shrink (k={mk} v={mv})");
680            // V rows match the kept set: the heaviest positions
681            // (tail 60..79) must be present in the attend output.
682            let (out, _) = c.attend(&[1.0, 1.0, 1.0, 1.0], 0);
683            assert!(out[0] > 30.0,
684                    "V from the kept tail, not the stale head (k={mk} v={mv}, out {})",
685                    out[0]);
686        }
687    }
688
689    #[test]
690    fn born_eviction_keeps_high_mass_position() {
691        let mut cache = KvCache::new(1, 1, 2, 16);
692        cache.policy = EvictionPolicy::Born { sink: 1 };
693        for l in &mut cache.layers { l.mode = KvMode::F32; }
694        let layer = &mut cache.layers[0];
695        // 8 positions; keys carry the position index so we can verify
696        // exactly which positions survive the gather.
697        for pos in 0..8 {
698            let k = vec![pos as f32; 2];
699            let v = vec![pos as f32 + 100.0; 2];
700            layer.append(&k, &v, &[true]);
701        }
702        // Position 3 carries the most attention mass (Born importance).
703        let mut imp = vec![0.05f32; 8];
704        imp[3] = 5.0;
705        layer.accumulate_imp(&imp);
706
707        cache.evict(4); // sink 1 + recent 2 + 1 top-mass slot
708        let layer = &cache.layers[0];
709        assert_eq!(layer.seq_len, 4);
710        let kept_keys: Vec<f32> = (0..4).map(|i| layer.head_keys(0)[i * 2]).collect();
711        assert_eq!(
712            kept_keys,
713            vec![0.0, 3.0, 6.0, 7.0],
714            "kept = sink(0) + Born-top(3) + recent(6,7)"
715        );
716        // imp stays aligned with the gathered positions.
717        assert_eq!(layer.head_len(0), 4);
718    }
719}