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/// Source of `LayerKvCache::generation`. Process-wide so two caches (or a
52/// cache and a restored clone of itself from another moment) never share a
53/// value by accident: a device mirror keyed by (generation, rows) can then tell
54/// "the rows I hold" from "the same number of different rows".
55static KV_GEN: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(1);
56
57fn next_gen() -> u64 {
58    KV_GEN.fetch_add(1, std::sync::atomic::Ordering::Relaxed)
59}
60
61/// Per-layer O(1) Nyström attention state (runtime `attn_type`
62/// override — spec §7 presence-driven pattern, no format change).
63///
64/// Collecting: the prompt pass still runs EXACT cache attention (the
65/// prefill outputs feed the residual stream, so they cannot be
66/// deferred) while the per-position rotated queries are buffered;
67/// `o1_seal()` then freezes landmarks + M from the full prompt, replays
68/// it into per-KV-group streaming states, and DROPS the full KV.
69/// Sealed: decode replaces cache attention with
70/// `NystromState::step_group()`.
71#[derive(Debug, Clone)]
72pub enum O1State {
73    Collecting {
74        m: usize,
75        w: usize,
76        sink: usize,
77        rect: crate::nystrom::O1Rect,
78        /// Optional completed-row barrier. `None` means the caller will
79        /// request a full-prompt seal; `Some(B)` keeps a short prompt exact
80        /// until the first skeleton-safe boundary B.
81        seal_at: Option<usize>,
82        /// Rotated post-norm queries, `[pos × num_heads × head_dim]`.
83        q_buf: Vec<f32>,
84    },
85    /// One state per KV GROUP, each holding its group's Q heads. The
86    /// exact window / sinks / K̃ are stored once per group (every Q head
87    /// of the group reads the same k/v rows); only the far field, Q̃ and
88    /// M — the query-dependent pieces — stay per Q head. See
89    /// `NystromState` for which piece is which and why.
90    Sealed {
91        groups: Vec<crate::nystrom::NystromState>,
92    },
93}
94
95/// KV cache for a single layer, head-major.
96#[derive(Debug, Clone)]
97pub struct LayerKvCache {
98    pub mode: KvMode,
99    /// Per-KV-head keys: `k[h]` is `[seq_len × head_dim]` (empty if head is dead).
100    k: Vec<Vec<f32>>,
101    /// Per-KV-head values, same layout.
102    v: Vec<Vec<f32>>,
103    /// q8 storage (mode == Q8_2F): int8 rows + f32 scale per row.
104    kq: Vec<Vec<i8>>,
105    ks: Vec<Vec<f32>>,
106    vq: Vec<Vec<i8>>,
107    vs: Vec<Vec<f32>>,
108    /// Per-channel scale fields per head [head_dim]; empty until frozen.
109    kcol: Vec<Vec<f32>>,
110    vcol: Vec<Vec<f32>>,
111    /// Accumulated attention mass per stored position: importance of a
112    /// position is how much probability mass reads it.
113    imp: Vec<f32>,
114    /// Rows stored (grows once per token, dead heads included). Stored
115    /// row 0 is absolute position `base`, so this equals the context depth
116    /// only while `base == 0`; [`Self::pos_len`] is the depth. Every attend
117    /// indexes rows, never positions (eviction made this a row count long
118    /// before `trim_window` did).
119    pub seq_len: usize,
120    /// Absolute position of stored row 0. Only `trim_window` advances it;
121    /// 0 on every layer that was never trimmed.
122    base: usize,
123    /// The attend window of a layer that keeps only its tail
124    /// (`trim_window`), set by the first call. A property of the layer,
125    /// not of the conversation: `clear()` keeps it. Such a layer bounds
126    /// itself, so the cache-wide eviction leaves it alone.
127    tail: Option<usize>,
128    /// Storage generation: a fresh value on every mutation other than
129    /// `append` and `truncate_last` (clear, trim, eviction, wire import,
130    /// resets). A device mirror that copies these rows keys its validity on
131    /// (generation, rows): a trimmed layer's row count is periodic (576..639 for a
132    /// 512 window), so the count alone can match a different set of rows.
133    /// Appends and rollbacks show up in the count, and the speculative
134    /// paths re-point their mirrors themselves, so they keep the value.
135    generation: u64,
136    pub num_kv_heads: usize,
137    pub head_dim: usize,
138    /// Linear-core recurrent state S (vmf_phase), f64; empty on full layers.
139    pub linear_state: Vec<f32>,
140    /// Tentative lane-2 state during speculative verify.
141    pub linear_scratch: Vec<f32>,
142    /// The legacy cache wire has no operator tag. Delta operators therefore
143    /// bind this layer to a fail-closed boundary until a versioned wire
144    /// schema can carry the linear-core identity.
145    linear_wire_allowed: bool,
146    /// O(1) Nyström override (None = plain cache attention).
147    pub o1: Option<O1State>,
148    /// A deferred seal failure is terminal for the current request. The
149    /// attention functions return only a hidden vector, so the pipeline
150    /// consumes this side channel at its next forward boundary.
151    o1_error: Option<String>,
152    /// Set when a collecting state actually becomes sealed. Pipeline owns
153    /// the epoch bump and consumes this bit after a complete forward.
154    o1_transitioned: bool,
155    /// Learned per-Q-head attention-sink logits of this layer (gpt-oss /
156    /// MiMo-V2 `self_attn.sinks`, one f32 per Q head). The sink is an
157    /// extra softmax column with no value: it joins the max and the
158    /// denominator of every head's softmax and so lets a head attend to
159    /// "nothing". These are WEIGHTS, not sequence state — `clear()` and
160    /// the wire import keep them. None = an ordinary softmax.
161    pub sinks: Option<Vec<f32>>,
162    /// Natively bounded anchor (`swa_sink_v1`): a fixed-size ring
163    /// installed from the header at load, never per prompt. A layer that
164    /// carries it stores NOTHING per position (`k`/`v` stay empty).
165    pub bounded: Option<crate::bounded::BoundedState>,
166    /// Which state record this layer exchanges on the wire (v2).
167    pub wire_kind: WireKind,
168    /// This layer's index in the stack (the wire header names it).
169    pub wire_layer: u32,
170    /// hash64 of the model's operator identity
171    /// (`ModelArch::linear_core_identity` JSON); 0 = no operator record.
172    pub wire_identity: u64,
173    /// `kv_abs_max`'s running answer: max |x| over the f32 K rows and over
174    /// the V rows `[0, kv_amax_rows)` of every head (+inf once one was inf
175    /// or NaN).
176    kv_amax: [f32; 2],
177    kv_amax_rows: usize,
178}
179
180/// State-record kind of the versioned cache wire (`export_wire` v2).
181#[derive(Debug, Clone, Copy, PartialEq, Eq)]
182#[repr(u8)]
183pub enum WireKind {
184    /// Per-position K/V (+ importance) — the legacy body.
185    Full = 0,
186    /// Recurrent state vector (S + conv ring), f32.
187    Linear = 1,
188    /// Bounded anchor: insert counter + ring K/V.
189    Bounded = 2,
190    /// A full layer that stores only its tail (`trim_window`): `base` as
191    /// u64, then the `Full` body of the stored rows. Sent only when
192    /// `base > 0`, so an untrimmed record is byte-identical to `Full`; a
193    /// peer that predates it refuses it as an unknown kind instead of
194    /// reading the tail as a whole history.
195    FullTail = 3,
196}
197
198impl WireKind {
199    fn from_u8(v: u8) -> Option<Self> {
200        match v {
201            0 => Some(WireKind::Full),
202            1 => Some(WireKind::Linear),
203            2 => Some(WireKind::Bounded),
204            3 => Some(WireKind::FullTail),
205            _ => None,
206        }
207    }
208}
209
210/// Magic of the versioned state wire.
211pub const WIRE_MAGIC: &[u8; 4] = b"CMFS";
212/// Current wire version.
213pub const WIRE_VERSION: u32 = 2;
214
215impl LayerKvCache {
216    pub fn new(num_kv_heads: usize, head_dim: usize) -> Self {
217        Self {
218            sinks: None,
219            mode: KvMode::from_env(),
220            k: vec![Vec::new(); num_kv_heads],
221            v: vec![Vec::new(); num_kv_heads],
222            kq: vec![Vec::new(); num_kv_heads],
223            ks: vec![Vec::new(); num_kv_heads],
224            vq: vec![Vec::new(); num_kv_heads],
225            vs: vec![Vec::new(); num_kv_heads],
226            kcol: vec![Vec::new(); num_kv_heads],
227            vcol: vec![Vec::new(); num_kv_heads],
228            imp: Vec::new(),
229            seq_len: 0,
230            base: 0,
231            tail: None,
232            generation: next_gen(),
233            num_kv_heads,
234            head_dim,
235            linear_state: Vec::new(),
236            linear_scratch: Vec::new(),
237            linear_wire_allowed: true,
238            o1: None,
239            o1_error: None,
240            o1_transitioned: false,
241            bounded: None,
242            wire_kind: WireKind::Full,
243            wire_layer: 0,
244            wire_identity: 0,
245            kv_amax: [0.0; 2],
246            kv_amax_rows: 0,
247        }
248    }
249
250    /// (max |k|, max |v|) over every K and over every V entry the cache
251    /// stores in f32, +inf when one is inf or NaN: what the wgpu prefill
252    /// compares with f16's range before its matrix-unit attention casts
253    /// these rows to f16. Lazy — it scans only the rows appended since the
254    /// last call; every other change to the rows (truncate, evict, clear,
255    /// import, an o1 seal) moves the scanned mark back
256    /// (`kv_rows_changed`), and a scan from the first row starts afresh.
257    pub(crate) fn kv_abs_max(&mut self) -> (f32, f32) {
258        let hd = self.head_dim.max(1);
259        let rows = self.kv_rows();
260        if rows < self.kv_amax_rows {
261            self.kv_amax_rows = 0;
262        }
263        if self.kv_amax_rows == 0 {
264            self.kv_amax = [0.0; 2];
265        }
266        let from = self.kv_amax_rows * hd;
267        for (m, heads) in self.kv_amax.iter_mut().zip([&self.k, &self.v]) {
268            for x in heads.iter().filter(|x| x.len() > from) {
269                *m = m.max(crate::gpu::abs_max_or_inf(&x[from..]));
270            }
271        }
272        self.kv_amax_rows = rows;
273        (self.kv_amax[0], self.kv_amax[1])
274    }
275
276    /// `kv_abs_max` for a caller that took the maxima of the rows it just
277    /// appended itself: `chunk` = (from, upto, max |k|, max |v|) over the
278    /// rows `[from, upto)`. Folded in without a scan when the cache holds
279    /// exactly `upto` rows and every row before `from` was scanned; else
280    /// the plain scan.
281    pub(crate) fn kv_abs_max_after(&mut self, chunk: (usize, usize, f32, f32)) -> (f32, f32) {
282        let (from, upto, km, vm) = chunk;
283        if self.kv_amax_rows != from || self.kv_rows() != upto || from > upto {
284            return self.kv_abs_max();
285        }
286        if from == 0 {
287            self.kv_amax = [0.0; 2];
288        }
289        self.kv_amax = [self.kv_amax[0].max(km), self.kv_amax[1].max(vm)];
290        self.kv_amax_rows = upto;
291        (self.kv_amax[0], self.kv_amax[1])
292    }
293
294    /// Rows any head stores.
295    fn kv_rows(&self) -> usize {
296        let hd = self.head_dim.max(1);
297        self.k
298            .iter()
299            .chain(&self.v)
300            .map(|x| x.len() / hd)
301            .max()
302            .unwrap_or(0)
303    }
304
305    /// Rows from `first` on may no longer be the ones `kv_abs_max` scanned.
306    fn kv_rows_changed(&mut self, first: usize) {
307        self.kv_amax_rows = self.kv_amax_rows.min(first);
308    }
309
310    // ── Natively bounded anchor (swa_sink_v1) ──
311
312    /// Give this layer its fixed-size ring (`[kvh][window][hd]` K and V).
313    /// Called once at load from the header; the record is zeroed on
314    /// `clear()` and never reallocated.
315    pub fn install_bounded(&mut self, window: usize) {
316        self.bounded = Some(crate::bounded::BoundedState::new(
317            self.num_kv_heads,
318            self.head_dim,
319            window,
320        ));
321        self.wire_kind = WireKind::Bounded;
322    }
323
324    /// One position of the bounded operator: insert the raw `k, v`
325    /// (`[kvh][hd]`) into slot `t mod W`, then attend every Q head of
326    /// `q` (`[nh][hd]`, raw) over sinks ∪ window into `out` (`[nh][hd]`).
327    /// Nothing is appended per position.
328    #[allow(clippy::too_many_arguments)]
329    pub fn bounded_step(
330        &mut self,
331        q: &[f32],
332        k: &[f32],
333        v: &[f32],
334        w: &crate::bounded::BoundedWeights,
335        rope: &crate::bounded::BoundedRope,
336        scale: f32,
337        num_heads: usize,
338        out: &mut [f32],
339    ) {
340        let st = self
341            .bounded
342            .as_mut()
343            .expect("bounded_step on a layer without an installed ring");
344        st.insert(k, v);
345        st.attend(q, num_heads, &w.sink_k, &w.sink_v, w.sink, rope, scale, out);
346        // Honest context depth for the memory/seq report — nothing is
347        // stored per position.
348        self.seq_len += 1;
349    }
350
351    /// Bytes of the bounded ring (0 on other layers).
352    pub fn bounded_state_bytes(&self) -> usize {
353        self.bounded.as_ref().map(|b| b.state_bytes()).unwrap_or(0)
354    }
355
356    /// Bit-for-bit copy of the ring for speculation (None on other layers).
357    pub fn bounded_snapshot(&self) -> Option<crate::bounded::BoundedSnapshot> {
358        self.bounded.as_ref().map(|b| b.snapshot())
359    }
360
361    /// Restore a ring snapshot taken on this layer; `seq_len` follows the
362    /// restored insert counter.
363    pub fn bounded_restore(&mut self, s: &crate::bounded::BoundedSnapshot) {
364        if let Some(b) = self.bounded.as_mut() {
365            b.restore(s);
366            self.seq_len = b.seen;
367            self.generation = next_gen();
368        }
369    }
370
371    /// Bind this layer to the legacy cache-wire policy. Delta layers must
372    /// refuse untagged state exchange rather than risk a plausible additive
373    /// interpretation on the peer.
374    pub fn set_linear_wire_allowed(&mut self, allowed: bool) {
375        self.linear_wire_allowed = allowed;
376    }
377
378    /// Discard tentative recurrent state after a speculative rejection or
379    /// any other path that abandons the lane-2 result.
380    pub fn discard_linear_scratch(&mut self) {
381        self.linear_scratch.clear();
382    }
383
384    /// Per-KV-head stored keys `[seq_len × head_dim]` (GPU token graph
385    /// sync): the stored tail, row 0 = position `base()`.
386    pub fn k_heads(&self) -> &[Vec<f32>] {
387        &self.k
388    }
389    /// Per-KV-head stored values `[seq_len × head_dim]`, row 0 = `base()`.
390    pub fn v_heads(&self) -> &[Vec<f32>] {
391        &self.v
392    }
393
394    // ── Sliding-window tail (trim_window) ──
395
396    /// Absolute position of stored row 0 (0 unless `trim_window` dropped
397    /// a prefix).
398    pub fn base(&self) -> usize {
399        self.base
400    }
401
402    /// Absolute context depth: the position the next append will hold.
403    pub fn pos_len(&self) -> usize {
404        self.base + self.seq_len
405    }
406
407    /// Storage generation (see the field): with the row count, the key a
408    /// device copy of these rows is valid under.
409    pub fn generation(&self) -> u64 {
410        self.generation
411    }
412
413    /// The window this layer trims to, once `trim_window` ran on it.
414    pub fn tail_window(&self) -> Option<usize> {
415        self.tail
416    }
417
418    /// Drop the rows a sliding-window layer can never read again. A query
419    /// at position p reads p−w+1..=p only, so once more than `2w` rows are
420    /// stored the front is drained down to `w + slack` rows, the cut
421    /// rounded DOWN to an absolute position that is a multiple of `align`
422    /// (so `w + slack ..= w + slack + align − 1` rows stay). Returns the
423    /// rows dropped.
424    ///
425    /// The rows that stay keep their order and values (`Vec::drain` of the
426    /// front), and every attend indexes rows relative to the newest one
427    /// (`first = rows − w`), so each later query sees the very rows it saw
428    /// untrimmed — bit for bit. `slack` rows past the window are what a
429    /// rollback (`truncate_last`) may take back without losing a row the
430    /// window needs. `align` keeps device GEMMs that reduce over the rows
431    /// from row 0 in fixed tiles on the tile grid they had untrimmed (the
432    /// dropped rows were whole all-zero tiles there).
433    ///
434    /// Callers trim only at a safe point: never between a call's appends
435    /// and its attend or importance accumulation (a prefill chunk appends
436    /// all its rows first and indexes them from the count it started at).
437    ///
438    /// No-op on an O(1) or bounded layer (they own their bound), and on a
439    /// q8 layer whose per-channel field is not frozen yet: the freeze
440    /// reads every stored row, so dropping one before it would change it.
441    pub fn trim_window(&mut self, w: usize, slack: usize, align: usize) -> usize {
442        if self.o1.is_some() || self.bounded.is_some() || w == 0 {
443            return 0;
444        }
445        self.tail = Some(w);
446        if self.seq_len <= 2 * w {
447            return 0;
448        }
449        if matches!(self.mode, KvMode::Q8 { .. })
450            && self.kcol.iter().all(Vec::is_empty)
451            && self.vcol.iter().all(Vec::is_empty)
452        {
453            return 0;
454        }
455        let align = align.max(1);
456        let keep_from = self.pos_len().saturating_sub(w + slack);
457        let new_base = keep_from / align * align;
458        if new_base <= self.base {
459            return 0;
460        }
461        let d = new_base - self.base;
462        let hd = self.head_dim;
463        fn drop_front<T>(v: &mut Vec<T>, n: usize) {
464            let n = n.min(v.len());
465            v.drain(..n);
466        }
467        for h in 0..self.num_kv_heads {
468            drop_front(&mut self.k[h], d * hd);
469            drop_front(&mut self.v[h], d * hd);
470            drop_front(&mut self.kq[h], d * hd);
471            drop_front(&mut self.vq[h], d * hd);
472            drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
473            drop_front(&mut self.vs[h], d);
474        }
475        drop_front(&mut self.imp, d);
476        self.base = new_base;
477        self.seq_len -= d;
478        self.generation = next_gen();
479        // `kv_abs_max` counts its scanned mark from the old row 0. Moved
480        // with the rows, the next scan still starts at the first row it has
481        // not seen; left in place, rows appended after the trim up to the
482        // old mark would never be scanned. The maxima may keep a dropped
483        // row's value, which only over-states.
484        self.kv_amax_rows = self.kv_amax_rows.saturating_sub(d);
485        d
486    }
487
488    // ── O(1) Nyström override ──
489
490    /// Arm query collection for a fresh prompt pass (a cleared cache).
491    pub fn o1_begin(&mut self, m: usize, w: usize, sink: usize, rect: crate::nystrom::O1Rect) {
492        self.o1_begin_with_boundary(m, w, sink, rect, None);
493    }
494
495    /// Arm query collection with an optional completed-row seal barrier.
496    /// The barrier is deliberately part of the existing collecting state:
497    /// no second history or scheduler is introduced for short prompts.
498    pub(crate) fn o1_begin_with_boundary(
499        &mut self,
500        m: usize,
501        w: usize,
502        sink: usize,
503        rect: crate::nystrom::O1Rect,
504        seal_at: Option<usize>,
505    ) {
506        self.o1 = Some(O1State::Collecting {
507            m,
508            w,
509            sink,
510            rect,
511            seal_at,
512            q_buf: Vec::new(),
513        });
514        self.o1_error = None;
515        self.o1_transitioned = false;
516    }
517
518    /// Record one position's rotated queries (`[num_heads × head_dim]`)
519    /// during the exact prompt pass. No-op unless collecting — the hook
520    /// sits inside qwen_attention so every prefill flavor (sequential,
521    /// batched) feeds the same trace.
522    pub fn o1_push_q(&mut self, q_all: &[f32]) {
523        if let Some(O1State::Collecting { q_buf, .. }) = &mut self.o1 {
524            q_buf.extend_from_slice(q_all);
525        }
526    }
527
528    pub fn o1_sealed(&self) -> bool {
529        matches!(self.o1, Some(O1State::Sealed { .. }))
530    }
531
532    /// Pending completed-row barrier, if any. A plain full-prompt seal has
533    /// no barrier until the caller asks to seal.
534    pub(crate) fn o1_pending_boundary(&self) -> Option<usize> {
535        match &self.o1 {
536            Some(O1State::Collecting { seal_at, .. }) => *seal_at,
537            _ => None,
538        }
539    }
540
541    /// Whether a batch of `count` exact rows would cross the deferred
542    /// boundary. This lets batched/pair callers split before row B rather
543    /// than appending exact KV past the point where conversion is required.
544    pub(crate) fn o1_boundary_crossed_by(&self, count: usize) -> bool {
545        let Some(target) = self.o1_pending_boundary() else {
546            return false;
547        };
548        target <= self.seq_len
549            || self
550                .seq_len
551                .checked_add(count)
552                .map_or(true, |next| next >= target)
553    }
554
555    pub(crate) fn take_o1_transition(&mut self) -> bool {
556        std::mem::take(&mut self.o1_transitioned)
557    }
558
559    pub(crate) fn take_o1_error(&self) -> Option<String> {
560        // Error observation is deliberately non-consuming.  The error is the
561        // append guard for this request; removing it would let a caller that
562        // ignored the returned Err resume ordinary KV growth after the
563        // bounded transition dropped its overlay.  `clear()` is the explicit
564        // reset boundary that clears the latch.
565        self.o1_error.clone()
566    }
567
568    /// Abort a malformed deferred transition after attention has already
569    /// produced its current row. Clearing the exact storage and dropping
570    /// the overlay makes the state unrecoverable by continued decode; the
571    /// pipeline then routes through its normal graph/cancel cleanup.
572    pub(crate) fn o1_abort(&mut self, err: String) {
573        self.k.iter_mut().for_each(Vec::clear);
574        self.v.iter_mut().for_each(Vec::clear);
575        self.kv_rows_changed(0);
576        self.kq.iter_mut().for_each(Vec::clear);
577        self.ks.iter_mut().for_each(Vec::clear);
578        self.vq.iter_mut().for_each(Vec::clear);
579        self.vs.iter_mut().for_each(Vec::clear);
580        self.kcol.iter_mut().for_each(Vec::clear);
581        self.vcol.iter_mut().for_each(Vec::clear);
582        self.imp.clear();
583        self.o1 = None;
584        self.seq_len = 0;
585        self.base = 0;
586        self.generation = next_gen();
587        self.o1_transitioned = false;
588        self.o1_error = Some(err);
589    }
590
591    /// Freeze the prompt into per-KV-group Nyström states and drop this
592    /// layer's full KV. Returns false while a short collecting layer is
593    /// below its deferred boundary; malformed prerequisites abort the
594    /// layer instead of silently resuming exact KV growth. The seal needs
595    /// f32 KV rows (`CMF_KV=q8` stores int8), every group densely stored, a
596    /// full q trace, and a GQA fan-out that actually divides.
597    pub fn o1_seal(&mut self, num_heads: usize) -> bool {
598        match self.o1_seal_checked(num_heads) {
599            Ok(sealed) => sealed,
600            Err(err) => {
601                tracing::error!("o1: seal aborted: {err}");
602                self.o1_abort(err);
603                false
604            }
605        }
606    }
607
608    /// Checked seal implementation. Validation happens while the collecting
609    /// state and full KV are still intact; only a valid completed boundary
610    /// is allowed to destructively convert them.
611    pub(crate) fn o1_seal_checked(&mut self, num_heads: usize) -> Result<bool, String> {
612        if let Some(err) = self.o1_error.clone() {
613            return Err(err);
614        }
615        // Idempotent: sealing a sealed (or plain) layer must not disturb its
616        // state, and a plain layer is not an O(1) participant.
617        if !matches!(self.o1, Some(O1State::Collecting { .. })) {
618            return Ok(self.o1_sealed());
619        }
620        let (m, w, sink, requested_boundary, q_len) = match &self.o1 {
621            Some(O1State::Collecting {
622                m,
623                w,
624                sink,
625                rect: _,
626                seal_at,
627                q_buf,
628            }) => (*m, *w, *sink, *seal_at, q_buf.len()),
629            _ => unreachable!("checked above"),
630        };
631        let floor = crate::nystrom::o1_deferred_boundary(w, sink)
632            .ok_or_else(|| "o1 seal: w + sink + slack + 1 overflow".to_string())?;
633        let target = requested_boundary.unwrap_or(floor).max(floor);
634        let t = self.seq_len;
635        if t < target {
636            if let Some(O1State::Collecting { seal_at, .. }) = &mut self.o1 {
637                if *seal_at != Some(target) {
638                    *seal_at = Some(target);
639                    tracing::info!(
640                        "o1 deferred seal: current rows={t}, boundary={target} (floor={floor})"
641                    );
642                }
643            }
644            return Ok(false);
645        }
646
647        let hd = self.head_dim;
648        if t == 0 {
649            return Err("o1 seal: cannot seal an empty layer".into());
650        }
651        if self.mode != KvMode::F32 {
652            return Err("o1 seal: requires dense F32 KV storage".into());
653        }
654        if self.num_kv_heads == 0 || num_heads == 0 || num_heads % self.num_kv_heads != 0 {
655            return Err(format!(
656                "o1 seal: invalid GQA geometry num_heads={num_heads} num_kv_heads={}",
657                self.num_kv_heads
658            ));
659        }
660        let hpk = num_heads / self.num_kv_heads;
661        let expected_k = t
662            .checked_mul(hd)
663            .ok_or_else(|| "o1 seal: KV row length overflow".to_string())?;
664        let expected_q = expected_k
665            .checked_mul(num_heads)
666            .ok_or_else(|| "o1 seal: query trace length overflow".to_string())?;
667        if q_len != expected_q {
668            return Err(format!(
669                "o1 seal: query trace has {q_len} values, expected {expected_q}"
670            ));
671        }
672        if (0..self.num_kv_heads)
673            .any(|g| self.k[g].len() != expected_k || self.v[g].len() != expected_k)
674        {
675            return Err("o1 seal: KV heads are not densely populated".into());
676        }
677        if m < 4 || w == 0 {
678            return Err(format!("o1 seal: invalid geometry m={m} w={w}"));
679        }
680
681        let Some(O1State::Collecting {
682            m,
683            w,
684            sink,
685            rect,
686            q_buf,
687            ..
688        }) = self.o1.take()
689        else {
690            unreachable!("collecting state disappeared after validation");
691        };
692        let mut groups = Vec::with_capacity(self.num_kv_heads);
693        // Query trace is position-major; the state wants each head's
694        // queries contiguous, so transpose one group at a time.
695        let mut qh = vec![0.0f32; hpk * t * hd];
696        for g in 0..self.num_kv_heads {
697            for hh in 0..hpk {
698                let h = g * hpk + hh;
699                for p in 0..t {
700                    let src = (p * num_heads + h) * hd;
701                    let dst = (hh * t + p) * hd;
702                    qh[dst..dst + hd].copy_from_slice(&q_buf[src..src + hd]);
703                }
704            }
705            let qs: Vec<&[f32]> = (0..hpk)
706                .map(|hh| &qh[hh * t * hd..(hh + 1) * t * hd])
707                .collect();
708            let mut st = crate::nystrom::NystromState::new_group(m, w, sink, hpk).with_rect(rect);
709            st.prefill_group(&qs, &self.k[g], &self.v[g], t, hd, hd);
710            groups.push(st);
711        }
712        // The states now carry everything decode needs — release the
713        // O(context) storage (this is the memory claim, not a cosmetic).
714        for h in 0..self.num_kv_heads {
715            self.k[h] = Vec::new();
716            self.v[h] = Vec::new();
717        }
718        self.kv_rows_changed(0);
719        self.imp = Vec::new();
720        self.o1 = Some(O1State::Sealed { groups });
721        self.o1_transitioned = true;
722        Ok(true)
723    }
724
725    /// One decode step on a sealed layer: per KV group, insert the
726    /// group's fresh (k, v) ONCE and read every Q head's attention
727    /// output. Returns `[num_heads × head_dim]`. Head h belongs to group
728    /// h/hpk, so a group's Q heads are contiguous in `q_all`/`out` —
729    /// same math as the shared KV row the exact path appends once.
730    /// Device views of the sealed o1 groups, or None when o1 is not
731    /// sealed on this layer (or any group is in the degenerate
732    /// exact-only mode the GPU path does not carry).
733    pub fn o1_views(&self) -> Option<Vec<crate::nystrom::O1DeviceView<'_>>> {
734        let Some(O1State::Sealed { groups }) = &self.o1 else {
735            return None;
736        };
737        let views: Vec<_> = groups.iter().map(|g| g.device_view()).collect();
738        if views.iter().any(|v| v.exact_only) {
739            return None;
740        }
741        Some(views)
742    }
743
744    pub fn o1_step(
745        &mut self,
746        q_all: &[f32],
747        k_new: &[f32],
748        v_new: &[f32],
749        num_heads: usize,
750    ) -> Vec<f32> {
751        let hd = self.head_dim;
752        let hpk = num_heads / self.num_kv_heads.max(1);
753        let mut out = vec![0.0f32; num_heads * hd];
754        let Some(O1State::Sealed { groups }) = &mut self.o1 else {
755            debug_assert!(false, "o1_step on an unsealed layer");
756            return out;
757        };
758        for (g, st) in groups.iter_mut().enumerate() {
759            let (lo, hi) = (g * hpk * hd, (g + 1) * hpk * hd);
760            st.step_group(
761                &q_all[lo..hi],
762                &k_new[g * hd..(g + 1) * hd],
763                &v_new[g * hd..(g + 1) * hd],
764                &mut out[lo..hi],
765            );
766        }
767        // Track the true context depth for the honest memory/seq report
768        // (nothing is stored per position — the state is O(1)).
769        self.seq_len += 1;
770        out
771    }
772
773    /// Bytes held by the O(1) override (query trace while collecting,
774    /// per-KV-group states once sealed).
775    pub fn o1_memory_bytes(&self) -> usize {
776        match &self.o1 {
777            Some(O1State::Collecting { q_buf, .. }) => q_buf.len() * std::mem::size_of::<f32>(),
778            Some(O1State::Sealed { groups }) => groups.iter().map(|s| s.memory_bytes()).sum(),
779            None => 0,
780        }
781    }
782
783    /// Quantize one row against the per-channel field (empty col = 1);
784    /// `group` — elements per scale (the whole row or KV_K_GROUP).
785    fn quant_row(row: &[f32], col: &[f32], q: &mut Vec<i8>, sc: &mut Vec<f32>, group: usize) {
786        let mut resid = vec![0.0f32; row.len()];
787        for (d, &x) in row.iter().enumerate() {
788            resid[d] = if col.is_empty() { x } else { x / col[d] };
789        }
790        for g0 in (0..row.len()).step_by(group) {
791            let g1 = (g0 + group).min(row.len());
792            let mut absmax = 0.0f32;
793            for &r in &resid[g0..g1] {
794                absmax = absmax.max(r.abs());
795            }
796            let s = (absmax / 127.0).max(1e-12);
797            sc.push(s);
798            for &r in &resid[g0..g1] {
799                q.push((r / s).round().clamp(-127.0, 127.0) as i8);
800            }
801        }
802    }
803
804    /// Freeze the 2f field: col = RMS of channels over stored rows, old
805    /// rows are requantized against the new field (once per conversation).
806    fn freeze_cols(&mut self) {
807        let hd = self.head_dim;
808        let ngk = hd.div_ceil(KV_K_GROUP);
809        for h in 0..self.num_kv_heads {
810            for (qv, sv, colv, group) in [
811                (
812                    &mut self.kq[h],
813                    &mut self.ks[h],
814                    &mut self.kcol[h],
815                    KV_K_GROUP,
816                ),
817                (&mut self.vq[h], &mut self.vs[h], &mut self.vcol[h], hd),
818            ] {
819                let spp = if group == hd { 1 } else { ngk }; // scales per position
820                let n = sv.len() / spp;
821                if n == 0 {
822                    continue;
823                }
824                // Dequantize to f32, RMS over channels, requantize.
825                let mut rows = vec![0.0f32; n * hd];
826                for p in 0..n {
827                    for d in 0..hd {
828                        rows[p * hd + d] = qv[p * hd + d] as f32 * sv[p * spp + d / group];
829                    }
830                }
831                let mut col = vec![0.0f32; hd];
832                for p in 0..n {
833                    for d in 0..hd {
834                        col[d] += rows[p * hd + d] * rows[p * hd + d];
835                    }
836                }
837                for c in col.iter_mut() {
838                    *c = (*c / n as f32).sqrt().max(1e-6);
839                }
840                qv.clear();
841                sv.clear();
842                for p in 0..n {
843                    Self::quant_row(&rows[p * hd..(p + 1) * hd], &col, qv, sv, group);
844                }
845                *colv = col;
846            }
847        }
848    }
849
850    /// Append K/V for one position. `k_new`/`v_new` are
851    /// `[num_kv_heads × head_dim]`; heads with `alive[h] == false` are
852    /// skipped (their slices stay empty).
853    pub fn append(&mut self, k_new: &[f32], v_new: &[f32], alive: &[bool]) {
854        // A failed bounded transition is terminal until the sequence is
855        // cleared. Do not let an ignored boolean/result resume plain KV
856        // growth after the O(1) collector has aborted.
857        if self.o1_error.is_some() {
858            return;
859        }
860        debug_assert_eq!(k_new.len(), self.num_kv_heads * self.head_dim);
861        debug_assert_eq!(v_new.len(), self.num_kv_heads * self.head_dim);
862        // Freeze the 2f field AT THE START of append: only rows that
863        // survived verify are visible (a rejected lane-2 draft does not
864        // pollute the field — found in review), and the threshold uses >=
865        // rather than strict equality (in small windows eviction may
866        // oscillate across 64).
867        if matches!(self.mode, KvMode::Q8 { .. })
868            && self.seq_len >= KV_COL_WARMUP
869            && self.kcol.iter().all(Vec::is_empty)
870            && self.vcol.iter().all(Vec::is_empty)
871        {
872            self.freeze_cols();
873        }
874        for h in 0..self.num_kv_heads {
875            if !alive.get(h).copied().unwrap_or(true) {
876                continue;
877            }
878            let s = h * self.head_dim;
879            if self.mode.quant_k() {
880                Self::quant_row(
881                    &k_new[s..s + self.head_dim],
882                    &self.kcol[h],
883                    &mut self.kq[h],
884                    &mut self.ks[h],
885                    KV_K_GROUP,
886                );
887            } else {
888                self.k[h].extend_from_slice(&k_new[s..s + self.head_dim]);
889            }
890            if self.mode.quant_v() {
891                Self::quant_row(
892                    &v_new[s..s + self.head_dim],
893                    &self.vcol[h],
894                    &mut self.vq[h],
895                    &mut self.vs[h],
896                    self.head_dim,
897                );
898            } else {
899                self.v[h].extend_from_slice(&v_new[s..s + self.head_dim]);
900            }
901        }
902        self.imp.push(0.0);
903        self.seq_len += 1;
904    }
905
906    /// Per-head attention over its own storage: the f32 branch is
907    /// bit-for-bit equal to attention_head() over slices; the q8 branch
908    /// computes score = s_k·⟨q⊙col_k, k_q⟩ and the weighted sum of V in i8
909    /// with f32 accumulation. Returns (output [head_dim], probs [stored]).
910    pub fn attend(&self, q: &[f32], kv_head: usize) -> (Vec<f32>, Vec<f32>) {
911        // Full-context layers only: a trimmed tail is not the whole row.
912        debug_assert!(self.tail.is_none(), "attend() on a sliding-window tail");
913        let hd = self.head_dim;
914        if self.mode == KvMode::F32 {
915            let stored = self.k[kv_head].len() / hd;
916            return crate::attention::attention_head(
917                q,
918                &self.k[kv_head],
919                &self.v[kv_head],
920                hd,
921                stored,
922            );
923        }
924        let stored = self.head_len(kv_head);
925        let scale = 1.0 / (hd as f32).sqrt();
926        let mut scores = vec![0.0f32; stored];
927        if self.mode.quant_k() {
928            let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
929            // q ⊙ col_k — once per call.
930            let kcol = &self.kcol[kv_head];
931            let mut qc = vec![0.0f32; hd];
932            for d in 0..hd {
933                qc[d] = if kcol.is_empty() {
934                    q[d]
935                } else {
936                    q[d] * kcol[d]
937                };
938            }
939            let ng = hd.div_ceil(KV_K_GROUP);
940            for p in 0..stored {
941                let row = &kq[p * hd..(p + 1) * hd];
942                // SAFETY: i8 and u8 share layout; dot_i8_f32 reads the
943                // bytes back as i8.
944                let row_u8 =
945                    unsafe { std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len()) };
946                let mut dot = 0.0f32;
947                for g in 0..ng {
948                    let g0 = g * KV_K_GROUP;
949                    let g1 = (g0 + KV_K_GROUP).min(hd);
950                    dot +=
951                        crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qc[g0..g1]) * ks[p * ng + g];
952                }
953                scores[p] = dot * scale;
954            }
955        } else {
956            let k = &self.k[kv_head];
957            for p in 0..stored {
958                let row = &k[p * hd..(p + 1) * hd];
959                scores[p] = crate::attention::dot_f32(q, row) * scale;
960            }
961        }
962        let max_score = scores.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
963        let mut sum = 0.0f32;
964        for s in scores.iter_mut() {
965            *s = (*s - max_score).exp();
966            sum += *s;
967        }
968        if sum > 0.0 {
969            for s in scores.iter_mut() {
970                *s /= sum;
971            }
972        }
973        let mut acc = vec![0.0f32; hd];
974        if self.mode.quant_v() {
975            let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
976            for p in 0..stored {
977                let w = scores[p] * vs[p];
978                if w.abs() < 1e-12 {
979                    continue;
980                }
981                crate::qtensor::axpy_i8_f32(&mut acc, &vq[p * hd..(p + 1) * hd], w);
982            }
983            let vcol = &self.vcol[kv_head];
984            if !vcol.is_empty() {
985                for d in 0..hd {
986                    acc[d] *= vcol[d];
987                }
988            }
989        } else {
990            let v = &self.v[kv_head];
991            for p in 0..stored {
992                let w = scores[p];
993                if w.abs() < 1e-12 {
994                    continue;
995                }
996                crate::attention::axpy_f32(&mut acc, &v[p * hd..(p + 1) * hd], w);
997            }
998        }
999        (acc, scores)
1000    }
1001
1002    /// Grouped GQA attention: all Q-heads of one KV group in a single
1003    /// pass over the stored K rows and a single pass over the V rows
1004    /// (per-head `attend` re-read the shared group storage
1005    /// heads_per_kv times — roadmap §3 P1). Per-head score order,
1006    /// softmax and V accumulation are IDENTICAL to `attend`, so each
1007    /// head's output is bit-for-bit the same.
1008    ///
1009    /// `q_group`: `[n_heads_in_group × head_dim]` (global head order);
1010    /// `out`: same shape; `imp_acc[0..stored]` accumulates the probabilities
1011    /// of every head (attention importance), matching the caller's former loop.
1012    /// `scale` is the score scale (1/√hd unless the arch overrides);
1013    /// `first` is the earliest visible position — sliding-window layers
1014    /// pass `stored − window` so older rows get zero probability.
1015    /// `sinks` holds one learned sink logit per head of `q_group` (the
1016    /// caller slices `self.sinks` to the group's heads) or is empty for
1017    /// an ordinary softmax.
1018    #[allow(clippy::too_many_arguments)]
1019    pub fn attend_group(
1020        &self,
1021        q_group: &[f32],
1022        kv_head: usize,
1023        out: &mut [f32],
1024        imp_acc: &mut [f32],
1025        scale: f32,
1026        first: usize,
1027        softcap: f32,
1028        sinks: &[f32],
1029    ) {
1030        self.attend_group_upto(
1031            q_group,
1032            kv_head,
1033            out,
1034            imp_acc,
1035            scale,
1036            first,
1037            softcap,
1038            usize::MAX,
1039            sinks,
1040        )
1041    }
1042
1043    /// `attend_group` over the first `upto` stored rows only — what the
1044    /// same call saw when the cache held exactly `upto` rows. A prefill
1045    /// chunk appends all its rows first and then attends every position
1046    /// in parallel; position `i` passes `upto = s0 + i + 1`, which makes
1047    /// its result bit-identical to the sequential append-then-attend.
1048    ///
1049    /// Only the visible rows `[first, stored)` are scored, so a
1050    /// sliding-window decode step costs O(window), not O(context). The
1051    /// rows before `first` get probability exactly 0 — what the former
1052    /// −inf-filled score row produced (exp(−inf) = 0 adds nothing to the
1053    /// max, the sum, V or the importance), so the result is bit-identical
1054    /// to scoring the whole row.
1055    ///
1056    /// Learned sinks (`sinks` non-empty, one per head): the sink logit
1057    /// joins the softmax max and adds exp(sink − max) to the denominator;
1058    /// it has no value row. Identical to appending a value-less column,
1059    /// which is how gpt-oss and MiMo-V2 define it.
1060    #[allow(clippy::too_many_arguments)]
1061    pub fn attend_group_upto(
1062        &self,
1063        q_group: &[f32],
1064        kv_head: usize,
1065        out: &mut [f32],
1066        imp_acc: &mut [f32],
1067        scale: f32,
1068        first: usize,
1069        softcap: f32,
1070        upto: usize,
1071        sinks: &[f32],
1072    ) {
1073        let hd = self.head_dim;
1074        let nheads = q_group.len() / hd;
1075        debug_assert_eq!(out.len(), nheads * hd);
1076        assert!(
1077            sinks.is_empty() || sinks.len() == nheads,
1078            "attend_group: {} sink logits for {nheads} heads",
1079            sinks.len()
1080        );
1081        let stored = if self.mode == KvMode::F32 {
1082            self.k[kv_head].len() / hd
1083        } else {
1084            self.head_len(kv_head)
1085        }
1086        .min(upto);
1087        if stored == 0 {
1088            out.fill(0.0);
1089            return;
1090        }
1091        let first = first.min(stored.saturating_sub(1));
1092        // Visible span: row p ∈ [first, stored) lives at column p − first.
1093        let span = stored - first;
1094
1095        thread_local! {
1096            /// scores [nheads × span] — reused across layers/tokens.
1097            static GQA_SCORES: std::cell::RefCell<Vec<f32>> =
1098                const { std::cell::RefCell::new(Vec::new()) };
1099            /// q ⊙ col_k per head (q8 K mode).
1100            static GQA_QC: std::cell::RefCell<Vec<f32>> =
1101                const { std::cell::RefCell::new(Vec::new()) };
1102        }
1103
1104        GQA_SCORES.with(|sc| {
1105            let mut scores = sc.borrow_mut();
1106            scores.resize(nheads * span, 0.0);
1107
1108            // ── score pass: each stored K row is read ONCE for all heads.
1109            if self.mode.quant_k() {
1110                let (kq, ks) = (&self.kq[kv_head], &self.ks[kv_head]);
1111                let kcol = &self.kcol[kv_head];
1112                let ng = hd.div_ceil(KV_K_GROUP);
1113                GQA_QC.with(|qc| {
1114                    let mut qcb = qc.borrow_mut();
1115                    qcb.resize(nheads * hd, 0.0);
1116                    for h in 0..nheads {
1117                        for d in 0..hd {
1118                            let qv = q_group[h * hd + d];
1119                            qcb[h * hd + d] = if kcol.is_empty() { qv } else { qv * kcol[d] };
1120                        }
1121                    }
1122                    for p in first..stored {
1123                        let row = &kq[p * hd..(p + 1) * hd];
1124                        // SAFETY: i8 and u8 share layout; dot_i8_f32 reads
1125                        // the bytes back as i8.
1126                        let row_u8 = unsafe {
1127                            std::slice::from_raw_parts(row.as_ptr() as *const u8, row.len())
1128                        };
1129                        for h in 0..nheads {
1130                            let qch = &qcb[h * hd..(h + 1) * hd];
1131                            let mut dot = 0.0f32;
1132                            for g in 0..ng {
1133                                let g0 = g * KV_K_GROUP;
1134                                let g1 = (g0 + KV_K_GROUP).min(hd);
1135                                dot += crate::qtensor::dot_i8_f32(&row_u8[g0..g1], &qch[g0..g1])
1136                                    * ks[p * ng + g];
1137                            }
1138                            scores[h * span + (p - first)] = dot * scale;
1139                        }
1140                    }
1141                });
1142            } else {
1143                let k = &self.k[kv_head];
1144                for p in first..stored {
1145                    let row = &k[p * hd..(p + 1) * hd];
1146                    for h in 0..nheads {
1147                        scores[h * span + (p - first)] =
1148                            crate::attention::dot_f32(&q_group[h * hd..(h + 1) * hd], row) * scale;
1149                    }
1150                }
1151            }
1152
1153            // Gemma-2 attention-logit soft-capping: tanh-squash the
1154            // COMPUTED scores before the softmax (every scored row is in
1155            // the window; a learned sink is not a score and is not capped).
1156            if softcap > 0.0 {
1157                for v in scores.iter_mut() {
1158                    *v = softcap * (*v / softcap).tanh();
1159                }
1160            }
1161
1162            // ── per-head softmax (identical to attend / attention_head;
1163            // a sink joins the max and the denominator, never the rows).
1164            for h in 0..nheads {
1165                let s = &mut scores[h * span..(h + 1) * span];
1166                let row_max = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
1167                let sink = sinks.get(h).copied();
1168                let max_score = match sink {
1169                    Some(z) => row_max.max(z),
1170                    None => row_max,
1171                };
1172                let mut sum = 0.0f32;
1173                for v in s.iter_mut() {
1174                    *v = (*v - max_score).exp();
1175                    sum += *v;
1176                }
1177                if let Some(z) = sink {
1178                    sum += (z - max_score).exp();
1179                }
1180                if sum > 0.0 {
1181                    for v in s.iter_mut() {
1182                        *v /= sum;
1183                    }
1184                }
1185            }
1186
1187            // ── value pass: each stored V row is read ONCE for all heads.
1188            out.fill(0.0);
1189            if self.mode.quant_v() {
1190                let (vq, vs) = (&self.vq[kv_head], &self.vs[kv_head]);
1191                for p in first..stored {
1192                    let row = &vq[p * hd..(p + 1) * hd];
1193                    for h in 0..nheads {
1194                        let w = scores[h * span + (p - first)] * vs[p];
1195                        if w.abs() < 1e-12 {
1196                            continue;
1197                        }
1198                        crate::qtensor::axpy_i8_f32(&mut out[h * hd..(h + 1) * hd], row, w);
1199                    }
1200                }
1201                let vcol = &self.vcol[kv_head];
1202                if !vcol.is_empty() {
1203                    for h in 0..nheads {
1204                        for d in 0..hd {
1205                            out[h * hd + d] *= vcol[d];
1206                        }
1207                    }
1208                }
1209            } else {
1210                let v = &self.v[kv_head];
1211                for p in first..stored {
1212                    let row = &v[p * hd..(p + 1) * hd];
1213                    for h in 0..nheads {
1214                        let w = scores[h * span + (p - first)];
1215                        if w.abs() < 1e-12 {
1216                            continue;
1217                        }
1218                        crate::attention::axpy_f32(&mut out[h * hd..(h + 1) * hd], row, w);
1219                    }
1220                }
1221            }
1222
1223            // ── Attention-importance accumulation (Σ probs over heads), same
1224            // head order as the caller's former per-head loop. Rows before
1225            // `first` carry probability 0 and are left untouched.
1226            let n = imp_acc.len().min(stored);
1227            if n > first {
1228                for h in 0..nheads {
1229                    let s = &scores[h * span..(h + 1) * span];
1230                    for (dst, &p) in imp_acc[first..n].iter_mut().zip(s) {
1231                        *dst += p;
1232                    }
1233                }
1234            }
1235        });
1236    }
1237
1238    /// Batched causal attend for a prefill chunk (macOS/AArch64): the
1239    /// cache already holds every chunk row (`s0` old + `b` new). Per
1240    /// Q-head the scores GEMM `Q·Kᵀ` rides the AMX, the causal softmax
1241    /// zeroes the not-yet-visible tail so the `P·V` GEMM needs no
1242    /// mask, and attention importance takes the masked column sums. Same
1243    /// math as the per-position attend; summation order differs
1244    /// (tolerance-class, like the projection GEMMs).
1245    #[cfg(target_arch = "aarch64")]
1246    #[allow(clippy::too_many_arguments)]
1247    pub fn attend_chunk(
1248        &mut self,
1249        q_all: &[f32],
1250        b: usize,
1251        s0: usize,
1252        nh: usize,
1253        heads_per_kv: usize,
1254        hd: usize,
1255        out: &mut [f32],
1256        pool: Option<&crate::pool::Pool>,
1257        scale: f32,
1258        window: Option<usize>,
1259    ) {
1260        let n = s0 + b;
1261        struct SendPtr(*mut f32);
1262        unsafe impl Send for SendPtr {}
1263        unsafe impl Sync for SendPtr {}
1264        impl SendPtr {
1265            fn at(&self, i: usize) -> *mut f32 {
1266                // Method receiver keeps the closure capturing &SendPtr
1267                // (2021 disjoint capture would grab the raw field).
1268                unsafe { self.0.add(i) }
1269            }
1270        }
1271        thread_local! {
1272            static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>)> =
1273                const { std::cell::RefCell::new((Vec::new(), Vec::new(), Vec::new(), Vec::new())) };
1274        }
1275        // The portable NEON GEMM pays dearly for the gathered Bᵀ loads
1276        // of the scores multiply — pack Kᵀ once per (group, chunk) and
1277        // hand it the sequential-B fast path instead. Accelerate keeps
1278        // the no-copy transposed call.
1279        let neon_gemm = cfg!(not(target_os = "macos"))
1280            || std::env::var("CMF_FORCE_NEON_GEMM")
1281                .map(|v| v == "1")
1282                .unwrap_or(false);
1283        SCRATCH.with(|s| {
1284            let mut s = s.borrow_mut();
1285            let (qpanel, scores, aopanel, ktpack) = &mut *s;
1286            // The whole KV-group attends in one GEMM pair: the group's
1287            // Q-heads stack head-major into one tall panel [hpk·b, hd]
1288            // (row hl·b + bi), so each layer costs 2 sgemm calls per
1289            // group instead of 2 per head — fat M keeps the AMX fed.
1290            let m = heads_per_kv * b;
1291            qpanel.resize(m * hd, 0.0);
1292            scores.resize(m * n, 0.0);
1293            aopanel.resize(m * hd, 0.0);
1294            for g in 0..self.num_kv_heads {
1295                let kmat = &self.k[g];
1296                let vmat = &self.v[g];
1297                debug_assert_eq!(kmat.len(), n * hd);
1298                for hl in 0..heads_per_kv {
1299                    let hh = g * heads_per_kv + hl;
1300                    for bi in 0..b {
1301                        qpanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]
1302                            .copy_from_slice(&q_all[bi * nh * hd + hh * hd..][..hd]);
1303                    }
1304                }
1305                if neon_gemm {
1306                    ktpack.resize(hd * n, 0.0);
1307                    for p in 0..n {
1308                        let row = &kmat[p * hd..(p + 1) * hd];
1309                        for (d, &v) in row.iter().enumerate() {
1310                            ktpack[d * n + p] = v;
1311                        }
1312                    }
1313                    // Accelerate threads its own GEMM; the NEON kernel
1314                    // splits the m rows across the pool instead.
1315                    let sp_q = SendPtr(qpanel.as_ptr() as *mut f32);
1316                    let sp_s = SendPtr(scores.as_mut_ptr());
1317                    let kt = &*ktpack;
1318                    let run = |start: usize, end: usize| {
1319                        if end > start {
1320                            // SAFETY: workers write disjoint score rows.
1321                            let a = unsafe {
1322                                std::slice::from_raw_parts(sp_q.at(start * hd), (end - start) * hd)
1323                            };
1324                            let c = unsafe {
1325                                std::slice::from_raw_parts_mut(
1326                                    sp_s.at(start * n),
1327                                    (end - start) * n,
1328                                )
1329                            };
1330                            crate::qtensor::neon_gemm_rm(
1331                                end - start,
1332                                n,
1333                                hd,
1334                                scale,
1335                                a,
1336                                hd,
1337                                kt,
1338                                n,
1339                                false,
1340                                c,
1341                                n,
1342                            );
1343                        }
1344                    };
1345                    match pool {
1346                        Some(p) if m >= 64 => p.run_rows(m, &run),
1347                        _ => run(0, m),
1348                    }
1349                } else {
1350                    crate::qtensor::sgemm_rm(
1351                        m, n, hd, scale, qpanel, hd, kmat, hd, true, scores, n,
1352                    );
1353                }
1354                // Causal softmax, row-parallel (rows are disjoint).
1355                let sp = SendPtr(scores.as_mut_ptr());
1356                let run = |start: usize, end: usize| {
1357                    for r in start..end {
1358                        let allowed = s0 + (r % b) + 1;
1359                        // Sliding-window layers see only the last W of
1360                        // the causal range; the zeroed head contributes
1361                        // nothing to P·V or attention importance.
1362                        let lo = window.map(|w| allowed.saturating_sub(w)).unwrap_or(0);
1363                        // SAFETY: workers cover disjoint row ranges.
1364                        let row = unsafe { std::slice::from_raw_parts_mut(sp.at(r * n), n) };
1365                        crate::attention::softmax_row(&mut row[lo..allowed]);
1366                        row[..lo].fill(0.0);
1367                        row[allowed..].fill(0.0);
1368                    }
1369                };
1370                match pool {
1371                    Some(p) if m >= 64 => p.run_rows(m, &run),
1372                    _ => run(0, m),
1373                }
1374                // Attention importance: masked column sums (probs of the
1375                // zeroed tail contribute nothing, same as the CPU
1376                // per-position accumulate).
1377                let ni = self.imp.len().min(n);
1378                for r in 0..m {
1379                    let al = (s0 + (r % b) + 1).min(ni);
1380                    for (dst, &p) in self.imp[..al].iter_mut().zip(&scores[r * n..r * n + al]) {
1381                        *dst += p;
1382                    }
1383                }
1384                if neon_gemm {
1385                    let sp_s = SendPtr(scores.as_mut_ptr());
1386                    let sp_o = SendPtr(aopanel.as_mut_ptr());
1387                    let run = |start: usize, end: usize| {
1388                        if end > start {
1389                            // SAFETY: workers write disjoint output rows.
1390                            let a = unsafe {
1391                                std::slice::from_raw_parts(sp_s.at(start * n), (end - start) * n)
1392                            };
1393                            let c = unsafe {
1394                                std::slice::from_raw_parts_mut(
1395                                    sp_o.at(start * hd),
1396                                    (end - start) * hd,
1397                                )
1398                            };
1399                            crate::qtensor::neon_gemm_rm(
1400                                end - start,
1401                                hd,
1402                                n,
1403                                1.0,
1404                                a,
1405                                n,
1406                                vmat,
1407                                hd,
1408                                false,
1409                                c,
1410                                hd,
1411                            );
1412                        }
1413                    };
1414                    match pool {
1415                        Some(p) if m >= 64 => p.run_rows(m, &run),
1416                        _ => run(0, m),
1417                    }
1418                } else {
1419                    crate::qtensor::sgemm_rm(
1420                        m, hd, n, 1.0, scores, n, vmat, hd, false, aopanel, hd,
1421                    );
1422                }
1423                for hl in 0..heads_per_kv {
1424                    let hh = g * heads_per_kv + hl;
1425                    for bi in 0..b {
1426                        out[bi * nh * hd + hh * hd..][..hd]
1427                            .copy_from_slice(&aopanel[(hl * b + bi) * hd..(hl * b + bi + 1) * hd]);
1428                    }
1429                }
1430            }
1431        });
1432    }
1433
1434    /// Roll back the last `n_drop` positions (speculative-decode reject).
1435    pub fn truncate_last(&mut self, n_drop: usize) {
1436        self.discard_linear_scratch();
1437        let d = n_drop.min(self.seq_len);
1438        if let Some(b) = self.bounded.as_mut() {
1439            // The ring rolls back through its undo rows; nothing is
1440            // stored per position, so there is nothing else to drop.
1441            let rolled = b.rollback(d);
1442            if rolled < d {
1443                tracing::warn!(
1444                    "bounded anchor: rollback of {d} exceeds the undo depth ({rolled} restored)"
1445                );
1446            }
1447            self.seq_len = b.seen;
1448            return;
1449        }
1450        // A trimmed tail keeps `slack` rows past its window for exactly
1451        // this; a deeper rollback leaves the next query short of rows its
1452        // window reads (no Spark path rolls back more than 2).
1453        if let Some(w) = self.tail
1454            && self.base > 0
1455            && self.seq_len - d < w.saturating_sub(1)
1456        {
1457            tracing::warn!(
1458                "sliding tail: rollback of {d} leaves {} rows under the {w}-row window",
1459                self.seq_len - d
1460            );
1461        }
1462        let mut first_changed = usize::MAX;
1463        for h in 0..self.num_kv_heads {
1464            let keep = self.k[h].len().saturating_sub(d * self.head_dim);
1465            if !self.k[h].is_empty() || !self.v[h].is_empty() {
1466                first_changed = first_changed.min(keep / self.head_dim.max(1));
1467            }
1468            self.k[h].truncate(keep);
1469            self.v[h].truncate(keep);
1470            let ngk = self.head_dim.div_ceil(KV_K_GROUP);
1471            let keep_q = self.kq[h].len().saturating_sub(d * self.head_dim);
1472            self.kq[h].truncate(keep_q);
1473            let keep_vq = self.vq[h].len().saturating_sub(d * self.head_dim);
1474            self.vq[h].truncate(keep_vq);
1475            let keep_ks = self.ks[h].len().saturating_sub(d * ngk);
1476            self.ks[h].truncate(keep_ks);
1477            let keep_vs = self.vs[h].len().saturating_sub(d);
1478            self.vs[h].truncate(keep_vs);
1479        }
1480        self.imp.truncate(self.imp.len().saturating_sub(d));
1481        self.seq_len -= d;
1482        self.kv_rows_changed(first_changed);
1483    }
1484
1485    /// Accumulate attention mass per stored position (summed over heads).
1486    pub fn accumulate_imp(&mut self, probs: &[f32]) {
1487        for (dst, &p) in self.imp.iter_mut().zip(probs) {
1488            *dst += p;
1489        }
1490    }
1491
1492    /// Contiguous keys of one head: `[stored_len × head_dim]`.
1493    pub fn head_keys(&self, kv_head: usize) -> &[f32] {
1494        &self.k[kv_head]
1495    }
1496
1497    pub fn head_values(&self, kv_head: usize) -> &[f32] {
1498        &self.v[kv_head]
1499    }
1500
1501    /// Number of positions actually stored for a head (0 for dead heads).
1502    pub fn head_len(&self, kv_head: usize) -> usize {
1503        let ng = self.head_dim.div_ceil(KV_K_GROUP);
1504        (self.k[kv_head].len() / self.head_dim)
1505            .max(self.ks[kv_head].len() / ng)
1506            .max(self.vs[kv_head].len())
1507    }
1508
1509    /// Clear cache (e.g. on new conversation or task switch).
1510    pub fn clear(&mut self) {
1511        self.kv_rows_changed(0);
1512        for h in 0..self.num_kv_heads {
1513            self.k[h].clear();
1514            self.v[h].clear();
1515            self.kq[h].clear();
1516            self.ks[h].clear();
1517            self.vq[h].clear();
1518            self.vs[h].clear();
1519            self.kcol[h].clear();
1520            self.vcol[h].clear();
1521        }
1522        self.imp.clear();
1523        self.linear_state.clear();
1524        self.discard_linear_scratch();
1525        // Fresh conversation → the pipeline re-arms collection if the
1526        // layer is o1-flagged (landmarks are per-prompt, never reused).
1527        self.o1 = None;
1528        self.o1_error = None;
1529        self.o1_transitioned = false;
1530        // The bounded ring is zeroed in place: its size is a property of
1531        // the file, not of the conversation.
1532        if let Some(b) = self.bounded.as_mut() {
1533            b.clear();
1534        }
1535        self.seq_len = 0;
1536        // `tail` stays: the window is the layer's, not the conversation's.
1537        self.base = 0;
1538        self.generation = next_gen();
1539    }
1540
1541    /// Serialize this layer's state for the wire (versioned, v2): a
1542    /// fixed header `{magic "CMFS", version, operator identity hash64,
1543    /// layer, kind, f16 flag, position}` followed by one record whose
1544    /// shape the kind fixes — per-position K/V (+ importance) for a full
1545    /// layer, the recurrent vector for a linear layer, the insert counter
1546    /// + ring K/V for a bounded anchor. `f16` halves the K/V payloads and
1547    /// is the caller's explicit choice, exactly like the hidden-state
1548    /// wire; recurrent vectors stay f32 whatever the wire dtype (they are
1549    /// the ONLY state a linear layer has — rounding them rounds the whole
1550    /// history).
1551    ///
1552    /// REFUSES rather than travelling half-complete. A cache carrying
1553    /// frozen columns, a Nyström overlay or q8 storage holds state this
1554    /// format does not describe, and shipping the rest would land a
1555    /// plausible-looking cache that answers differently — the failure
1556    /// mode this whole format exists to avoid.
1557    pub fn export_wire(&self, f16: bool) -> Result<Vec<u8>, String> {
1558        if !matches!(self.mode, KvMode::F32) {
1559            return Err("kv export: only the F32 cache is described by this                         format (CMF_KV=q8 stores int8 rows and per-row scales)"
1560                .into());
1561        }
1562        if self.o1.is_some() {
1563            return Err("kv export: an O(1) Nyström overlay is not part of                         this format — the skeletons are irreversible and                         would have to travel with it"
1564                .into());
1565        }
1566        // Frozen columns only exist under q8 storage, which is refused
1567        // above. If one shows up under an F32 cache the format is lying
1568        // about something and the transfer must not proceed.
1569        if self.kcol.iter().any(|c| !c.is_empty()) || self.vcol.iter().any(|c| !c.is_empty()) {
1570            return Err(
1571                "kv export: frozen columns under an F32 cache — refusing to ship \
1572                        a state this format does not describe"
1573                    .into(),
1574            );
1575        }
1576        let mut out = Vec::with_capacity(self.memory_bytes() / if f16 { 2 } else { 1 } + 64);
1577        let u = |v: u32, o: &mut Vec<u8>| o.extend_from_slice(&v.to_le_bytes());
1578        out.extend_from_slice(WIRE_MAGIC);
1579        u(WIRE_VERSION, &mut out);
1580        out.extend_from_slice(&self.wire_identity.to_le_bytes());
1581        u(self.wire_layer, &mut out);
1582        // A trimmed full layer travels as its tail and says where it starts.
1583        let kind = match self.wire_kind {
1584            WireKind::Full if self.base > 0 => WireKind::FullTail,
1585            k => k,
1586        };
1587        out.push(kind as u8);
1588        out.push(u8::from(f16));
1589        out.extend_from_slice(&0u16.to_le_bytes());
1590        // The header carries the absolute depth (== seq_len untrimmed).
1591        out.extend_from_slice(&(self.pos_len() as u64).to_le_bytes());
1592        let push = |xs: &[f32], o: &mut Vec<u8>| {
1593            if f16 {
1594                for &x in xs {
1595                    o.extend_from_slice(&cortiq_core::quant::f32_to_f16(x).to_le_bytes());
1596                }
1597            } else {
1598                for &x in xs {
1599                    o.extend_from_slice(&x.to_le_bytes());
1600                }
1601            }
1602        };
1603        match kind {
1604            WireKind::Full => self.export_full_body(f16, &mut out),
1605            WireKind::FullTail => {
1606                out.extend_from_slice(&(self.base as u64).to_le_bytes());
1607                self.export_full_body(f16, &mut out);
1608            }
1609            WireKind::Linear => {
1610                u(self.num_kv_heads as u32, &mut out);
1611                u(self.head_dim as u32, &mut out);
1612                u(self.linear_state.len() as u32, &mut out);
1613                for &x in &self.linear_state {
1614                    out.extend_from_slice(&x.to_le_bytes());
1615                }
1616            }
1617            WireKind::Bounded => {
1618                let b = self
1619                    .bounded
1620                    .as_ref()
1621                    .ok_or("kv export: bounded wire kind without an installed ring")?;
1622                u(self.num_kv_heads as u32, &mut out);
1623                u(self.head_dim as u32, &mut out);
1624                u(b.window as u32, &mut out);
1625                u(b.len() as u32, &mut out);
1626                u(b.head() as u32, &mut out);
1627                push(&b.ring_k, &mut out);
1628                push(&b.ring_v, &mut out);
1629            }
1630        }
1631        Ok(out)
1632    }
1633
1634    /// The per-position record (the whole legacy wire): f16 flag,
1635    /// seq_len, geometry, recurrent vector, importance, per-head K/V.
1636    fn export_full_body(&self, f16: bool, out: &mut Vec<u8>) {
1637        let u = |v: u32, o: &mut Vec<u8>| o.extend_from_slice(&v.to_le_bytes());
1638        u(u8::from(f16) as u32, out);
1639        u(self.seq_len as u32, out);
1640        u(self.num_kv_heads as u32, out);
1641        u(self.head_dim as u32, out);
1642        u(self.linear_state.len() as u32, out);
1643        // Attention importance is ordinary state: every attention call
1644        // accumulates it and eviction reads it. Leaving it behind would
1645        // hand the far side a cache that forgets the RIGHT positions
1646        // later — a divergence that shows up only under pressure.
1647        u(self.imp.len() as u32, out);
1648        let push = |xs: &[f32], o: &mut Vec<u8>| {
1649            if f16 {
1650                for &x in xs {
1651                    o.extend_from_slice(&cortiq_core::quant::f32_to_f16(x).to_le_bytes());
1652                }
1653            } else {
1654                for &x in xs {
1655                    o.extend_from_slice(&x.to_le_bytes());
1656                }
1657            }
1658        };
1659        for &x in &self.linear_state {
1660            out.extend_from_slice(&x.to_le_bytes());
1661        }
1662        for &x in &self.imp {
1663            out.extend_from_slice(&x.to_le_bytes());
1664        }
1665        for h in 0..self.num_kv_heads {
1666            u(self.k[h].len() as u32, out);
1667            push(&self.k[h], out);
1668            u(self.v[h].len() as u32, out);
1669            push(&self.v[h], out);
1670        }
1671    }
1672
1673    /// Install a peer's state over this layer. The geometry must match the
1674    /// model both sides hold — it is checked, not assumed. Accepts the
1675    /// versioned wire (magic "CMFS") and, for full/linear layers, the old
1676    /// unversioned per-position body.
1677    pub fn import_wire(&mut self, buf: &[u8]) -> Result<(), String> {
1678        if buf.len() >= 4 && &buf[..4] == WIRE_MAGIC {
1679            return self.import_wire_v2(buf);
1680        }
1681        // Legacy (unversioned) wire: no operator tag travels with it.
1682        if !self.linear_wire_allowed {
1683            return Err(
1684                "kv import: Delta linear state cannot use the unversioned cache wire; refusing until the wire carries operator identity".into(),
1685            );
1686        }
1687        if self.bounded.is_some() {
1688            return Err(
1689                "kv import: a bounded anchor takes only the versioned wire (v2) — the \
1690                 unversioned body has no ring record"
1691                    .into(),
1692            );
1693        }
1694        let n = self.import_full_body(buf)?;
1695        if n != buf.len() {
1696            return Err(format!(
1697                "kv import: {} trailing byte(s) after the record",
1698                buf.len() - n
1699            ));
1700        }
1701        Ok(())
1702    }
1703
1704    fn import_wire_v2(&mut self, buf: &[u8]) -> Result<(), String> {
1705        let need = |n: usize, o: usize| -> Result<(), String> {
1706            if o + n > buf.len() {
1707                Err("kv import: truncated header".into())
1708            } else {
1709                Ok(())
1710            }
1711        };
1712        need(28, 0)?;
1713        let version = u32::from_le_bytes(buf[4..8].try_into().unwrap());
1714        if version != WIRE_VERSION {
1715            return Err(format!(
1716                "kv import: wire version {version}, this runtime speaks {WIRE_VERSION}"
1717            ));
1718        }
1719        let identity = u64::from_le_bytes(buf[8..16].try_into().unwrap());
1720        let layer = u32::from_le_bytes(buf[16..20].try_into().unwrap());
1721        let kind = WireKind::from_u8(buf[20])
1722            .ok_or_else(|| format!("kv import: unknown state kind {}", buf[20]))?;
1723        let f16 = buf[21] != 0;
1724        let position = u64::from_le_bytes(buf[24..32].try_into().unwrap()) as usize;
1725        if identity != self.wire_identity {
1726            return Err(format!(
1727                "kv import: peer operator identity {identity:016x} != mine {:016x} — \
1728                 the two sides do not hold the same operator",
1729                self.wire_identity
1730            ));
1731        }
1732        if layer != self.wire_layer {
1733            return Err(format!(
1734                "kv import: record is for layer {layer}, this is layer {}",
1735                self.wire_layer
1736            ));
1737        }
1738        // A full layer takes its trimmed tail too.
1739        let fits = kind == self.wire_kind
1740            || (kind == WireKind::FullTail && self.wire_kind == WireKind::Full);
1741        if !fits {
1742            return Err(format!(
1743                "kv import: record kind {kind:?} does not match this layer's {:?}",
1744                self.wire_kind
1745            ));
1746        }
1747        let mut o = 32usize;
1748        let u32_at = |o: &mut usize| -> Result<u32, String> {
1749            if *o + 4 > buf.len() {
1750                return Err("kv import: truncated record".into());
1751            }
1752            let v = u32::from_le_bytes(buf[*o..*o + 4].try_into().unwrap());
1753            *o += 4;
1754            Ok(v)
1755        };
1756        let need_payload = |n: usize, o: usize| -> Result<(), String> {
1757            if o + n > buf.len() {
1758                Err("kv import: truncated payload".into())
1759            } else {
1760                Ok(())
1761            }
1762        };
1763        match kind {
1764            WireKind::Full => {
1765                let n = self.import_full_body(&buf[o..])?;
1766                o += n;
1767                if self.seq_len != position {
1768                    return Err(format!(
1769                        "kv import: header position {position} != record seq_len {}",
1770                        self.seq_len
1771                    ));
1772                }
1773            }
1774            WireKind::FullTail => {
1775                need_payload(8, o)?;
1776                let base = u64::from_le_bytes(buf[o..o + 8].try_into().unwrap()) as usize;
1777                o += 8;
1778                let n = self.import_full_body(&buf[o..])?;
1779                o += n;
1780                if base.checked_add(self.seq_len) != Some(position) {
1781                    let rows = self.seq_len;
1782                    self.reset_per_position_storage();
1783                    self.seq_len = 0;
1784                    return Err(format!(
1785                        "kv import: tail base {base} + {rows} rows != header position {position}"
1786                    ));
1787                }
1788                self.base = base;
1789            }
1790            WireKind::Linear => {
1791                let heads = u32_at(&mut o)? as usize;
1792                let hd = u32_at(&mut o)? as usize;
1793                if heads != self.num_kv_heads || hd != self.head_dim {
1794                    return Err(format!(
1795                        "kv import: peer sent {heads}×{hd} per position, this layer is {}×{}",
1796                        self.num_kv_heads, self.head_dim
1797                    ));
1798                }
1799                let lin = u32_at(&mut o)? as usize;
1800                need_payload(lin * 4, o)?;
1801                self.linear_state = (0..lin)
1802                    .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1803                    .collect();
1804                o += lin * 4;
1805                self.reset_per_position_storage();
1806                self.seq_len = position;
1807            }
1808            WireKind::Bounded => {
1809                let heads = u32_at(&mut o)? as usize;
1810                let hd = u32_at(&mut o)? as usize;
1811                let window = u32_at(&mut o)? as usize;
1812                let len = u32_at(&mut o)? as usize;
1813                let head = u32_at(&mut o)? as usize;
1814                let b = self
1815                    .bounded
1816                    .as_mut()
1817                    .ok_or("kv import: bounded record for a layer without a ring")?;
1818                if heads != b.num_kv_heads || hd != b.head_dim || window != b.window {
1819                    return Err(format!(
1820                        "kv import: bounded record {heads}×{window}×{hd} does not fit this \
1821                         layer's ring {}×{}×{}",
1822                        b.num_kv_heads, b.window, b.head_dim
1823                    ));
1824                }
1825                if len != position.min(window) || head != position % window {
1826                    return Err(format!(
1827                        "kv import: bounded record len/head {len}/{head} inconsistent with \
1828                         position {position} (window {window})"
1829                    ));
1830                }
1831                let n = b.ring_k.len();
1832                let w = if f16 { 2 } else { 4 };
1833                need_payload(2 * n * w, o)?;
1834                let read = |o: usize, dst: &mut [f32]| {
1835                    for (i, d) in dst.iter_mut().enumerate() {
1836                        let at = o + i * w;
1837                        *d = if f16 {
1838                            cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
1839                                buf[at..at + 2].try_into().unwrap(),
1840                            ))
1841                        } else {
1842                            f32::from_le_bytes(buf[at..at + 4].try_into().unwrap())
1843                        };
1844                    }
1845                };
1846                read(o, &mut b.ring_k);
1847                o += n * w;
1848                read(o, &mut b.ring_v);
1849                o += n * w;
1850                b.seen = position;
1851                // The undo rows describe inserts this side never made.
1852                let snap = b.snapshot();
1853                b.restore(&snap);
1854                self.linear_state = Vec::new();
1855                self.reset_per_position_storage();
1856                self.seq_len = position;
1857            }
1858        }
1859        if o != buf.len() {
1860            return Err(format!(
1861                "kv import: {} trailing byte(s) after the record",
1862                buf.len() - o
1863            ));
1864        }
1865        Ok(())
1866    }
1867
1868    /// Empty every per-position store (keeps the recurrent vector and
1869    /// the bounded ring untouched).
1870    fn reset_per_position_storage(&mut self) {
1871        let heads = self.num_kv_heads;
1872        self.mode = KvMode::F32;
1873        self.k = vec![Vec::new(); heads];
1874        self.v = vec![Vec::new(); heads];
1875        self.kv_rows_changed(0);
1876        self.kq = vec![Vec::new(); heads];
1877        self.ks = vec![Vec::new(); heads];
1878        self.vq = vec![Vec::new(); heads];
1879        self.vs = vec![Vec::new(); heads];
1880        self.kcol = vec![Vec::new(); heads];
1881        self.vcol = vec![Vec::new(); heads];
1882        self.imp = Vec::new();
1883        self.discard_linear_scratch();
1884        self.o1 = None;
1885        self.o1_error = None;
1886        self.o1_transitioned = false;
1887        self.base = 0;
1888        self.generation = next_gen();
1889    }
1890
1891    /// Parse the per-position record (legacy wire body); returns the
1892    /// bytes consumed.
1893    fn import_full_body(&mut self, buf: &[u8]) -> Result<usize, String> {
1894        let mut o = 0usize;
1895        let u32_at = |o: &mut usize| -> Result<u32, String> {
1896            if *o + 4 > buf.len() {
1897                return Err("kv import: truncated header".into());
1898            }
1899            let v = u32::from_le_bytes(buf[*o..*o + 4].try_into().unwrap());
1900            *o += 4;
1901            Ok(v)
1902        };
1903        let f16 = u32_at(&mut o)? != 0;
1904        let seq_len = u32_at(&mut o)? as usize;
1905        let heads = u32_at(&mut o)? as usize;
1906        let hd = u32_at(&mut o)? as usize;
1907        let lin = u32_at(&mut o)? as usize;
1908        let nimp = u32_at(&mut o)? as usize;
1909        if heads != self.num_kv_heads || hd != self.head_dim {
1910            return Err(format!(
1911                "kv import: peer sent {heads}×{hd} per position, this layer is {}×{}",
1912                self.num_kv_heads, self.head_dim
1913            ));
1914        }
1915        let w = if f16 { 2 } else { 4 };
1916        let need = |n: usize, o: usize| -> Result<(), String> {
1917            if o + n > buf.len() {
1918                Err("kv import: truncated payload".into())
1919            } else {
1920                Ok(())
1921            }
1922        };
1923        need(lin * 4, o)?;
1924        self.linear_state = (0..lin)
1925            .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1926            .collect();
1927        o += lin * 4;
1928        need(nimp * 4, o)?;
1929        let imp: Vec<f32> = (0..nimp)
1930            .map(|i| f32::from_le_bytes(buf[o + i * 4..o + i * 4 + 4].try_into().unwrap()))
1931            .collect();
1932        o += nimp * 4;
1933        let mut k: Vec<Vec<f32>> = Vec::with_capacity(heads);
1934        let mut v: Vec<Vec<f32>> = Vec::with_capacity(heads);
1935        for _ in 0..heads {
1936            for which in 0..2 {
1937                let n = u32_at(&mut o)? as usize;
1938                need(n * w, o)?;
1939                let xs: Vec<f32> = (0..n)
1940                    .map(|i| {
1941                        let at = o + i * w;
1942                        if f16 {
1943                            cortiq_core::quant::f16_to_f32(u16::from_le_bytes(
1944                                buf[at..at + 2].try_into().unwrap(),
1945                            ))
1946                        } else {
1947                            f32::from_le_bytes(buf[at..at + 4].try_into().unwrap())
1948                        }
1949                    })
1950                    .collect();
1951                o += n * w;
1952                if which == 0 { k.push(xs) } else { v.push(xs) }
1953            }
1954        }
1955        self.reset_per_position_storage();
1956        self.k = k;
1957        self.v = v;
1958        self.imp = imp;
1959        self.seq_len = seq_len;
1960        Ok(o)
1961    }
1962
1963    pub fn memory_bytes(&self) -> usize {
1964        let floats: usize = self.k.iter().map(Vec::len).sum::<usize>()
1965            + self.v.iter().map(Vec::len).sum::<usize>()
1966            + self.ks.iter().map(Vec::len).sum::<usize>()
1967            + self.vs.iter().map(Vec::len).sum::<usize>()
1968            + self.kcol.iter().map(Vec::len).sum::<usize>()
1969            + self.vcol.iter().map(Vec::len).sum::<usize>();
1970        let bytes: usize = self.kq.iter().map(Vec::len).sum::<usize>()
1971            + self.vq.iter().map(Vec::len).sum::<usize>();
1972        floats * std::mem::size_of::<f32>()
1973            + bytes
1974            // O(1) recurrent state of linear-core layers (vmf_phase/GDN):
1975            // constant in context, but real memory — the honest "KV+state"
1976            // line must count it (a pure-linear model reported 0 before).
1977            + self.linear_state.len() * std::mem::size_of::<f32>()
1978            // O(1) Nyström state (window + sinks + skeleton) — same
1979            // discipline: constant in context, but real memory.
1980            + self.o1_memory_bytes()
1981            // Natively bounded anchor: the ring is the whole state of the
1982            // layer, fixed by the header.
1983            + self.bounded_state_bytes()
1984    }
1985
1986    /// Drop oldest positions, keeping the last `keep_last`.
1987    fn evict(&mut self, keep_last: usize) {
1988        // A collecting o1 layer owns a still-needed exact prefix and query
1989        // trace.  Evicting it would lower the effective seal boundary while
1990        // leaving q_buf untouched, so conversion could never match its KV
1991        // rows.  A sealed layer stores nothing per position — the Nyström
1992        // state IS the eviction policy; resetting seq_len here would lie
1993        // about the context depth.  Both states therefore bypass ordinary
1994        // eviction until the transition or explicit reset completes.  A
1995        // bounded anchor likewise: the ring evicts itself every token.
1996        // A sliding-window tail (`trim_window`) bounds itself and must stay
1997        // the exact contiguous window its attends index.
1998        if self.o1.is_some()
1999            || self.bounded.is_some()
2000            || self.tail.is_some()
2001            || self.seq_len <= keep_last
2002        {
2003            return;
2004        }
2005        self.generation = next_gen();
2006        let drop = self.seq_len - keep_last;
2007        for h in 0..self.num_kv_heads {
2008            // Dead heads store fewer positions; drop proportionally.
2009            let stored = self.head_len(h);
2010            let d = drop.min(stored);
2011            let hd = self.head_dim;
2012            fn drop_front<T>(v: &mut Vec<T>, n: usize) {
2013                let n = n.min(v.len());
2014                v.drain(..n);
2015            }
2016            drop_front(&mut self.k[h], d * hd);
2017            drop_front(&mut self.v[h], d * hd);
2018            self.kv_rows_changed(0);
2019            drop_front(&mut self.kq[h], d * hd);
2020            drop_front(&mut self.vq[h], d * hd);
2021            drop_front(&mut self.ks[h], d * hd.div_ceil(KV_K_GROUP));
2022            drop_front(&mut self.vs[h], d);
2023        }
2024        let d = drop.min(self.imp.len());
2025        self.imp.drain(..d);
2026        self.seq_len = keep_last;
2027    }
2028
2029    /// Mass-based eviction: keep `sink` earliest positions (attention sinks),
2030    /// the `recent` latest, and fill the rest of the `keep_last` budget
2031    /// with the positions carrying the highest accumulated attention
2032    /// mass (vmfcore: PPL 8.342 vs 8.687 for recency-only, full 8.295).
2033    fn evict_born(&mut self, keep_last: usize, sink: usize, recent: usize) {
2034        if self.o1.is_some() || self.bounded.is_some() || self.tail.is_some() {
2035            // See evict(): collecting must retain the exact prefix as well as
2036            // sealed O(1) state must retain its own bounded representation;
2037            // a bounded anchor's ring is its own eviction, and a sliding
2038            // tail its own (a gather would make its rows non-contiguous).
2039            return;
2040        }
2041        let stored = self.imp.len();
2042        if stored <= keep_last {
2043            return;
2044        }
2045        self.generation = next_gen();
2046        // Budget discipline: sinks first, recents next, both clamped so
2047        // the total never exceeds keep_last.
2048        let sink_n = sink.min(keep_last);
2049        let recent_n = recent.min(keep_last - sink_n);
2050        let mut keep = vec![false; stored];
2051        for k in keep.iter_mut().take(sink_n) {
2052            *k = true;
2053        }
2054        for k in keep.iter_mut().skip(stored.saturating_sub(recent_n)) {
2055            *k = true;
2056        }
2057        let mut budget = keep_last.saturating_sub(keep.iter().filter(|&&x| x).count());
2058        // Highest accumulated mass first among the middle positions.
2059        let mut order: Vec<usize> = (0..stored).filter(|&i| !keep[i]).collect();
2060        order.sort_by(|&a, &b| {
2061            self.imp[b]
2062                .partial_cmp(&self.imp[a])
2063                .unwrap_or(std::cmp::Ordering::Equal)
2064        });
2065        for i in order {
2066            if budget == 0 {
2067                break;
2068            }
2069            keep[i] = true;
2070            budget -= 1;
2071        }
2072
2073        let kept: Vec<usize> = (0..stored).filter(|&i| keep[i]).collect();
2074        let hd = self.head_dim;
2075        fn gather<T: Copy>(src: &[T], kept: &[usize], step: usize) -> Vec<T> {
2076            let mut out = Vec::with_capacity(kept.len() * step);
2077            for &i in kept {
2078                out.extend_from_slice(&src[i * step..(i + 1) * step]);
2079            }
2080            out
2081        }
2082        // Each storage is gathered INDEPENDENTLY: in mixed modes
2083        // (q8k/q8v) K and V live in different storages — the paired branch
2084        // panicked (q8v) or silently left V uncompressed (q8k);
2085        // found by adversarial review, closed by regression tests.
2086        for h in 0..self.num_kv_heads {
2087            if !self.k[h].is_empty() {
2088                self.k[h] = gather(&self.k[h], &kept, hd);
2089                self.kv_rows_changed(0);
2090            }
2091            if !self.v[h].is_empty() {
2092                self.v[h] = gather(&self.v[h], &kept, hd);
2093                self.kv_rows_changed(0);
2094            }
2095            if !self.kq[h].is_empty() {
2096                self.kq[h] = gather(&self.kq[h], &kept, hd);
2097                self.ks[h] = gather(&self.ks[h], &kept, hd.div_ceil(KV_K_GROUP));
2098            }
2099            if !self.vq[h].is_empty() {
2100                self.vq[h] = gather(&self.vq[h], &kept, hd);
2101                self.vs[h] = gather(&self.vs[h], &kept, 1);
2102            }
2103        }
2104        self.imp = kept.iter().map(|&i| self.imp[i]).collect();
2105        self.seq_len = kept.len();
2106    }
2107}
2108
2109/// Eviction policy for a bounded cache.
2110#[derive(Debug, Clone, Copy, PartialEq, Eq)]
2111pub enum EvictionPolicy {
2112    /// Sliding window: keep only the most recent positions.
2113    Recent,
2114    /// Mass-based eviction: sinks + recents + top accumulated attention mass.
2115    Born { sink: usize },
2116}
2117
2118/// Full KV cache for all layers.
2119#[derive(Debug)]
2120pub struct KvCache {
2121    pub layers: Vec<LayerKvCache>,
2122    pub max_seq_len: usize,
2123    pub policy: EvictionPolicy,
2124}
2125
2126impl KvCache {
2127    pub fn new(
2128        num_layers: usize,
2129        num_kv_heads: usize,
2130        head_dim: usize,
2131        max_seq_len: usize,
2132    ) -> Self {
2133        let layers = (0..num_layers)
2134            .map(|li| {
2135                let mut l = LayerKvCache::new(num_kv_heads, head_dim);
2136                l.wire_layer = li as u32;
2137                l
2138            })
2139            .collect();
2140        Self {
2141            layers,
2142            max_seq_len,
2143            policy: EvictionPolicy::Born { sink: 4 },
2144        }
2145    }
2146
2147    pub fn clear(&mut self) {
2148        for layer in &mut self.layers {
2149            layer.clear();
2150        }
2151    }
2152
2153    pub fn total_memory_bytes(&self) -> usize {
2154        self.layers.iter().map(|l| l.memory_bytes()).sum()
2155    }
2156
2157    /// Bytes owned by linear-core recurrent state (including the tentative
2158    /// speculative scratch).  This is reported separately from attention KV
2159    /// so a serving slot's O(1) capacity can be compared with its context
2160    /// cache without guessing from model geometry.
2161    pub fn recurrent_state_bytes(&self) -> usize {
2162        let floats: usize = self
2163            .layers
2164            .iter()
2165            .map(|l| l.linear_state.len() + l.linear_scratch.len())
2166            .sum();
2167        floats * std::mem::size_of::<f32>()
2168    }
2169
2170    /// Attention KV (or sealed O(1) attention state) bytes, excluding the
2171    /// linear recurrent vectors returned by [`recurrent_state_bytes`].
2172    pub fn attention_state_bytes(&self) -> usize {
2173        self.total_memory_bytes()
2174            .saturating_sub(self.recurrent_state_bytes())
2175    }
2176
2177    /// Current sequence length (max across layers — dead layers may lag):
2178    /// the absolute depth, which a trimmed sliding layer keeps in
2179    /// `pos_len()` while it stores only its tail.
2180    pub fn seq_len(&self) -> usize {
2181        self.layers.iter().map(|l| l.pos_len()).max().unwrap_or(0)
2182    }
2183
2184    /// Bytes owned by bounded-anchor rings (constant in context).
2185    pub fn bounded_state_bytes(&self) -> usize {
2186        self.layers.iter().map(|l| l.bounded_state_bytes()).sum()
2187    }
2188
2189    /// True when some layer with PER-POSITION storage reached the cap.
2190    /// Bounded anchors hold nothing per position and never need it: a
2191    /// model whose every layer is O(1) has no eviction cliff at all. A
2192    /// sliding tail (`trim_window`) bounds itself and is not counted.
2193    pub fn needs_eviction(&self) -> bool {
2194        self.layers
2195            .iter()
2196            .filter(|l| l.bounded.is_none() && l.tail.is_none())
2197            .map(|l| l.seq_len)
2198            .max()
2199            .unwrap_or(0)
2200            >= self.max_seq_len
2201    }
2202
2203    /// Evict down to `keep_last` positions according to the policy.
2204    pub fn evict(&mut self, keep_last: usize) {
2205        match self.policy {
2206            EvictionPolicy::Recent => {
2207                for layer in &mut self.layers {
2208                    layer.evict(keep_last);
2209                }
2210            }
2211            EvictionPolicy::Born { sink } => {
2212                let recent = (keep_last / 2).max(1);
2213                for layer in &mut self.layers {
2214                    layer.evict_born(keep_last, sink, recent);
2215                }
2216            }
2217        }
2218    }
2219}
2220
2221#[cfg(test)]
2222mod tests {
2223    use super::*;
2224
2225    /// The lazy K/V maximum the wgpu prefill's f16 guard reads: rows a
2226    /// truncation drops and later appends rewrite are scanned again, an
2227    /// inf or NaN reports +inf, `clear` starts afresh, and a caller's own
2228    /// maxima of the rows it just appended fold in only when they continue
2229    /// the scanned prefix.
2230    #[test]
2231    fn kv_abs_max_follows_appends_truncation_and_clear() {
2232        let (heads, hd) = (2usize, 4usize);
2233        let mut c = LayerKvCache::new(heads, hd);
2234        c.mode = KvMode::F32;
2235        let row = |x: f32| vec![x; heads * hd];
2236        for _ in 0..4 {
2237            c.append(&row(0.5), &row(-0.25), &[]);
2238        }
2239        assert_eq!(c.kv_abs_max(), (0.5, 0.25));
2240        c.truncate_last(2);
2241        c.append(&row(-9.0), &row(3.0), &[]);
2242        c.append(&row(0.5), &row(0.25), &[]);
2243        assert_eq!(c.kv_abs_max(), (9.0, 3.0), "rewritten rows were not rescanned");
2244        c.append(&row(f32::NAN), &row(0.0), &[]);
2245        assert_eq!(c.kv_abs_max().0, f32::INFINITY);
2246        c.clear();
2247        c.append(&row(1.0), &row(2.0), &[]);
2248        assert_eq!(c.kv_abs_max(), (1.0, 2.0));
2249        c.append(&row(4.0), &row(1.0), &[]);
2250        // Rows [1, 2) continue the scanned prefix: the caller's maxima.
2251        assert_eq!(c.kv_abs_max_after((1, 2, 4.0, 1.0)), (4.0, 2.0));
2252        c.append(&row(8.0), &row(1.0), &[]);
2253        // A stale `from` falls back to the scan.
2254        assert_eq!(c.kv_abs_max_after((0, 3, 0.0, 0.0)), (8.0, 2.0));
2255    }
2256
2257    /// A trim drops rows from the front: `kv_abs_max` must still scan every
2258    /// row appended after it, also when more are appended than were dropped
2259    /// before the next call.
2260    #[test]
2261    fn kv_abs_max_scans_rows_appended_after_a_trim() {
2262        let (heads, hd, w) = (1usize, 4usize, 8usize);
2263        let mut c = LayerKvCache::new(heads, hd);
2264        c.mode = KvMode::F32;
2265        let row = |x: f32| vec![x; heads * hd];
2266        for _ in 0..(2 * w + 4) {
2267            c.append(&row(0.5), &row(0.5), &[]);
2268        }
2269        assert_eq!(c.kv_abs_max(), (0.5, 0.5));
2270        let d = c.trim_window(w, 2, 2);
2271        assert!(d > 0, "the trim must drop rows for this test");
2272        // The first row after the trim is large, then more rows than were
2273        // dropped: without moving the mark, rows [stored - d, old mark) —
2274        // the large one among them — would go unscanned.
2275        let first_new = c.kv_rows();
2276        c.append(&row(-7.0), &row(6.0), &[]);
2277        for _ in 0..(d + 3) {
2278            c.append(&row(0.5), &row(0.5), &[]);
2279        }
2280        assert!(c.kv_rows() > 2 * w + 4 && first_new < 2 * w + 4);
2281        assert_eq!(c.kv_abs_max(), (7.0, 6.0));
2282    }
2283
2284    #[test]
2285    fn memory_breakdown_separates_recurrent_and_attention_state() {
2286        let mut cache = KvCache::new(1, 1, 4, 16);
2287        cache.layers[0].linear_state = vec![0.0; 8];
2288        cache.layers[0].linear_scratch = vec![0.0; 4];
2289        cache.layers[0].append(&[0.0; 4], &[1.0; 4], &[true]);
2290        let recurrent = cache.recurrent_state_bytes();
2291        assert_eq!(recurrent, 12 * std::mem::size_of::<f32>());
2292        assert_eq!(
2293            cache.attention_state_bytes() + recurrent,
2294            cache.total_memory_bytes()
2295        );
2296        assert!(cache.attention_state_bytes() > 0);
2297    }
2298
2299    #[test]
2300    fn wire_round_trip_reproduces_attention() {
2301        // The state has to arrive as state, not as something that looks
2302        // like it: the oracle is what the layer ANSWERS, not what it
2303        // stores. Same query, same output, bit for bit.
2304        let (heads, hd) = (2usize, 4usize);
2305        let mut a = LayerKvCache::new(heads, hd);
2306        for p in 0..5 {
2307            let k: Vec<f32> = (0..heads * hd)
2308                .map(|i| (p * 10 + i) as f32 * 0.031)
2309                .collect();
2310            let v: Vec<f32> = (0..heads * hd)
2311                .map(|i| (p * 7 + i) as f32 * -0.017)
2312                .collect();
2313            a.append(&k, &v, &[true, true]);
2314        }
2315        a.linear_state = vec![0.5, -0.25, 1.0];
2316        let q: Vec<f32> = (0..hd).map(|i| 0.1 * (i as f32 + 1.0)).collect();
2317
2318        let bytes = a.export_wire(false).expect("f32 cache exports");
2319        let mut b = LayerKvCache::new(heads, hd);
2320        b.linear_scratch = vec![9.0; 3];
2321        b.import_wire(&bytes).expect("import");
2322        assert!(
2323            b.linear_scratch.is_empty(),
2324            "import must discard tentative state"
2325        );
2326
2327        assert_eq!(b.seq_len, a.seq_len);
2328        assert_eq!(b.linear_state, a.linear_state);
2329        for h in 0..heads {
2330            let (oa, sa) = a.attend(&q, h);
2331            let (ob, sb) = b.attend(&q, h);
2332            assert_eq!(oa, ob, "head {h} attention output diverged");
2333            assert_eq!(sa, sb, "head {h} attention scores diverged");
2334        }
2335    }
2336
2337    #[test]
2338    fn wire_refuses_what_it_cannot_describe() {
2339        // A refusal is the feature: a cache whose extra state this format
2340        // does not carry must not travel looking complete.
2341        let mut c = LayerKvCache::new(1, 4);
2342        c.mode = KvMode::Q8 { k: true, v: true };
2343        let err = c.export_wire(false).unwrap_err();
2344        assert!(err.contains("F32"), "{err}");
2345    }
2346
2347    #[test]
2348    fn wire_refuses_unversioned_delta_state() {
2349        // The versioned wire carries the operator identity, so a delta
2350        // layer exports; only the OLD unversioned body is refused.
2351        let mut c = LayerKvCache::new(1, 4);
2352        c.set_linear_wire_allowed(false);
2353        let bytes = c.export_wire(false).expect("v2 export carries identity");
2354        assert_eq!(&bytes[..4], WIRE_MAGIC);
2355        let err = c.import_wire(&[0, 0, 0, 0]).unwrap_err();
2356        assert!(err.contains("operator identity"), "{err}");
2357    }
2358
2359    #[test]
2360    fn wire_v2_round_trips_linear_and_bounded_records() {
2361        // Linear record: the recurrent vector travels f32 whatever the
2362        // wire dtype and the per-position stores come back empty.
2363        let mut a = LayerKvCache::new(1, 4);
2364        a.wire_kind = WireKind::Linear;
2365        a.wire_identity = 0xC0FFEE;
2366        a.linear_state = vec![0.5, -0.25, 1.0, 3.5];
2367        a.seq_len = 9;
2368        let bytes = a.export_wire(true).unwrap();
2369        let mut b = LayerKvCache::new(1, 4);
2370        b.wire_kind = WireKind::Linear;
2371        b.wire_identity = 0xC0FFEE;
2372        b.import_wire(&bytes).unwrap();
2373        assert_eq!(b.linear_state, a.linear_state);
2374        assert_eq!(b.seq_len, 9);
2375        // Identity mismatch is a refusal, not a warning.
2376        let mut c = LayerKvCache::new(1, 4);
2377        c.wire_kind = WireKind::Linear;
2378        let err = c.import_wire(&bytes).unwrap_err();
2379        assert!(err.contains("operator identity"), "{err}");
2380
2381        // Bounded record, f32 and f16 rings.
2382        let (kvh, hd, w) = (2, 4, 8);
2383        let mut a = LayerKvCache::new(kvh, hd);
2384        a.install_bounded(w);
2385        let k: Vec<f32> = (0..kvh * hd).map(|i| i as f32 * 0.125).collect();
2386        for p in 0..11 {
2387            a.bounded.as_mut().unwrap().insert(&k, &k);
2388            a.seq_len = p + 1;
2389        }
2390        for f16 in [false, true] {
2391            let bytes = a.export_wire(f16).unwrap();
2392            let mut b = LayerKvCache::new(kvh, hd);
2393            b.install_bounded(w);
2394            b.import_wire(&bytes).unwrap();
2395            let (ra, rb) = (a.bounded.as_ref().unwrap(), b.bounded.as_ref().unwrap());
2396            assert_eq!(rb.seen, 11);
2397            assert_eq!(b.seq_len, 11);
2398            // 0.125 multiples are exact in f16, so both dtypes round-trip bit for bit.
2399            assert!(ra.same_state(rb), "f16={f16}");
2400            // A ring of another width refuses the record.
2401            let mut c = LayerKvCache::new(kvh, hd);
2402            c.install_bounded(w * 2);
2403            assert!(c.import_wire(&bytes).is_err());
2404        }
2405    }
2406
2407    #[test]
2408    fn wire_import_checks_geometry() {
2409        let a = LayerKvCache::new(2, 4);
2410        let bytes = a.export_wire(false).unwrap();
2411        let mut wrong = LayerKvCache::new(2, 8);
2412        let err = wrong.import_wire(&bytes).unwrap_err();
2413        assert!(err.contains("2×4"), "{err}");
2414    }
2415
2416    #[test]
2417    fn append_tracks_seq_len_and_layout() {
2418        let mut cache = LayerKvCache::new(4, 8);
2419        cache.mode = KvMode::F32;
2420        assert_eq!(cache.seq_len, 0);
2421
2422        let k: Vec<f32> = (0..32).map(|i| i as f32).collect();
2423        let v = vec![2.0f32; 32];
2424        cache.append(&k, &v, &[true; 4]);
2425
2426        assert_eq!(cache.seq_len, 1);
2427        assert_eq!(cache.head_len(0), 1);
2428        // head 1 slice is contiguous and equals its part of k_new
2429        assert_eq!(cache.head_keys(1), &k[8..16]);
2430        assert_eq!(cache.memory_bytes(), 256);
2431    }
2432
2433    #[test]
2434    fn dead_head_stores_nothing() {
2435        let mut cache = LayerKvCache::new(2, 4);
2436        cache.mode = KvMode::F32;
2437        let k = vec![1.0f32; 8];
2438        let v = vec![2.0f32; 8];
2439        cache.append(&k, &v, &[true, false]);
2440        cache.append(&k, &v, &[true, false]);
2441
2442        assert_eq!(cache.seq_len, 2);
2443        assert_eq!(cache.head_len(0), 2);
2444        assert_eq!(cache.head_len(1), 0, "dead head must not store KV");
2445        assert_eq!(cache.memory_bytes(), 2 * 2 * 4 * 4);
2446    }
2447
2448    #[test]
2449    fn eviction_keeps_recent() {
2450        let mut cache = KvCache::new(2, 4, 8, 10);
2451        cache.policy = EvictionPolicy::Recent;
2452        for l in &mut cache.layers {
2453            l.mode = KvMode::F32;
2454        }
2455        let k = vec![1.0f32; 32];
2456        let v = vec![2.0f32; 32];
2457        for _ in 0..8 {
2458            for layer in &mut cache.layers {
2459                layer.append(&k, &v, &[true; 4]);
2460            }
2461        }
2462        assert_eq!(cache.seq_len(), 8);
2463        assert!(!cache.needs_eviction());
2464
2465        cache.evict(4);
2466        assert_eq!(cache.seq_len(), 4);
2467        assert_eq!(cache.layers[0].head_len(0), 4);
2468    }
2469
2470    #[test]
2471    fn collecting_o1_eviction_retains_exact_storage_until_boundary() {
2472        const B: usize = 19;
2473        let q = vec![0.1f32; 8];
2474        let k = vec![0.2f32; 4];
2475        let v = vec![0.3f32; 4];
2476
2477        for policy in [EvictionPolicy::Recent, EvictionPolicy::Born { sink: 2 }] {
2478            let mut cache = KvCache::new(1, 1, 4, 6);
2479            cache.policy = policy;
2480            cache.layers[0].mode = KvMode::F32;
2481            cache.layers[0].o1_begin_with_boundary(
2482                4,
2483                8,
2484                2,
2485                crate::nystrom::O1Rect::Aggregate,
2486                Some(B),
2487            );
2488
2489            for pos in 0..B {
2490                {
2491                    let layer = &mut cache.layers[0];
2492                    layer.o1_push_q(&q);
2493                    layer.append(&k, &v, &[]);
2494                }
2495                if pos + 1 < B {
2496                    cache.evict(3);
2497                }
2498            }
2499
2500            let layer = &cache.layers[0];
2501            let rows = B * layer.head_dim;
2502            assert_eq!(layer.seq_len, B, "policy {policy:?} retained depth");
2503            assert_eq!(layer.k[0].len(), rows, "policy {policy:?} K rows");
2504            assert_eq!(layer.v[0].len(), rows, "policy {policy:?} V rows");
2505            assert!(
2506                layer.k[0].capacity() >= rows,
2507                "policy {policy:?} K capacity"
2508            );
2509            assert!(
2510                layer.v[0].capacity() >= rows,
2511                "policy {policy:?} V capacity"
2512            );
2513            let q_capacity = match layer.o1.as_ref() {
2514                Some(O1State::Collecting { q_buf, .. }) => q_buf.capacity(),
2515                other => panic!("policy {policy:?} changed state early: {other:?}"),
2516            };
2517            assert!(
2518                q_capacity >= B * 8,
2519                "policy {policy:?} Q capacity must cover the exact prefix"
2520            );
2521
2522            assert!(cache.layers[0].o1_seal_checked(2).unwrap());
2523            assert_eq!(cache.layers[0].k[0].capacity(), 0, "K released after seal");
2524            assert_eq!(cache.layers[0].v[0].capacity(), 0, "V released after seal");
2525        }
2526    }
2527
2528    #[test]
2529    fn truncate_rolls_back_speculative_positions() {
2530        let mut cache = LayerKvCache::new(2, 4);
2531        cache.mode = KvMode::F32;
2532        for pos in 0..5 {
2533            let k = vec![pos as f32; 8];
2534            let v = vec![pos as f32; 8];
2535            cache.append(&k, &v, &[true; 2]);
2536        }
2537        cache.truncate_last(2);
2538        assert_eq!(cache.seq_len, 3);
2539        assert_eq!(cache.head_len(0), 3);
2540        assert_eq!(cache.head_keys(0)[2 * 4], 2.0, "position 2 survives");
2541    }
2542
2543    /// q8_2f-attend ≈ f32-attend: 100 positions (crosses the field freeze
2544    /// at the 64th), pseudo-random vectors, relative tolerance of the
2545    /// int8 grid. Plus rollback and mass-based eviction on the q8 storage.
2546    #[test]
2547    fn q8_attend_matches_f32_within_grid() {
2548        let (heads, hd) = (2, 32);
2549        let mut f = LayerKvCache::new(heads, hd);
2550        f.mode = KvMode::F32;
2551        let mut q8 = LayerKvCache::new(heads, hd);
2552        q8.mode = KvMode::Q8 { k: true, v: true };
2553
2554        let synth = |p: usize, salt: usize| -> Vec<f32> {
2555            (0..heads * hd)
2556                .map(|i| {
2557                    let x = ((i * 31 + p * 17 + salt * 7 + 3) % 97) as f32 / 97.0 - 0.5;
2558                    // channel structure: even channels ×4 (checks the 2f field)
2559                    if i % 2 == 0 { x * 4.0 } else { x * 0.25 }
2560                })
2561                .collect()
2562        };
2563        for p in 0..100 {
2564            let k = synth(p, 1);
2565            let v = synth(p, 2);
2566            f.append(&k, &v, &[true; 2]);
2567            q8.append(&k, &v, &[true; 2]);
2568        }
2569        let q: Vec<f32> = (0..hd)
2570            .map(|i| ((i * 13 + 5) % 89) as f32 / 89.0 - 0.5)
2571            .collect();
2572        for g in 0..heads {
2573            let (of, pf) = f.attend(&q, g);
2574            let (o8, p8) = q8.attend(&q, g);
2575            let scale = of.iter().fold(0f32, |m, x| m.max(x.abs())).max(1e-6);
2576            for d in 0..hd {
2577                assert!(
2578                    (of[d] - o8[d]).abs() <= scale * 0.03 + 1e-3,
2579                    "g{g} d{d}: f32 {} vs q8 {}",
2580                    of[d],
2581                    o8[d]
2582                );
2583            }
2584            for p in 0..100 {
2585                assert!((pf[p] - p8[p]).abs() < 0.02, "prob p{p}");
2586            }
2587        }
2588        // rollback + eviction live on the q8 storage
2589        q8.truncate_last(30);
2590        assert_eq!(q8.head_len(0), 70);
2591        let imp: Vec<f32> = (0..70).map(|i| i as f32).collect();
2592        q8.accumulate_imp(&imp);
2593        q8.evict_born(20, 2, 8);
2594        assert_eq!(q8.head_len(0), 20);
2595        let (o, _) = q8.attend(&q, 0);
2596        assert!(o.iter().all(|x| x.is_finite()));
2597        // memory: q8 ≈ 1 byte/element + scale per row (vs 4 for f32)
2598        assert!(q8.memory_bytes() * 3 < f.memory_bytes());
2599    }
2600
2601    /// Grouped GQA attend must be bit-identical to per-head attend in
2602    /// every KV mode (it is the same math with rows streamed once).
2603    #[test]
2604    fn attend_group_equals_per_head_attend_bitexact() {
2605        let (kv_heads, hd, hpk) = (2usize, 32usize, 3usize); // 6 Q-heads
2606        for mode in [KvMode::F32, KvMode::Q8 { k: true, v: true }] {
2607            let mut c = LayerKvCache::new(kv_heads, hd);
2608            c.mode = mode;
2609            for p in 0..70 {
2610                let k: Vec<f32> = (0..kv_heads * hd)
2611                    .map(|i| ((i * 31 + p * 17 + 3) % 97) as f32 / 97.0 - 0.5)
2612                    .collect();
2613                let v: Vec<f32> = (0..kv_heads * hd)
2614                    .map(|i| ((i * 13 + p * 29 + 7) % 89) as f32 / 89.0 - 0.5)
2615                    .collect();
2616                c.append(&k, &v, &[true; 2]);
2617            }
2618            let q: Vec<f32> = (0..kv_heads * hpk * hd)
2619                .map(|i| ((i * 11 + 5) % 83) as f32 / 83.0 - 0.5)
2620                .collect();
2621            for g in 0..kv_heads {
2622                let span = g * hpk * hd..(g + 1) * hpk * hd;
2623                let mut out = vec![0f32; hpk * hd];
2624                let mut imp = vec![0f32; 70];
2625                c.attend_group(
2626                    &q[span.clone()],
2627                    g,
2628                    &mut out,
2629                    &mut imp,
2630                    1.0 / (hd as f32).sqrt(),
2631                    0,
2632                    0.0,
2633                    &[],
2634                );
2635                let mut imp_ref = vec![0f32; 70];
2636                for h in 0..hpk {
2637                    let qh = &q[span.start + h * hd..span.start + (h + 1) * hd];
2638                    let (o, probs) = c.attend(qh, g);
2639                    assert_eq!(
2640                        &out[h * hd..(h + 1) * hd],
2641                        &o[..],
2642                        "mode {mode:?} g{g} h{h}: grouped attend must be bit-identical"
2643                    );
2644                    for (dst, &p) in imp_ref.iter_mut().zip(&probs) {
2645                        *dst += p;
2646                    }
2647                }
2648                assert_eq!(
2649                    imp, imp_ref,
2650                    "mode {mode:?} g{g}: attention mass must match"
2651                );
2652            }
2653        }
2654    }
2655
2656    /// Learned sinks (gpt-oss / MiMo-V2): the grouped softmax must equal a
2657    /// softmax over the visible rows PLUS an explicit value-less sink
2658    /// column, computed independently in f64 — at position 0 (one stored
2659    /// row), mid-sequence, and with the window truncating the rows. The
2660    /// importance row must be the row probabilities (the sink's share is
2661    /// nobody's importance).
2662    #[test]
2663    fn sink_attend_matches_explicit_sink_column() {
2664        let (nkv, hd, hpk) = (2usize, 8usize, 3usize);
2665        let rows = 9usize;
2666        let mut c = LayerKvCache::new(nkv, hd);
2667        c.mode = KvMode::F32;
2668        let kv = |r: usize, i: usize, a: usize, m: usize| {
2669            (((r * a + i * 7 + 3) % m) as f32 / m as f32 - 0.5) * 2.0
2670        };
2671        let mut ks = Vec::new();
2672        let mut vs = Vec::new();
2673        for r in 0..rows {
2674            let k: Vec<f32> = (0..nkv * hd).map(|i| kv(r, i, 31, 97)).collect();
2675            let v: Vec<f32> = (0..nkv * hd).map(|i| kv(r, i, 17, 89)).collect();
2676            c.append(&k, &v, &[]);
2677            ks.push(k);
2678            vs.push(v);
2679        }
2680        let q: Vec<f32> = (0..nkv * hpk * hd)
2681            .map(|i| (((i * 11 + 5) % 83) as f32 / 83.0 - 0.5) * 3.0)
2682            .collect();
2683        // Mixed signs and one sink that dominates its head.
2684        let sinks = [0.7f32, -1.3, 2.5, 0.0, -4.0, 6.0];
2685        let scale = 1.0 / (hd as f32).sqrt();
2686        let mut checked = 0usize;
2687        for upto in [1usize, 2, 5, 9] {
2688            for window in [None, Some(3usize), Some(1)] {
2689                let first = window.map(|w| upto.saturating_sub(w)).unwrap_or(0);
2690                for g in 0..nkv {
2691                    let qg = &q[g * hpk * hd..(g + 1) * hpk * hd];
2692                    let sg = &sinks[g * hpk..(g + 1) * hpk];
2693                    let mut out = vec![0f32; hpk * hd];
2694                    let mut imp = vec![0f32; upto];
2695                    c.attend_group_upto(qg, g, &mut out, &mut imp, scale, first, 0.0, upto, sg);
2696                    let mut imp_ref = vec![0f64; upto];
2697                    for h in 0..hpk {
2698                        let qh = &qg[h * hd..(h + 1) * hd];
2699                        // Logits of the visible rows, then the sink column.
2700                        let mut z: Vec<f64> = (first..upto)
2701                            .map(|p| {
2702                                let k = &ks[p][g * hd..(g + 1) * hd];
2703                                qh.iter()
2704                                    .zip(k)
2705                                    .map(|(&a, &b)| a as f64 * b as f64)
2706                                    .sum::<f64>()
2707                                    * scale as f64
2708                            })
2709                            .collect();
2710                        z.push(sg[h] as f64);
2711                        let m = z.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
2712                        let e: Vec<f64> = z.iter().map(|&x| (x - m).exp()).collect();
2713                        let s: f64 = e.iter().sum();
2714                        let p: Vec<f64> = e.iter().map(|&x| x / s).collect();
2715                        // The sink column carries a zero value vector.
2716                        for d in 0..hd {
2717                            let want: f64 = (first..upto)
2718                                .map(|r| p[r - first] * vs[r][g * hd + d] as f64)
2719                                .sum();
2720                            let got = out[h * hd + d] as f64;
2721                            assert!(
2722                                (got - want).abs() < 1e-6,
2723                                "upto {upto} window {window:?} g{g} h{h} d{d}: {got} vs {want}"
2724                            );
2725                        }
2726                        for r in first..upto {
2727                            imp_ref[r] += p[r - first];
2728                        }
2729                        checked += 1;
2730                    }
2731                    for r in 0..upto {
2732                        assert!(
2733                            (imp[r] as f64 - imp_ref[r]).abs() < 1e-6,
2734                            "imp upto {upto} window {window:?} g{g} row {r}: {} vs {}",
2735                            imp[r],
2736                            imp_ref[r]
2737                        );
2738                    }
2739                    // The sink takes real mass: rows sum below 1.
2740                    let row_mass: f32 = imp.iter().sum();
2741                    assert!(row_mass < hpk as f32, "sinks must absorb some mass");
2742                }
2743            }
2744        }
2745        assert_eq!(checked, 4 * 3 * nkv * hpk);
2746    }
2747
2748    /// Scoring only the window's rows must be bit-identical to the former
2749    /// whole-row scoring (−inf outside the window) — with and without a
2750    /// sink, output and importance. The reference is the plain per-head
2751    /// softmax over the same rows written out here.
2752    #[test]
2753    fn windowed_attend_equals_masked_full_row() {
2754        let (nkv, hd, hpk) = (1usize, 16usize, 2usize);
2755        let rows = 40usize;
2756        let mut c = LayerKvCache::new(nkv, hd);
2757        c.mode = KvMode::F32;
2758        for r in 0..rows {
2759            let k: Vec<f32> = (0..hd)
2760                .map(|i| ((r * 13 + i * 5) % 29) as f32 / 29.0 - 0.5)
2761                .collect();
2762            let v: Vec<f32> = (0..hd)
2763                .map(|i| ((r * 7 + i * 3) % 31) as f32 / 31.0 - 0.5)
2764                .collect();
2765            c.append(&k, &v, &[]);
2766        }
2767        let q: Vec<f32> = (0..hpk * hd)
2768            .map(|i| ((i * 19) % 23) as f32 / 23.0 - 0.5)
2769            .collect();
2770        let scale = 0.25f32;
2771        for w in [1usize, 7, 39, 40, 100] {
2772            let first = rows.saturating_sub(w);
2773            let mut out = vec![0f32; hpk * hd];
2774            let mut imp = vec![0f32; rows];
2775            c.attend_group(&q, 0, &mut out, &mut imp, scale, first, 0.0, &[]);
2776            // Reference: the historical full-row −inf-masked kernel.
2777            let mut out_ref = vec![0f32; hpk * hd];
2778            let mut imp_ref = vec![0f32; rows];
2779            for h in 0..hpk {
2780                let mut s = vec![f32::NEG_INFINITY; rows];
2781                for p in first..rows {
2782                    s[p] = crate::attention::dot_f32(
2783                        &q[h * hd..(h + 1) * hd],
2784                        &c.head_keys(0)[p * hd..(p + 1) * hd],
2785                    ) * scale;
2786                }
2787                let m = s.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
2788                let mut sum = 0f32;
2789                for v in s.iter_mut() {
2790                    *v = (*v - m).exp();
2791                    sum += *v;
2792                }
2793                for v in s.iter_mut() {
2794                    *v /= sum;
2795                }
2796                for p in first..rows {
2797                    if s[p].abs() < 1e-12 {
2798                        continue;
2799                    }
2800                    crate::attention::axpy_f32(
2801                        &mut out_ref[h * hd..(h + 1) * hd],
2802                        &c.head_values(0)[p * hd..(p + 1) * hd],
2803                        s[p],
2804                    );
2805                }
2806                for (d, &p) in imp_ref.iter_mut().zip(&s) {
2807                    *d += p;
2808                }
2809            }
2810            let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
2811            assert_eq!(bits(&out), bits(&out_ref), "out window {w}");
2812            assert_eq!(bits(&imp), bits(&imp_ref), "imp window {w}");
2813        }
2814    }
2815
2816    /// Sinks are weights: a new conversation (clear) and a wire import
2817    /// must keep them.
2818    #[test]
2819    fn sinks_survive_clear_and_wire_import() {
2820        let mut c = LayerKvCache::new(1, 4);
2821        c.mode = KvMode::F32;
2822        c.sinks = Some(vec![0.5, -0.5]);
2823        c.append(&[1.0; 4], &[2.0; 4], &[]);
2824        let wire = c.export_wire(false).unwrap();
2825        c.clear();
2826        assert_eq!(c.sinks.as_deref(), Some(&[0.5f32, -0.5][..]));
2827        c.import_wire(&wire).unwrap();
2828        assert_eq!(c.sinks.as_deref(), Some(&[0.5f32, -0.5][..]));
2829        assert_eq!(c.seq_len, 1);
2830    }
2831
2832    /// Review regression: mass-based eviction in MIXED modes. q8v used to
2833    /// panic (gather over an empty v[h]), q8k silently left raw V
2834    /// uncompressed (stale rows under kept keys + memory leak).
2835    #[test]
2836    fn born_eviction_mixed_modes_stay_consistent() {
2837        for (mk, mv) in [(false, true), (true, false)] {
2838            let mut c = LayerKvCache::new(1, 4);
2839            c.mode = KvMode::Q8 { k: mk, v: mv };
2840            for p in 0..80 {
2841                let k = vec![p as f32 * 0.01; 4];
2842                let v = vec![p as f32; 4];
2843                c.append(&k, &v, &[true]);
2844            }
2845            let imp: Vec<f32> = (0..80).map(|i| i as f32).collect();
2846            c.accumulate_imp(&imp);
2847            let before = c.memory_bytes();
2848            c.evict_born(20, 4, 8); // q8v: used to panic here
2849            assert_eq!(c.head_len(0), 20, "k={mk} v={mv}");
2850            assert!(
2851                c.memory_bytes() < before / 2,
2852                "memory must shrink (k={mk} v={mv})"
2853            );
2854            // V rows match the kept set: the heaviest positions
2855            // (tail 60..79) must be present in the attend output.
2856            let (out, _) = c.attend(&[1.0, 1.0, 1.0, 1.0], 0);
2857            assert!(
2858                out[0] > 30.0,
2859                "V from the kept tail, not the stale head (k={mk} v={mv}, out {})",
2860                out[0]
2861            );
2862        }
2863    }
2864
2865    #[test]
2866    fn born_eviction_keeps_high_mass_position() {
2867        let mut cache = KvCache::new(1, 1, 2, 16);
2868        cache.policy = EvictionPolicy::Born { sink: 1 };
2869        for l in &mut cache.layers {
2870            l.mode = KvMode::F32;
2871        }
2872        let layer = &mut cache.layers[0];
2873        // 8 positions; keys carry the position index so we can verify
2874        // exactly which positions survive the gather.
2875        for pos in 0..8 {
2876            let k = vec![pos as f32; 2];
2877            let v = vec![pos as f32 + 100.0; 2];
2878            layer.append(&k, &v, &[true]);
2879        }
2880        // Position 3 carries the most attention mass.
2881        let mut imp = vec![0.05f32; 8];
2882        imp[3] = 5.0;
2883        layer.accumulate_imp(&imp);
2884
2885        cache.evict(4); // sink 1 + recent 2 + 1 top-mass slot
2886        let layer = &cache.layers[0];
2887        assert_eq!(layer.seq_len, 4);
2888        let kept_keys: Vec<f32> = (0..4).map(|i| layer.head_keys(0)[i * 2]).collect();
2889        assert_eq!(
2890            kept_keys,
2891            vec![0.0, 3.0, 6.0, 7.0],
2892            "kept = sink(0) + mass-top(3) + recent(6,7)"
2893        );
2894        // imp stays aligned with the gathered positions.
2895        assert_eq!(layer.head_len(0), 4);
2896    }
2897
2898    // ── Sliding-window tail (trim_window) ──
2899
2900    /// Deterministic row of `n` values in [-1, 1) for position `p`.
2901    fn swa_row(p: usize, n: usize, salt: usize) -> Vec<f32> {
2902        (0..n)
2903            .map(|i| ((p * 7919 + i * 104_729 + salt * 1_299_709) % 2003) as f32 / 1001.5 - 1.0)
2904            .collect()
2905    }
2906
2907    fn bits(v: &[f32]) -> Vec<u32> {
2908        v.iter().map(|x| x.to_bits()).collect()
2909    }
2910
2911    /// The stored tail is the untrimmed cache's rows from `base` on, in
2912    /// order, with the bounds the trigger promises — a dead head included.
2913    #[test]
2914    fn swa_trim_keeps_contiguous_tail() {
2915        let (nkv, hd, w, slack, align) = (2usize, 8usize, 6usize, 2usize, 4usize);
2916        let alive = [true, false];
2917        let mut full = LayerKvCache::new(nkv, hd);
2918        full.mode = KvMode::F32;
2919        let mut t = full.clone();
2920        let mut trims = 0;
2921        for p in 0..100 {
2922            let (k, v) = (swa_row(p, nkv * hd, 1), swa_row(p, nkv * hd, 2));
2923            full.append(&k, &v, &alive);
2924            t.append(&k, &v, &alive);
2925            if t.trim_window(w, slack, align) > 0 {
2926                trims += 1;
2927                assert!(
2928                    (w + slack..w + slack + align).contains(&t.seq_len),
2929                    "rows after a trim: {}",
2930                    t.seq_len
2931                );
2932            }
2933            assert!(t.seq_len <= 2 * w, "rows {} past the trigger", t.seq_len);
2934            assert_eq!(t.pos_len(), p + 1);
2935            assert_eq!(t.base() % align, 0);
2936            assert_eq!(t.base() + t.seq_len, full.seq_len);
2937            let b = t.base();
2938            assert_eq!(t.head_keys(0), &full.head_keys(0)[b * hd..]);
2939            assert_eq!(t.head_values(0), &full.head_values(0)[b * hd..]);
2940            assert!(t.head_keys(1).is_empty() && t.head_values(1).is_empty());
2941            assert_eq!(t.imp.len(), t.seq_len);
2942        }
2943        assert!(trims > 5, "only {trims} trims in 100 positions");
2944        assert_eq!(t.tail_window(), Some(w));
2945        assert_eq!(full.tail_window(), None);
2946    }
2947
2948    /// Two caches fed the same rows, one trimmed: every windowed attend
2949    /// (`first = rows − w`) gives the same output and importance, bit for
2950    /// bit, under F32 and the q8 cache — the q8 field freezes at 64 rows,
2951    /// before the first trim.
2952    #[test]
2953    fn swa_trimmed_attend_equals_untrimmed() {
2954        let (nkv, hd, hpk) = (2usize, 16usize, 2usize);
2955        for mode in [KvMode::F32, KvMode::Q8 { k: true, v: true }] {
2956            for (w, slack, align) in [(40usize, 2usize, 4usize), (37, 5, 16), (64, 64, 64)] {
2957                let mut full = LayerKvCache::new(nkv, hd);
2958                full.mode = mode;
2959                let mut t = full.clone();
2960                let scale = 0.3f32;
2961                let mut trims = 0;
2962                for p in 0..400 {
2963                    let (k, v) = (swa_row(p, nkv * hd, 3), swa_row(p, nkv * hd, 4));
2964                    full.append(&k, &v, &[]);
2965                    t.append(&k, &v, &[]);
2966                    let q = swa_row(p, nkv * hpk * hd, 5);
2967                    for g in 0..nkv {
2968                        let qg = &q[g * hpk * hd..(g + 1) * hpk * hd];
2969                        let (nf, nt) = (full.head_len(g), t.head_len(g));
2970                        let (mut of, mut ot) = (vec![0f32; hpk * hd], vec![0f32; hpk * hd]);
2971                        let (mut imf, mut imt) = (vec![0f32; nf], vec![0f32; nt]);
2972                        full.attend_group(
2973                            qg,
2974                            g,
2975                            &mut of,
2976                            &mut imf,
2977                            scale,
2978                            nf.saturating_sub(w),
2979                            0.0,
2980                            &[],
2981                        );
2982                        t.attend_group(
2983                            qg,
2984                            g,
2985                            &mut ot,
2986                            &mut imt,
2987                            scale,
2988                            nt.saturating_sub(w),
2989                            0.0,
2990                            &[],
2991                        );
2992                        assert_eq!(bits(&of), bits(&ot), "{mode:?} w {w} pos {p} group {g}");
2993                        assert_eq!(bits(&imf[t.base()..]), bits(&imt), "{mode:?} w {w} pos {p}");
2994                        full.accumulate_imp(&imf);
2995                        t.accumulate_imp(&imt);
2996                    }
2997                    if t.trim_window(w, slack, align) > 0 {
2998                        trims += 1;
2999                    }
3000                    assert_eq!(bits(&full.imp[t.base()..]), bits(&t.imp));
3001                }
3002                assert!(trims > 0, "{mode:?} w {w}: never trimmed");
3003            }
3004        }
3005    }
3006
3007    /// A rollback no deeper than the slack, then new rows: still the
3008    /// untrimmed answer.
3009    #[test]
3010    fn swa_truncate_after_trim() {
3011        let (nkv, hd, w) = (1usize, 8usize, 10usize);
3012        let mut full = LayerKvCache::new(nkv, hd);
3013        full.mode = KvMode::F32;
3014        let mut t = full.clone();
3015        let mut p = 0usize;
3016        for round in 0..60 {
3017            for _ in 0..3 {
3018                let (k, v) = (swa_row(p, hd, 6), swa_row(p, hd, 7));
3019                full.append(&k, &v, &[]);
3020                t.append(&k, &v, &[]);
3021                p += 1;
3022            }
3023            t.trim_window(w, 2, 4);
3024            // Speculative-style rollback of 2, then replacement rows.
3025            full.truncate_last(2);
3026            t.truncate_last(2);
3027            p -= 2;
3028            assert_eq!(t.pos_len(), full.seq_len);
3029            for _ in 0..2 {
3030                let (k, v) = (swa_row(p, hd, 8 + round), swa_row(p, hd, 9 + round));
3031                full.append(&k, &v, &[]);
3032                t.append(&k, &v, &[]);
3033                p += 1;
3034            }
3035            let q = swa_row(p, hd, 10);
3036            let (nf, nt) = (full.head_len(0), t.head_len(0));
3037            let (mut of, mut ot) = (vec![0f32; hd], vec![0f32; hd]);
3038            let (mut imf, mut imt) = (vec![0f32; nf], vec![0f32; nt]);
3039            full.attend_group(&q, 0, &mut of, &mut imf, 0.5, nf - w.min(nf), 0.0, &[]);
3040            t.attend_group(&q, 0, &mut ot, &mut imt, 0.5, nt - w.min(nt), 0.0, &[]);
3041            assert_eq!(bits(&of), bits(&ot), "round {round}");
3042        }
3043        assert!(t.base() > 0);
3044    }
3045
3046    /// The cache-wide eviction leaves a sliding tail alone, and a tail
3047    /// does not count toward the eviction trigger.
3048    #[test]
3049    fn swa_evict_skips_tail() {
3050        for policy in [EvictionPolicy::Recent, EvictionPolicy::Born { sink: 1 }] {
3051            let mut c = KvCache::new(2, 1, 4, 30);
3052            c.policy = policy;
3053            for p in 0..25 {
3054                for l in &mut c.layers {
3055                    l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3056                }
3057                c.layers[0].trim_window(8, 2, 4);
3058            }
3059            let t_rows = c.layers[0].seq_len;
3060            let t_base = c.layers[0].base();
3061            assert!(t_base > 0);
3062            assert!(!c.needs_eviction(), "25 rows < cap 30");
3063            for p in 25..30 {
3064                for l in &mut c.layers {
3065                    l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3066                }
3067            }
3068            // Layer 0 (tail, 5 rows past its last trim) is not what triggers.
3069            assert!(c.needs_eviction());
3070            let gen0 = c.layers[0].generation();
3071            c.evict(10);
3072            assert_eq!(c.layers[0].seq_len, t_rows + 5, "{policy:?}: tail evicted");
3073            assert_eq!(c.layers[0].base(), t_base);
3074            assert_eq!(c.layers[0].generation(), gen0);
3075            assert_eq!(
3076                c.layers[1].seq_len, 10,
3077                "{policy:?}: full layer not evicted"
3078            );
3079            assert_eq!(c.seq_len(), 30, "absolute depth from the tail");
3080        }
3081    }
3082
3083    /// `generation` moves on every mutation a device copy cannot follow by its
3084    /// row count, and stays on appends and rollbacks.
3085    #[test]
3086    fn swa_gen_bumps() {
3087        let mut l = LayerKvCache::new(1, 4);
3088        l.mode = KvMode::F32;
3089        let g0 = l.generation();
3090        assert_ne!(
3091            LayerKvCache::new(1, 4).generation(),
3092            g0,
3093            "fresh caches differ"
3094        );
3095        for p in 0..20 {
3096            l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3097        }
3098        assert_eq!(l.generation(), g0, "append");
3099        l.truncate_last(2);
3100        assert_eq!(l.generation(), g0, "truncate_last");
3101        assert_eq!(l.trim_window(10, 2, 4), 0, "18 rows ≤ 2w: no trim");
3102        assert_eq!(l.generation(), g0, "a no-op trim");
3103        for p in 18..21 {
3104            l.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3105        }
3106        assert!(l.trim_window(10, 2, 4) > 0);
3107        let g1 = l.generation();
3108        assert_ne!(g1, g0, "trim");
3109        let snap = l.clone();
3110        l.clear();
3111        let g2 = l.generation();
3112        assert_ne!(g2, g1, "clear");
3113        assert_eq!((l.base(), l.pos_len(), l.tail_window()), (0, 0, Some(10)));
3114        let mut e = LayerKvCache::new(1, 4);
3115        e.mode = KvMode::F32;
3116        for p in 0..12 {
3117            e.append(&swa_row(p, 4, 1), &swa_row(p, 4, 2), &[]);
3118        }
3119        let ge = e.generation();
3120        e.evict(20);
3121        assert_eq!(e.generation(), ge, "an eviction that does nothing");
3122        e.evict(6);
3123        assert_ne!(e.generation(), ge, "evict");
3124        let ge = e.generation();
3125        e.evict_born(3, 1, 1);
3126        assert_ne!(e.generation(), ge, "evict_born");
3127        let mut i = LayerKvCache::new(1, 4);
3128        let gi = i.generation();
3129        i.import_wire(&snap.export_wire(false).unwrap()).unwrap();
3130        assert_ne!(i.generation(), gi, "import");
3131        assert_ne!(
3132            i.generation(),
3133            snap.generation(),
3134            "an import is new storage"
3135        );
3136    }
3137
3138    /// A trimmed layer travels as `FullTail` and answers the same on the
3139    /// far side; an untrimmed one is still byte-for-byte `Full`; a layer
3140    /// with another kind refuses the tail.
3141    #[test]
3142    fn wire_full_tail_roundtrip() {
3143        let (nkv, hd, w) = (2usize, 4usize, 6usize);
3144        let mut a = LayerKvCache::new(nkv, hd);
3145        a.mode = KvMode::F32;
3146        for p in 0..5 {
3147            a.append(&swa_row(p, nkv * hd, 1), &swa_row(p, nkv * hd, 2), &[]);
3148        }
3149        let plain = a.export_wire(false).unwrap();
3150        assert_eq!(plain[20], WireKind::Full as u8, "untrimmed stays Full");
3151        assert_eq!(u64::from_le_bytes(plain[24..32].try_into().unwrap()), 5);
3152        for p in 5..40 {
3153            a.append(&swa_row(p, nkv * hd, 1), &swa_row(p, nkv * hd, 2), &[]);
3154            a.trim_window(w, 2, 4);
3155        }
3156        assert!(a.base() > 0);
3157        for f16 in [false, true] {
3158            let bytes = a.export_wire(f16).unwrap();
3159            assert_eq!(bytes[20], WireKind::FullTail as u8);
3160            assert_eq!(u64::from_le_bytes(bytes[24..32].try_into().unwrap()), 40);
3161            let mut b = LayerKvCache::new(nkv, hd);
3162            b.import_wire(&bytes).expect("tail import");
3163            assert_eq!(
3164                (b.base(), b.seq_len, b.pos_len()),
3165                (a.base(), a.seq_len, 40)
3166            );
3167            if !f16 {
3168                let q = swa_row(99, 2 * hd, 3);
3169                for g in 0..nkv {
3170                    let n = a.head_len(g);
3171                    let (mut oa, mut ob) = (vec![0f32; 2 * hd], vec![0f32; 2 * hd]);
3172                    let (mut ia, mut ib) = (vec![0f32; n], vec![0f32; n]);
3173                    a.attend_group(&q, g, &mut oa, &mut ia, 0.5, n - w, 0.0, &[]);
3174                    b.attend_group(&q, g, &mut ob, &mut ib, 0.5, n - w, 0.0, &[]);
3175                    assert_eq!(bits(&oa), bits(&ob));
3176                }
3177                // ... and a Full import over it puts row 0 back at 0.
3178                b.import_wire(&plain).unwrap();
3179                assert_eq!((b.base(), b.pos_len()), (0, 5));
3180            }
3181        }
3182        // A position that disagrees with base + rows is refused.
3183        let mut bad = a.export_wire(false).unwrap();
3184        bad[24..32].copy_from_slice(&41u64.to_le_bytes());
3185        assert!(LayerKvCache::new(nkv, hd).import_wire(&bad).is_err());
3186        // A bounded layer does not take a tail.
3187        let mut r = LayerKvCache::new(nkv, hd);
3188        r.install_bounded(8);
3189        let err = r.import_wire(&a.export_wire(false).unwrap()).unwrap_err();
3190        assert!(err.contains("does not match"), "{err}");
3191    }
3192}