#[derive(Debug, Clone)]
pub struct SparseMatrix {
rows: usize,
cols: usize,
row_starts: Vec<usize>,
col_indices: Vec<usize>,
values: Vec<f64>,
}
impl SparseMatrix {
#[must_use]
pub fn from_triplets(rows: usize, cols: usize, triplets: &[(usize, usize, f64)]) -> Self {
let mut sorted: Vec<(usize, usize, f64)> = triplets
.iter()
.inspect(|(r, c, _)| {
assert!(*r < rows && *c < cols, "triplet outside the matrix");
})
.copied()
.collect();
sorted.sort_by_key(|&(r, c, _)| (r, c));
let mut row_starts = vec![0usize; rows + 1];
let mut col_indices = Vec::with_capacity(sorted.len());
let mut values: Vec<f64> = Vec::with_capacity(sorted.len());
let mut next = sorted.into_iter().peekable();
for (r, start) in row_starts.iter_mut().enumerate().take(rows) {
let row_begin = col_indices.len();
*start = row_begin;
while let Some(&(tr, c, v)) = next.peek() {
if tr != r {
break;
}
next.next();
if col_indices.len() > row_begin && col_indices.last() == Some(&c) {
let last = values.len() - 1;
values[last] += v;
} else {
col_indices.push(c);
values.push(v);
}
}
}
row_starts[rows] = col_indices.len();
Self {
rows,
cols,
row_starts,
col_indices,
values,
}
}
#[must_use]
pub fn shape(&self) -> (usize, usize) {
(self.rows, self.cols)
}
#[must_use]
pub fn stored(&self) -> usize {
self.values.len()
}
#[must_use]
pub fn multiply(&self, x: &[f64]) -> Vec<f64> {
assert_eq!(x.len(), self.cols, "vector length must match columns");
let mut y = vec![0.0; self.rows];
for (y_r, window) in y.iter_mut().zip(self.row_starts.windows(2)) {
let mut sum = 0.0;
for i in window[0]..window[1] {
sum = self.values[i].mul_add(x[self.col_indices[i]], sum);
}
*y_r = sum;
}
y
}
#[must_use]
pub fn transpose_multiply(&self, x: &[f64]) -> Vec<f64> {
assert_eq!(x.len(), self.rows, "vector length must match rows");
let mut y = vec![0.0; self.cols];
for (window, xr) in self.row_starts.windows(2).zip(x) {
for i in window[0]..window[1] {
y[self.col_indices[i]] = self.values[i].mul_add(*xr, y[self.col_indices[i]]);
}
}
y
}
}
fn dot(a: &[f64], b: &[f64]) -> f64 {
a.iter().zip(b).fold(0.0, |acc, (x, y)| x.mul_add(*y, acc))
}
#[must_use]
pub fn least_squares_cgnr(
a: &SparseMatrix,
b: &[f64],
tolerance: f64,
max_iterations: usize,
) -> Option<Vec<f64>> {
let (rows, cols) = a.shape();
assert_eq!(b.len(), rows, "right-hand side must match rows");
let mut x = vec![0.0; cols];
let mut r = b.to_vec();
let mut z = a.transpose_multiply(&r);
let target = tolerance * dot(&z, &z).sqrt().max(f64::MIN_POSITIVE);
let mut p = z.clone();
let mut zz = dot(&z, &z);
if zz.sqrt() <= target {
return Some(x);
}
for _ in 0..max_iterations {
let w = a.multiply(&p);
let ww = dot(&w, &w);
if ww <= 0.0 {
return Some(x);
}
let alpha = zz / ww;
for (xi, pi) in x.iter_mut().zip(&p) {
*xi = alpha.mul_add(*pi, *xi);
}
for (ri, wi) in r.iter_mut().zip(&w) {
*ri = alpha.mul_add(-wi, *ri);
}
z = a.transpose_multiply(&r);
let zz_next = dot(&z, &z);
if zz_next.sqrt() <= target {
return Some(x);
}
let beta = zz_next / zz;
zz = zz_next;
for (pi, zi) in p.iter_mut().zip(&z) {
*pi = beta.mul_add(*pi, *zi);
}
}
None
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
#[test]
fn triplets_assemble_sum_and_multiply() {
let a = SparseMatrix::from_triplets(
2,
3,
&[(0, 0, 1.0), (0, 2, 1.0), (1, 1, 3.0), (0, 0, 1.0)],
);
assert_eq!(a.stored(), 3);
assert_eq!(a.multiply(&[1.0, 2.0, 3.0]), vec![5.0, 6.0]);
assert_eq!(a.transpose_multiply(&[1.0, 1.0]), vec![2.0, 3.0, 1.0]);
}
#[test]
fn an_overdetermined_system_lands_on_the_normal_equation_answer() {
let a = SparseMatrix::from_triplets(
3,
2,
&[
(0, 0, 1.0),
(0, 1, 0.0),
(1, 0, 1.0),
(1, 1, 1.0),
(2, 0, 1.0),
(2, 1, 2.0),
],
);
let x = least_squares_cgnr(&a, &[1.0, 2.0, 4.0], 1e-14, 100).unwrap();
assert!((x[0] - 5.0 / 6.0).abs() < 1e-10, "{x:?}");
assert!((x[1] - 1.5).abs() < 1e-10, "{x:?}");
}
#[test]
fn a_rank_deficient_system_returns_the_minimum_norm_solution() {
let a = SparseMatrix::from_triplets(1, 2, &[(0, 0, 1.0), (0, 1, 1.0)]);
let x = least_squares_cgnr(&a, &[2.0], 1e-14, 50).unwrap();
assert!(
(x[0] - 1.0).abs() < 1e-12 && (x[1] - 1.0).abs() < 1e-12,
"{x:?}"
);
}
#[test]
fn a_zero_matrix_answers_zero_rather_than_spinning() {
let a = SparseMatrix::from_triplets(2, 2, &[]);
let x = least_squares_cgnr(&a, &[1.0, 1.0], 1e-12, 10).unwrap();
assert_eq!(x, vec![0.0, 0.0]);
}
}