Skip to main content

cortiq_engine/
nystrom.rs

1//! Nyström (landmark) attention kernel — streaming per-GQA-group runtime
2//! for long-context `attn_type: nystrom` layers.
3//!
4//! Attention splits into an EXACT sliding window (last `w` keys) and a
5//! landmark-skeleton far field sharing ONE joint denominator:
6//!
7//! ```text
8//! out(q_t) = (Σ_{j>t-w} e_j·v_j + F·M·T_far) / (Σ_{j>t-w} e_j + F·M·Z_far)
9//! e_j   = exp(q_t·k_j/√d)                      exact near weights
10//! F_i   = exp(q_t·k̃_i/√d)                      scores vs landmark keys
11//! M     = pinv_reg(exp(Q̃·K̃ᵀ/√d))               fixed after prefill
12//! T_far = Σ_{j≤t-w} exp(Q̃·k_j/√d)·v_jᵀ         [m × dv]
13//! Z_far = Σ_{j≤t-w} exp(Q̃·k_j/√d)              [m]
14//! ```
15//!
16//! exp(q·k) is a PSD kernel, so the UNNORMALIZED skeleton (classic
17//! Nyström/CUR) is legal.  Do NOT row-softmax the factors and do NOT
18//! normalize the key scores over landmarks — both "simplifications"
19//! measurably collapse quality (validated in the torch matrix probes).
20//!
21//! Boundary discipline: key j enters T/Z at the exact step it LEAVES
22//! the window (t = j+w) — delayed insertion, no overlap, no hole; the
23//! near mass stays exact rather than Nyström-estimated.
24//!
25//! Sink tokens (spec §5b, StreamingLLM discipline): the first `sink`
26//! keys of the sequence are PERMANENT exact keys — the near mask is
27//! (t-j < w) OR (j < sink) — and must never enter the far accumulators.
28//! Here they never enter the ring window in the first place (they live
29//! in a dedicated buffer), so delayed insertion cannot see them: no
30//! double count, no gap.  Measured: sinks make the full 28/28-layer
31//! O(1) conversion viable — the default mode.
32//!
33//! Quality of THIS kernel, measured through it (`cortiq ppl --o1 all`,
34//! Qwen3-0.6B, all 28 layers, m=32 W=128 sink=4, wikitext-2 val, 12×512
35//! windows, landmarks frozen at a 256-token prefill): ×1.296 vs exact
36//! attention over the same scored tokens (28.04 vs 21.63).
37//!
38//! The older ×1.177 figure is NOT this operator: it comes from the torch
39//! matrix probe, which (a) rectifies every per-(t,j) weight — impossible
40//! to stream, the weights are never materialized — (b) builds landmarks
41//! from the FULL sequence rather than the prefill, and (c) averages in
42//! the first W positions, which are pure-exact and cost nothing.  Quote
43//! ×1.296 for the runtime; ×1.177 is an upper bound the runtime cannot
44//! reach by construction.
45//!
46//! fp32 numerics: raw exp overflows on real logits, so shifts are
47//! absorbed into diagonals.  T̂[i]/Ẑ[i] live at scale e^{-m_i} with a
48//! per-landmark running max m_i (flash-style rescale on growth); each
49//! token's landmark row uses its own shift f; near and far are brought
50//! to one common scale before the single joint division.
51
52/// Which rectifier keeps the skeleton's estimated far mass non-negative.
53///
54/// pinv(exp(Q̃K̃ᵀ/√d)) is violently ill-conditioned, so M is indefinite
55/// and the raw skeleton estimates negative weights for a large minority
56/// of keys (measured on Qwen3-0.6B: 23.5% of far weights negative,
57/// carrying 24.5% of the absolute far mass).  Unrectified, the joint
58/// denominator goes near-zero/negative and the model collapses (×510).
59///
60/// The matrix probe rectifies every estimated weight — `west =
61/// ((Fu@Mu)@E).clamp_min(0)`.  A STREAMING kernel cannot do that: the
62/// per-(t, j) weights are never materialized, they exist only already
63/// contracted against the accumulators.  Two streaming-legal stand-ins:
64///
65/// MEASURED (Qwen3-0.6B, all 28 layers, W=128, sink=4, wikitext-2 val,
66/// 12×512 windows, landmarks frozen at a 256-token prefill — i.e. the
67/// runtime's real discipline, `cortiq ppl --o1`):
68///
69/// ```text
70///        m=8              m=16             m=32
71/// agg    28.51 (×1.318)   28.82 (×1.332)   28.04 (×1.296)  ← default
72/// fm     28.97 (×1.340)   29.69 (×1.373)   30.58 (×1.414)
73/// ```
74///
75/// `Aggregate` wins at every m, so it stays the default.  `Fm` is kept
76/// selectable because its per-key guarantee is the intuitively "correct"
77/// fix and someone will re-derive it: this table is the evidence that it
78/// costs quality HERE, and the reason is that the guarantee is bought by
79/// destroying signal — clamping a landmark's coefficient zeroes its
80/// contribution to EVERY far key, including the majority where the
81/// weighted sum was already positive and accurate.
82#[derive(Clone, Copy, Debug, PartialEq, Eq)]
83pub enum O1Rect {
84    /// Clamp only the AGGREGATE far denominator: a row whose skeleton
85    /// denominator comes out negative drops its far field entirely.
86    /// Coarse — negative per-key mass survives untouched whenever the
87    /// row sum happens to stay positive — but measured BEST (see above):
88    /// the surviving negatives are apparently error-cancelling, not
89    /// error-causing.
90    Aggregate,
91    /// Clamp FM = F_u·M_u (an m-vector, per query row) at zero.
92    /// ŵ(t,j) = Σ_b FM[b]·E[b,j] and E = exp(·) ≥ 0 ELEMENTWISE, so
93    /// FM ≥ 0 is SUFFICIENT for every far weight to be non-negative —
94    /// a per-key guarantee bought with O(m) work on a vector the row
95    /// already materializes, state untouched.  It is strictly stronger
96    /// than the probe's clamp (a negative landmark is dropped for every
97    /// key, not only where the sum would go negative), so this is a
98    /// DIFFERENT operator, not an emulation of the matrix reference —
99    /// and, measured, a worse one.  Opt in with `--o1-rect fm`.
100    Fm,
101}
102
103/// Ridge factor for the regularized pseudo-inverse of the landmark
104/// kernel: λ = RIDGE_REL · mean(diag(AᵀA)).
105const RIDGE_REL: f64 = 1e-6;
106/// Floor for the joint denominator (mirrors the reference probe).
107const DEN_EPS: f32 = 1e-30;
108/// Prompts of length ≤ w + EXACT_SLACK skip the skeleton entirely:
109/// tiny prefills duplicate segment-mean landmarks (singular Au).
110const EXACT_SLACK: usize = 8;
111
112/// Streaming Nyström attention state for ONE GQA group.
113///
114/// State splits along the GQA grain, because the operator does:
115///
116/// * SHARED per KV group (`NystromGroup`) — the exact window ring, the
117///   sink buffer and the key landmarks K̃.  Under GQA every Q head of a
118///   group reads the SAME k/v rows, so all three are bit-identical
119///   across the group; storing them once per group instead of once per
120///   Q head is the point of this split (identical arithmetic,
121///   ×heads_per_kv less window memory).  K̃ = seg_means(ks, t, d, m_eff)
122///   is a pure function of the group's keys and of `t` (which fixes
123///   m_eff), so it is shareable for the same reason the keys are.
124/// * PRIVATE per Q head (`NystromHead`) — the far accumulators T̂/Ẑ and
125///   their per-landmark running maxima, the QUERY landmarks Q̃, and the
126///   mixing matrix M = pinv(exp(Q̃K̃ᵀ/√d)).  Q̃ is built from that head's
127///   own queries, so M and the far field it drives are per-Q-head and
128///   cannot be shared: the far mass a head accumulates is contracted
129///   against its own query landmarks.
130///
131/// Lifecycle: `new(m, w, sink)` → `prefill(prompt)` once → `step()` per
132/// decode token (single-head façade), or `new_group`/`prefill_group`/
133/// `step_group` for a whole GQA group at once.  All buffers are flat
134/// `Vec<f32>`, row-major; the skeleton path performs no allocations
135/// inside `step()`.
136#[derive(Clone, Debug)]
137pub struct NystromState {
138    group: NystromGroup,
139    heads: Vec<NystromHead>,
140}
141
142/// The part of the state a GQA group shares: everything derived from
143/// the group's KEYS and VALUES alone (see `NystromState`).
144#[derive(Clone, Debug)]
145struct NystromGroup {
146    /// Landmark budget (m) — effective count may be lower (`m_eff`).
147    m: usize,
148    /// Exact-window width in keys.
149    w: usize,
150    /// Permanent exact sink keys at positions 0..sink (spec §5b).
151    sink: usize,
152    d: usize,
153    dv: usize,
154    /// Effective landmark count: clamp(t/8, 4, m) at prefill.  Derived
155    /// from the prompt length, hence equal for every head of the group.
156    m_eff: usize,
157    /// Short-prompt mode: window holds ALL keys, no skeleton.  The
158    /// buffer grows on decode, so this mode may allocate in `step()` —
159    /// acceptable for the ≤ w+8-token degenerate case.
160    exact_only: bool,
161    scale: f32,
162    /// Window keys `[cap][d]` — ring buffer in skeleton mode (cap = w),
163    /// append-only in exact-only mode.
164    win_k: Vec<f32>,
165    /// Window values `[cap][dv]`.
166    win_v: Vec<f32>,
167    win_len: usize,
168    /// Ring slot of the OLDEST window entry (0 while not yet full).
169    win_head: usize,
170    /// Sink keys `[sink_len][d]` — filled once at prefill, immutable.
171    sink_k: Vec<f32>,
172    /// Sink values `[sink_len][dv]`.
173    sink_v: Vec<f32>,
174    /// Number of stored sink tokens (0 in exact-only mode, where every
175    /// key is permanent-exact anyway).
176    sink_len: usize,
177    /// Key landmarks `[m_eff][d]` — segment means of the group's keys.
178    k_tilde: Vec<f32>,
179}
180
181/// The part of the state that is private to one Q head: everything that
182/// touches that head's QUERIES (see `NystromState`).
183#[derive(Clone, Debug)]
184struct NystromHead {
185    /// How the indefinite skeleton is rectified (see `O1Rect`).
186    rect: O1Rect,
187    /// Far numerator `[m_eff][dv]`, stored at scale e^{-m_max[i]}.
188    t_hat: Vec<f32>,
189    /// Far denominator `[m_eff]`, same scale.
190    z_hat: Vec<f32>,
191    /// Per-landmark running max of far logits q̃_i·k_j/√d.
192    m_max: Vec<f32>,
193    /// Number of keys absorbed into the far field.
194    far_len: usize,
195    /// Query landmarks `[m_eff][d]` (segment means of the prefill).
196    q_tilde: Vec<f32>,
197    /// Regularized pseudo-inverse of Au = exp(Q̃·K̃ᵀ/√d), `[m_eff][m_eff]`.
198    mu: Vec<f32>,
199    // Scratch preallocated at prefill so skeleton-mode step() is
200    // allocation-free.  Per head rather than per group: the heads of a
201    // group write it independently, and it is ~0.5 KB.
202    scr_s: Vec<f32>,
203    scr_fh: Vec<f32>,
204    scr_u: Vec<f32>,
205    scr_l: Vec<f32>,
206}
207
208impl NystromState {
209    /// Single-head state (`heads_per_kv == 1`, and the shape the kernel
210    /// unit tests use).
211    ///
212    /// `m` — landmark budget (≥ 4; see `O1_DEFAULT_M`),
213    /// `w` — exact window width (validated setting is 128),
214    /// `sink` — permanent exact sink keys (validated default is 4;
215    /// 0 reproduces the sink-free kernel bit-for-bit).
216    /// Rectifier defaults to `O1_DEFAULT_RECT`; override with
217    /// `with_rect` (the golden-parity test pins it explicitly).
218    pub fn new(m: usize, w: usize, sink: usize) -> Self {
219        Self::new_group(m, w, sink, 1)
220    }
221
222    /// State for one GQA group of `q_heads` query heads sharing a KV
223    /// head.  The window/sink/K̃ are stored ONCE for the group; each Q
224    /// head keeps its own far field, Q̃ and M.
225    pub fn new_group(m: usize, w: usize, sink: usize, q_heads: usize) -> Self {
226        assert!(m >= 4, "landmark budget must be at least 4");
227        assert!(w >= 1, "window must hold at least one key");
228        assert!(q_heads >= 1, "a GQA group needs at least one query head");
229        NystromState {
230            group: NystromGroup {
231                m,
232                w,
233                sink,
234                d: 0,
235                dv: 0,
236                m_eff: 0,
237                exact_only: true,
238                scale: 0.0,
239                win_k: Vec::new(),
240                win_v: Vec::new(),
241                win_len: 0,
242                win_head: 0,
243                sink_k: Vec::new(),
244                sink_v: Vec::new(),
245                sink_len: 0,
246                k_tilde: Vec::new(),
247            },
248            heads: (0..q_heads).map(|_| NystromHead::new()).collect(),
249        }
250    }
251
252    /// Select the skeleton rectifier for every head of the group
253    /// (builder; see `O1Rect`).
254    pub fn with_rect(mut self, rect: O1Rect) -> Self {
255        for h in &mut self.heads {
256            h.rect = rect;
257        }
258        self
259    }
260
261    /// Query heads in this group.
262    pub fn num_q_heads(&self) -> usize {
263        self.heads.len()
264    }
265
266    /// Keys absorbed into head `head`'s far field.  Exposed for the
267    /// delayed-insertion invariant test: eviction is a GROUP event, but
268    /// each head must absorb the evicted key EXACTLY once, so this must
269    /// equal the number of evictions — never a multiple of it.
270    pub fn far_len(&self, head: usize) -> usize {
271        self.heads[head].far_len
272    }
273
274    /// Absorb the whole prompt for a single-head state — see
275    /// `prefill_group`.
276    pub fn prefill(&mut self, qs: &[f32], ks: &[f32], vs: &[f32], t: usize, d: usize, dv: usize) {
277        assert_eq!(self.heads.len(), 1, "use prefill_group for a GQA group");
278        self.prefill_group(&[qs], ks, vs, t, d, dv);
279    }
280
281    /// Absorb the whole prompt for a GQA group: freeze each head's
282    /// landmarks and M, then replay the prompt through the step() state
283    /// semantics (window fill + delayed far insertion).  `qs[h]` is that
284    /// head's `[t][d]` query block; `ks` is `[t][d]` and `vs` is
285    /// `[t][dv]` — the group's shared keys/values, row-major.
286    pub fn prefill_group(
287        &mut self,
288        qs: &[&[f32]],
289        ks: &[f32],
290        vs: &[f32],
291        t: usize,
292        d: usize,
293        dv: usize,
294    ) {
295        assert_eq!(qs.len(), self.heads.len(), "one query block per head");
296        for q in qs {
297            assert_eq!(q.len(), t * d);
298        }
299        assert_eq!(ks.len(), t * d);
300        assert_eq!(vs.len(), t * dv);
301
302        let Some(k_tilde64) = self.group.prefill_shared(ks, vs, t, d, dv) else {
303            // exact-only: no skeleton, no far field — nothing per head
304            // beyond the score scratch.
305            for h in &mut self.heads {
306                h.seal_exact(t);
307            }
308            return;
309        };
310        for (h, q) in self.heads.iter_mut().zip(qs) {
311            h.seal(&self.group, q, t, &k_tilde64);
312        }
313        // Replay the post-sink prompt ONCE for the group: each key
314        // enters the shared window, evicting the (j-w)-th into every
315        // head's far field.
316        for j in self.group.sink..t {
317            Self::advance(
318                &mut self.group,
319                &mut self.heads,
320                &ks[j * d..(j + 1) * d],
321                &vs[j * dv..(j + 1) * dv],
322            );
323        }
324    }
325
326    /// One decode step for a single-head state — see `step_group`.
327    pub fn step(&mut self, q: &[f32], k: &[f32], v: &[f32], out: &mut [f32]) {
328        assert_eq!(self.heads.len(), 1, "use step_group for a GQA group");
329        self.step_group(q, k, v, out);
330    }
331
332    /// One decode step for the whole GQA group.  Inserts the group's
333    /// (k, v) ONCE, evicting the oldest window key into every head's far
334    /// accumulators, then writes each head's attention output.
335    /// `q_all` is `[q_heads][d]`, `out_all` is `[q_heads][dv]`.
336    pub fn step_group(&mut self, q_all: &[f32], k: &[f32], v: &[f32], out_all: &mut [f32]) {
337        let (d, dv) = (self.group.d, self.group.dv);
338        assert!(d > 0, "prefill() must run before step()");
339        let nh = self.heads.len();
340        assert_eq!(q_all.len(), nh * d);
341        assert_eq!(k.len(), d);
342        assert_eq!(v.len(), dv);
343        assert_eq!(out_all.len(), nh * dv);
344        // The current token is part of its own near window (t-j = 0),
345        // so insertion happens BEFORE any output is computed.
346        Self::advance(&mut self.group, &mut self.heads, k, v);
347        for (h, head) in self.heads.iter_mut().enumerate() {
348            head.step(
349                &self.group,
350                &q_all[h * d..(h + 1) * d],
351                &mut out_all[h * dv..(h + 1) * dv],
352            );
353        }
354    }
355
356    /// Heap bytes held by this group's state (shared window + sinks +
357    /// K̃, plus each head's skeleton and scratch) — feeds the honest
358    /// "KV+state" memory line, same discipline as counting
359    /// `linear_state` for the linear core.
360    pub fn memory_bytes(&self) -> usize {
361        self.group.memory_bytes()
362            + self
363                .heads
364                .iter()
365                .map(NystromHead::memory_bytes)
366                .sum::<usize>()
367    }
368
369    /// Push the group's (k, v) into the shared window.  In skeleton mode
370    /// a full ring first evicts its oldest key (delayed insertion — the
371    /// key leaves the exact window at this very step).
372    ///
373    /// The eviction is a GROUP event: the window is shared, so there is
374    /// exactly ONE eviction per position, not one per Q head.  The far
375    /// accumulators are per head, though, so that single evicted key is
376    /// absorbed once into EACH head — one eviction, `q_heads`
377    /// insertions.  Getting this wrong in either direction breaks the
378    /// boundary invariant (a key enters the far field at exactly the
379    /// step it leaves the window: no double count, no hole).
380    fn advance(g: &mut NystromGroup, heads: &mut [NystromHead], k: &[f32], v: &[f32]) {
381        let (d, dv) = (g.d, g.dv);
382        if !g.exact_only && g.win_len == g.w {
383            let slot = g.win_head;
384            // Every head absorbs the outgoing key BEFORE the slot is
385            // overwritten by the incoming one.
386            for h in heads.iter_mut() {
387                h.far_insert(g, slot);
388            }
389            g.win_k[slot * d..(slot + 1) * d].copy_from_slice(k);
390            g.win_v[slot * dv..(slot + 1) * dv].copy_from_slice(v);
391            g.win_head = (g.win_head + 1) % g.w;
392        } else if g.exact_only {
393            g.win_k.extend_from_slice(k);
394            g.win_v.extend_from_slice(v);
395            g.win_len += 1;
396        } else {
397            g.win_k[g.win_len * d..(g.win_len + 1) * d].copy_from_slice(k);
398            g.win_v[g.win_len * dv..(g.win_len + 1) * dv].copy_from_slice(v);
399            g.win_len += 1;
400        }
401    }
402}
403
404impl NystromGroup {
405    /// Freeze the group-shared geometry from the prompt's keys/values.
406    /// Returns the f64 key landmarks (which the heads need at full
407    /// precision to build Au), or None in exact-only mode.
408    fn prefill_shared(
409        &mut self,
410        ks: &[f32],
411        vs: &[f32],
412        t: usize,
413        d: usize,
414        dv: usize,
415    ) -> Option<Vec<f64>> {
416        self.d = d;
417        self.dv = dv;
418        self.scale = 1.0 / (d as f32).sqrt();
419        self.win_len = 0;
420        self.win_head = 0;
421        self.sink_len = 0;
422        self.exact_only = t <= self.w + self.sink + EXACT_SLACK;
423
424        if self.exact_only {
425            // Everything fits in the exact window (plus slack for a few
426            // decode steps before Vec growth); no skeleton is built and
427            // no separate sink buffer is needed — every key is already
428            // a permanent exact key in this mode.
429            self.win_k = Vec::with_capacity((t + 64) * d);
430            self.win_v = Vec::with_capacity((t + 64) * dv);
431            self.win_k.extend_from_slice(ks);
432            self.win_v.extend_from_slice(vs);
433            self.win_len = t;
434            return None;
435        }
436
437        // Sink tokens: positions 0..sink become permanent exact keys.
438        // They bypass the ring window entirely, so the delayed-insertion
439        // path can never move them into the far accumulators.
440        self.sink_len = self.sink; // skeleton mode guarantees t > sink
441        self.sink_k = ks[..self.sink * d].to_vec();
442        self.sink_v = vs[..self.sink * dv].to_vec();
443
444        // Landmarks: contiguous segment means of the prompt.  The
445        // integer split (i·t)/m matches the reference probe; the clamp
446        // keeps tiny prompts from producing duplicate landmarks.
447        let m_eff = (t / 8).clamp(4, self.m);
448        self.m_eff = m_eff;
449        let k_tilde64 = seg_means(ks, t, d, m_eff);
450        self.k_tilde = k_tilde64.iter().map(|&x| x as f32).collect();
451
452        self.win_k = vec![0.0; self.w * d];
453        self.win_v = vec![0.0; self.w * dv];
454        Some(k_tilde64)
455    }
456
457    fn memory_bytes(&self) -> usize {
458        (self.win_k.len()
459            + self.win_v.len()
460            + self.sink_k.len()
461            + self.sink_v.len()
462            + self.k_tilde.len())
463            * std::mem::size_of::<f32>()
464    }
465}
466
467impl NystromHead {
468    fn new() -> Self {
469        NystromHead {
470            rect: O1_DEFAULT_RECT,
471            t_hat: Vec::new(),
472            z_hat: Vec::new(),
473            m_max: Vec::new(),
474            far_len: 0,
475            q_tilde: Vec::new(),
476            mu: Vec::new(),
477            scr_s: Vec::new(),
478            scr_fh: Vec::new(),
479            scr_u: Vec::new(),
480            scr_l: Vec::new(),
481        }
482    }
483
484    /// exact-only mode: no skeleton state at all, just room to score the
485    /// growing window.
486    fn seal_exact(&mut self, t: usize) {
487        self.far_len = 0;
488        self.scr_s = Vec::with_capacity(t + 64);
489    }
490
491    /// Freeze this head's query landmarks and mixing matrix against the
492    /// group's (already frozen) key landmarks.
493    fn seal(&mut self, g: &NystromGroup, qs: &[f32], t: usize, k_tilde64: &[f64]) {
494        let (d, dv, m_eff) = (g.d, g.dv, g.m_eff);
495        self.far_len = 0;
496        let q_tilde64 = seg_means(qs, t, d, m_eff);
497        self.q_tilde = q_tilde64.iter().map(|&x| x as f32).collect();
498
499        // Au and its regularized pseudo-inverse in f64 — one-off m×m
500        // work at prefill only; the hot path stays f32.
501        let mut au = vec![0.0f64; m_eff * m_eff];
502        for i in 0..m_eff {
503            for j in 0..m_eff {
504                let mut s = 0.0f64;
505                for c in 0..d {
506                    s += q_tilde64[i * d + c] * k_tilde64[j * d + c];
507                }
508                au[i * m_eff + j] = (s * g.scale as f64).exp();
509            }
510        }
511        let mu64 = ridge_pinv(&au, m_eff);
512        self.mu = mu64.iter().map(|&x| x as f32).collect();
513
514        self.t_hat = vec![0.0; m_eff * dv];
515        self.z_hat = vec![0.0; m_eff];
516        self.m_max = vec![f32::NEG_INFINITY; m_eff];
517        self.scr_s = vec![0.0; g.sink + g.w];
518        self.scr_fh = vec![0.0; m_eff];
519        self.scr_u = vec![0.0; m_eff];
520        self.scr_l = vec![0.0; m_eff];
521    }
522
523    /// This head's output for `q` against the group's current window and
524    /// sinks and its own far field.  The window insertion for this
525    /// position already happened at group level (`NystromState::advance`).
526    fn step(&mut self, g: &NystromGroup, q: &[f32], out: &mut [f32]) {
527        let (d, dv) = (g.d, g.dv);
528        assert_eq!(q.len(), d);
529        assert_eq!(out.len(), dv);
530
531        // Near field: exact logits over sinks + window, one shared
532        // shift.  Sinks are permanent exact keys (near mask §5b:
533        // t-j < w OR j < sink); sink_len = 0 in exact-only mode.
534        let ns = g.sink_len;
535        let n = ns + g.win_len;
536        self.scr_s.resize(n, 0.0);
537        let mut c = f32::NEG_INFINITY;
538        for s in 0..ns {
539            let lg = dot(q, &g.sink_k[s * d..(s + 1) * d]) * g.scale;
540            self.scr_s[s] = lg;
541            c = c.max(lg);
542        }
543        // Window scores are the decode hot loop — NEON dot (same
544        // products, regrouped sums; parity-gated by the golden tests).
545        for s in 0..g.win_len {
546            let lg = crate::attention::dot_f32(q, &g.win_k[s * d..(s + 1) * d]) * g.scale;
547            self.scr_s[ns + s] = lg;
548            c = c.max(lg);
549        }
550
551        // Far field: shifted skeleton (spec §3).  All exp arguments are
552        // ≤ 0 relative to the joint shift c_all, so nothing overflows.
553        let mut far_den = 0.0f32;
554        let mut c_all = c;
555        let mut have_far = false;
556        if self.far_len > 0 {
557            // Per-token row shift f over landmark scores.
558            let mut f = f32::NEG_INFINITY;
559            for a in 0..g.m_eff {
560                let s = crate::attention::dot_f32(q, &g.k_tilde[a * d..(a + 1) * d]) * g.scale;
561                self.scr_fh[a] = s;
562                f = f.max(s);
563            }
564            for a in 0..g.m_eff {
565                self.scr_fh[a] = (self.scr_fh[a] - f).exp();
566            }
567            // u = (F·e^{-f}) · M — the landmark mixing row (= FM, up to
568            // the positive factor e^{-f}).
569            for b in 0..g.m_eff {
570                let mut s = 0.0f32;
571                for a in 0..g.m_eff {
572                    s += self.scr_fh[a] * self.mu[a * g.m_eff + b];
573                }
574                // FM rectifier: every far weight is Σ_b FM[b]·E[b,j]
575                // with E ≥ 0 elementwise, so clamping this m-vector is
576                // enough to make all of them non-negative — the per-key
577                // guarantee the streaming form otherwise cannot state.
578                // The row shift e^{-f} and the flash factors below are
579                // strictly positive, so clamping here or after the
580                // rescale is the same predicate.
581                self.scr_u[b] = if self.rect == O1Rect::Fm {
582                    s.max(0.0)
583                } else {
584                    s
585                };
586            }
587            // Joint scale: the far term b carries e^{f + m_max[b]}, the
588            // near term e^{c}; take the max so every factor is ≤ 1.
589            for b in 0..g.m_eff {
590                c_all = c_all.max(f + self.m_max[b]);
591            }
592            for b in 0..g.m_eff {
593                let gain = self.scr_u[b] * (f + self.m_max[b] - c_all).exp();
594                self.scr_u[b] = gain;
595                far_den += gain * self.z_hat[b];
596            }
597            // Aggregate guard — the rectifier of `O1Rect::Aggregate`,
598            // and a second line of defence under `Fm` (where far_den is
599            // a sum of non-negative terms, so this can only fire on
600            // rounding): a negative denominator means the skeleton
601            // estimate is unusable for this row — drop the far field.
602            if far_den >= 0.0 {
603                have_far = true;
604            } else {
605                far_den = 0.0;
606            }
607        }
608
609        for o in out.iter_mut() {
610            *o = 0.0;
611        }
612        if have_far {
613            for b in 0..g.m_eff {
614                crate::attention::axpy_f32(out, &self.t_hat[b * dv..(b + 1) * dv], self.scr_u[b]);
615            }
616        }
617        let mut den = far_den;
618        for s in 0..n {
619            let p = (self.scr_s[s] - c_all).exp();
620            den += p;
621            // scr_s rows 0..ns are sinks, the rest are window entries.
622            let vv = if s < ns {
623                &g.sink_v[s * dv..(s + 1) * dv]
624            } else {
625                &g.win_v[(s - ns) * dv..(s - ns + 1) * dv]
626            };
627            crate::attention::axpy_f32(out, vv, p);
628        }
629        let den = den.max(DEN_EPS);
630        for o in out.iter_mut() {
631            *o /= den;
632        }
633    }
634
635    /// Absorb the group's window slot into THIS head's far accumulators
636    /// with the per-landmark flash shift: T̂[i]/Ẑ[i] live at scale
637    /// e^{-m_max[i]}; when a new logit raises the max, existing mass is
638    /// rescaled by e^{old-new} (exactly 0 on first insertion, since
639    /// m_max = -inf).
640    fn far_insert(&mut self, g: &NystromGroup, slot: usize) {
641        let (d, dv) = (g.d, g.dv);
642        // Runs once per evicted key per head — NEON dot/axpy like the
643        // decode loop (same products, regrouped sums).
644        for i in 0..g.m_eff {
645            self.scr_l[i] = crate::attention::dot_f32(
646                &self.q_tilde[i * d..(i + 1) * d],
647                &g.win_k[slot * d..(slot + 1) * d],
648            ) * g.scale;
649        }
650        for i in 0..g.m_eff {
651            let l = self.scr_l[i];
652            if l > self.m_max[i] {
653                let r = (self.m_max[i] - l).exp();
654                self.z_hat[i] *= r;
655                for e in self.t_hat[i * dv..(i + 1) * dv].iter_mut() {
656                    *e *= r;
657                }
658                self.m_max[i] = l;
659            }
660            let e = (l - self.m_max[i]).exp();
661            self.z_hat[i] += e;
662            crate::attention::axpy_f32(
663                &mut self.t_hat[i * dv..(i + 1) * dv],
664                &g.win_v[slot * dv..(slot + 1) * dv],
665                e,
666            );
667        }
668        self.far_len += 1;
669    }
670
671    fn memory_bytes(&self) -> usize {
672        (self.t_hat.len()
673            + self.z_hat.len()
674            + self.m_max.len()
675            + self.q_tilde.len()
676            + self.mu.len()
677            + self.scr_s.len()
678            + self.scr_fh.len()
679            + self.scr_u.len()
680            + self.scr_l.len())
681            * std::mem::size_of::<f32>()
682    }
683}
684
685/// Contiguous segment means (the Nyströmformer landmark recipe), f64
686/// accumulation.  The split (i·t)/m matches the Python reference.
687fn seg_means(xs: &[f32], t: usize, d: usize, m: usize) -> Vec<f64> {
688    let mut out = vec![0.0f64; m * d];
689    for i in 0..m {
690        let lo = i * t / m;
691        let hi = (i + 1) * t / m;
692        for j in lo..hi {
693            for c in 0..d {
694                out[i * d + c] += xs[j * d + c] as f64;
695            }
696        }
697        let inv = 1.0 / (hi - lo) as f64;
698        for c in 0..d {
699            out[i * d + c] *= inv;
700        }
701    }
702    out
703}
704
705fn dot(a: &[f32], b: &[f32]) -> f32 {
706    let mut s = 0.0f32;
707    for (x, y) in a.iter().zip(b) {
708        s += x * y;
709    }
710    s
711}
712
713/// Regularized pseudo-inverse M = (AᵀA + λI)⁻¹ Aᵀ of a square matrix,
714/// λ = RIDGE_REL·mean(diag(AᵀA)), solved via Cholesky.  f64 internal —
715/// this runs once per prefill on an m×m matrix (m ≤ 32).  If Cholesky
716/// fails (Au numerically singular despite the m_eff clamp), λ grows
717/// tenfold — the jitter fallback of the reference probe.
718/// pub(crate): the FCD polish trainer builds its (constant-in-backward)
719/// mixing matrix with the SAME solver the runtime seals with.
720pub(crate) fn ridge_pinv(a: &[f64], n: usize) -> Vec<f64> {
721    let mut ata = vec![0.0f64; n * n];
722    for i in 0..n {
723        for j in 0..n {
724            let mut s = 0.0;
725            for k in 0..n {
726                s += a[k * n + i] * a[k * n + j];
727            }
728            ata[i * n + j] = s;
729        }
730    }
731    let mean_diag: f64 = (0..n).map(|i| ata[i * n + i]).sum::<f64>() / n as f64;
732    let mut lambda = RIDGE_REL * mean_diag.max(f64::MIN_POSITIVE);
733    for _ in 0..12 {
734        let mut g = ata.clone();
735        for i in 0..n {
736            g[i * n + i] += lambda;
737        }
738        if let Some(l) = cholesky(&mut g, n) {
739            // Solve G·M = Aᵀ column by column; column j of Aᵀ is row j
740            // of A.
741            let mut m_out = vec![0.0f64; n * n];
742            let mut x = vec![0.0f64; n];
743            for j in 0..n {
744                let rhs = &a[j * n..(j + 1) * n];
745                // Forward: L·y = rhs.
746                for i in 0..n {
747                    let mut s = rhs[i];
748                    for k in 0..i {
749                        s -= l[i * n + k] * x[k];
750                    }
751                    x[i] = s / l[i * n + i];
752                }
753                // Backward: Lᵀ·x = y.
754                for i in (0..n).rev() {
755                    let mut s = x[i];
756                    for k in i + 1..n {
757                        s -= l[k * n + i] * x[k];
758                    }
759                    x[i] = s / l[i * n + i];
760                }
761                for i in 0..n {
762                    m_out[i * n + j] = x[i];
763                }
764            }
765            return m_out;
766        }
767        lambda *= 10.0;
768    }
769    // Unreachable in practice: λ eventually dominates the diagonal.
770    // Degrade to a scaled identity rather than poison the output.
771    let mut fallback = vec![0.0f64; n * n];
772    for i in 0..n {
773        fallback[i * n + i] = 1.0 / mean_diag.max(f64::MIN_POSITIVE);
774    }
775    fallback
776}
777
778// ── Runtime configuration (v1: runtime-level, NOT a format change) ──
779//
780// A layer set + {m, w, sink}, resolved in priority order:
781//   1. CLI flag (`--o1` on run/serve/bench) — explicit user intent;
782//   2. env `CMF_O1` (all | deepN | i,j,k | off) with CMF_O1_M /
783//      CMF_O1_WINDOW / CMF_O1_SINK parameter overrides;
784//   3. converter hint in the header JSON (`provenance.o1_attn`,
785//      written by `cortiq convert --o1`) — additive metadata, the
786//      binary envelope is untouched.
787
788/// Validated defaults (spec: m=32, W=128, sink=4; sink ablation ×2.39).
789pub const O1_DEFAULT_M: usize = 32;
790pub const O1_DEFAULT_W: usize = 128;
791pub const O1_DEFAULT_SINK: usize = 4;
792/// Rectifier default (see `O1Rect`).
793pub const O1_DEFAULT_RECT: O1Rect = O1Rect::Aggregate;
794
795/// Which layers run the O(1) kernel.
796#[derive(Clone, Debug, PartialEq, Eq)]
797pub enum O1Layers {
798    All,
799    /// The N deepest layers (deep-N ladder of the price map; the
800    /// early stack is the most sink-dependent, depth converts best).
801    Deep(usize),
802    /// Explicit layer indices.
803    List(Vec<usize>),
804}
805
806/// Per-model O(1)-attention setting.
807#[derive(Clone, Debug)]
808pub struct O1Cfg {
809    pub layers: O1Layers,
810    /// Landmark budget (≥ 4; m=64 measured WORSE — collinear segment
811    /// means poison the pinv, so don't "help" by raising it).
812    pub m: usize,
813    /// Exact-window width — the main quality lever.
814    pub w: usize,
815    /// Permanent exact sink keys (StreamingLLM discipline, spec §5b).
816    pub sink: usize,
817    /// Skeleton rectifier (see `O1Rect`).
818    pub rect: O1Rect,
819}
820
821/// Three-state env reading: unset falls through to the header hint,
822/// `off`/`0` force-disables even a header hint (the escape hatch).
823pub enum O1Env {
824    Unset,
825    Off,
826    On(O1Cfg),
827}
828
829impl O1Cfg {
830    /// Parse a layer spec: `all` | `deepN` | `i,j,k`. None = not a spec
831    /// (also used for `off`/`0`/empty).
832    pub fn parse_layers(spec: &str) -> Option<O1Layers> {
833        let s = spec.trim();
834        match s {
835            "" | "off" | "0" | "none" => None,
836            "all" => Some(O1Layers::All),
837            _ => {
838                if let Some(n) = s.strip_prefix("deep") {
839                    return n
840                        .parse::<usize>()
841                        .ok()
842                        .filter(|&n| n > 0)
843                        .map(O1Layers::Deep);
844                }
845                let idx: Result<Vec<usize>, _> =
846                    s.split(',').map(|p| p.trim().parse::<usize>()).collect();
847                idx.ok().filter(|v| !v.is_empty()).map(O1Layers::List)
848            }
849        }
850    }
851
852    /// Parse a rectifier spec: `agg`/`aggregate` | `fm`. None = not a
853    /// spec.
854    pub fn parse_rect(spec: &str) -> Option<O1Rect> {
855        match spec.trim() {
856            "agg" | "aggregate" => Some(O1Rect::Aggregate),
857            "fm" => Some(O1Rect::Fm),
858            _ => None,
859        }
860    }
861
862    /// Rectifier from an explicit value, else `CMF_O1_RECT`, else the
863    /// default.
864    fn rect_or_env(rect: Option<O1Rect>) -> O1Rect {
865        rect.or_else(|| {
866            std::env::var("CMF_O1_RECT")
867                .ok()
868                .as_deref()
869                .and_then(Self::parse_rect)
870        })
871        .unwrap_or(O1_DEFAULT_RECT)
872    }
873
874    /// Build from an explicit spec (CLI path). None = `off` or malformed.
875    /// Explicit m/w/sink/rect beat env overrides beat validated defaults.
876    pub fn from_spec(
877        spec: &str,
878        m: Option<usize>,
879        w: Option<usize>,
880        sink: Option<usize>,
881        rect: Option<O1Rect>,
882    ) -> Option<O1Cfg> {
883        let layers = Self::parse_layers(spec)?;
884        let env = |k: &str| std::env::var(k).ok().and_then(|v| v.parse::<usize>().ok());
885        Some(O1Cfg {
886            layers,
887            // NystromState asserts m ≥ 4 and w ≥ 1 — clamp rather than
888            // panic deep in the first prefill.
889            m: m.or_else(|| env("CMF_O1_M")).unwrap_or(O1_DEFAULT_M).max(4),
890            w: w.or_else(|| env("CMF_O1_WINDOW"))
891                .unwrap_or(O1_DEFAULT_W)
892                .max(1),
893            sink: sink
894                .or_else(|| env("CMF_O1_SINK"))
895                .unwrap_or(O1_DEFAULT_SINK),
896            rect: Self::rect_or_env(rect),
897        })
898    }
899
900    /// Converter hint from the header JSON: `{"layers": "all"|[i,…],
901    /// "m": …, "w": …, "sink": …}`. Env parameter overrides still apply
902    /// (the operator's knob wins over the file's suggestion).
903    pub fn from_json(v: &serde_json::Value) -> Option<O1Cfg> {
904        let layers = match v.get("layers") {
905            Some(serde_json::Value::String(s)) => Self::parse_layers(s)?,
906            Some(serde_json::Value::Array(a)) => O1Layers::List(
907                a.iter()
908                    .filter_map(|x| x.as_u64().map(|n| n as usize))
909                    .collect(),
910            ),
911            _ => return None,
912        };
913        let f = |k: &str| v.get(k).and_then(|x| x.as_u64()).map(|n| n as usize);
914        let env = |k: &str| std::env::var(k).ok().and_then(|s| s.parse::<usize>().ok());
915        Some(O1Cfg {
916            layers,
917            m: env("CMF_O1_M")
918                .or_else(|| f("m"))
919                .unwrap_or(O1_DEFAULT_M)
920                .max(4),
921            w: env("CMF_O1_WINDOW")
922                .or_else(|| f("w"))
923                .unwrap_or(O1_DEFAULT_W)
924                .max(1),
925            sink: env("CMF_O1_SINK")
926                .or_else(|| f("sink"))
927                .unwrap_or(O1_DEFAULT_SINK),
928            // The rectifier is a runtime property of the kernel, not a
929            // property of the weights — a file hint cannot pin it.
930            rect: Self::rect_or_env(None),
931        })
932    }
933
934    /// Per-layer flags over `num_layers` (indices past the end are
935    /// silently dropped; the pipeline additionally filters non-Full
936    /// layers — a linear layer keeps its own operator).
937    pub fn layer_flags(&self, num_layers: usize) -> Vec<bool> {
938        let mut flags = vec![false; num_layers];
939        match &self.layers {
940            O1Layers::All => flags.iter_mut().for_each(|f| *f = true),
941            O1Layers::Deep(n) => {
942                for f in flags.iter_mut().skip(num_layers.saturating_sub(*n)) {
943                    *f = true;
944                }
945            }
946            O1Layers::List(idx) => {
947                for &i in idx {
948                    if i < num_layers {
949                        flags[i] = true;
950                    }
951                }
952            }
953        }
954        flags
955    }
956}
957
958/// Read `CMF_O1` (+ parameter overrides) — the embedding-friendly path
959/// for hosts that don't go through the CLI flags.
960pub fn o1_from_env() -> O1Env {
961    match std::env::var("CMF_O1") {
962        Err(_) => O1Env::Unset,
963        Ok(s) => match O1Cfg::from_spec(&s, None, None, None, None) {
964            Some(cfg) => O1Env::On(cfg),
965            None => O1Env::Off,
966        },
967    }
968}
969
970/// In-place lower Cholesky of an SPD matrix; None if a pivot fails.
971fn cholesky(g: &mut [f64], n: usize) -> Option<&[f64]> {
972    for i in 0..n {
973        for j in 0..=i {
974            let mut s = g[i * n + j];
975            for k in 0..j {
976                s -= g[i * n + k] * g[j * n + k];
977            }
978            if i == j {
979                if s <= 0.0 || !s.is_finite() {
980                    return None;
981                }
982                g[i * n + i] = s.sqrt();
983            } else {
984                g[i * n + j] = s / g[j * n + j];
985            }
986        }
987    }
988    Some(g)
989}