use ndarray::{Array1, Array2, ArrayView1, ArrayView2};
#[derive(Clone, Debug)]
pub struct SparseCode {
pub indices: Vec<u32>,
pub codes: Vec<f32>,
}
pub fn solve_row_codes(
row: ArrayView1<'_, f32>,
decoder: ArrayView2<'_, f32>,
active: &[(u32, f32)],
s: usize,
ridge: f32,
) -> SparseCode {
assert!(s > 0, "sparse-code support width must be positive");
assert!(
!ridge.is_nan() && ridge >= 0.0,
"active-set ridge must be nonnegative, got {ridge}"
);
assert_eq!(
row.len(),
decoder.ncols(),
"row width must equal decoder width"
);
let m = active.len().min(s);
if m == 0 {
return SparseCode {
indices: vec![0u32; s],
codes: vec![0.0f32; s],
};
}
assert!(
active
.iter()
.take(m)
.all(|&(atom, _)| (atom as usize) < decoder.nrows()),
"active atom index is out of range"
);
if ridge == f32::INFINITY {
return SparseCode {
indices: active
.iter()
.take(m)
.map(|&(atom, _)| atom)
.chain(std::iter::repeat_n(active[0].0, s - m))
.collect(),
codes: vec![0.0; s],
};
}
let p = row.len();
let mut gram = Array2::<f64>::zeros((m, m));
let mut rhs = Array1::<f64>::zeros(m);
for i in 0..m {
let ai = active[i].0 as usize;
let di = decoder.row(ai);
let mut proj = 0.0f64;
for c in 0..p {
proj += di[c] as f64 * row[c] as f64;
}
rhs[i] = proj;
for j in i..m {
let aj = active[j].0 as usize;
let dj = decoder.row(aj);
let mut g = 0.0f64;
for c in 0..p {
g += di[c] as f64 * dj[c] as f64;
}
gram[[i, j]] = g;
gram[[j, i]] = g;
}
gram[[i, i]] += ridge as f64;
}
let solution = solve_spd(&gram, &rhs);
let mut indices = Vec::with_capacity(s);
let mut codes = Vec::with_capacity(s);
for i in 0..m {
indices.push(active[i].0);
codes.push(solution[i] as f32);
}
while indices.len() < s {
indices.push(active[0].0);
codes.push(0.0f32);
}
SparseCode { indices, codes }
}
fn solve_spd(gram: &Array2<f64>, rhs: &Array1<f64>) -> Array1<f64> {
use faer::Side;
use gam_linalg::faer_ndarray::{FaerCholesky, FaerEigh};
let m = rhs.len();
if let Ok(factor) = gram.cholesky(Side::Lower) {
return factor.solvevec(rhs);
}
let (eigenvalues, eigenvectors) = gram
.eigh(Side::Lower)
.expect("an active Gram matrix must admit a symmetric eigendecomposition");
let spectral_radius = eigenvalues
.iter()
.map(|value| value.abs())
.fold(0.0_f64, f64::max);
if spectral_radius == 0.0 {
return Array1::<f64>::zeros(m);
}
let cutoff = f64::EPSILON * (m as f64) * spectral_radius;
let mut out = Array1::<f64>::zeros(m);
for eigen_index in 0..m {
let eigenvalue = eigenvalues[eigen_index];
assert!(
eigenvalue >= -cutoff,
"active Gram matrix is not positive semidefinite: eigenvalue {eigenvalue:e}, cutoff {cutoff:e}"
);
if eigenvalue <= cutoff {
continue;
}
let eigenvector = eigenvectors.column(eigen_index);
let projection = eigenvector.dot(rhs) / eigenvalue;
for coordinate in 0..m {
out[coordinate] += projection * eigenvector[coordinate];
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
#[test]
fn zero_prior_variance_has_zero_codes_with_the_requested_sparse_support() {
let row = array![3.0_f32, -2.0];
let decoder = array![[1.0_f32, 0.0], [0.0, 1.0]];
let codes = solve_row_codes(row.view(), decoder.view(), &[(1, 2.0)], 2, f32::INFINITY);
assert_eq!(codes.indices, vec![1, 1]);
assert_eq!(codes.codes, vec![0.0, 0.0]);
}
#[test]
fn duplicate_selected_atoms_use_joint_minimum_norm_least_squares() {
let row = array![1.0_f32, 0.0];
let decoder = array![[1.0_f32, 0.0], [1.0_f32, 0.0]];
let active = vec![(0_u32, 0.0_f32), (1_u32, 0.0_f32)];
let code = solve_row_codes(row.view(), decoder.view(), &active, 2, 0.0);
assert!((code.codes[0] - 0.5).abs() < 1.0e-6);
assert!((code.codes[1] - 0.5).abs() < 1.0e-6);
let reconstructed = code.codes[0] * decoder[[0, 0]] + code.codes[1] * decoder[[1, 0]];
assert!((reconstructed - 1.0).abs() < 1.0e-6);
}
}