use crate::infer::GraphExt;
use crate::op::{Activation, Op};
use crate::{DType, Graph, NodeId, Shape};
pub const SPD_JACOBI_SWEEPS: u32 = 6;
pub fn spd_jacobi_sweeps() -> u32 {
std::env::var("RLX_SPD_JACOBI_SWEEPS")
.ok()
.and_then(|s| s.parse::<u32>().ok())
.filter(|&v| v > 0)
.unwrap_or(SPD_JACOBI_SWEEPS)
}
fn const_bytes(xs: &[f64], dtype: DType) -> Vec<u8> {
match dtype {
DType::F64 => xs.iter().flat_map(|v| v.to_le_bytes()).collect(),
_ => xs.iter().flat_map(|v| (*v as f32).to_le_bytes()).collect(),
}
}
fn cmat(g: &mut Graph, xs: &[f64], dims: &[usize], dtype: DType) -> NodeId {
g.add_node(
Op::Constant {
data: const_bytes(xs, dtype),
},
vec![],
Shape::new(dims, dtype),
)
}
fn cscalar(g: &mut Graph, v: f64, dtype: DType) -> NodeId {
cmat(g, &[v], &[1, 1], dtype)
}
fn round_robin_rounds(ne: usize) -> Vec<Vec<(usize, usize)>> {
let mut players: Vec<usize> = (0..ne).collect();
let mut rounds = Vec::with_capacity(ne - 1);
for _ in 0..ne - 1 {
let pairs: Vec<(usize, usize)> = (0..ne / 2)
.map(|i| {
let (a, b) = (players[i], players[ne - 1 - i]);
(a.min(b), a.max(b))
})
.collect();
rounds.push(pairs);
let last = players[ne - 1];
for i in (2..ne).rev() {
players[i] = players[i - 1];
}
players[1] = last;
}
rounds
}
fn diag_k(g: &mut Graph, m: NodeId, ident_k: NodeId) -> NodeId {
let d = g.mul(m, ident_k);
g.sum(d, vec![1], true) }
#[allow(clippy::too_many_arguments)]
fn one_round(
g: &mut Graph,
av: NodeId,
vv: NodeId,
spp: NodeId,
sqq: NodeId,
ident_n: NodeId,
ident_k: NodeId,
dt: DType,
) -> (NodeId, NodeId) {
let one = cscalar(g, 1.0, dt);
let two = cscalar(g, 2.0, dt);
let tiny = cscalar(g, 1e-30, dt);
let sbias = cscalar(g, 1e-12, dt);
let floor = cscalar(g, 1e-6, dt);
let sppt = g.transpose_(spp, vec![1, 0]); let sqqt = g.transpose_(sqq, vec![1, 0]);
let tpp = g.mm(sppt, av); let m_pp = g.mm(tpp, spp); let app = diag_k(g, m_pp, ident_k); let tqq = g.mm(sqqt, av);
let m_qq = g.mm(tqq, sqq);
let aqq = diag_k(g, m_qq, ident_k);
let m_pq = g.mm(tpp, sqq); let apq = diag_k(g, m_pq, ident_k);
let num = g.sub(aqq, app);
let den0 = g.mul(two, apq);
let den0b = g.add(den0, sbias); let den0sq = g.mul(den0b, den0b);
let den0abs = g.sqrt(den0sq);
let den0abst = g.add(den0abs, tiny);
let sgnden = g.div(den0b, den0abst);
let sfloor = g.mul(sgnden, floor);
let den = g.add(den0, sfloor);
let tau = g.div(num, den);
let tau2 = g.mul(tau, tau);
let atau = g.sqrt(tau2);
let satau = g.add(atau, tiny);
let sgn = g.div(tau, satau);
let tau2p1 = g.add(tau2, one);
let rt = g.sqrt(tau2p1);
let tden = g.add(atau, rt);
let t = g.div(sgn, tden);
let t2 = g.mul(t, t);
let t2p1 = g.add(t2, one);
let rc = g.sqrt(t2p1);
let c = g.div(one, rc);
let s = g.mul(t, c);
let cm1 = g.sub(c, one);
let negs = g.neg(s);
let d_cm1 = g.mul(cm1, ident_k); let d_s = g.mul(s, ident_k);
let d_negs = g.mul(negs, ident_k);
let pp_diag = {
let l = g.mm(spp, d_cm1);
g.mm(l, sppt)
};
let qq_diag = {
let l = g.mm(sqq, d_cm1);
g.mm(l, sqqt)
};
let pq = {
let l = g.mm(spp, d_s);
g.mm(l, sqqt)
};
let qp = {
let l = g.mm(sqq, d_negs);
g.mm(l, sppt)
};
let j1 = g.add(ident_n, pp_diag);
let j2 = g.add(j1, qq_diag);
let j3 = g.add(j2, pq);
let jj = g.add(j3, qp);
let jt = g.transpose_(jj, vec![1, 0]);
let jta = g.mm(jt, av);
let av2 = g.mm(jta, jj);
let vv2 = g.mm(vv, jj);
(av2, vv2)
}
fn round_body(n: usize, k: usize, dt: DType) -> Graph {
let mut body = Graph::new("jacobi_round");
let carry = body.input("carry", Shape::new(&[2 * n, n], dt));
let spp = body.input("spp", Shape::new(&[n, k], dt));
let sqq = body.input("sqq", Shape::new(&[n, k], dt));
let a_in = body.narrow_(carry, 0, 0, n);
let v_in = body.narrow_(carry, 0, n, n);
let mut ident = vec![0f64; n * n];
for i in 0..n {
ident[i * n + i] = 1.0;
}
let ident_n = cmat(&mut body, &ident, &[n, n], dt);
let mut idk = vec![0f64; k * k];
for i in 0..k {
idk[i * k + i] = 1.0;
}
let ident_k = cmat(&mut body, &idk, &[k, k], dt);
let (a_out, v_out) = one_round(&mut body, a_in, v_in, spp, sqq, ident_n, ident_k, dt);
let out = body.concat_(vec![a_out, v_out], 0);
body.set_outputs(vec![out]);
body
}
fn eigensolve(g: &mut Graph, a: NodeId, n: usize, sweeps: u32, dt: DType) -> (NodeId, NodeId) {
let mut ident = vec![0f64; n * n];
for i in 0..n {
ident[i * n + i] = 1.0;
}
let iv = cmat(g, &ident, &[n, n], dt);
if n < 2 {
return (a, iv); }
let mut carry = g.concat_(vec![a, iv], 0);
let ne = if n.is_multiple_of(2) { n } else { n + 1 };
let k = ne / 2;
let rounds = round_robin_rounds(ne);
let nr = rounds.len();
if std::env::var("RLX_SPD_UNROLL").as_deref() == Ok("1") {
let mut idk = vec![0f64; k * k];
for i in 0..k {
idk[i * k + i] = 1.0;
}
let ident_k = cmat(g, &idk, &[k, k], dt);
let sel: Vec<(NodeId, NodeId)> = rounds
.iter()
.map(|round| {
let mut sp = vec![0f64; n * k];
let mut sq = vec![0f64; n * k];
for (i, &(p, q)) in round.iter().enumerate() {
if p < n {
sp[p * k + i] = 1.0;
}
if q < n {
sq[q * k + i] = 1.0;
}
}
(cmat(g, &sp, &[n, k], dt), cmat(g, &sq, &[n, k], dt))
})
.collect();
let (mut av, mut vv) = (a, iv);
for _ in 0..sweeps {
for &(spp_r, sqq_r) in &sel {
let (a2, v2) = one_round(g, av, vv, spp_r, sqq_r, iv, ident_k, dt);
av = a2;
vv = v2;
}
}
return (av, vv);
}
let total = nr * sweeps as usize;
let mut spp = vec![0f64; total * n * k];
let mut sqq = vec![0f64; total * n * k];
for s in 0..sweeps as usize {
for (r, round) in rounds.iter().enumerate() {
let base = (s * nr + r) * n * k;
for (i, &(p, q)) in round.iter().enumerate() {
if p < n {
spp[base + p * k + i] = 1.0;
}
if q < n {
sqq[base + q * k + i] = 1.0;
}
}
}
}
let spp_c = cmat(g, &spp, &[total, n, k], dt);
let sqq_c = cmat(g, &sqq, &[total, n, k], dt);
let body = round_body(n, k, dt);
carry = g.scan_with_xs(carry, &[spp_c, sqq_c], body, total as u32);
let av = g.narrow_(carry, 0, 0, n);
let vv = g.narrow_(carry, 0, n, n);
(av, vv)
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub enum SpectralFn {
Re,
Log,
Sqrt,
InvSqrt,
}
fn spectral_fmat(
g: &mut Graph,
av: NodeId,
n: usize,
eps: f64,
f: SpectralFn,
dt: DType,
) -> (NodeId, Vec<NodeId>) {
let epsn = cscalar(g, eps, dt);
let half = cscalar(g, 0.5, dt);
let one = cscalar(g, 1.0, dt);
let mut terms: Vec<NodeId> = Vec::with_capacity(n);
let mut lam_parts: Vec<NodeId> = Vec::with_capacity(n);
for i in 0..n {
let ri = g.narrow_(av, 0, i, 1); let lam = g.narrow_(ri, 1, i, 1); lam_parts.push(g.reshape_(lam, vec![1])); let sumv = g.add(lam, epsn);
let diff = g.sub(lam, epsn);
let d2 = g.mul(diff, diff);
let ad = g.sqrt(d2);
let sad = g.add(sumv, ad);
let mx = g.mul(half, sad);
let fl = match f {
SpectralFn::Re => mx,
SpectralFn::Log => g.activation(Activation::Log, mx, Shape::new(&[1, 1], dt)),
SpectralFn::Sqrt => g.sqrt(mx),
SpectralFn::InvSqrt => {
let s = g.sqrt(mx);
g.div(one, s)
}
};
let mut eii = vec![0f64; n * n];
eii[i * n + i] = 1.0;
let eii = cmat(g, &eii, &[n, n], dt);
terms.push(g.mul(fl, eii));
}
let mut fmat = terms[0];
for t in &terms[1..] {
fmat = g.add(fmat, *t);
}
(fmat, lam_parts)
}
pub fn spectral_matfn(
g: &mut Graph,
a: NodeId,
n: usize,
sweeps: u32,
eps: f64,
f: SpectralFn,
) -> NodeId {
let dt = g.shape(a).dtype();
let (av, vv) = eigensolve(g, a, n, sweeps, dt);
let (fmat, _) = spectral_fmat(g, av, n, eps, f, dt);
let vt = g.transpose_(vv, vec![1, 0]);
let vf = g.mm(vv, fmat);
g.mm(vf, vt)
}
fn ident_mat3(g: &mut Graph, n: usize, dt: DType) -> NodeId {
let mut m = vec![0f64; n * n];
for i in 0..n {
m[i * n + i] = 1.0;
}
cmat(g, &m, &[1, n, n], dt)
}
fn jacobi_cs(
g: &mut Graph,
app: NodeId,
aqq: NodeId,
apq: NodeId,
dt: DType,
) -> (NodeId, NodeId, NodeId) {
let one = cscalar(g, 1.0, dt);
let two = cscalar(g, 2.0, dt);
let tiny = cscalar(g, 1e-30, dt);
let sbias = cscalar(g, 1e-12, dt);
let floor = cscalar(g, 1e-6, dt);
let num = g.sub(aqq, app);
let den0 = g.mul(two, apq);
let den0b = g.add(den0, sbias);
let den0sq = g.mul(den0b, den0b);
let den0abs = g.sqrt(den0sq);
let den0abst = g.add(den0abs, tiny);
let sgnden = g.div(den0b, den0abst);
let sfloor = g.mul(sgnden, floor);
let den = g.add(den0, sfloor);
let tau = g.div(num, den);
let tau2 = g.mul(tau, tau);
let atau = g.sqrt(tau2);
let satau = g.add(atau, tiny);
let sgn = g.div(tau, satau);
let tau2p1 = g.add(tau2, one);
let rt = g.sqrt(tau2p1);
let tden = g.add(atau, rt);
let t = g.div(sgn, tden);
let t2 = g.mul(t, t);
let t2p1 = g.add(t2, one);
let rc = g.sqrt(t2p1);
let c = g.div(one, rc);
let s = g.mul(t, c);
let cm1 = g.sub(c, one);
let negs = g.neg(s);
(cm1, s, negs)
}
#[allow(clippy::too_many_arguments)]
fn one_round_b(
g: &mut Graph,
av: NodeId,
vv: NodeId,
spp: NodeId,
sqq: NodeId,
ident_n: NodeId,
ident_k: NodeId,
dt: DType,
) -> (NodeId, NodeId) {
let sppt = g.transpose_(spp, vec![0, 2, 1]); let sqqt = g.transpose_(sqq, vec![0, 2, 1]);
let tpp = g.mm(sppt, av); let m_pp = g.mm(tpp, spp); let app = {
let d = g.mul(m_pp, ident_k);
g.sum(d, vec![2], true)
}; let tqq = g.mm(sqqt, av);
let m_qq = g.mm(tqq, sqq);
let aqq = {
let d = g.mul(m_qq, ident_k);
g.sum(d, vec![2], true)
};
let m_pq = g.mm(tpp, sqq);
let apq = {
let d = g.mul(m_pq, ident_k);
g.sum(d, vec![2], true)
};
let (cm1, s, negs) = jacobi_cs(g, app, aqq, apq, dt);
let d_cm1 = g.mul(cm1, ident_k); let d_s = g.mul(s, ident_k);
let d_negs = g.mul(negs, ident_k);
let pp_diag = {
let l = g.mm(spp, d_cm1);
g.mm(l, sppt)
};
let qq_diag = {
let l = g.mm(sqq, d_cm1);
g.mm(l, sqqt)
};
let pq = {
let l = g.mm(spp, d_s);
g.mm(l, sqqt)
};
let qp = {
let l = g.mm(sqq, d_negs);
g.mm(l, sppt)
};
let j1 = g.add(ident_n, pp_diag);
let j2 = g.add(j1, qq_diag);
let j3 = g.add(j2, pq);
let jj = g.add(j3, qp);
let jt = g.transpose_(jj, vec![0, 2, 1]); let jta = g.mm(jt, av);
let av2 = g.mm(jta, jj);
let vv2 = g.mm(vv, jj);
(av2, vv2)
}
fn round_body_b(batch: usize, n: usize, k: usize, dt: DType) -> Graph {
let mut body = Graph::new("jacobi_round_b");
let carry = body.input("carry", Shape::new(&[batch, 2 * n, n], dt));
let spp = body.input("spp", Shape::new(&[n, k], dt));
let sqq = body.input("sqq", Shape::new(&[n, k], dt));
let a_in = body.narrow_(carry, 1, 0, n); let v_in = body.narrow_(carry, 1, n, n);
let _ = batch;
let spp3 = body.reshape_(spp, vec![1, n as i64, k as i64]);
let sqq3 = body.reshape_(sqq, vec![1, n as i64, k as i64]);
let ident_n = ident_mat3(&mut body, n, dt);
let ident_k = ident_mat3(&mut body, k, dt);
let (a_out, v_out) = one_round_b(&mut body, a_in, v_in, spp3, sqq3, ident_n, ident_k, dt);
let out = body.concat_(vec![a_out, v_out], 1); body.set_outputs(vec![out]);
body
}
fn eigensolve_b(
g: &mut Graph,
a: NodeId,
batch: usize,
n: usize,
sweeps: u32,
dt: DType,
) -> (NodeId, NodeId) {
let mut ivdata = vec![0f64; batch * n * n];
for b in 0..batch {
for i in 0..n {
ivdata[b * n * n + i * n + i] = 1.0;
}
}
let iv = cmat(g, &ivdata, &[batch, n, n], dt);
if n < 2 {
return (a, iv);
}
let mut carry = g.concat_(vec![a, iv], 1);
let ne = if n.is_multiple_of(2) { n } else { n + 1 };
let k = ne / 2;
let rounds = round_robin_rounds(ne);
let nr = rounds.len();
let mut spp = vec![0f64; nr * n * k];
let mut sqq = vec![0f64; nr * n * k];
for (r, round) in rounds.iter().enumerate() {
for (i, &(p, q)) in round.iter().enumerate() {
if p < n {
spp[r * n * k + p * k + i] = 1.0;
}
if q < n {
sqq[r * n * k + q * k + i] = 1.0;
}
}
}
let spp_c = cmat(g, &spp, &[nr, n, k], dt);
let sqq_c = cmat(g, &sqq, &[nr, n, k], dt);
let xs = [spp_c, sqq_c];
for _ in 0..sweeps {
let body = round_body_b(batch, n, k, dt);
carry = g.scan_with_xs(carry, &xs, body, nr as u32);
}
let av = g.narrow_(carry, 1, 0, n);
let vv = g.narrow_(carry, 1, n, n);
(av, vv)
}
pub fn spectral_matfn_batched(
g: &mut Graph,
a: NodeId,
batch: usize,
n: usize,
sweeps: u32,
eps: f64,
f: SpectralFn,
) -> NodeId {
let dt = g.shape(a).dtype();
let ident_n = ident_mat3(g, n, dt); let (av, vv) = eigensolve_b(g, a, batch, n, sweeps, dt);
let epsn = cscalar(g, eps, dt);
let half = cscalar(g, 0.5, dt);
let one = cscalar(g, 1.0, dt);
let diag_av = {
let d = g.mul(av, ident_n);
g.sum(d, vec![2], true)
}; let sumv = g.add(diag_av, epsn);
let diff = g.sub(diag_av, epsn);
let d2 = g.mul(diff, diff);
let ad = g.sqrt(d2);
let sad = g.add(sumv, ad);
let mx = g.mul(half, sad); let fl = match f {
SpectralFn::Re => mx,
SpectralFn::Log => g.activation(Activation::Log, mx, Shape::new(&[batch, n, 1], dt)),
SpectralFn::Sqrt => g.sqrt(mx),
SpectralFn::InvSqrt => {
let s = g.sqrt(mx);
g.div(one, s)
}
};
let fmat = g.mul(fl, ident_n); let vt = g.transpose_(vv, vec![0, 2, 1]);
let vf = g.mm(vv, fmat);
g.mm(vf, vt)
}
impl Graph {
pub fn spd_reeig_batched(
&mut self,
x: NodeId,
batch: usize,
n: usize,
sweeps: u32,
eps: f64,
) -> NodeId {
spectral_matfn_batched(self, x, batch, n, sweeps, eps, SpectralFn::Re)
}
pub fn spd_logeig_batched(
&mut self,
x: NodeId,
batch: usize,
n: usize,
sweeps: u32,
eps: f64,
) -> NodeId {
spectral_matfn_batched(self, x, batch, n, sweeps, eps, SpectralFn::Log)
}
}
pub fn bimap(g: &mut Graph, w: NodeId, x: NodeId) -> NodeId {
let wt = g.transpose_(w, vec![1, 0]);
let wx = g.mm(w, x);
g.mm(wx, wt)
}
pub fn spd_batch_norm_transport(
g: &mut Graph,
x: NodeId,
mean: NodeId,
g_: NodeId,
n: usize,
batch: usize,
sweeps: u32,
eps: f64,
) -> NodeId {
let ms = spectral_matfn(g, mean, n, sweeps, eps, SpectralFn::InvSqrt); let gs = spectral_matfn(g, g_, n, sweeps, eps, SpectralFn::Sqrt); let mut slices = Vec::with_capacity(batch);
for bi in 0..batch {
let xi = g.narrow_(x, 0, bi, 1);
let xi = g.reshape_(xi, vec![n as i64, n as i64]);
let mx = g.mm(ms, xi); let ci = g.mm(mx, ms); let gc = g.mm(gs, ci); let yi = g.mm(gc, gs); slices.push(g.reshape_(yi, vec![1, n as i64, n as i64]));
}
g.concat_(slices, 0)
}
pub fn spectral_reeig(g: &mut Graph, a: NodeId, n: usize, sweeps: u32, eps: f64) -> NodeId {
spectral_matfn(g, a, n, sweeps, eps, SpectralFn::Re)
}
pub fn spectral_logeig(g: &mut Graph, a: NodeId, n: usize, sweeps: u32, eps: f64) -> NodeId {
spectral_matfn(g, a, n, sweeps, eps, SpectralFn::Log)
}
pub fn spectral_packed(
g: &mut Graph,
a: NodeId,
n: usize,
sweeps: u32,
eps: f64,
log: bool,
) -> NodeId {
let dt = g.shape(a).dtype();
let (av, vv) = eigensolve(g, a, n, sweeps, dt);
let f = if log { SpectralFn::Log } else { SpectralFn::Re };
let (fmat, lam_parts) = spectral_fmat(g, av, n, eps, f, dt);
let vt = g.transpose_(vv, vec![1, 0]);
let vf = g.mm(vv, fmat);
let y = g.mm(vf, vt); let y_flat = g.reshape_(y, vec![(n * n) as i64]); let lam = g.concat_(lam_parts, 0); let u_flat = g.reshape_(vv, vec![(n * n) as i64]); g.concat_(vec![y_flat, lam, u_flat], 0) }
impl Graph {
pub fn spectral_reeig(&mut self, x: NodeId, sweeps: u32, eps: f64) -> NodeId {
let n = self.shape(x).dim(0).unwrap_static();
spectral_reeig(self, x, n, sweeps, eps)
}
pub fn spectral_logeig(&mut self, x: NodeId, sweeps: u32, eps: f64) -> NodeId {
let n = self.shape(x).dim(0).unwrap_static();
spectral_logeig(self, x, n, sweeps, eps)
}
}