use rayon::prelude::*;
use super::kernel::Kernel;
use super::linalg::{backward_solve, cholesky_in_place, dot, forward_solve, pivoted_cholesky};
use crate::error::{HessboostError, Result};
use crate::rng::Rng;
const GRAM_BLOCK: usize = 1024;
const LANDMARK_TOL: f64 = 1e-10;
const NYSTROM_SALT: u64 = 0x4E59_5354;
pub(super) enum RidgeSolver {
Exact { factor: Vec<f64>, n: usize },
Nystrom {
f: Vec<f64>,
n: usize,
r: usize,
m_factor: Vec<f64>,
f_sums: Vec<f64>,
},
}
pub(super) struct Solved {
pub(super) gram: Vec<f64>,
pub(super) sums: Vec<f64>,
}
fn not_positive_definite() -> HessboostError {
HessboostError::invalid_param(
"solver",
"the kernel ridge system is not numerically positive definite",
)
}
impl RidgeSolver {
pub(super) fn exact(kernel: &impl Kernel, c: f64) -> Result<Self> {
let n = kernel.n();
let mut a = kernel.dense();
for i in 0..n {
a[i * n + i] += c;
}
cholesky_in_place(&mut a, n).map_err(|_| not_positive_definite())?;
Ok(RidgeSolver::Exact { factor: a, n })
}
pub(super) fn nystrom(
kernel: &impl Kernel,
c: f64,
landmarks: usize,
seed: u64,
) -> Result<Self> {
let n = kernel.n();
let mut rows: Vec<usize> = (0..n).collect();
if landmarks < n {
Rng::new(seed ^ NYSTROM_SALT).shuffle(&mut rows);
rows.truncate(landmarks);
rows.sort_unstable();
}
let s = rows.len();
let mut w = vec![0.0; s * s];
w.par_chunks_mut(s.max(1)).zip(&rows).for_each_init(
|| (vec![0.0; n], kernel.scratch()),
|(buffer, scratch), (out, &row)| {
buffer.fill(0.0);
kernel.add_row(row, scratch, buffer);
for (o, &j) in out.iter_mut().zip(&rows) {
*o = buffer[j];
}
},
);
let pivoted = pivoted_cholesky(&w, s, LANDMARK_TOL);
let r = pivoted.rank();
let pivot_rows: Vec<usize> = pivoted.pivots.iter().map(|&p| rows[p]).collect();
let mut f = vec![0.0; n * r];
let mut buffer = vec![0.0; n];
let mut scratch = kernel.scratch();
for (t, &row) in pivot_rows.iter().enumerate() {
buffer.fill(0.0);
kernel.add_row(row, &mut scratch, &mut buffer);
for (i, &v) in buffer.iter().enumerate() {
f[i * r + t] = v;
}
}
let lw = &pivoted.factor;
f.par_chunks_mut(r.max(1)).for_each(|row| {
forward_solve(lw, r, row, 1);
});
let partials: Vec<(Vec<f64>, Vec<f64>)> = f
.par_chunks(GRAM_BLOCK * r.max(1))
.map(|block| {
let mut g = vec![0.0; r * r];
let mut sums = vec![0.0; r];
for row in block.chunks_exact(r.max(1)) {
for (a, &fa) in row.iter().enumerate() {
sums[a] += fa;
if fa != 0.0 {
super::linalg::axpy(fa, &row[..=a], &mut g[a * r..=a * r + a]);
}
}
}
(g, sums)
})
.collect();
let mut m = vec![0.0; r * r];
let mut f_sums = vec![0.0; r];
for (g, sums) in partials {
for (a, b) in m.iter_mut().zip(&g) {
*a += b;
}
for (a, b) in f_sums.iter_mut().zip(&sums) {
*a += b;
}
}
for a in 0..r {
m[a * r + a] += c;
}
cholesky_in_place(&mut m, r).map_err(|_| not_positive_definite())?;
Ok(RidgeSolver::Nystrom {
f,
n,
r,
m_factor: m,
f_sums,
})
}
pub(super) fn solve(&self, kvecs: &[f64], m: usize, c: f64) -> Solved {
match self {
RidgeSolver::Exact { factor, n } => {
let n = *n;
let mut u = vec![0.0; n * m];
for (a, k) in kvecs.chunks_exact(n).enumerate() {
for (i, &v) in k.iter().enumerate() {
u[i * m + a] = v;
}
}
forward_solve(factor, n, &mut u, m);
backward_solve(factor, n, &mut u, m);
let mut gram = vec![0.0; m * m];
let mut sums = vec![0.0; m];
for row in u.chunks_exact(m) {
for a in 0..m {
sums[a] += row[a];
for b in 0..=a {
gram[a * m + b] += row[a] * row[b];
}
}
}
symmetrize(&mut gram, m);
Solved { gram, sums }
}
RidgeSolver::Nystrom {
f,
n,
r,
m_factor,
f_sums,
} => {
let (n, r) = (*n, *r);
let mut v = vec![0.0; m * r];
for (a, k) in kvecs.chunks_exact(n).enumerate() {
let va = &mut v[a * r..(a + 1) * r];
for (i, &ki) in k.iter().enumerate() {
if ki != 0.0 {
super::linalg::axpy(ki, &f[i * r..(i + 1) * r], va);
}
}
}
let mut z = vec![0.0; r * m];
for a in 0..m {
for t in 0..r {
z[t * m + a] = v[a * r + t];
}
}
forward_solve(m_factor, r, &mut z, m);
backward_solve(m_factor, r, &mut z, m);
let mut zq = vec![0.0; m * r];
for t in 0..r {
for a in 0..m {
zq[a * r + t] = z[t * m + a];
}
}
let mut gram = vec![0.0; m * m];
let mut sums = vec![0.0; m];
for a in 0..m {
let ka = &kvecs[a * n..(a + 1) * n];
let za = &zq[a * r..(a + 1) * r];
let k_sum: f64 = ka.iter().sum();
sums[a] = (k_sum - dot(za, f_sums)) / c;
for b in 0..=a {
let kb = &kvecs[b * n..(b + 1) * n];
let zb = &zq[b * r..(b + 1) * r];
let vz = dot(&v[a * r..(a + 1) * r], zb);
gram[a * m + b] = (dot(ka, kb) - vz - c * dot(za, zb)) / (c * c);
}
}
symmetrize(&mut gram, m);
Solved { gram, sums }
}
}
}
pub(super) fn solve_vectors(&self, rhs: &mut [f64], m: usize, c: f64) {
match self {
RidgeSolver::Exact { factor, n } => {
let n = *n;
let mut u = vec![0.0; n * m];
for (a, k) in rhs.chunks_exact(n).enumerate() {
for (i, &v) in k.iter().enumerate() {
u[i * m + a] = v;
}
}
forward_solve(factor, n, &mut u, m);
backward_solve(factor, n, &mut u, m);
for (a, k) in rhs.chunks_exact_mut(n).enumerate() {
for (i, v) in k.iter_mut().enumerate() {
*v = u[i * m + a];
}
}
}
RidgeSolver::Nystrom {
f, n, r, m_factor, ..
} => {
let (n, r) = (*n, *r);
for k in rhs.chunks_exact_mut(n) {
let mut z = vec![0.0; r];
for (i, &ki) in k.iter().enumerate() {
if ki != 0.0 {
super::linalg::axpy(ki, &f[i * r..(i + 1) * r], &mut z);
}
}
forward_solve(m_factor, r, &mut z, 1);
backward_solve(m_factor, r, &mut z, 1);
for (i, ki) in k.iter_mut().enumerate() {
*ki = (*ki - dot(&f[i * r..(i + 1) * r], &z)) / c;
}
}
}
}
}
}
fn symmetrize(g: &mut [f64], m: usize) {
for a in 0..m {
for b in 0..a {
g[b * m + a] = g[a * m + b];
}
}
}