uncertain_numerics/
conditioning.rs1use nalgebra::{DMatrix, DVector, Dyn, linalg::Cholesky};
4
5use crate::ConditioningError;
6
7#[derive(Debug, Clone)]
14pub struct GaussianConditioner {
15 dimension: usize,
16 jitter: f64,
17 cholesky: Cholesky<f64, Dyn>,
18}
19
20impl GaussianConditioner {
21 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 #[must_use]
69 pub const fn dimension(&self) -> usize {
70 self.dimension
71 }
72
73 #[must_use]
75 pub const fn jitter(&self) -> f64 {
76 self.jitter
77 }
78
79 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)] mod 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}