#[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;
fn reseal_every() -> usize {
static R: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*R.get_or_init(|| {
std::env::var("CMF_O1_RESEAL")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0)
})
}
const RESEAL_CAP: usize = 256;
fn far_only() -> bool {
static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ON.get_or_init(|| std::env::var("CMF_O1_FARONLY").as_deref() == Ok("1"))
}
#[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>,
samp_k: Vec<f32>,
samp_v: Vec<f32>,
samp_len: usize,
samp_head: usize,
since_reseal: usize,
}
#[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>,
samp_q: Vec<f32>,
samp_q_len: usize,
samp_q_head: usize,
}
pub struct O1Snapshot {
win_k: Vec<f32>,
win_v: Vec<f32>,
win_len: usize,
win_head: usize,
heads: Vec<(Vec<f32>, Vec<f32>, Vec<f32>, usize)>,
}
impl NystromState {
pub fn snapshot(&self) -> O1Snapshot {
O1Snapshot {
win_k: self.group.win_k.clone(),
win_v: self.group.win_v.clone(),
win_len: self.group.win_len,
win_head: self.group.win_head,
heads: self
.heads
.iter()
.map(|h| (h.t_hat.clone(), h.z_hat.clone(), h.m_max.clone(), h.far_len))
.collect(),
}
}
pub fn restore(&mut self, s: &O1Snapshot) {
self.group.win_k = s.win_k.clone();
self.group.win_v = s.win_v.clone();
self.group.win_len = s.win_len;
self.group.win_head = s.win_head;
debug_assert_eq!(self.heads.len(), s.heads.len());
for (h, (t, z, m, fl)) in self.heads.iter_mut().zip(&s.heads) {
h.t_hat = t.clone();
h.z_hat = z.clone();
h.m_max = m.clone();
h.far_len = *fl;
}
}
}
pub struct O1DeviceView<'a> {
pub m_eff: usize,
pub w: usize,
pub sink_len: usize,
pub d: usize,
pub dv: usize,
pub exact_only: bool,
pub scale: f32,
pub win_len: usize,
pub win_head: usize,
pub far_len: usize,
pub win_k: &'a [f32],
pub win_v: &'a [f32],
pub sink_k: &'a [f32],
pub sink_v: &'a [f32],
pub k_tilde: &'a [f32],
pub heads: Vec<O1HeadView<'a>>,
}
pub struct O1HeadView<'a> {
pub rect_fm: bool,
pub t_hat: &'a [f32],
pub z_hat: &'a [f32],
pub m_max: &'a [f32],
pub q_tilde: &'a [f32],
pub mu: &'a [f32],
}
impl NystromState {
pub fn device_view(&self) -> O1DeviceView<'_> {
let g = &self.group;
O1DeviceView {
m_eff: g.m_eff,
w: g.w,
sink_len: g.sink_len,
d: g.d,
dv: g.dv,
exact_only: g.exact_only,
scale: g.scale,
win_len: g.win_len,
win_head: g.win_head,
far_len: self.heads.first().map_or(0, |h| h.far_len),
win_k: &g.win_k,
win_v: &g.win_v,
sink_k: &g.sink_k,
sink_v: &g.sink_v,
k_tilde: &g.k_tilde,
heads: self
.heads
.iter()
.map(|h| O1HeadView {
rect_fm: h.rect == O1Rect::Fm,
t_hat: &h.t_hat,
z_hat: &h.z_hat,
m_max: &h.m_max,
q_tilde: &h.q_tilde,
mu: &h.mu,
})
.collect(),
}
}
}
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(),
samp_k: Vec::new(),
samp_v: Vec::new(),
samp_len: 0,
samp_head: 0,
since_reseal: 0,
},
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);
let rs = reseal_every();
for (h, head) in self.heads.iter_mut().enumerate() {
let qh = &q_all[h * d..(h + 1) * d];
if rs > 0 && !self.group.exact_only {
if head.samp_q.is_empty() {
head.samp_q = vec![0.0; RESEAL_CAP * d];
}
let sp = head.samp_q_head;
head.samp_q[sp * d..(sp + 1) * d].copy_from_slice(qh);
head.samp_q_head = (sp + 1) % RESEAL_CAP;
head.samp_q_len = (head.samp_q_len + 1).min(RESEAL_CAP);
}
head.step(&self.group, qh, &mut out_all[h * dv..(h + 1) * dv]);
}
if rs > 0 && self.group.since_reseal >= rs && self.group.samp_len >= 2 * self.group.m_eff {
self.reseal();
}
}
fn reseal(&mut self) {
let g = &mut self.group;
let (d, dv, m_eff) = (g.d, g.dv, g.m_eff);
let n = g.samp_len;
let start = if n == RESEAL_CAP { g.samp_head } else { 0 };
let mut ks = vec![0.0f32; n * d];
let mut vs = vec![0.0f32; n * dv];
for i in 0..n {
let idx = (start + i) % RESEAL_CAP;
ks[i * d..(i + 1) * d].copy_from_slice(&g.samp_k[idx * d..(idx + 1) * d]);
vs[i * dv..(i + 1) * dv].copy_from_slice(&g.samp_v[idx * dv..(idx + 1) * dv]);
}
let k_tilde64 = seg_means(&ks, n, d, m_eff);
g.k_tilde = k_tilde64.iter().map(|&x| x as f32).collect();
g.since_reseal = 0;
for head in &mut self.heads {
let qn = head.samp_q_len;
if qn < m_eff {
continue; }
let qstart = if qn == RESEAL_CAP {
head.samp_q_head
} else {
0
};
let mut qs = vec![0.0f32; qn * d];
for i in 0..qn {
let idx = (qstart + i) % RESEAL_CAP;
qs[i * d..(i + 1) * d].copy_from_slice(&head.samp_q[idx * d..(idx + 1) * d]);
}
let q_tilde64 = seg_means(&qs, qn, d, m_eff);
head.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);
head.mu = mu64.iter().map(|&x| x as f32).collect();
head.t_hat.iter_mut().for_each(|x| *x = 0.0);
head.z_hat.iter_mut().for_each(|x| *x = 0.0);
head.m_max.iter_mut().for_each(|x| *x = f32::NEG_INFINITY);
head.far_len = 0;
for i in 0..n {
head.far_absorb(
m_eff,
d,
dv,
g.scale,
&ks[i * d..(i + 1) * d],
&vs[i * dv..(i + 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);
}
if reseal_every() > 0 {
if g.samp_k.is_empty() {
g.samp_k = vec![0.0; RESEAL_CAP * d];
g.samp_v = vec![0.0; RESEAL_CAP * dv];
}
let sp = g.samp_head;
g.samp_k[sp * d..(sp + 1) * d].copy_from_slice(&g.win_k[slot * d..(slot + 1) * d]);
g.samp_v[sp * dv..(sp + 1) * dv]
.copy_from_slice(&g.win_v[slot * dv..(slot + 1) * dv]);
g.samp_head = (sp + 1) % RESEAL_CAP;
g.samp_len = (g.samp_len + 1).min(RESEAL_CAP);
g.since_reseal += 1;
}
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 {
tracing::info!(
"o1 seal: exact-only (prompt t={t} <= w {} + sink {} + slack {}) — \
not graph-portable; longer prompt or smaller --o1-window lifts it",
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);
if m_eff < self.m {
use std::sync::atomic::{AtomicBool, Ordering};
static SAID: AtomicBool = AtomicBool::new(false);
if !SAID.swap(true, Ordering::Relaxed) {
tracing::warn!(
"o1: landmark budget m={} clamped to m_eff={} — the prefill is {t} tokens \
and the skeleton takes t/8. Prefill at least {} tokens to use the budget \
you asked for.",
self.m,
m_eff,
self.m * 8
);
}
}
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(),
samp_q: Vec::new(),
samp_q_len: 0,
samp_q_head: 0,
}
}
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 skip_win = far_only() && !g.exact_only && self.far_len > 0;
let n = if skip_win { ns } else { 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);
}
if !skip_win {
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);
let k = &g.win_k[slot * d..(slot + 1) * d];
let v = &g.win_v[slot * dv..(slot + 1) * dv];
let (m_eff, scale) = (g.m_eff, g.scale);
self.far_absorb_slices(m_eff, d, dv, scale, k, v);
}
fn far_absorb(&mut self, m_eff: usize, d: usize, dv: usize, scale: f32, k: &[f32], v: &[f32]) {
self.far_absorb_slices(m_eff, d, dv, scale, k, v);
}
fn far_absorb_slices(
&mut self,
m_eff: usize,
d: usize,
dv: usize,
scale: f32,
k: &[f32],
v: &[f32],
) {
if self.scr_l.len() < m_eff {
self.scr_l.resize(m_eff, 0.0);
}
for i in 0..m_eff {
self.scr_l[i] = crate::attention::dot_f32(&self.q_tilde[i * d..(i + 1) * d], k) * scale;
}
for i in 0..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], v, 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)
}