const RIDGE_REL: f64 = 1e-6;
const DEN_EPS: f32 = 1e-30;
const EXACT_SLACK: usize = 8;
#[derive(Clone, Debug)]
pub struct NystromState {
m: usize,
w: usize,
sink: usize,
d: usize,
dv: usize,
m_eff: usize,
exact_only: bool,
scale: f32,
win_k: Vec<f32>,
win_v: Vec<f32>,
win_len: usize,
win_head: usize,
sink_k: Vec<f32>,
sink_v: Vec<f32>,
sink_len: usize,
t_hat: Vec<f32>,
z_hat: Vec<f32>,
m_max: Vec<f32>,
far_len: usize,
q_tilde: Vec<f32>,
k_tilde: Vec<f32>,
mu: Vec<f32>,
scr_s: Vec<f32>,
scr_fh: Vec<f32>,
scr_u: Vec<f32>,
scr_l: Vec<f32>,
}
impl NystromState {
pub fn new(m: usize, w: usize, sink: usize) -> Self {
assert!(m >= 4, "landmark budget must be at least 4");
assert!(w >= 1, "window must hold at least one key");
NystromState {
m,
w,
sink,
d: 0,
dv: 0,
m_eff: 0,
exact_only: true,
scale: 0.0,
win_k: Vec::new(),
win_v: Vec::new(),
win_len: 0,
win_head: 0,
sink_k: Vec::new(),
sink_v: Vec::new(),
sink_len: 0,
t_hat: Vec::new(),
z_hat: Vec::new(),
m_max: Vec::new(),
far_len: 0,
q_tilde: Vec::new(),
k_tilde: Vec::new(),
mu: Vec::new(),
scr_s: Vec::new(),
scr_fh: Vec::new(),
scr_u: Vec::new(),
scr_l: Vec::new(),
}
}
pub fn prefill(&mut self, qs: &[f32], ks: &[f32], vs: &[f32], t: usize, d: usize, dv: usize) {
assert_eq!(qs.len(), t * d);
assert_eq!(ks.len(), t * d);
assert_eq!(vs.len(), t * dv);
self.d = d;
self.dv = dv;
self.scale = 1.0 / (d as f32).sqrt();
self.win_len = 0;
self.win_head = 0;
self.far_len = 0;
self.sink_len = 0;
self.exact_only = t <= self.w + self.sink + EXACT_SLACK;
if self.exact_only {
self.win_k = Vec::with_capacity((t + 64) * d);
self.win_v = Vec::with_capacity((t + 64) * dv);
self.win_k.extend_from_slice(ks);
self.win_v.extend_from_slice(vs);
self.win_len = t;
self.scr_s = Vec::with_capacity(t + 64);
return;
}
self.sink_len = self.sink; self.sink_k = ks[..self.sink * d].to_vec();
self.sink_v = vs[..self.sink * dv].to_vec();
let m_eff = (t / 8).clamp(4, self.m);
self.m_eff = m_eff;
let q_tilde64 = seg_means(qs, t, d, m_eff);
let k_tilde64 = seg_means(ks, t, d, m_eff);
self.q_tilde = q_tilde64.iter().map(|&x| x as f32).collect();
self.k_tilde = k_tilde64.iter().map(|&x| x as f32).collect();
let mut au = vec![0.0f64; m_eff * m_eff];
for i in 0..m_eff {
for j in 0..m_eff {
let mut s = 0.0f64;
for c in 0..d {
s += q_tilde64[i * d + c] * k_tilde64[j * d + c];
}
au[i * m_eff + j] = (s * self.scale as f64).exp();
}
}
let mu64 = ridge_pinv(&au, m_eff);
self.mu = mu64.iter().map(|&x| x as f32).collect();
self.t_hat = vec![0.0; m_eff * dv];
self.z_hat = vec![0.0; m_eff];
self.m_max = vec![f32::NEG_INFINITY; m_eff];
self.win_k = vec![0.0; self.w * d];
self.win_v = vec![0.0; self.w * dv];
self.scr_s = vec![0.0; self.sink + self.w];
self.scr_fh = vec![0.0; m_eff];
self.scr_u = vec![0.0; m_eff];
self.scr_l = vec![0.0; m_eff];
for j in self.sink..t {
self.advance_window(&ks[j * d..(j + 1) * d], &vs[j * dv..(j + 1) * dv]);
}
}
pub fn step(&mut self, q: &[f32], k: &[f32], v: &[f32], out: &mut [f32]) {
let (d, dv) = (self.d, self.dv);
assert!(d > 0, "prefill() must run before step()");
assert_eq!(q.len(), d);
assert_eq!(k.len(), d);
assert_eq!(v.len(), dv);
assert_eq!(out.len(), dv);
self.advance_window(k, v);
let ns = self.sink_len;
let n = ns + self.win_len;
self.scr_s.resize(n, 0.0);
let mut c = f32::NEG_INFINITY;
for s in 0..ns {
let lg = dot(q, &self.sink_k[s * d..(s + 1) * d]) * self.scale;
self.scr_s[s] = lg;
c = c.max(lg);
}
for s in 0..self.win_len {
let lg =
crate::attention::dot_f32(q, &self.win_k[s * d..(s + 1) * d]) * self.scale;
self.scr_s[ns + s] = lg;
c = c.max(lg);
}
let mut far_den = 0.0f32;
let mut c_all = c;
let mut have_far = false;
if self.far_len > 0 {
let mut f = f32::NEG_INFINITY;
for a in 0..self.m_eff {
let s = crate::attention::dot_f32(q, &self.k_tilde[a * d..(a + 1) * d])
* self.scale;
self.scr_fh[a] = s;
f = f.max(s);
}
for a in 0..self.m_eff {
self.scr_fh[a] = (self.scr_fh[a] - f).exp();
}
for b in 0..self.m_eff {
let mut s = 0.0f32;
for a in 0..self.m_eff {
s += self.scr_fh[a] * self.mu[a * self.m_eff + b];
}
self.scr_u[b] = s;
}
for b in 0..self.m_eff {
c_all = c_all.max(f + self.m_max[b]);
}
for b in 0..self.m_eff {
let g = self.scr_u[b] * (f + self.m_max[b] - c_all).exp();
self.scr_u[b] = g;
far_den += g * self.z_hat[b];
}
if far_den >= 0.0 {
have_far = true;
} else {
far_den = 0.0;
}
}
for o in out.iter_mut() {
*o = 0.0;
}
if have_far {
for b in 0..self.m_eff {
crate::attention::axpy_f32(
out,
&self.t_hat[b * dv..(b + 1) * dv],
self.scr_u[b],
);
}
}
let mut den = far_den;
for s in 0..n {
let p = (self.scr_s[s] - c_all).exp();
den += p;
let vv = if s < ns {
&self.sink_v[s * dv..(s + 1) * dv]
} else {
&self.win_v[(s - ns) * dv..(s - ns + 1) * dv]
};
crate::attention::axpy_f32(out, vv, p);
}
let den = den.max(DEN_EPS);
for o in out.iter_mut() {
*o /= den;
}
}
pub fn memory_bytes(&self) -> usize {
(self.win_k.len()
+ self.win_v.len()
+ self.sink_k.len()
+ self.sink_v.len()
+ self.t_hat.len()
+ self.z_hat.len()
+ self.m_max.len()
+ self.q_tilde.len()
+ self.k_tilde.len()
+ self.mu.len()
+ self.scr_s.len()
+ self.scr_fh.len()
+ self.scr_u.len()
+ self.scr_l.len())
* std::mem::size_of::<f32>()
}
fn advance_window(&mut self, k: &[f32], v: &[f32]) {
let (d, dv) = (self.d, self.dv);
if !self.exact_only && self.win_len == self.w {
let slot = self.win_head;
self.far_insert(slot);
self.win_k[slot * d..(slot + 1) * d].copy_from_slice(k);
self.win_v[slot * dv..(slot + 1) * dv].copy_from_slice(v);
self.win_head = (self.win_head + 1) % self.w;
} else if self.exact_only {
self.win_k.extend_from_slice(k);
self.win_v.extend_from_slice(v);
self.win_len += 1;
} else {
self.win_k[self.win_len * d..(self.win_len + 1) * d].copy_from_slice(k);
self.win_v[self.win_len * dv..(self.win_len + 1) * dv].copy_from_slice(v);
self.win_len += 1;
}
}
fn far_insert(&mut self, slot: usize) {
let (d, dv) = (self.d, self.dv);
for i in 0..self.m_eff {
self.scr_l[i] = crate::attention::dot_f32(
&self.q_tilde[i * d..(i + 1) * d],
&self.win_k[slot * d..(slot + 1) * d],
) * self.scale;
}
for i in 0..self.m_eff {
let l = self.scr_l[i];
if l > self.m_max[i] {
let r = (self.m_max[i] - l).exp();
self.z_hat[i] *= r;
for e in self.t_hat[i * dv..(i + 1) * dv].iter_mut() {
*e *= r;
}
self.m_max[i] = l;
}
let e = (l - self.m_max[i]).exp();
self.z_hat[i] += e;
crate::attention::axpy_f32(
&mut self.t_hat[i * dv..(i + 1) * dv],
&self.win_v[slot * dv..(slot + 1) * dv],
e,
);
}
self.far_len += 1;
}
}
fn seg_means(xs: &[f32], t: usize, d: usize, m: usize) -> Vec<f64> {
let mut out = vec![0.0f64; m * d];
for i in 0..m {
let lo = i * t / m;
let hi = (i + 1) * t / m;
for j in lo..hi {
for c in 0..d {
out[i * d + c] += xs[j * d + c] as f64;
}
}
let inv = 1.0 / (hi - lo) as f64;
for c in 0..d {
out[i * d + c] *= inv;
}
}
out
}
fn dot(a: &[f32], b: &[f32]) -> f32 {
let mut s = 0.0f32;
for (x, y) in a.iter().zip(b) {
s += x * y;
}
s
}
pub(crate) fn ridge_pinv(a: &[f64], n: usize) -> Vec<f64> {
let mut ata = vec![0.0f64; n * n];
for i in 0..n {
for j in 0..n {
let mut s = 0.0;
for k in 0..n {
s += a[k * n + i] * a[k * n + j];
}
ata[i * n + j] = s;
}
}
let mean_diag: f64 = (0..n).map(|i| ata[i * n + i]).sum::<f64>() / n as f64;
let mut lambda = RIDGE_REL * mean_diag.max(f64::MIN_POSITIVE);
for _ in 0..12 {
let mut g = ata.clone();
for i in 0..n {
g[i * n + i] += lambda;
}
if let Some(l) = cholesky(&mut g, n) {
let mut m_out = vec![0.0f64; n * n];
let mut x = vec![0.0f64; n];
for j in 0..n {
let rhs = &a[j * n..(j + 1) * n];
for i in 0..n {
let mut s = rhs[i];
for k in 0..i {
s -= l[i * n + k] * x[k];
}
x[i] = s / l[i * n + i];
}
for i in (0..n).rev() {
let mut s = x[i];
for k in i + 1..n {
s -= l[k * n + i] * x[k];
}
x[i] = s / l[i * n + i];
}
for i in 0..n {
m_out[i * n + j] = x[i];
}
}
return m_out;
}
lambda *= 10.0;
}
let mut fallback = vec![0.0f64; n * n];
for i in 0..n {
fallback[i * n + i] = 1.0 / mean_diag.max(f64::MIN_POSITIVE);
}
fallback
}
pub const O1_DEFAULT_M: usize = 32;
pub const O1_DEFAULT_W: usize = 128;
pub const O1_DEFAULT_SINK: usize = 4;
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum O1Layers {
All,
Deep(usize),
List(Vec<usize>),
}
#[derive(Clone, Debug)]
pub struct O1Cfg {
pub layers: O1Layers,
pub m: usize,
pub w: usize,
pub sink: usize,
}
pub enum O1Env {
Unset,
Off,
On(O1Cfg),
}
impl O1Cfg {
pub fn parse_layers(spec: &str) -> Option<O1Layers> {
let s = spec.trim();
match s {
"" | "off" | "0" | "none" => None,
"all" => Some(O1Layers::All),
_ => {
if let Some(n) = s.strip_prefix("deep") {
return n.parse::<usize>().ok().filter(|&n| n > 0).map(O1Layers::Deep);
}
let idx: Result<Vec<usize>, _> =
s.split(',').map(|p| p.trim().parse::<usize>()).collect();
idx.ok().filter(|v| !v.is_empty()).map(O1Layers::List)
}
}
}
pub fn from_spec(
spec: &str,
m: Option<usize>,
w: Option<usize>,
sink: Option<usize>,
) -> Option<O1Cfg> {
let layers = Self::parse_layers(spec)?;
let env = |k: &str| std::env::var(k).ok().and_then(|v| v.parse::<usize>().ok());
Some(O1Cfg {
layers,
m: m.or_else(|| env("CMF_O1_M")).unwrap_or(O1_DEFAULT_M).max(4),
w: w.or_else(|| env("CMF_O1_WINDOW")).unwrap_or(O1_DEFAULT_W).max(1),
sink: sink.or_else(|| env("CMF_O1_SINK")).unwrap_or(O1_DEFAULT_SINK),
})
}
pub fn from_json(v: &serde_json::Value) -> Option<O1Cfg> {
let layers = match v.get("layers") {
Some(serde_json::Value::String(s)) => Self::parse_layers(s)?,
Some(serde_json::Value::Array(a)) => O1Layers::List(
a.iter().filter_map(|x| x.as_u64().map(|n| n as usize)).collect(),
),
_ => return None,
};
let f = |k: &str| v.get(k).and_then(|x| x.as_u64()).map(|n| n as usize);
let env = |k: &str| std::env::var(k).ok().and_then(|s| s.parse::<usize>().ok());
Some(O1Cfg {
layers,
m: env("CMF_O1_M").or_else(|| f("m")).unwrap_or(O1_DEFAULT_M).max(4),
w: env("CMF_O1_WINDOW").or_else(|| f("w")).unwrap_or(O1_DEFAULT_W).max(1),
sink: env("CMF_O1_SINK").or_else(|| f("sink")).unwrap_or(O1_DEFAULT_SINK),
})
}
pub fn layer_flags(&self, num_layers: usize) -> Vec<bool> {
let mut flags = vec![false; num_layers];
match &self.layers {
O1Layers::All => flags.iter_mut().for_each(|f| *f = true),
O1Layers::Deep(n) => {
for f in flags.iter_mut().skip(num_layers.saturating_sub(*n)) {
*f = true;
}
}
O1Layers::List(idx) => {
for &i in idx {
if i < num_layers {
flags[i] = true;
}
}
}
}
flags
}
}
pub fn o1_from_env() -> O1Env {
match std::env::var("CMF_O1") {
Err(_) => O1Env::Unset,
Ok(s) => match O1Cfg::from_spec(&s, None, None, None) {
Some(cfg) => O1Env::On(cfg),
None => O1Env::Off,
},
}
}
fn cholesky(g: &mut [f64], n: usize) -> Option<&[f64]> {
for i in 0..n {
for j in 0..=i {
let mut s = g[i * n + j];
for k in 0..j {
s -= g[i * n + k] * g[j * n + k];
}
if i == j {
if s <= 0.0 || !s.is_finite() {
return None;
}
g[i * n + i] = s.sqrt();
} else {
g[i * n + j] = s / g[j * n + j];
}
}
}
Some(g)
}