use deep_causality_num::{FromPrimitive, RealField};
#[derive(Debug, Clone)]
pub(crate) struct CgFailure<R: RealField> {
pub iterations: usize,
pub residual: R,
}
pub(crate) fn cg_solve<R, Apply>(
apply: Apply,
b: &[R],
tolerance: R,
max_iterations: usize,
) -> Result<Vec<R>, CgFailure<R>>
where
R: RealField + FromPrimitive + PartialEq,
Apply: Fn(&[R]) -> Vec<R>,
{
let n = b.len();
let mut x = vec![R::zero(); n];
let mut r: Vec<R> = b.to_vec();
let mut p = r.clone();
let mut rsold: R = dot(&r, &r);
let b_norm = dot(b, b).sqrt();
let abs_tol = if b_norm == R::zero() {
tolerance
} else {
tolerance * b_norm
};
if rsold.sqrt() < abs_tol {
return Ok(x);
}
for iter in 0..max_iterations {
let ap = apply(&p);
if ap.len() != n {
return Err(CgFailure {
iterations: iter,
residual: rsold.sqrt(),
});
}
let pap = dot(&p, &ap);
if pap == R::zero() {
return Err(CgFailure {
iterations: iter,
residual: rsold.sqrt(),
});
}
let alpha = rsold / pap;
for i in 0..n {
x[i] += alpha * p[i];
r[i] -= alpha * ap[i];
}
let rsnew = dot(&r, &r);
if rsnew.sqrt() < abs_tol {
return Ok(x);
}
let beta = rsnew / rsold;
for i in 0..n {
p[i] = r[i] + beta * p[i];
}
rsold = rsnew;
}
Err(CgFailure {
iterations: max_iterations,
residual: rsold.sqrt(),
})
}
pub(crate) fn subtract_mean_in_place<R>(v: &mut [R])
where
R: RealField + FromPrimitive,
{
if v.is_empty() {
return;
}
let n = <R as FromPrimitive>::from_usize(v.len())
.expect("v.len() is representable in every supported RealField");
let sum: R = v.iter().copied().fold(R::zero(), |a, b| a + b);
let mean = sum / n;
for entry in v.iter_mut() {
*entry -= mean;
}
}
#[inline]
fn dot<R>(a: &[R], b: &[R]) -> R
where
R: RealField,
{
a.iter()
.zip(b.iter())
.map(|(&x, &y)| x * y)
.fold(R::zero(), |acc, t| acc + t)
}
#[cfg(test)]
mod tests {
use super::*;
fn dense_apply(a: &[[f64; 3]; 3], v: &[f64]) -> Vec<f64> {
(0..3)
.map(|i| (0..3).map(|j| a[i][j] * v[j]).sum())
.collect()
}
#[test]
fn cg_solves_2x2_spd_system() {
let a = [[4.0_f64, 1.0], [1.0, 3.0]];
let b = vec![1.0_f64, 2.0];
let apply = |v: &[f64]| -> Vec<f64> {
(0..2)
.map(|i| (0..2).map(|j| a[i][j] * v[j]).sum())
.collect()
};
let x = cg_solve(apply, &b, 1e-12_f64, 100).expect("CG converges");
assert!((x[0] - 1.0 / 11.0).abs() < 1e-10);
assert!((x[1] - 7.0 / 11.0).abs() < 1e-10);
}
#[test]
fn cg_solves_3x3_spd_system() {
let a = [[2.0_f64, 0.0, 0.0], [0.0, 3.0, 0.0], [0.0, 0.0, 5.0]];
let b = vec![4.0_f64, 9.0, 25.0];
let apply = |v: &[f64]| dense_apply(&a, v);
let x = cg_solve(apply, &b, 1e-12_f64, 100).expect("CG converges");
assert!((x[0] - 2.0).abs() < 1e-10);
assert!((x[1] - 3.0).abs() < 1e-10);
assert!((x[2] - 5.0).abs() < 1e-10);
}
#[test]
fn cg_returns_zero_for_zero_rhs() {
let a = [[2.0_f64, 0.0, 0.0], [0.0, 3.0, 0.0], [0.0, 0.0, 5.0]];
let b = vec![0.0_f64; 3];
let apply = |v: &[f64]| dense_apply(&a, v);
let x = cg_solve(apply, &b, 1e-12_f64, 100).expect("CG converges");
for &xi in &x {
assert!(xi.abs() < 1e-15);
}
}
#[test]
fn cg_reports_nonconvergence_at_iteration_cap() {
let n = 100;
let apply = |v: &[f64]| -> Vec<f64> {
v.iter()
.enumerate()
.map(|(i, &vi)| (i as f64 + 1.0) * vi)
.collect()
};
let b: Vec<f64> = (0..n)
.map(|i| (i as f64 + 1.0) * (i as f64 + 1.0))
.collect();
let result = cg_solve(apply, &b, 1e-12_f64, 1);
let err = result.expect_err("CG must fail with iteration cap 1");
assert_eq!(err.iterations, 1);
assert!(err.residual > 0.0);
}
#[test]
fn cg_reports_nonconvergence_with_zero_iteration_budget() {
let apply = |v: &[f64]| v.to_vec();
let b = vec![1.0_f64, 2.0, 3.0];
let result = cg_solve(apply, &b, 1e-12_f64, 0);
let err = result.expect_err("CG must fail with iteration cap 0");
assert_eq!(err.iterations, 0);
assert!((err.residual - dot(&b, &b).sqrt()).abs() < 1e-14);
}
#[test]
fn cg_converges_at_f32_precision() {
let a = [[4.0_f32, 1.0], [1.0, 3.0]];
let b = vec![1.0_f32, 2.0];
let apply = |v: &[f32]| -> Vec<f32> {
(0..2)
.map(|i| (0..2).map(|j| a[i][j] * v[j]).sum())
.collect()
};
let x = cg_solve(apply, &b, 1e-5_f32, 100).expect("CG converges at f32");
assert!((x[0] - 1.0 / 11.0).abs() < 1e-4);
assert!((x[1] - 7.0 / 11.0).abs() < 1e-4);
}
#[test]
fn subtract_mean_in_place_zeros_a_constant_vector() {
let mut v = vec![5.0_f64; 7];
subtract_mean_in_place(&mut v);
for &x in &v {
assert!(x.abs() < 1e-15);
}
}
#[test]
fn subtract_mean_in_place_preserves_zero_mean_input() {
let mut v = vec![-2.0_f64, -1.0, 0.0, 1.0, 2.0];
let original = v.clone();
subtract_mean_in_place(&mut v);
for (a, b) in v.iter().zip(original.iter()) {
assert!((a - b).abs() < 1e-15);
}
}
#[test]
fn subtract_mean_in_place_handles_empty_slice() {
let mut v: Vec<f64> = vec![];
subtract_mean_in_place(&mut v);
assert!(v.is_empty());
}
#[test]
fn subtract_mean_in_place_subtracts_correct_mean_from_arbitrary_input() {
let mut v = vec![1.0_f64, 2.0, 3.0, 4.0]; subtract_mean_in_place(&mut v);
let expected = [-1.5, -0.5, 0.5, 1.5];
for (a, e) in v.iter().zip(expected.iter()) {
assert!((a - e).abs() < 1e-15);
}
}
#[test]
fn cg_failure_struct_is_clonable_and_debug_printable() {
let f = CgFailure {
iterations: 42_usize,
residual: 1.23_f64,
};
let f2 = f.clone();
assert_eq!(f2.iterations, 42);
assert!((f2.residual - 1.23).abs() < 1e-15);
let s = format!("{:?}", f);
assert!(s.contains("42"));
}
}