use crate::qtensor::QTensor;
pub const UNDO_DEPTH: usize = 64;
pub struct BoundedWeights {
pub wq: QTensor,
pub wk: QTensor,
pub wv: QTensor,
pub wo: QTensor,
pub sink_k: Vec<f32>,
pub sink_v: Vec<f32>,
pub sink: usize,
pub window: usize,
}
#[derive(Debug, Clone)]
pub struct BoundedRope {
pub window: usize,
pub half: usize,
pub cos: Vec<f32>,
pub sin: Vec<f32>,
}
impl BoundedRope {
pub fn new(window: usize, inv_freq: &[f32], rope_scale: f32) -> Self {
let half = inv_freq.len();
let s2 = (rope_scale as f64) * (rope_scale as f64);
let mut cos = Vec::with_capacity(window * half);
let mut sin = Vec::with_capacity(window * half);
for delta in 0..window {
for &f in inv_freq {
let (sn, cs) = ((delta as f64) * (f as f64)).sin_cos();
cos.push((cs * s2) as f32);
sin.push((sn * s2) as f32);
}
}
Self {
window,
half,
cos,
sin,
}
}
#[inline]
pub fn rotate(&self, delta: usize, x: &[f32], out: &mut [f32]) {
let half = self.half;
let c = &self.cos[delta * half..(delta + 1) * half];
let s = &self.sin[delta * half..(delta + 1) * half];
for i in 0..half {
let x0 = x[i];
let x1 = x[i + half];
out[i] = x0 * c[i] - x1 * s[i];
out[i + half] = x0 * s[i] + x1 * c[i];
}
let r = 2 * half;
if x.len() > r {
out[r..x.len()].copy_from_slice(&x[r..]);
}
}
}
#[derive(Debug, Clone)]
pub struct BoundedSnapshot {
ring_k: Vec<f32>,
ring_v: Vec<f32>,
seen: usize,
}
#[derive(Debug, Clone)]
pub struct BoundedState {
pub window: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
pub ring_k: Vec<f32>,
pub ring_v: Vec<f32>,
pub seen: usize,
undo_k: Vec<f32>,
undo_v: Vec<f32>,
undo_len: usize,
undo_head: usize,
}
impl BoundedState {
pub fn new(num_kv_heads: usize, head_dim: usize, window: usize) -> Self {
let n = num_kv_heads * window * head_dim;
let u = UNDO_DEPTH * num_kv_heads * head_dim;
Self {
window,
num_kv_heads,
head_dim,
ring_k: vec![0.0; n],
ring_v: vec![0.0; n],
seen: 0,
undo_k: vec![0.0; u],
undo_v: vec![0.0; u],
undo_len: 0,
undo_head: 0,
}
}
#[inline]
pub fn len(&self) -> usize {
self.seen.min(self.window)
}
#[inline]
pub fn is_empty(&self) -> bool {
self.seen == 0
}
#[inline]
pub fn head(&self) -> usize {
self.seen % self.window
}
pub fn state_bytes(&self) -> usize {
(self.ring_k.len() + self.ring_v.len()) * std::mem::size_of::<f32>()
+ std::mem::size_of::<u64>()
}
pub fn clear(&mut self) {
self.ring_k.fill(0.0);
self.ring_v.fill(0.0);
self.seen = 0;
self.undo_len = 0;
self.undo_head = 0;
}
pub fn insert(&mut self, k: &[f32], v: &[f32]) {
let (kvh, hd, w) = (self.num_kv_heads, self.head_dim, self.window);
debug_assert_eq!(k.len(), kvh * hd);
debug_assert_eq!(v.len(), kvh * hd);
let slot = self.head();
let u = self.undo_head;
for h in 0..kvh {
let r = (h * w + slot) * hd;
let uo = (u * kvh + h) * hd;
self.undo_k[uo..uo + hd].copy_from_slice(&self.ring_k[r..r + hd]);
self.undo_v[uo..uo + hd].copy_from_slice(&self.ring_v[r..r + hd]);
self.ring_k[r..r + hd].copy_from_slice(&k[h * hd..(h + 1) * hd]);
self.ring_v[r..r + hd].copy_from_slice(&v[h * hd..(h + 1) * hd]);
}
self.undo_head = (u + 1) % UNDO_DEPTH;
self.undo_len = (self.undo_len + 1).min(UNDO_DEPTH);
self.seen += 1;
}
pub fn rollback(&mut self, n: usize) -> usize {
let (kvh, hd, w) = (self.num_kv_heads, self.head_dim, self.window);
let n = n.min(self.undo_len).min(self.seen);
for _ in 0..n {
self.seen -= 1;
let slot = self.seen % w;
let u = (self.undo_head + UNDO_DEPTH - 1) % UNDO_DEPTH;
for h in 0..kvh {
let r = (h * w + slot) * hd;
let uo = (u * kvh + h) * hd;
self.ring_k[r..r + hd].copy_from_slice(&self.undo_k[uo..uo + hd]);
self.ring_v[r..r + hd].copy_from_slice(&self.undo_v[uo..uo + hd]);
}
self.undo_head = u;
self.undo_len -= 1;
}
n
}
pub fn snapshot(&self) -> BoundedSnapshot {
BoundedSnapshot {
ring_k: self.ring_k.clone(),
ring_v: self.ring_v.clone(),
seen: self.seen,
}
}
pub fn restore(&mut self, s: &BoundedSnapshot) {
debug_assert_eq!(s.ring_k.len(), self.ring_k.len());
self.ring_k.copy_from_slice(&s.ring_k);
self.ring_v.copy_from_slice(&s.ring_v);
self.seen = s.seen;
self.undo_len = 0;
self.undo_head = 0;
}
pub fn same_state(&self, other: &BoundedState) -> bool {
self.window == other.window
&& self.seen == other.seen
&& self.ring_k == other.ring_k
&& self.ring_v == other.ring_v
}
#[allow(clippy::too_many_arguments)]
pub fn attend(
&self,
q: &[f32],
num_heads: usize,
sink_k: &[f32],
sink_v: &[f32],
sink: usize,
rope: &BoundedRope,
scale: f32,
out: &mut [f32],
) {
let (kvh, hd, w) = (self.num_kv_heads, self.head_dim, self.window);
let nh = num_heads;
let hpk = nh / kvh.max(1);
debug_assert_eq!(hpk * kvh, nh);
debug_assert_eq!(q.len(), nh * hd);
debug_assert_eq!(out.len(), nh * hd);
debug_assert_eq!(sink_k.len(), kvh * sink * hd);
debug_assert_eq!(rope.window, w);
let m = self.len();
let n = sink + m;
let head = self.head();
thread_local! {
static SCRATCH: std::cell::RefCell<(Vec<f32>, Vec<f32>)> =
const { std::cell::RefCell::new((Vec::new(), Vec::new())) };
}
SCRATCH.with(|s| {
let mut s = s.borrow_mut();
let (scores, qrot) = &mut *s;
scores.clear();
scores.resize(nh * n, 0.0);
qrot.clear();
qrot.resize(nh * hd, 0.0);
for h in 0..nh {
let g = h / hpk;
let qh = &q[h * hd..(h + 1) * hd];
for s in 0..sink {
let kr = &sink_k[(g * sink + s) * hd..(g * sink + s + 1) * hd];
scores[h * n + s] = crate::attention::dot_f32(qh, kr) * scale;
}
}
for d in 0..m {
let slot = (head + w - 1 - d) % w;
for h in 0..nh {
rope.rotate(d, &q[h * hd..(h + 1) * hd], &mut qrot[h * hd..(h + 1) * hd]);
}
for h in 0..nh {
let g = h / hpk;
let kr = &self.ring_k[(g * w + slot) * hd..(g * w + slot + 1) * hd];
scores[h * n + sink + d] =
crate::attention::dot_f32(&qrot[h * hd..(h + 1) * hd], kr) * scale;
}
}
for h in 0..nh {
let sc = &mut scores[h * n..(h + 1) * n];
let max = sc.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for v in sc.iter_mut() {
*v = (*v - max).exp();
sum += *v;
}
if sum > 0.0 {
for v in sc.iter_mut() {
*v /= sum;
}
}
}
out.fill(0.0);
for h in 0..nh {
let g = h / hpk;
let oh = &mut out[h * hd..(h + 1) * hd];
let p = &scores[h * n..(h + 1) * n];
for s in 0..sink {
let vr = &sink_v[(g * sink + s) * hd..(g * sink + s + 1) * hd];
if p[s].abs() >= 1e-12 {
crate::attention::axpy_f32(oh, vr, p[s]);
}
}
for d in 0..m {
let slot = (head + w - 1 - d) % w;
let vr = &self.ring_v[(g * w + slot) * hd..(g * w + slot + 1) * hd];
let pw = p[sink + d];
if pw.abs() >= 1e-12 {
crate::attention::axpy_f32(oh, vr, pw);
}
}
}
});
}
}
pub struct BoundedAttnCfg<'a> {
pub num_heads: usize,
pub num_kv_heads: usize,
pub head_dim: usize,
pub hidden_size: usize,
pub scale: f32,
pub rope: &'a BoundedRope,
pub pool: Option<&'a crate::pool::Pool>,
}
pub fn bounded_attention(
hidden: &[f32],
w: &BoundedWeights,
cache: &mut crate::kv_cache::LayerKvCache,
cfg: &BoundedAttnCfg,
) -> Vec<f32> {
let (nh, nkv, hd) = (cfg.num_heads, cfg.num_kv_heads, cfg.head_dim);
let mut q = crate::attention::take_buf(nh * hd);
let mut k = crate::attention::take_buf(nkv * hd);
let mut v = crate::attention::take_buf(nkv * hd);
w.wq.matvec(hidden, &mut q, cfg.pool);
w.wk.matvec(hidden, &mut k, cfg.pool);
w.wv.matvec(hidden, &mut v, cfg.pool);
let mut ao = crate::attention::take_buf(nh * hd);
cache.bounded_step(&q, &k, &v, w, cfg.rope, cfg.scale, nh, &mut ao);
let mut out = crate::attention::take_buf(cfg.hidden_size);
w.wo.matvec(&ao, &mut out, cfg.pool);
crate::attention::recycle_buf(&mut q);
crate::attention::recycle_buf(&mut k);
crate::attention::recycle_buf(&mut v);
crate::attention::recycle_buf(&mut ao);
out
}
pub fn bounded_attention_batch(
normed_all: &[f32],
b: usize,
w: &BoundedWeights,
cache: &mut crate::kv_cache::LayerKvCache,
cfg: &BoundedAttnCfg,
) -> Vec<f32> {
let (nh, nkv, hd, hs) = (cfg.num_heads, cfg.num_kv_heads, cfg.head_dim, cfg.hidden_size);
debug_assert_eq!(normed_all.len(), b * hs);
let mut q_all = crate::attention::take_buf(b * nh * hd);
let mut k_all = crate::attention::take_buf(b * nkv * hd);
let mut v_all = crate::attention::take_buf(b * nkv * hd);
w.wq.matmat(normed_all, b, &mut q_all, cfg.pool);
w.wk.matmat(normed_all, b, &mut k_all, cfg.pool);
w.wv.matmat(normed_all, b, &mut v_all, cfg.pool);
let mut ao_all = crate::attention::take_buf(b * nh * hd);
for bi in 0..b {
let q = &q_all[bi * nh * hd..(bi + 1) * nh * hd];
let k = &k_all[bi * nkv * hd..(bi + 1) * nkv * hd];
let v = &v_all[bi * nkv * hd..(bi + 1) * nkv * hd];
let ao = &mut ao_all[bi * nh * hd..(bi + 1) * nh * hd];
cache.bounded_step(q, k, v, w, cfg.rope, cfg.scale, nh, ao);
}
let mut out = crate::attention::take_buf(b * hs);
w.wo.matmat(&ao_all, b, &mut out, cfg.pool);
crate::attention::recycle_buf(&mut q_all);
crate::attention::recycle_buf(&mut k_all);
crate::attention::recycle_buf(&mut v_all);
crate::attention::recycle_buf(&mut ao_all);
out
}
#[cfg(test)]
mod tests {
use super::*;
fn synth(n: usize, salt: u64, scale: f32) -> Vec<f32> {
(0..n)
.map(|i| {
let x = (i as u64)
.wrapping_mul(6364136223846793005)
.wrapping_add(salt.wrapping_mul(1442695040888963407) ^ 0x9E3779B97F4A7C15);
let x = (x ^ (x >> 31)).wrapping_mul(0xBF58476D1CE4E5B9);
(((x >> 11) as f64 / (1u64 << 53) as f64 - 0.5) as f32) * scale
})
.collect()
}
#[test]
fn relative_rotation_equals_absolute_pair() {
let hd = 16;
let inv = crate::attention::rope_inv_freq(hd, 10_000.0);
let rope = BoundedRope::new(64, &inv, 1.0);
for (t, j) in [(0usize, 0usize), (5, 5), (7, 3), (63, 0), (300, 250), (1000, 990)] {
let q = synth(hd, t as u64 + 1, 1.0);
let k = synth(hd, j as u64 + 77, 1.0);
let mut qa = q.clone();
let mut ka = k.clone();
crate::attention::rope_rotate_scaled(&mut qa, t, &inv, 1.0);
crate::attention::rope_rotate_scaled(&mut ka, j, &inv, 1.0);
let absolute: f64 = qa.iter().zip(&ka).map(|(a, b)| (*a as f64) * (*b as f64)).sum();
let mut qr = vec![0.0; hd];
rope.rotate(t - j, &q, &mut qr);
let relative: f64 = qr.iter().zip(&k).map(|(a, b)| (*a as f64) * (*b as f64)).sum();
assert!(
(absolute - relative).abs() < 2e-4,
"t={t} j={j}: absolute {absolute} vs relative {relative}"
);
}
}
#[test]
fn rollback_restores_overwritten_slots_bit_for_bit() {
let (kvh, hd, w) = (2, 4, 8);
let mut st = BoundedState::new(kvh, hd, w);
for p in 0..20 {
st.insert(&synth(kvh * hd, p, 1.0), &synth(kvh * hd, 100 + p, 1.0));
}
let snap = st.snapshot();
for p in 20..25 {
st.insert(&synth(kvh * hd, p, 1.0), &synth(kvh * hd, 100 + p, 1.0));
}
assert_eq!(st.rollback(5), 5);
assert!(st.same_state(&{
let mut s = BoundedState::new(kvh, hd, w);
s.restore(&snap);
s
}));
assert_eq!(st.seen, 20);
assert_eq!(st.len(), w);
}
}