conspire 0.7.2

The Rust interface to conspire.
Documentation
use super::{SparseSolver, Vector};
use crate::math::assert::Assert;
use crate::math::assert::AssertionError;

const N: usize = 100;

fn pattern() -> Vec<(usize, usize)> {
    let mut pattern: Vec<(usize, usize)> = (0..N).map(|i| (i, i)).collect();
    (0..N).for_each(|i| {
        pattern.push((i, (3 * i + 1) % N));
        pattern.push(((5 * i + 2) % N, i));
    });
    pattern
}

fn values(scale: f64) -> impl Fn(usize, usize) -> f64 {
    move |i, j| {
        if i == j {
            8.0
        } else {
            scale * (((i * 7 + j * 3) % 5) as f64 / 5.0 - 0.4)
        }
    }
}

fn residual(
    source: impl Fn(usize, usize) -> f64,
    x: &Vector,
    b: &Vector,
) -> Result<(), AssertionError> {
    let mut product = Vector::zero(N);
    pattern()
        .into_iter()
        .for_each(|(i, j)| product[i] += source(i, j) * x[j]);
    Assert::default().eq_within_tols(&product, b)
}

#[test]
fn solve_factor_then_refactor() -> Result<(), AssertionError> {
    let solver = SparseSolver::from_pattern(N, pattern(), false);
    let b: Vector = (0..N).map(|i| (i % 13) as f64 - 6.0).collect();
    residual(values(1.0), &solver.solve(values(1.0), &b)?, &b)?;
    residual(values(-2.0), &solver.solve(values(-2.0), &b)?, &b)
}

#[test]
fn clones_share_factorization() -> Result<(), AssertionError> {
    let solver = SparseSolver::from_pattern(N, pattern(), false);
    let b: Vector = (0..N).map(|i| (i % 13) as f64 - 6.0).collect();
    residual(values(1.0), &solver.clone().solve(values(1.0), &b)?, &b)?;
    assert!(solver.lu.borrow().is_some());
    residual(values(3.0), &solver.clone().solve(values(3.0), &b)?, &b)
}

#[test]
fn recovers_from_degraded_pivot() -> Result<(), AssertionError> {
    let solver = SparseSolver::from_pattern(2, vec![(0, 0), (0, 1), (1, 0), (1, 1)], false);
    let b = Vector::from([1.0, 1.0]);
    let x = solver.solve(|i, j| ((2 * i + j) % 3) as f64 + 1.0, &b)?;
    Assert::default().eq_within_tols(Vector::from([x[0] + 2.0 * x[1], 3.0 * x[0] + x[1]]), &b)?;
    let x = solver.solve(
        |i, j| {
            if (i, j) == (1, 0) {
                0.0
            } else {
                j as f64 + 1.0
            }
        },
        &b,
    )?;
    Assert::default().eq_within_tols(Vector::from([x[0] + 2.0 * x[1], 2.0 * x[1]]), &b)
}

#[test]
fn symmetric_uses_ldl() -> Result<(), AssertionError> {
    let n = 30;
    let mut pattern: Vec<(usize, usize)> = (0..n).map(|i| (i, i)).collect();
    (0..n - 3).for_each(|i| {
        pattern.push((i, i + 3));
        pattern.push((i + 3, i));
    });
    (0..4).for_each(|c| {
        pattern.push((n + c, 7 * c));
        pattern.push((7 * c, n + c));
    });
    let solver = SparseSolver::from_pattern(n + 4, pattern, true);
    let b: Vector = (0..n + 4).map(|i| (i % 7) as f64 - 3.0).collect();
    let source = |scale: f64| {
        move |i: usize, j: usize| {
            if i == j {
                12.0
            } else if i >= n || j >= n {
                2.0
            } else {
                scale * ((i.min(j) % 3) as f64 - 1.0)
            }
        }
    };
    let x = solver.solve(source(1.0), &b)?;
    let mut residual = Vector::zero(n + 4);
    solver
        .pattern()
        .iter()
        .for_each(|&(i, j)| residual[i] += source(1.0)(i, j) * x[j]);
    Assert::default().eq_within_tols(&residual, &b)?;
    assert!(solver.ldl.borrow().is_some());
    let x = solver.solve(source(-2.0), &b)?;
    let mut residual = Vector::zero(n + 4);
    solver
        .pattern()
        .iter()
        .for_each(|&(i, j)| residual[i] += source(-2.0)(i, j) * x[j]);
    Assert::default().eq_within_tols(&residual, &b)
}

#[test]
fn asymmetric_falls_back_to_lu() -> Result<(), AssertionError> {
    let n = 20;
    let mut pattern: Vec<(usize, usize)> = (0..n).map(|i| (i, i)).collect();
    (0..n - 3).for_each(|i| {
        pattern.push((i, i + 3));
        pattern.push((i + 3, i));
    });
    let solver = SparseSolver::from_pattern(n, pattern, false);
    let b: Vector = (0..n).map(|i| (i % 5) as f64 - 2.0).collect();
    let source = |i: usize, j: usize| {
        if i == j {
            12.0
        } else {
            (i % 3) as f64 - (j % 2) as f64
        }
    };
    let x = solver.solve(source, &b)?;
    assert!(solver.ldl.borrow().is_none());
    assert!(solver.lu.borrow().is_some());
    let mut residual = Vector::zero(n);
    solver
        .pattern()
        .iter()
        .for_each(|&(i, j)| residual[i] += source(i, j) * x[j]);
    Assert::default().eq_within_tols(&residual, &b)
}