qualia_core_db/solvers/linear_algebra/
cholesky.rs1use crate::solvers::SolversError;
11
12pub fn cholesky_factor(n: usize, a: &[f64], l: &mut [f64]) -> Result<(), SolversError> {
22 if a.len() != n * n || l.len() != n * n {
23 return Err(SolversError::InvalidDimension);
24 }
25 for x in l.iter_mut() {
26 *x = 0.0;
27 }
28 for j in 0..n {
29 let mut diag = a[j * n + j];
31 for k in 0..j {
32 diag -= l[j * n + k] * l[j * n + k];
33 }
34 if !(diag > 0.0) {
35 return Err(SolversError::SingularMatrix);
36 }
37 let ljj = diag.sqrt();
38 l[j * n + j] = ljj;
39
40 for i in (j + 1)..n {
42 let mut s = a[i * n + j];
43 for k in 0..j {
44 s -= l[i * n + k] * l[j * n + k];
45 }
46 l[i * n + j] = s / ljj;
47 }
48 }
49 Ok(())
50}
51
52pub fn cholesky_solve(n: usize, l: &[f64], b: &[f64], x: &mut [f64]) -> Result<(), SolversError> {
57 if l.len() != n * n || b.len() != n || x.len() != n {
58 return Err(SolversError::InvalidDimension);
59 }
60 for i in 0..n {
62 let mut s = b[i];
63 for k in 0..i {
64 s -= l[i * n + k] * x[k];
65 }
66 x[i] = s / l[i * n + i];
67 }
68 for i in (0..n).rev() {
70 let mut s = x[i];
71 for k in (i + 1)..n {
72 s -= l[k * n + i] * x[k];
73 }
74 x[i] = s / l[i * n + i];
75 }
76 Ok(())
77}
78
79pub fn cholesky_determinant(n: usize, l: &[f64]) -> f64 {
81 let mut prod = 1.0;
82 for i in 0..n {
83 let d = l[i * n + i];
84 prod *= d * d;
85 }
86 prod
87}
88
89#[cfg(test)]
90mod tests {
91 use super::*;
92
93 const EPS: f64 = 1e-9;
94
95 const A3: [f64; 9] = [4.0, 12.0, -16.0, 12.0, 37.0, -43.0, -16.0, -43.0, 98.0];
99
100 #[test]
101 fn factor_matches_known_lower_triangle() {
102 let mut l = [0.0; 9];
103 cholesky_factor(3, &A3, &mut l).unwrap();
104 let expect = [2.0, 0.0, 0.0, 6.0, 1.0, 0.0, -8.0, 5.0, 3.0];
105 for i in 0..9 {
106 assert!(
107 (l[i] - expect[i]).abs() < EPS,
108 "l[{i}]={} != {}",
109 l[i],
110 expect[i]
111 );
112 }
113 }
114
115 #[test]
116 fn reconstructs_a() {
117 let mut l = [0.0; 9];
118 cholesky_factor(3, &A3, &mut l).unwrap();
119 for i in 0..3 {
121 for j in 0..3 {
122 let mut s = 0.0;
123 for k in 0..3 {
124 s += l[i * 3 + k] * l[j * 3 + k];
125 }
126 assert!((s - A3[i * 3 + j]).abs() < 1e-6);
127 }
128 }
129 }
130
131 #[test]
132 fn solves_linear_system() {
133 let mut l = [0.0; 9];
134 cholesky_factor(3, &A3, &mut l).unwrap();
135 let b = [1.0, 2.0, 3.0];
136 let mut x = [0.0; 3];
137 cholesky_solve(3, &l, &b, &mut x).unwrap();
138 for i in 0..3 {
140 let mut s = 0.0;
141 for j in 0..3 {
142 s += A3[i * 3 + j] * x[j];
143 }
144 assert!((s - b[i]).abs() < 1e-6, "row {i}: {} != {}", s, b[i]);
145 }
146 }
147
148 #[test]
149 fn determinant_via_factor() {
150 let mut l = [0.0; 9];
151 cholesky_factor(3, &A3, &mut l).unwrap();
152 assert!((cholesky_determinant(3, &l) - 36.0).abs() < 1e-6);
154 }
155
156 #[test]
157 fn rejects_non_positive_definite() {
158 let a = [1.0, 2.0, 2.0, 1.0];
160 let mut l = [0.0; 4];
161 assert_eq!(
162 cholesky_factor(2, &a, &mut l),
163 Err(SolversError::SingularMatrix)
164 );
165 }
166
167 #[test]
168 fn rejects_bad_dims() {
169 let mut l = [0.0; 9];
170 assert_eq!(
171 cholesky_factor(2, &A3, &mut l),
172 Err(SolversError::InvalidDimension)
173 );
174 }
175}