Skip to main content

cortiq_engine/
bounded.rs

1//! Natively bounded softmax anchor `swa_sink_v1` (Embryo-O1).
2//!
3//! Contract: `docs/EMBRYO_BOUNDED_ANCHOR.md`. The operator is a TRAINED
4//! bounded attention, not a masked full attention: per token `t`
5//!
6//! ```text
7//! window keys:  j ∈ (t − W, t]                (W keys INCLUDING the token)
8//! window score: s_j = ( R(t − j) q̂_t ) · k̂_j / √hd
9//! sink score:   s_s = q̂_t · k̂ˢ_s / √hd,   s ∈ [0, S)      (NoPE)
10//! one softmax over {s_s} ∪ {s_j};   out_t = Σ_s p_s v̂ˢ_s + Σ_j p_j v_j
11//! ```
12//!
13//! The ring stores RAW (unrotated) keys; the query is rotated by the
14//! distance `Δ = t − j ∈ [0, W)` through a `[W][rd/2]` cos/sin table.
15//! With absolute RoPE `q_rot(t) = R(t) q̂`, `k_rot(j) = R(j) k̂` one has
16//! `q_rot(t)·k_rot(j) = q̂ᵀ R(t)ᵀ R(j) k̂ = (R(t−j) q̂)·k̂` (every 2-D block
17//! of `R` is a plane rotation, so `R(t)ᵀ R(j) = R(j − t)`), which is why
18//! the trainer may keep rotating by absolute positions and only mask,
19//! while the served operator never sees an absolute position at all.
20//!
21//! State per layer: `ring_k, ring_v [kvh][W][hd]` + the insert counter —
22//! a record of fixed size derived from the header, identical in prefill,
23//! decode and across turns. Nothing here grows with the context.
24
25use crate::qtensor::QTensor;
26
27/// Rows of rollback history kept beside the ring so a speculative reject
28/// (`truncate_last`) can restore the slots the rejected tokens overwrote.
29/// Bounded by construction; not part of the wire state.
30pub const UNDO_DEPTH: usize = 64;
31
32/// Weights of one bounded-anchor layer.
33pub struct BoundedWeights {
34    pub wq: QTensor,
35    pub wk: QTensor,
36    pub wv: QTensor,
37    pub wo: QTensor,
38    /// Trained NoPE sink keys `[kvh][sink][hd]` (weights, not positions).
39    pub sink_k: Vec<f32>,
40    /// Trained sink values `[kvh][sink][hd]`.
41    pub sink_v: Vec<f32>,
42    pub sink: usize,
43    pub window: usize,
44}
45
46/// Relative-rotation table: `cos/sin[Δ][i] = cos/sin(Δ · inv_freq[i])`
47/// for `Δ ∈ [0, W)`, built once per model from the layer's `inv_freq`
48/// (same convention as `attention::rope_rotate_scaled`, angles computed
49/// in f64 so the table carries only the final f32 rounding).
50#[derive(Debug, Clone)]
51pub struct BoundedRope {
52    pub window: usize,
53    /// `rotary_dim / 2`: dims `[0, half)` pair with `[half, 2·half)`;
54    /// dims past `2·half` are copied unrotated (partial rotary).
55    pub half: usize,
56    pub cos: Vec<f32>,
57    pub sin: Vec<f32>,
58}
59
60impl BoundedRope {
61    /// `rope_scale` is the YaRN attention factor the absolute path applies
62    /// to BOTH q and k; the relative path rotates only q, so the window
63    /// score carries `scale²` (`(s·R(t)q)·(s·R(j)k) = s²·(R(t−j)q)·k`).
64    pub fn new(window: usize, inv_freq: &[f32], rope_scale: f32) -> Self {
65        let half = inv_freq.len();
66        let s2 = (rope_scale as f64) * (rope_scale as f64);
67        let mut cos = Vec::with_capacity(window * half);
68        let mut sin = Vec::with_capacity(window * half);
69        for delta in 0..window {
70            for &f in inv_freq {
71                let (sn, cs) = ((delta as f64) * (f as f64)).sin_cos();
72                cos.push((cs * s2) as f32);
73                sin.push((sn * s2) as f32);
74            }
75        }
76        Self {
77            window,
78            half,
79            cos,
80            sin,
81        }
82    }
83
84    /// `out = R(Δ) x` (first `2·half` dims rotated, the rest copied).
85    #[inline]
86    pub fn rotate(&self, delta: usize, x: &[f32], out: &mut [f32]) {
87        let half = self.half;
88        let c = &self.cos[delta * half..(delta + 1) * half];
89        let s = &self.sin[delta * half..(delta + 1) * half];
90        for i in 0..half {
91            let x0 = x[i];
92            let x1 = x[i + half];
93            out[i] = x0 * c[i] - x1 * s[i];
94            out[i + half] = x0 * s[i] + x1 * c[i];
95        }
96        let r = 2 * half;
97        if x.len() > r {
98            out[r..x.len()].copy_from_slice(&x[r..]);
99        }
100    }
101}
102
103/// Bit-for-bit copy of everything the operator mutates (speculation).
104#[derive(Debug, Clone)]
105pub struct BoundedSnapshot {
106    ring_k: Vec<f32>,
107    ring_v: Vec<f32>,
108    seen: usize,
109}
110
111/// Per-layer bounded state: the ring of the last `W` raw keys/values per
112/// KV head plus the insert counter. `len = min(seen, W)`, next write slot
113/// `head = seen mod W`, and the key of position `j` lives in slot
114/// `j mod W` — so no absolute position is ever stored.
115#[derive(Debug, Clone)]
116pub struct BoundedState {
117    pub window: usize,
118    pub num_kv_heads: usize,
119    pub head_dim: usize,
120    /// `[kvh][W][hd]`, raw (unrotated) keys.
121    pub ring_k: Vec<f32>,
122    /// `[kvh][W][hd]`.
123    pub ring_v: Vec<f32>,
124    /// Tokens inserted since the last clear (= the next position).
125    pub seen: usize,
126    /// Rollback rows: the slot contents each of the last `UNDO_DEPTH`
127    /// inserts overwrote (`[UNDO_DEPTH][kvh][hd]` for k and v).
128    undo_k: Vec<f32>,
129    undo_v: Vec<f32>,
130    undo_len: usize,
131    undo_head: usize,
132}
133
134impl BoundedState {
135    pub fn new(num_kv_heads: usize, head_dim: usize, window: usize) -> Self {
136        let n = num_kv_heads * window * head_dim;
137        let u = UNDO_DEPTH * num_kv_heads * head_dim;
138        Self {
139            window,
140            num_kv_heads,
141            head_dim,
142            ring_k: vec![0.0; n],
143            ring_v: vec![0.0; n],
144            seen: 0,
145            undo_k: vec![0.0; u],
146            undo_v: vec![0.0; u],
147            undo_len: 0,
148            undo_head: 0,
149        }
150    }
151
152    /// Filled slots, `min(seen, W)`.
153    #[inline]
154    pub fn len(&self) -> usize {
155        self.seen.min(self.window)
156    }
157
158    #[inline]
159    pub fn is_empty(&self) -> bool {
160        self.seen == 0
161    }
162
163    /// Next write slot, `seen mod W`.
164    #[inline]
165    pub fn head(&self) -> usize {
166        self.seen % self.window
167    }
168
169    /// Bytes of the wire-visible state (ring + counter). The undo rows
170    /// are scratch of fixed size and are not state.
171    pub fn state_bytes(&self) -> usize {
172        (self.ring_k.len() + self.ring_v.len()) * std::mem::size_of::<f32>()
173            + std::mem::size_of::<u64>()
174    }
175
176    /// Zero the record (fresh sequence). Capacity never changes.
177    pub fn clear(&mut self) {
178        self.ring_k.fill(0.0);
179        self.ring_v.fill(0.0);
180        self.seen = 0;
181        self.undo_len = 0;
182        self.undo_head = 0;
183    }
184
185    /// Write `k, v` (`[kvh][hd]`, raw) into slot `seen mod W`, saving
186    /// the overwritten rows for rollback.
187    pub fn insert(&mut self, k: &[f32], v: &[f32]) {
188        let (kvh, hd, w) = (self.num_kv_heads, self.head_dim, self.window);
189        debug_assert_eq!(k.len(), kvh * hd);
190        debug_assert_eq!(v.len(), kvh * hd);
191        let slot = self.head();
192        let u = self.undo_head;
193        for h in 0..kvh {
194            let r = (h * w + slot) * hd;
195            let uo = (u * kvh + h) * hd;
196            self.undo_k[uo..uo + hd].copy_from_slice(&self.ring_k[r..r + hd]);
197            self.undo_v[uo..uo + hd].copy_from_slice(&self.ring_v[r..r + hd]);
198            self.ring_k[r..r + hd].copy_from_slice(&k[h * hd..(h + 1) * hd]);
199            self.ring_v[r..r + hd].copy_from_slice(&v[h * hd..(h + 1) * hd]);
200        }
201        self.undo_head = (u + 1) % UNDO_DEPTH;
202        self.undo_len = (self.undo_len + 1).min(UNDO_DEPTH);
203        self.seen += 1;
204    }
205
206    /// Undo the last `n` inserts exactly (restores the overwritten slots).
207    /// Returns how many were rolled back — fewer than `n` only when the
208    /// rollback history (`UNDO_DEPTH`) or the inserted count is shorter.
209    pub fn rollback(&mut self, n: usize) -> usize {
210        let (kvh, hd, w) = (self.num_kv_heads, self.head_dim, self.window);
211        let n = n.min(self.undo_len).min(self.seen);
212        for _ in 0..n {
213            self.seen -= 1;
214            let slot = self.seen % w;
215            let u = (self.undo_head + UNDO_DEPTH - 1) % UNDO_DEPTH;
216            for h in 0..kvh {
217                let r = (h * w + slot) * hd;
218                let uo = (u * kvh + h) * hd;
219                self.ring_k[r..r + hd].copy_from_slice(&self.undo_k[uo..uo + hd]);
220                self.ring_v[r..r + hd].copy_from_slice(&self.undo_v[uo..uo + hd]);
221            }
222            self.undo_head = u;
223            self.undo_len -= 1;
224        }
225        n
226    }
227
228    pub fn snapshot(&self) -> BoundedSnapshot {
229        BoundedSnapshot {
230            ring_k: self.ring_k.clone(),
231            ring_v: self.ring_v.clone(),
232            seen: self.seen,
233        }
234    }
235
236    /// Restore a snapshot taken on THIS state (same geometry). The undo
237    /// history is discarded: it described inserts that no longer exist.
238    pub fn restore(&mut self, s: &BoundedSnapshot) {
239        debug_assert_eq!(s.ring_k.len(), self.ring_k.len());
240        self.ring_k.copy_from_slice(&s.ring_k);
241        self.ring_v.copy_from_slice(&s.ring_v);
242        self.seen = s.seen;
243        self.undo_len = 0;
244        self.undo_head = 0;
245    }
246
247    /// Bit-for-bit equality of the wire-visible state.
248    pub fn same_state(&self, other: &BoundedState) -> bool {
249        self.window == other.window
250            && self.seen == other.seen
251            && self.ring_k == other.ring_k
252            && self.ring_v == other.ring_v
253    }
254
255    /// Attend every Q head of every KV group over `sink ∪ window` and
256    /// write `[nh][hd]` into `out`. `q` is `[nh][hd]` RAW (unrotated);
257    /// the current token must already be inserted (the window includes
258    /// it, matching the trainer's band `col == S + row`).
259    #[allow(clippy::too_many_arguments)]
260    pub fn attend(
261        &self,
262        q: &[f32],
263        num_heads: usize,
264        sink_k: &[f32],
265        sink_v: &[f32],
266        sink: usize,
267        rope: &BoundedRope,
268        scale: f32,
269        out: &mut [f32],
270    ) {
271        let (kvh, hd, w) = (self.num_kv_heads, self.head_dim, self.window);
272        let nh = num_heads;
273        let hpk = nh / kvh.max(1);
274        debug_assert_eq!(hpk * kvh, nh);
275        debug_assert_eq!(q.len(), nh * hd);
276        debug_assert_eq!(out.len(), nh * hd);
277        debug_assert_eq!(sink_k.len(), kvh * sink * hd);
278        debug_assert_eq!(rope.window, w);
279        let m = self.len();
280        let n = sink + m;
281        let head = self.head();
282        thread_local! {
283            static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<f32>)> =
284                const { std::cell::RefCell::new((Vec::new(), Vec::new())) };
285        }
286        SCRATCH.with(|s| {
287            let mut s = s.borrow_mut();
288            let (scores, qrot) = &mut *s;
289            scores.clear();
290            scores.resize(nh * n, 0.0);
291            qrot.clear();
292            qrot.resize(nh * hd, 0.0);
293            // Sink scores on the raw query (NoPE).
294            for h in 0..nh {
295                let g = h / hpk;
296                let qh = &q[h * hd..(h + 1) * hd];
297                for s in 0..sink {
298                    let kr = &sink_k[(g * sink + s) * hd..(g * sink + s + 1) * hd];
299                    scores[h * n + s] = crate::attention::dot_f32(qh, kr) * scale;
300                }
301            }
302            // Window scores: Δ = 0 is the token itself; slot(Δ) walks the
303            // ring backwards from the last written slot.
304            for d in 0..m {
305                let slot = (head + w - 1 - d) % w;
306                for h in 0..nh {
307                    rope.rotate(d, &q[h * hd..(h + 1) * hd], &mut qrot[h * hd..(h + 1) * hd]);
308                }
309                for h in 0..nh {
310                    let g = h / hpk;
311                    let kr = &self.ring_k[(g * w + slot) * hd..(g * w + slot + 1) * hd];
312                    scores[h * n + sink + d] =
313                        crate::attention::dot_f32(&qrot[h * hd..(h + 1) * hd], kr) * scale;
314                }
315            }
316            // One softmax per head over sinks ∪ window (attention_head order).
317            for h in 0..nh {
318                let sc = &mut scores[h * n..(h + 1) * n];
319                let max = sc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
320                let mut sum = 0.0f32;
321                for v in sc.iter_mut() {
322                    *v = (*v - max).exp();
323                    sum += *v;
324                }
325                if sum > 0.0 {
326                    for v in sc.iter_mut() {
327                        *v /= sum;
328                    }
329                }
330            }
331            out.fill(0.0);
332            for h in 0..nh {
333                let g = h / hpk;
334                let oh = &mut out[h * hd..(h + 1) * hd];
335                let p = &scores[h * n..(h + 1) * n];
336                for s in 0..sink {
337                    let vr = &sink_v[(g * sink + s) * hd..(g * sink + s + 1) * hd];
338                    if p[s].abs() >= 1e-12 {
339                        crate::attention::axpy_f32(oh, vr, p[s]);
340                    }
341                }
342                for d in 0..m {
343                    let slot = (head + w - 1 - d) % w;
344                    let vr = &self.ring_v[(g * w + slot) * hd..(g * w + slot + 1) * hd];
345                    let pw = p[sink + d];
346                    if pw.abs() >= 1e-12 {
347                        crate::attention::axpy_f32(oh, vr, pw);
348                    }
349                }
350            }
351        });
352    }
353}
354
355/// Per-token configuration of a bounded layer (geometry + the shared
356/// rotation table). No position: the operator has none.
357pub struct BoundedAttnCfg<'a> {
358    pub num_heads: usize,
359    pub num_kv_heads: usize,
360    pub head_dim: usize,
361    pub hidden_size: usize,
362    /// Score scale (1/√hd unless the arch overrides).
363    pub scale: f32,
364    pub rope: &'a BoundedRope,
365    pub pool: Option<&'a crate::pool::Pool>,
366}
367
368/// One position: `q̂ k̂ v = W x`, insert, attend, `W_o`. The cache's
369/// `bounded` record must exist (installed from the header at load).
370pub fn bounded_attention(
371    hidden: &[f32],
372    w: &BoundedWeights,
373    cache: &mut crate::kv_cache::LayerKvCache,
374    cfg: &BoundedAttnCfg,
375) -> Vec<f32> {
376    let (nh, nkv, hd) = (cfg.num_heads, cfg.num_kv_heads, cfg.head_dim);
377    let mut q = crate::attention::take_buf(nh * hd);
378    let mut k = crate::attention::take_buf(nkv * hd);
379    let mut v = crate::attention::take_buf(nkv * hd);
380    w.wq.matvec(hidden, &mut q, cfg.pool);
381    w.wk.matvec(hidden, &mut k, cfg.pool);
382    w.wv.matvec(hidden, &mut v, cfg.pool);
383    let mut ao = crate::attention::take_buf(nh * hd);
384    cache.bounded_step(&q, &k, &v, w, cfg.rope, cfg.scale, nh, &mut ao);
385    let mut out = crate::attention::take_buf(cfg.hidden_size);
386    w.wo.matvec(&ao, &mut out, cfg.pool);
387    crate::attention::recycle_buf(&mut q);
388    crate::attention::recycle_buf(&mut k);
389    crate::attention::recycle_buf(&mut v);
390    crate::attention::recycle_buf(&mut ao);
391    out
392}
393
394/// A prefill chunk of `b` positions: the projections run as chunk
395/// GEMMs (each weight row streams once per chunk), the operator runs
396/// per position over ring + chunk — the scores of a position are
397/// against at most `S + W` keys whatever the chunk or the context, and
398/// the per-position arithmetic is the same code as `bounded_attention`.
399pub fn bounded_attention_batch(
400    normed_all: &[f32],
401    b: usize,
402    w: &BoundedWeights,
403    cache: &mut crate::kv_cache::LayerKvCache,
404    cfg: &BoundedAttnCfg,
405) -> Vec<f32> {
406    let (nh, nkv, hd, hs) = (cfg.num_heads, cfg.num_kv_heads, cfg.head_dim, cfg.hidden_size);
407    debug_assert_eq!(normed_all.len(), b * hs);
408    let mut q_all = crate::attention::take_buf(b * nh * hd);
409    let mut k_all = crate::attention::take_buf(b * nkv * hd);
410    let mut v_all = crate::attention::take_buf(b * nkv * hd);
411    w.wq.matmat(normed_all, b, &mut q_all, cfg.pool);
412    w.wk.matmat(normed_all, b, &mut k_all, cfg.pool);
413    w.wv.matmat(normed_all, b, &mut v_all, cfg.pool);
414    let mut ao_all = crate::attention::take_buf(b * nh * hd);
415    for bi in 0..b {
416        let q = &q_all[bi * nh * hd..(bi + 1) * nh * hd];
417        let k = &k_all[bi * nkv * hd..(bi + 1) * nkv * hd];
418        let v = &v_all[bi * nkv * hd..(bi + 1) * nkv * hd];
419        let ao = &mut ao_all[bi * nh * hd..(bi + 1) * nh * hd];
420        cache.bounded_step(q, k, v, w, cfg.rope, cfg.scale, nh, ao);
421    }
422    let mut out = crate::attention::take_buf(b * hs);
423    w.wo.matmat(&ao_all, b, &mut out, cfg.pool);
424    crate::attention::recycle_buf(&mut q_all);
425    crate::attention::recycle_buf(&mut k_all);
426    crate::attention::recycle_buf(&mut v_all);
427    crate::attention::recycle_buf(&mut ao_all);
428    out
429}
430
431#[cfg(test)]
432mod tests {
433    use super::*;
434
435    fn synth(n: usize, salt: u64, scale: f32) -> Vec<f32> {
436        (0..n)
437            .map(|i| {
438                let x = (i as u64)
439                    .wrapping_mul(6364136223846793005)
440                    .wrapping_add(salt.wrapping_mul(1442695040888963407) ^ 0x9E3779B97F4A7C15);
441                let x = (x ^ (x >> 31)).wrapping_mul(0xBF58476D1CE4E5B9);
442                (((x >> 11) as f64 / (1u64 << 53) as f64 - 0.5) as f32) * scale
443            })
444            .collect()
445    }
446
447    /// `(R(t−j) q)·k == (R(t) q)·(R(j) k)` — the relative table against
448    /// the absolute `rope_rotate_scaled` of the full-attention path, so
449    /// the sign convention is the runtime's own, not assumed.
450    #[test]
451    fn relative_rotation_equals_absolute_pair() {
452        let hd = 16;
453        let inv = crate::attention::rope_inv_freq(hd, 10_000.0);
454        let rope = BoundedRope::new(64, &inv, 1.0);
455        for (t, j) in [(0usize, 0usize), (5, 5), (7, 3), (63, 0), (300, 250), (1000, 990)] {
456            let q = synth(hd, t as u64 + 1, 1.0);
457            let k = synth(hd, j as u64 + 77, 1.0);
458            let mut qa = q.clone();
459            let mut ka = k.clone();
460            crate::attention::rope_rotate_scaled(&mut qa, t, &inv, 1.0);
461            crate::attention::rope_rotate_scaled(&mut ka, j, &inv, 1.0);
462            let absolute: f64 = qa.iter().zip(&ka).map(|(a, b)| (*a as f64) * (*b as f64)).sum();
463            let mut qr = vec![0.0; hd];
464            rope.rotate(t - j, &q, &mut qr);
465            let relative: f64 = qr.iter().zip(&k).map(|(a, b)| (*a as f64) * (*b as f64)).sum();
466            assert!(
467                (absolute - relative).abs() < 2e-4,
468                "t={t} j={j}: absolute {absolute} vs relative {relative}"
469            );
470        }
471    }
472
473    #[test]
474    fn rollback_restores_overwritten_slots_bit_for_bit() {
475        let (kvh, hd, w) = (2, 4, 8);
476        let mut st = BoundedState::new(kvh, hd, w);
477        for p in 0..20 {
478            st.insert(&synth(kvh * hd, p, 1.0), &synth(kvh * hd, 100 + p, 1.0));
479        }
480        let snap = st.snapshot();
481        for p in 20..25 {
482            st.insert(&synth(kvh * hd, p, 1.0), &synth(kvh * hd, 100 + p, 1.0));
483        }
484        assert_eq!(st.rollback(5), 5);
485        assert!(st.same_state(&{
486            let mut s = BoundedState::new(kvh, hd, w);
487            s.restore(&snap);
488            s
489        }));
490        assert_eq!(st.seen, 20);
491        assert_eq!(st.len(), w);
492    }
493}