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 scale 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 {
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
41/// Positions before freezing the per-channel field (2f): before — col ≡ 1,
42/// after — col = RMS over channels of the stored rows, old rows are requantized.
43const KV_COL_WARMUP: usize = 64;
44
45/// K-rows are quantized in groups of 32 channels (scale per group):
46/// attention logits are sensitive to the dot-product error, per-group scales
47/// localize it along RoPE bands (35B: +4.6% PPL with a per-row scale
48/// → target <1% with a per-group one). V — per-row scale (measured +0.56%).
49const KV_K_GROUP: usize = 32;
50
51/// Per-layer O(1) Nyström attention state (runtime `attn_type`
52/// override — spec §7 presence-driven pattern, no format change).
53///
54/// Collecting: the prompt pass still runs EXACT cache attention (the
55/// prefill outputs feed the residual stream, so they cannot be
56/// deferred) while the per-position rotated queries are buffered;
57/// `o1_seal()` then freezes landmarks + M from the full prompt, replays
58/// it into per-KV-group streaming states, and DROPS the full KV.
59/// Sealed: decode replaces cache attention with
60/// `NystromState::step_group()`.
61#[derive(Debug, Clone)]
62pub enum O1State {
63    Collecting {
64        m: usize,
65        w: usize,
66        sink: usize,
67        rect: crate::nystrom::O1Rect,
68        /// Optional completed-row barrier. `None` means the caller will
69        /// request a full-prompt seal; `Some(B)` keeps a short prompt exact
70        /// until the first skeleton-safe boundary B.
71        seal_at: Option<usize>,
72        /// Rotated post-norm queries, `[pos × num_heads × head_dim]`.
73        q_buf: Vec<f32>,
74    },
75    /// One state per KV GROUP, each holding its group's Q heads. The
76    /// exact window / sinks / K̃ are stored once per group (every Q head
77    /// of the group reads the same k/v rows); only the far field, Q̃ and
78    /// M — the query-dependent pieces — stay per Q head. See
79    /// `NystromState` for which piece is which and why.
80    Sealed {
81        groups: Vec<crate::nystrom::NystromState>,
82    },
83}
84
85/// KV cache for a single layer, head-major.
86#[derive(Debug, Clone)]
87pub struct LayerKvCache {
88    pub mode: KvMode,
89    /// Per-KV-head keys: `k[h]` is `[seq_len × head_dim]` (empty if head is dead).
90    k: Vec<Vec<f32>>,
91    /// Per-KV-head values, same layout.
92    v: Vec<Vec<f32>>,
93    /// q8 storage (mode == Q8_2F): int8 rows + f32 scale per row.
94    kq: Vec<Vec<i8>>,
95    ks: Vec<Vec<f32>>,
96    vq: Vec<Vec<i8>>,
97    vs: Vec<Vec<f32>>,
98    /// Per-channel scale fields per head [head_dim]; empty until frozen.
99    kcol: Vec<Vec<f32>>,
100    vcol: Vec<Vec<f32>>,
101    /// Accumulated attention mass per stored position: importance of a
102    /// position is how much probability mass reads it.
103    imp: Vec<f32>,
104    /// Positions appended so far (grows once per token, dead heads included).
105    pub seq_len: usize,
106    pub num_kv_heads: usize,
107    pub head_dim: usize,
108    /// Linear-core recurrent state S (vmf_phase), f64; empty on full layers.
109    pub linear_state: Vec<f32>,
110    /// Tentative lane-2 state during speculative verify.
111    pub linear_scratch: Vec<f32>,
112    /// The legacy cache wire has no operator tag. Delta operators therefore
113    /// bind this layer to a fail-closed boundary until a versioned wire
114    /// schema can carry the linear-core identity.
115    linear_wire_allowed: bool,
116    /// O(1) Nyström override (None = plain cache attention).
117    pub o1: Option<O1State>,
118    /// A deferred seal failure is terminal for the current request. The
119    /// attention functions return only a hidden vector, so the pipeline
120    /// consumes this side channel at its next forward boundary.
121    o1_error: Option<String>,
122    /// Set when a collecting state actually becomes sealed. Pipeline owns
123    /// the epoch bump and consumes this bit after a complete forward.
124    o1_transitioned: bool,
125    /// Learned per-Q-head attention-sink logits of this layer (gpt-oss /
126    /// MiMo-V2 `self_attn.sinks`, one f32 per Q head). The sink is an
127    /// extra softmax column with no value: it joins the max and the
128    /// denominator of every head's softmax and so lets a head attend to
129    /// "nothing". These are WEIGHTS, not sequence state — `clear()` and
130    /// the wire import keep them. None = an ordinary softmax.
131    pub sinks: Option<Vec<f32>>,
132    /// Natively bounded anchor (`swa_sink_v1`): a fixed-size ring
133    /// installed from the header at load, never per prompt. A layer that
134    /// carries it stores NOTHING per position (`k`/`v` stay empty).
135    pub bounded: Option<crate::bounded::BoundedState>,
136    /// Which state record this layer exchanges on the wire (v2).
137    pub wire_kind: WireKind,
138    /// This layer's index in the stack (the wire header names it).
139    pub wire_layer: u32,
140    /// hash64 of the model's operator identity
141    /// (`ModelArch::linear_core_identity` JSON); 0 = no operator record.
142    pub wire_identity: u64,
143}
144
145/// State-record kind of the versioned cache wire (`export_wire` v2).
146#[derive(Debug, Clone, Copy, PartialEq, Eq)]
147#[repr(u8)]
148pub enum WireKind {
149    /// Per-position K/V (+ importance) — the legacy body.
150    Full = 0,
151    /// Recurrent state vector (S + conv ring), f32.
152    Linear = 1,
153    /// Bounded anchor: insert counter + ring K/V.
154    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
168/// Magic of the versioned state wire.
169pub const WIRE_MAGIC: &[u8; 4] = b"CMFS";
170/// Current wire version.
171pub 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    // ── Natively bounded anchor (swa_sink_v1) ──
204
205    /// Give this layer its fixed-size ring (`[kvh][window][hd]` K and V).
206    /// Called once at load from the header; the record is zeroed on
207    /// `clear()` and never reallocated.
208    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    /// One position of the bounded operator: insert the raw `k, v`
218    /// (`[kvh][hd]`) into slot `t mod W`, then attend every Q head of
219    /// `q` (`[nh][hd]`, raw) over sinks ∪ window into `out` (`[nh][hd]`).
220    /// Nothing is appended per position.
221    #[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        // Honest context depth for the memory/seq report — nothing is
240        // stored per position.
241        self.seq_len += 1;
242    }
243
244    /// Bytes of the bounded ring (0 on other layers).
245    pub fn bounded_state_bytes(&self) -> usize {
246        self.bounded.as_ref().map(|b| b.state_bytes()).unwrap_or(0)
247    }
248
249    /// Bit-for-bit copy of the ring for speculation (None on other layers).
250    pub fn bounded_snapshot(&self) -> Option<crate::bounded::BoundedSnapshot> {
251        self.bounded.as_ref().map(|b| b.snapshot())
252    }
253
254    /// Restore a ring snapshot taken on this layer; `seq_len` follows the
255    /// restored insert counter.
256    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    /// Bind this layer to the legacy cache-wire policy. Delta layers must
264    /// refuse untagged state exchange rather than risk a plausible additive
265    /// interpretation on the peer.
266    pub fn set_linear_wire_allowed(&mut self, allowed: bool) {
267        self.linear_wire_allowed = allowed;
268    }
269
270    /// Discard tentative recurrent state after a speculative rejection or
271    /// any other path that abandons the lane-2 result.
272    pub fn discard_linear_scratch(&mut self) {
273        self.linear_scratch.clear();
274    }
275
276    /// Per-KV-head stored keys `[seq_len × head_dim]` (GPU token graph sync).
277    pub fn k_heads(&self) -> &[Vec<f32>] {
278        &self.k
279    }
280    /// Per-KV-head stored values `[seq_len × head_dim]`.
281    pub fn v_heads(&self) -> &[Vec<f32>] {
282        &self.v
283    }
284
285    // ── O(1) Nyström override ──
286
287    /// Arm query collection for a fresh prompt pass (a cleared cache).
288    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    /// Arm query collection with an optional completed-row seal barrier.
293    /// The barrier is deliberately part of the existing collecting state:
294    /// no second history or scheduler is introduced for short prompts.
295    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    /// Record one position's rotated queries (`[num_heads × head_dim]`)
316    /// during the exact prompt pass. No-op unless collecting — the hook
317    /// sits inside qwen_attention so every prefill flavor (sequential,
318    /// batched) feeds the same trace.
319    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    /// Pending completed-row barrier, if any. A plain full-prompt seal has
330    /// no barrier until the caller asks to seal.
331    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    /// Whether a batch of `count` exact rows would cross the deferred
339    /// boundary. This lets batched/pair callers split before row B rather
340    /// than appending exact KV past the point where conversion is required.
341    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        // Error observation is deliberately non-consuming.  The error is the
358        // append guard for this request; removing it would let a caller that
359        // ignored the returned Err resume ordinary KV growth after the
360        // bounded transition dropped its overlay.  `clear()` is the explicit
361        // reset boundary that clears the latch.
362        self.o1_error.clone()
363    }
364
365    /// Abort a malformed deferred transition after attention has already
366    /// produced its current row. Clearing the exact storage and dropping
367    /// the overlay makes the state unrecoverable by continued decode; the
368    /// pipeline then routes through its normal graph/cancel cleanup.
369    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    /// Freeze the prompt into per-KV-group Nyström states and drop this
386    /// layer's full KV. Returns false while a short collecting layer is
387    /// below its deferred boundary; malformed prerequisites abort the
388    /// layer instead of silently resuming exact KV growth. The seal needs
389    /// f32 KV rows (`CMF_KV=q8` stores int8), every group densely stored, a
390    /// full q trace, and a GQA fan-out that actually divides.
391    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    /// Checked seal implementation. Validation happens while the collecting
403    /// state and full KV are still intact; only a valid completed boundary
404    /// is allowed to destructively convert them.
405    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        // Idempotent: sealing a sealed (or plain) layer must not disturb its
410        // state, and a plain layer is not an O(1) participant.
411        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        // Query trace is position-major; the state wants each head's
488        // queries contiguous, so transpose one group at a time.
489        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        // The states now carry everything decode needs — release the
507        // O(context) storage (this is the memory claim, not a cosmetic).
508        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    /// One decode step on a sealed layer: per KV group, insert the
519    /// group's fresh (k, v) ONCE and read every Q head's attention
520    /// output. Returns `[num_heads × head_dim]`. Head h belongs to group
521    /// h/hpk, so a group's Q heads are contiguous in `q_all`/`out` —
522    /// same math as the shared KV row the exact path appends once.
523    /// Device views of the sealed o1 groups, or None when o1 is not
524    /// sealed on this layer (or any group is in the degenerate
525    /// exact-only mode the GPU path does not carry).
526    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        // Track the true context depth for the honest memory/seq report
561        // (nothing is stored per position — the state is O(1)).
562        self.seq_len += 1;
563        out
564    }
565
566    /// Bytes held by the O(1) override (query trace while collecting,
567    /// per-KV-group states once sealed).
568    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    /// Quantize one row against the per-channel field (empty col = 1);
577    /// `group` — elements per scale (the whole row or KV_K_GROUP).
578    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    /// Freeze the 2f field: col = RMS of channels over stored rows, old
598    /// rows are requantized against the new field (once per conversation).
599    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 }; // scales per position
613                let n = sv.len() / spp;
614                if n == 0 {
615                    continue;
616                }
617                // Dequantize to f32, RMS over channels, requantize.
618                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    /// Append K/V for one position. `k_new`/`v_new` are
644    /// `[num_kv_heads × head_dim]`; heads with `alive[h] == false` are
645    /// skipped (their slices stay empty).
646    pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
647        // A failed bounded transition is terminal until the sequence is
648        // cleared. Do not let an ignored boolean/result resume plain KV
649        // growth after the O(1) collector has aborted.
650        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        // Freeze the 2f field AT THE START of append: only rows that
656        // survived verify are visible (a rejected lane-2 draft does not
657        // pollute the field — found in review), and the threshold uses >=
658        // rather than strict equality (in small windows eviction may
659        // oscillate across 64).
660        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    /// Per-head attention over its own storage: the f32 branch is
700    /// bit-for-bit equal to attention_head() over slices; the q8 branch
701    /// computes score = s_k·⟨q⊙col_k, k_q⟩ and the weighted sum of V in i8
702    /// with f32 accumulation. Returns (output [head_dim], probs [stored]).
703    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            // q ⊙ col_k — once per call.
721            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                // SAFETY: i8 and u8 share layout; dot_i8_f32 reads the
734                // bytes back as i8.
735                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    /// Grouped GQA attention: all Q-heads of one KV group in a single
794    /// pass over the stored K rows and a single pass over the V rows
795    /// (per-head `attend` re-read the shared group storage
796    /// heads_per_kv times — roadmap §3 P1). Per-head score order,
797    /// softmax and V accumulation are IDENTICAL to `attend`, so each
798    /// head's output is bit-for-bit the same.
799    ///
800    /// `q_group`: `[n_heads_in_group × head_dim]` (global head order);
801    /// `out`: same shape; `imp_acc[0..stored]` accumulates the probabilities
802    /// of every head (attention importance), matching the caller's former loop.
803    /// `scale` is the score scale (1/√hd unless the arch overrides);
804    /// `first` is the earliest visible position — sliding-window layers
805    /// pass `stored − window` so older rows get zero probability.
806    /// `sinks` holds one learned sink logit per head of `q_group` (the
807    /// caller slices `self.sinks` to the group's heads) or is empty for
808    /// an ordinary softmax.
809    #[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    /// `attend_group` over the first `upto` stored rows only — what the
835    /// same call saw when the cache held exactly `upto` rows. A prefill
836    /// chunk appends all its rows first and then attends every position
837    /// in parallel; position `i` passes `upto = s0 + i + 1`, which makes
838    /// its result bit-identical to the sequential append-then-attend.
839    ///
840    /// Only the visible rows `[first, stored)` are scored, so a
841    /// sliding-window decode step costs O(window), not O(context). The
842    /// rows before `first` get probability exactly 0 — what the former
843    /// −inf-filled score row produced (exp(−inf) = 0 adds nothing to the
844    /// max, the sum, V or the importance), so the result is bit-identical
845    /// to scoring the whole row.
846    ///
847    /// Learned sinks (`sinks` non-empty, one per head): the sink logit
848    /// joins the softmax max and adds exp(sink − max) to the denominator;
849    /// it has no value row. Identical to appending a value-less column,
850    /// which is how gpt-oss and MiMo-V2 define it.
851    #[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        // Visible span: row p ∈ [first, stored) lives at column p − first.
884        let span = stored - first;
885
886        thread_local! {
887            /// scores [nheads × span] — reused across layers/tokens.
888            static GQA_SCORES: std::cell::RefCell<Vec<f32>> =
889                const { std::cell::RefCell::new(Vec::new()) };
890            /// q ⊙ col_k per head (q8 K mode).
891            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            // ── score pass: each stored K row is read ONCE for all heads.
900            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                        // SAFETY: i8 and u8 share layout; dot_i8_f32 reads
916                        // the bytes back as i8.
917                        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            // Gemma-2 attention-logit soft-capping: tanh-squash the
945            // COMPUTED scores before the softmax (every scored row is in
946            // the window; a learned sink is not a score and is not capped).
947            if softcap > 0.0 {
948                for v in scores.iter_mut() {
949                    *v = softcap * (*v / softcap).tanh();
950                }
951            }
952
953            // ── per-head softmax (identical to attend / attention_head;
954            // a sink joins the max and the denominator, never the rows).
955            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            // ── value pass: each stored V row is read ONCE for all heads.
979            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            // ── Attention-importance accumulation (Σ probs over heads), same
1015            // head order as the caller's former per-head loop. Rows before
1016            // `first` carry probability 0 and are left untouched.
1017            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    /// Batched causal attend for a prefill chunk (macOS/AArch64): the
1030    /// cache already holds every chunk row (`s0` old + `b` new). Per
1031    /// Q-head the scores GEMM `Q·Kᵀ` rides the AMX, the causal softmax
1032    /// zeroes the not-yet-visible tail so the `P·V` GEMM needs no
1033    /// mask, and attention importance takes the masked column sums. Same
1034    /// math as the per-position attend; summation order differs
1035    /// (tolerance-class, like the projection GEMMs).
1036    #[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                // Method receiver keeps the closure capturing &SendPtr
1058                // (2021 disjoint capture would grab the raw field).
1059                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        // The portable NEON GEMM pays dearly for the gathered Bᵀ loads
1067        // of the scores multiply — pack Kᵀ once per (group, chunk) and
1068        // hand it the sequential-B fast path instead. Accelerate keeps
1069        // the no-copy transposed call.
1070        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            // The whole KV-group attends in one GEMM pair: the group's
1078            // Q-heads stack head-major into one tall panel [hpk·b, hd]
1079            // (row hl·b + bi), so each layer costs 2 sgemm calls per
1080            // group instead of 2 per head — fat M keeps the AMX fed.
1081            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                    // Accelerate threads its own GEMM; the NEON kernel
1105                    // splits the m rows across the pool instead.
1106                    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                            // SAFETY: workers write disjoint score rows.
1112                            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                // Causal softmax, row-parallel (rows are disjoint).
1146                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                        // Sliding-window layers see only the last W of
1151                        // the causal range; the zeroed head contributes
1152                        // nothing to P·V or attention importance.
1153                        let lo = window.map(|w| allowed.saturating_sub(w)).unwrap_or(0);
1154                        // SAFETY: workers cover disjoint row ranges.
1155                        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                // Attention importance: masked column sums (probs of the
1166                // zeroed tail contribute nothing, same as the CPU
1167                // per-position accumulate).
1168                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                            // SAFETY: workers write disjoint output rows.
1181                            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    /// Roll back the last `n_drop` positions (speculative-decode reject).
1226    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            // The ring rolls back through its undo rows; nothing is
1231            // stored per position, so there is nothing else to drop.
1232            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    /// Accumulate attention mass per stored position (summed over heads).
1260    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    /// Contiguous keys of one head: `[stored_len × head_dim]`.
1267    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    /// Number of positions actually stored for a head (0 for dead heads).
1276    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    /// Clear cache (e.g. on new conversation or task switch).
1284    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        // Fresh conversation → the pipeline re-arms collection if the
1299        // layer is o1-flagged (landmarks are per-prompt, never reused).
1300        self.o1 = None;
1301        self.o1_error = None;
1302        self.o1_transitioned = false;
1303        // The bounded ring is zeroed in place: its size is a property of
1304        // the file, not of the conversation.
1305        if let Some(b) = self.bounded.as_mut() {
1306            b.clear();
1307        }
1308        self.seq_len = 0;
1309    }
1310
1311    /// Serialize this layer's state for the wire (versioned, v2): a
1312    /// fixed header `{magic "CMFS", version, operator identity hash64,
1313    /// layer, kind, f16 flag, position}` followed by one record whose
1314    /// shape the kind fixes — per-position K/V (+ importance) for a full
1315    /// layer, the recurrent vector for a linear layer, the insert counter
1316    /// + ring K/V for a bounded anchor. `f16` halves the K/V payloads and
1317    /// is the caller's explicit choice, exactly like the hidden-state
1318    /// wire; recurrent vectors stay f32 whatever the wire dtype (they are
1319    /// the ONLY state a linear layer has — rounding them rounds the whole
1320    /// history).
1321    ///
1322    /// REFUSES rather than travelling half-complete. A cache carrying
1323    /// frozen columns, a Nyström overlay or q8 storage holds state this
1324    /// format does not describe, and shipping the rest would land a
1325    /// plausible-looking cache that answers differently — the failure
1326    /// mode this whole format exists to avoid.
1327    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        // Frozen columns only exist under q8 storage, which is refused
1337        // above. If one shows up under an F32 cache the format is lying
1338        // about something and the transfer must not proceed.
1339        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    /// The per-position record (the whole legacy wire): f16 flag,
1395    /// seq_len, geometry, recurrent vector, importance, per-head K/V.
1396    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        // Attention importance is ordinary state: every attention call
1404        // accumulates it and eviction reads it. Leaving it behind would
1405        // hand the far side a cache that forgets the RIGHT positions
1406        // later — a divergence that shows up only under pressure.
1407        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    /// Install a peer's state over this layer. The geometry must match the
1434    /// model both sides hold — it is checked, not assumed. Accepts the
1435    /// versioned wire (magic "CMFS") and, for full/linear layers, the old
1436    /// unversioned per-position body.
1437    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        // Legacy (unversioned) wire: no operator tag travels with it.
1442        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                // The undo rows describe inserts this side never made.
1593                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    /// Empty every per-position store (keeps the recurrent vector and
1610    /// the bounded ring untouched).
1611    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    /// Parse the per-position record (legacy wire body); returns the
1630    /// bytes consumed.
1631    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            // O(1) recurrent state of linear-core layers (vmf_phase/GDN):
1713            // constant in context, but real memory — the honest "KV+state"
1714            // line must count it (a pure-linear model reported 0 before).
1715            + self.linear_state.len() * std::mem::size_of::<f32>()
1716            // O(1) Nyström state (window + sinks + skeleton) — same
1717            // discipline: constant in context, but real memory.
1718            + self.o1_memory_bytes()
1719            // Natively bounded anchor: the ring is the whole state of the
1720            // layer, fixed by the header.
1721            + self.bounded_state_bytes()
1722    }
1723
1724    /// Drop oldest positions, keeping the last `keep_last`.
1725    fn evict(&mut self, keep_last: usize) {
1726        // A collecting o1 layer owns a still-needed exact prefix and query
1727        // trace.  Evicting it would lower the effective seal boundary while
1728        // leaving q_buf untouched, so conversion could never match its KV
1729        // rows.  A sealed layer stores nothing per position — the Nyström
1730        // state IS the eviction policy; resetting seq_len here would lie
1731        // about the context depth.  Both states therefore bypass ordinary
1732        // eviction until the transition or explicit reset completes.  A
1733        // bounded anchor likewise: the ring evicts itself every token.
1734        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            // Dead heads store fewer positions; drop proportionally.
1740            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    /// Mass-based eviction: keep `sink` earliest positions (attention sinks),
1760    /// the `recent` latest, and fill the rest of the `keep_last` budget
1761    /// with the positions carrying the highest accumulated attention
1762    /// mass (vmfcore: PPL 8.342 vs 8.687 for recency-only, full 8.295).
1763    fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
1764        if self.o1.is_some() || self.bounded.is_some() {
1765            // See evict(): collecting must retain the exact prefix as well as
1766            // sealed O(1) state must retain its own bounded representation;
1767            // a bounded anchor's ring is its own eviction.
1768            return;
1769        }
1770        let stored = self.imp.len();
1771        if stored <= keep_last {
1772            return;
1773        }
1774        // Budget discipline: sinks first, recents next, both clamped so
1775        // the total never exceeds keep_last.
1776        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        // Highest accumulated mass first among the middle positions.
1787        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        // Each storage is gathered INDEPENDENTLY: in mixed modes
1811        // (q8k/q8v) K and V live in different storages — the paired branch
1812        // panicked (q8v) or silently left V uncompressed (q8k);
1813        // found by adversarial review, closed by regression tests.
1814        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/// Eviction policy for a bounded cache.
1836#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1837pub enum EvictionPolicy {
1838    /// Sliding window: keep only the most recent positions.
1839    Recent,
1840    /// Mass-based eviction: sinks + recents + top accumulated attention mass.
1841    Born { sink: usize },
1842}
1843
1844/// Full KV cache for all layers.
1845#[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    /// Bytes owned by linear-core recurrent state (including the tentative
1884    /// speculative scratch).  This is reported separately from attention KV
1885    /// so a serving slot's O(1) capacity can be compared with its context
1886    /// cache without guessing from model geometry.
1887    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    /// Attention KV (or sealed O(1) attention state) bytes, excluding the
1897    /// linear recurrent vectors returned by [`recurrent_state_bytes`].
1898    pub fn attention_state_bytes(&self) -> usize {
1899        self.total_memory_bytes()
1900            .saturating_sub(self.recurrent_state_bytes())
1901    }
1902
1903    /// Current sequence length (max across layers — dead layers may lag).
1904    pub fn seq_len(&self) -> usize {
1905        self.layers.iter().map(|l| l.seq_len).max().unwrap_or(0)
1906    }
1907
1908    /// Bytes owned by bounded-anchor rings (constant in context).
1909    pub fn bounded_state_bytes(&self) -> usize {
1910        self.layers.iter().map(|l| l.bounded_state_bytes()).sum()
1911    }
1912
1913    /// True when some layer with PER-POSITION storage reached the cap.
1914    /// Bounded anchors hold nothing per position and never need it: a
1915    /// model whose every layer is O(1) has no eviction cliff at all.
1916    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    /// Evict down to `keep_last` positions according to the policy.
1927    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        // The state has to arrive as state, not as something that looks
1966        // like it: the oracle is what the layer ANSWERS, not what it
1967        // stores. Same query, same output, bit for bit.
1968        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        // A refusal is the feature: a cache whose extra state this format
2004        // does not carry must not travel looking complete.
2005        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        // The versioned wire carries the operator identity, so a delta
2014        // layer exports; only the OLD unversioned body is refused.
2015        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        // Linear record: the recurrent vector travels f32 whatever the
2026        // wire dtype and the per-position stores come back empty.
2027        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        // Identity mismatch is a refusal, not a warning.
2040        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        // Bounded record, f32 and f16 rings.
2046        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            // 0.125 multiples are exact in f16, so both dtypes round-trip bit for bit.
2063            assert!(ra.same_state(rb), "f16={f16}");
2064            // A ring of another width refuses the record.
2065            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        // head 1 slice is contiguous and equals its part of k_new
2093        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    /// q8_2f-attend ≈ f32-attend: 100 positions (crosses the field freeze
2208    /// at the 64th), pseudo-random vectors, relative tolerance of the
2209    /// int8 grid. Plus rollback and mass-based eviction on the q8 storage.
2210    #[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                    // channel structure: even channels ×4 (checks the 2f field)
2223                    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        // rollback + eviction live on the q8 storage
2253        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        // memory: q8 ≈ 1 byte/element + scale per row (vs 4 for f32)
2262        assert!(q8.memory_bytes() * 3 < f.memory_bytes());
2263    }
2264
2265    /// Grouped GQA attend must be bit-identical to per-head attend in
2266    /// every KV mode (it is the same math with rows streamed once).
2267    #[test]
2268    fn attend_group_equals_per_head_attend_bitexact() {
2269        let (kv_heads, hd, hpk) = (2usize, 32usize, 3usize); // 6 Q-heads
2270        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    /// Learned sinks (gpt-oss / MiMo-V2): the grouped softmax must equal a
2321    /// softmax over the visible rows PLUS an explicit value-less sink
2322    /// column, computed independently in f64 — at position 0 (one stored
2323    /// row), mid-sequence, and with the window truncating the rows. The
2324    /// importance row must be the row probabilities (the sink's share is
2325    /// nobody's importance).
2326    #[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        // Mixed signs and one sink that dominates its head.
2348        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                        // Logits of the visible rows, then the sink column.
2364                        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                        // The sink column carries a zero value vector.
2380                        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                    // The sink takes real mass: rows sum below 1.
2404                    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    /// Scoring only the window's rows must be bit-identical to the former
2413    /// whole-row scoring (−inf outside the window) — with and without a
2414    /// sink, output and importance. The reference is the plain per-head
2415    /// softmax over the same rows written out here.
2416    #[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            // Reference: the historical full-row −inf-masked kernel.
2441            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    /// Sinks are weights: a new conversation (clear) and a wire import
2481    /// must keep them.
2482    #[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    /// Review regression: mass-based eviction in MIXED modes. q8v used to
2497    /// panic (gather over an empty v[h]), q8k silently left raw V
2498    /// uncompressed (stale rows under kept keys + memory leak).
2499    #[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); // q8v: used to panic here
2513            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            // V rows match the kept set: the heaviest positions
2519            // (tail 60..79) must be present in the attend output.
2520            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        // 8 positions; keys carry the position index so we can verify
2538        // exactly which positions survive the gather.
2539        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        // Position 3 carries the most attention mass.
2545        let mut imp = vec![0.05f32; 8];
2546        imp[3] = 5.0;
2547        layer.accumulate_imp(&imp);
2548
2549        cache.evict(4); // sink 1 + recent 2 + 1 top-mass slot
2550        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        // imp stays aligned with the gathered positions.
2559        assert_eq!(layer.head_len(0), 4);
2560    }
2561}