Skip to main content

rusolver/
numerics.rs

1// SPDX-License-Identifier: Apache-2.0
2use crate::{
3    MatrixView, SolverError, Tolerance
4};
5pub(crate) fn zeros(n: usize) -> Result<Vec<f64>, SolverError> {
6    let mut out = Vec::new();
7    out.try_reserve_exact(n).map_err(|_| SolverError::Allocation)?;
8    out.resize(n, 0.0);
9    Ok(out)
10}
11pub(crate) fn finite(x: &[f64]) -> Result<(), SolverError> {
12    for (index, v) in x.iter().enumerate() {
13        if !v.is_finite() {
14            return Err(SolverError::NonFinite {
15                index
16            });
17        }
18    }
19    Ok(())
20}
21pub(crate) fn checked(x: f64, where_: &'static str) -> Result<f64, SolverError> {
22    if x.is_finite() {
23        Ok(x)
24    } else {
25        Err(SolverError::Arithmetic(where_))
26    }
27}
28/// Scaled sum of squares, avoiding overflow from squaring large finite entries.
29/// Returns an error if the norm itself exceeds the representable FP64 range.
30pub fn l2_norm(x: &[f64]) -> Result<f64, SolverError> {
31    norm_iter(x.iter().copied())
32}
33pub(crate) fn norm_iter(it: impl Iterator<Item=f64>) -> Result<f64, SolverError> {
34    let (mut scale, mut sum) = (0.0f64, 1.0f64);
35    for (index, value) in it.enumerate() {
36        if !value.is_finite() {
37            return Err(SolverError::NonFinite {
38                index
39            });
40        }
41        let a = value.abs();
42        if a != 0.0 {
43            if scale < a {
44                let q = scale/a;
45                sum = 1.0 + sum*q*q;
46                scale = a;
47            }
48            else {
49                let q = a/scale;
50                sum += q*q;
51            }
52        }
53    }
54    checked(scale * sum.sqrt(), "Euclidean norm")
55}
56/// Compensated sum with explicit failure on unrepresentable intermediates.
57/// It does not claim correctly-rounded arbitrary-precision dot products.
58pub(crate) fn sum_iter(it: impl Iterator<Item=f64>) -> Result<f64, SolverError> {
59    let (mut sum, mut correction) = (0.0f64, 0.0f64);
60    for value in it {
61        checked(value, "dot-product term")?;
62        let t = checked(sum + value, "dot-product sum")?;
63        let increment = if sum.abs() >= value.abs() {
64            (sum-t)+value
65        } else {
66            (value-t)+sum
67        };
68        correction = checked(correction + increment, "dot-product compensation")?;
69        sum = t;
70    }
71    checked(sum + correction, "compensated sum")
72}
73pub(crate) fn dot(x: &[f64], y: &[f64]) -> Result<f64, SolverError> {
74    if x.len()!=y.len() {
75        return Err(SolverError::Shape("dot-product lengths"));
76    }
77    sum_iter(x.iter().zip(y).map(|(x, y)| x*y))
78}
79pub(crate) fn symmetric(a: MatrixView<'_>, tolerance: Tolerance) -> Result<(), SolverError> {
80    let n = a.square()?;
81    tolerance.validate()?;
82    for i in 0..n {
83        for j in 0..i {
84            let (x, y) = (a.at(i, j), a.at(j, i));
85            // Scale each pair first so opposite, large finite values do not overflow.
86            let scale = x.abs().max(y.abs());
87            if scale != 0.0 && (x/scale-y/scale).abs() > (tolerance.absolute/scale).max(tolerance.relative) {
88                return Err(SolverError::NotSymmetric {
89                    row: i, column: j
90                });
91            }
92        }
93    }
94    Ok(())
95}
96/// `||A*x-b||_2 / ||b||_2`; for a zero RHS returns the absolute residual norm.
97/// This is a residual diagnostic, NOT a forward-error or condition estimate.
98pub fn relative_residual(a: MatrixView<'_>, x: &[f64], b: &[f64]) -> Result<f64, SolverError> {
99    if x.len()!=a.columns() || b.len()!=a.rows() {
100        return Err(SolverError::Shape("residual dimensions"));
101    }
102    finite(x)?;
103    finite(b)?;
104    let mut residual = zeros(b.len())?;
105    for i in 0..a.rows() {
106        residual[i] = sum_iter(std::iter::once(-b[i]).chain((0..a.columns()).map(|j| a.at(i, j)*x[j])))?;
107    }
108    let (r, n)=(l2_norm(&residual)?, l2_norm(b)?);
109    checked(if n == 0.0 {
110        r
111    } else {
112        r/n
113    }, "relative residual")
114}