1use crate::qtensor::QTensor;
26
27pub const UNDO_DEPTH: usize = 64;
31
32pub struct BoundedWeights {
34 pub wq: QTensor,
35 pub wk: QTensor,
36 pub wv: QTensor,
37 pub wo: QTensor,
38 pub sink_k: Vec<f32>,
40 pub sink_v: Vec<f32>,
42 pub sink: usize,
43 pub window: usize,
44}
45
46#[derive(Debug, Clone)]
51pub struct BoundedRope {
52 pub window: usize,
53 pub half: usize,
56 pub cos: Vec<f32>,
57 pub sin: Vec<f32>,
58}
59
60impl BoundedRope {
61 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 #[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#[derive(Debug, Clone)]
105pub struct BoundedSnapshot {
106 ring_k: Vec<f32>,
107 ring_v: Vec<f32>,
108 seen: usize,
109}
110
111#[derive(Debug, Clone)]
116pub struct BoundedState {
117 pub window: usize,
118 pub num_kv_heads: usize,
119 pub head_dim: usize,
120 pub ring_k: Vec<f32>,
122 pub ring_v: Vec<f32>,
124 pub seen: usize,
126 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 #[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 #[inline]
165 pub fn head(&self) -> usize {
166 self.seen % self.window
167 }
168
169 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 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 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 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 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 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 #[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 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 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 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
355pub 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 pub scale: f32,
364 pub rope: &'a BoundedRope,
365 pub pool: Option<&'a crate::pool::Pool>,
366}
367
368pub 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
394pub 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 #[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}