#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum O1Rect {
Aggregate,
Fm,
}
const RIDGE_REL: f64 = 1e-6;
const DEN_EPS: f32 = 1e-30;
const EXACT_SLACK: usize = 8;
#[derive(Clone, Debug)]
pub struct NystromState {
group: NystromGroup,
heads: Vec<NystromHead>,
}
#[derive(Clone, Debug)]
struct NystromGroup {
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,
k_tilde: Vec<f32>,
}
#[derive(Clone, Debug)]
struct NystromHead {
rect: O1Rect,
t_hat: Vec<f32>,
z_hat: Vec<f32>,
m_max: Vec<f32>,
far_len: usize,
q_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 {
Self::new_group(m, w, sink, 1)
}
pub fn new_group(m: usize, w: usize, sink: usize, q_heads: usize) -> Self {
assert!(m >= 4, "landmark budget must be at least 4");
assert!(w >= 1, "window must hold at least one key");
assert!(q_heads >= 1, "a GQA group needs at least one query head");
NystromState {
group: NystromGroup {
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,
k_tilde: Vec::new(),
},
heads: (0..q_heads).map(|_| NystromHead::new()).collect(),
}
}
pub fn with_rect(mut self, rect: O1Rect) -> Self {
for h in &mut self.heads {
h.rect = rect;
}
self
}
pub fn num_q_heads(&self) -> usize {
self.heads.len()
}
pub fn far_len(&self, head: usize) -> usize {
self.heads[head].far_len
}
pub fn prefill(&mut self, qs: &[f32], ks: &[f32], vs: &[f32], t: usize, d: usize, dv: usize) {
assert_eq!(self.heads.len(), 1, "use prefill_group for a GQA group");
self.prefill_group(&[qs], ks, vs, t, d, dv);
}
pub fn prefill_group(
&mut self,
qs: &[&[f32]],
ks: &[f32],
vs: &[f32],
t: usize,
d: usize,
dv: usize,
) {
assert_eq!(qs.len(), self.heads.len(), "one query block per head");
for q in qs {
assert_eq!(q.len(), t * d);
}
assert_eq!(ks.len(), t * d);
assert_eq!(vs.len(), t * dv);
let Some(k_tilde64) = self.group.prefill_shared(ks, vs, t, d, dv) else {
for h in &mut self.heads {
h.seal_exact(t);
}
return;
};
for (h, q) in self.heads.iter_mut().zip(qs) {
h.seal(&self.group, q, t, &k_tilde64);
}
for j in self.group.sink..t {
Self::advance(
&mut self.group,
&mut self.heads,
&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]) {
assert_eq!(self.heads.len(), 1, "use step_group for a GQA group");
self.step_group(q, k, v, out);
}
pub fn step_group(&mut self, q_all: &[f32], k: &[f32], v: &[f32], out_all: &mut [f32]) {
let (d, dv) = (self.group.d, self.group.dv);
assert!(d > 0, "prefill() must run before step()");
let nh = self.heads.len();
assert_eq!(q_all.len(), nh * d);
assert_eq!(k.len(), d);
assert_eq!(v.len(), dv);
assert_eq!(out_all.len(), nh * dv);
Self::advance(&mut self.group, &mut self.heads, k, v);
for (h, head) in self.heads.iter_mut().enumerate() {
head.step(
&self.group,
&q_all[h * d..(h + 1) * d],
&mut out_all[h * dv..(h + 1) * dv],
);
}
}
pub fn memory_bytes(&self) -> usize {
self.group.memory_bytes()
+ self.heads.iter().map(NystromHead::memory_bytes).sum::<usize>()
}
fn advance(g: &mut NystromGroup, heads: &mut [NystromHead], k: &[f32], v: &[f32]) {
let (d, dv) = (g.d, g.dv);
if !g.exact_only && g.win_len == g.w {
let slot = g.win_head;
for h in heads.iter_mut() {
h.far_insert(g, slot);
}
g.win_k[slot * d..(slot + 1) * d].copy_from_slice(k);
g.win_v[slot * dv..(slot + 1) * dv].copy_from_slice(v);
g.win_head = (g.win_head + 1) % g.w;
} else if g.exact_only {
g.win_k.extend_from_slice(k);
g.win_v.extend_from_slice(v);
g.win_len += 1;
} else {
g.win_k[g.win_len * d..(g.win_len + 1) * d].copy_from_slice(k);
g.win_v[g.win_len * dv..(g.win_len + 1) * dv].copy_from_slice(v);
g.win_len += 1;
}
}
}
impl NystromGroup {
fn prefill_shared(
&mut self,
ks: &[f32],
vs: &[f32],
t: usize,
d: usize,
dv: usize,
) -> Option<Vec<f64>> {
self.d = d;
self.dv = dv;
self.scale = 1.0 / (d as f32).sqrt();
self.win_len = 0;
self.win_head = 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;
return None;
}
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 k_tilde64 = seg_means(ks, t, d, m_eff);
self.k_tilde = k_tilde64.iter().map(|&x| x as f32).collect();
self.win_k = vec![0.0; self.w * d];
self.win_v = vec![0.0; self.w * dv];
Some(k_tilde64)
}
fn memory_bytes(&self) -> usize {
(self.win_k.len()
+ self.win_v.len()
+ self.sink_k.len()
+ self.sink_v.len()
+ self.k_tilde.len())
* std::mem::size_of::<f32>()
}
}
impl NystromHead {
fn new() -> Self {
NystromHead {
rect: O1_DEFAULT_RECT,
t_hat: Vec::new(),
z_hat: Vec::new(),
m_max: Vec::new(),
far_len: 0,
q_tilde: Vec::new(),
mu: Vec::new(),
scr_s: Vec::new(),
scr_fh: Vec::new(),
scr_u: Vec::new(),
scr_l: Vec::new(),
}
}
fn seal_exact(&mut self, t: usize) {
self.far_len = 0;
self.scr_s = Vec::with_capacity(t + 64);
}
fn seal(&mut self, g: &NystromGroup, qs: &[f32], t: usize, k_tilde64: &[f64]) {
let (d, dv, m_eff) = (g.d, g.dv, g.m_eff);
self.far_len = 0;
let q_tilde64 = seg_means(qs, t, d, m_eff);
self.q_tilde = q_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 * g.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.scr_s = vec![0.0; g.sink + g.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];
}
fn step(&mut self, g: &NystromGroup, q: &[f32], out: &mut [f32]) {
let (d, dv) = (g.d, g.dv);
assert_eq!(q.len(), d);
assert_eq!(out.len(), dv);
let ns = g.sink_len;
let n = ns + g.win_len;
self.scr_s.resize(n, 0.0);
let mut c = f32::NEG_INFINITY;
for s in 0..ns {
let lg = dot(q, &g.sink_k[s * d..(s + 1) * d]) * g.scale;
self.scr_s[s] = lg;
c = c.max(lg);
}
for s in 0..g.win_len {
let lg = crate::attention::dot_f32(q, &g.win_k[s * d..(s + 1) * d]) * g.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..g.m_eff {
let s = crate::attention::dot_f32(q, &g.k_tilde[a * d..(a + 1) * d]) * g.scale;
self.scr_fh[a] = s;
f = f.max(s);
}
for a in 0..g.m_eff {
self.scr_fh[a] = (self.scr_fh[a] - f).exp();
}
for b in 0..g.m_eff {
let mut s = 0.0f32;
for a in 0..g.m_eff {
s += self.scr_fh[a] * self.mu[a * g.m_eff + b];
}
self.scr_u[b] = if self.rect == O1Rect::Fm { s.max(0.0) } else { s };
}
for b in 0..g.m_eff {
c_all = c_all.max(f + self.m_max[b]);
}
for b in 0..g.m_eff {
let gain = self.scr_u[b] * (f + self.m_max[b] - c_all).exp();
self.scr_u[b] = gain;
far_den += gain * 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..g.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 {
&g.sink_v[s * dv..(s + 1) * dv]
} else {
&g.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;
}
}
fn far_insert(&mut self, g: &NystromGroup, slot: usize) {
let (d, dv) = (g.d, g.dv);
for i in 0..g.m_eff {
self.scr_l[i] = crate::attention::dot_f32(
&self.q_tilde[i * d..(i + 1) * d],
&g.win_k[slot * d..(slot + 1) * d],
) * g.scale;
}
for i in 0..g.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],
&g.win_v[slot * dv..(slot + 1) * dv],
e,
);
}
self.far_len += 1;
}
fn memory_bytes(&self) -> usize {
(self.t_hat.len()
+ self.z_hat.len()
+ self.m_max.len()
+ self.q_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 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;
pub const O1_DEFAULT_RECT: O1Rect = O1Rect::Aggregate;
#[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 rect: O1Rect,
}
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 parse_rect(spec: &str) -> Option<O1Rect> {
match spec.trim() {
"agg" | "aggregate" => Some(O1Rect::Aggregate),
"fm" => Some(O1Rect::Fm),
_ => None,
}
}
fn rect_or_env(rect: Option<O1Rect>) -> O1Rect {
rect.or_else(|| std::env::var("CMF_O1_RECT").ok().as_deref().and_then(Self::parse_rect))
.unwrap_or(O1_DEFAULT_RECT)
}
pub fn from_spec(
spec: &str,
m: Option<usize>,
w: Option<usize>,
sink: Option<usize>,
rect: Option<O1Rect>,
) -> 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),
rect: Self::rect_or_env(rect),
})
}
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),
rect: Self::rect_or_env(None),
})
}
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, 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)
}