Skip to main content

uncertain_numerics/
conditioning.rs

1//! Stable Gaussian conditioning through Cholesky factorization.
2
3use nalgebra::{DMatrix, DVector, Dyn, linalg::Cholesky};
4
5use crate::ConditioningError;
6
7/// Reusable factorization of a symmetric positive-definite linear system.
8///
9/// The constructor accepts a dense matrix in row-major order, optionally adds a
10/// fixed non-negative jitter value to its diagonal, and computes a Cholesky
11/// factorization. Solves reuse that factorization and never form an explicit
12/// matrix inverse.
13#[derive(Debug, Clone)]
14pub struct GaussianConditioner {
15    dimension: usize,
16    jitter: f64,
17    cholesky: Cholesky<f64, Dyn>,
18}
19
20impl GaussianConditioner {
21    /// Construct a Gaussian conditioning system from a flattened square matrix.
22    ///
23    /// `matrix` is interpreted in row-major order. The supplied `jitter` is
24    /// added exactly once to each diagonal entry before factorization.
25    ///
26    /// # Errors
27    ///
28    /// Returns [`ConditioningError`] when dimensions are inconsistent, values
29    /// are non-finite, jitter is negative, or Cholesky factorization fails.
30    pub fn new(matrix: &[f64], dimension: usize, jitter: f64) -> Result<Self, ConditioningError> {
31        if dimension == 0 {
32            return Err(ConditioningError::ZeroDimension);
33        }
34
35        let expected_len = dimension
36            .checked_mul(dimension)
37            .ok_or(ConditioningError::MatrixDimensionMismatch)?;
38        if matrix.len() != expected_len {
39            return Err(ConditioningError::MatrixDimensionMismatch);
40        }
41        if matrix.iter().any(|value| !value.is_finite()) {
42            return Err(ConditioningError::NonFiniteMatrixEntry);
43        }
44        if !jitter.is_finite() {
45            return Err(ConditioningError::NonFiniteJitter);
46        }
47        if jitter < 0.0 {
48            return Err(ConditioningError::NegativeJitter);
49        }
50
51        let mut regularized = DMatrix::from_row_slice(dimension, dimension, matrix);
52        if jitter > 0.0 {
53            for index in 0..dimension {
54                regularized[(index, index)] += jitter;
55            }
56        }
57
58        let cholesky = Cholesky::new(regularized).ok_or(ConditioningError::NotPositiveDefinite)?;
59
60        Ok(Self {
61            dimension,
62            jitter,
63            cholesky,
64        })
65    }
66
67    /// Return the system dimension.
68    #[must_use]
69    pub const fn dimension(&self) -> usize {
70        self.dimension
71    }
72
73    /// Return the fixed diagonal jitter used during factorization.
74    #[must_use]
75    pub const fn jitter(&self) -> f64 {
76        self.jitter
77    }
78
79    /// Solve `K x = rhs` using the stored Cholesky factorization.
80    ///
81    /// # Errors
82    ///
83    /// Returns [`ConditioningError`] when the right-hand side has the wrong
84    /// dimension or contains non-finite entries.
85    pub fn solve(&self, rhs: &[f64]) -> Result<Vec<f64>, ConditioningError> {
86        if rhs.len() != self.dimension {
87            return Err(ConditioningError::RightHandSideDimensionMismatch);
88        }
89        if rhs.iter().any(|value| !value.is_finite()) {
90            return Err(ConditioningError::NonFiniteRightHandSideEntry);
91        }
92
93        let right_hand_side = DVector::from_column_slice(rhs);
94        let solution = self.cholesky.solve(&right_hand_side);
95
96        Ok(solution.iter().copied().collect())
97    }
98}
99
100#[cfg(test)]
101#[allow(clippy::float_cmp)] // exact round-trips of constructor inputs are intended
102mod tests {
103    use super::GaussianConditioner;
104    use crate::ConditioningError;
105
106    const TOLERANCE: f64 = 1.0e-12;
107
108    fn assert_close(actual: f64, expected: f64) {
109        let scale = expected.abs().max(1.0);
110        assert!(
111            (actual - expected).abs() <= TOLERANCE * scale,
112            "expected {expected:.16e}, got {actual:.16e}"
113        );
114    }
115
116    #[test]
117    fn solves_known_spd_system_without_inverse() {
118        let conditioner =
119            GaussianConditioner::new(&[4.0, 1.0, 1.0, 3.0], 2, 0.0).expect("matrix is SPD");
120        let solution = conditioner.solve(&[1.0, 2.0]).expect("rhs is valid");
121
122        assert_close(solution[0], 1.0 / 11.0);
123        assert_close(solution[1], 7.0 / 11.0);
124    }
125
126    #[test]
127    fn factorization_is_reused_for_multiple_right_hand_sides() {
128        let conditioner =
129            GaussianConditioner::new(&[2.0, 0.5, 0.5, 1.5], 2, 0.0).expect("matrix is SPD");
130
131        let first = conditioner.solve(&[1.0, 0.0]).expect("rhs is valid");
132        let second = conditioner.solve(&[0.0, 1.0]).expect("rhs is valid");
133
134        assert_close(2.0 * first[0] + 0.5 * first[1], 1.0);
135        assert_close(0.5 * first[0] + 1.5 * first[1], 0.0);
136        assert_close(2.0 * second[0] + 0.5 * second[1], 0.0);
137        assert_close(0.5 * second[0] + 1.5 * second[1], 1.0);
138    }
139
140    #[test]
141    fn rejects_invalid_matrix_contracts() {
142        assert!(matches!(
143            GaussianConditioner::new(&[], 0, 0.0),
144            Err(ConditioningError::ZeroDimension)
145        ));
146        assert!(matches!(
147            GaussianConditioner::new(&[1.0, 0.0, 0.0], 2, 0.0),
148            Err(ConditioningError::MatrixDimensionMismatch)
149        ));
150        assert!(matches!(
151            GaussianConditioner::new(&[1.0, f64::NAN, 0.0, 1.0], 2, 0.0),
152            Err(ConditioningError::NonFiniteMatrixEntry)
153        ));
154    }
155
156    #[test]
157    fn rejects_invalid_jitter() {
158        assert!(matches!(
159            GaussianConditioner::new(&[1.0], 1, f64::NAN),
160            Err(ConditioningError::NonFiniteJitter)
161        ));
162        assert!(matches!(
163            GaussianConditioner::new(&[1.0], 1, -1.0e-6),
164            Err(ConditioningError::NegativeJitter)
165        ));
166    }
167
168    #[test]
169    fn singular_matrix_fails_without_jitter() {
170        assert!(matches!(
171            GaussianConditioner::new(&[1.0, 1.0, 1.0, 1.0], 2, 0.0),
172            Err(ConditioningError::NotPositiveDefinite)
173        ));
174    }
175
176    #[test]
177    fn explicit_jitter_can_regularize_singular_matrix() {
178        let conditioner = GaussianConditioner::new(&[1.0, 1.0, 1.0, 1.0], 2, 1.0e-6)
179            .expect("positive jitter makes matrix positive definite");
180        let solution = conditioner.solve(&[2.0, 2.0]).expect("rhs is valid");
181
182        assert_eq!(conditioner.jitter(), 1.0e-6);
183        assert_close((1.0 + 1.0e-6) * solution[0] + solution[1], 2.0);
184        assert_close(solution[0] + (1.0 + 1.0e-6) * solution[1], 2.0);
185    }
186
187    #[test]
188    fn rejects_invalid_right_hand_side() {
189        let conditioner =
190            GaussianConditioner::new(&[2.0, 0.0, 0.0, 3.0], 2, 0.0).expect("matrix is SPD");
191
192        assert_eq!(
193            conditioner.solve(&[1.0]),
194            Err(ConditioningError::RightHandSideDimensionMismatch)
195        );
196        assert_eq!(
197            conditioner.solve(&[1.0, f64::INFINITY]),
198            Err(ConditioningError::NonFiniteRightHandSideEntry)
199        );
200    }
201}