1use crate::solvers::SolversError;
21
22pub fn qr_factor(m: usize, n: usize, a: &mut [f64], tau: &mut [f64]) -> Result<(), SolversError> {
33 if m < n || a.len() != m * n || tau.len() != n {
34 return Err(SolversError::InvalidDimension);
35 }
36 for j in 0..n {
37 let mut norm_sq = 0.0;
39 for i in j..m {
40 let v = a[i * n + j];
41 norm_sq += v * v;
42 }
43 let s = norm_sq.sqrt();
44 if s == 0.0 {
45 tau[j] = 0.0;
46 continue;
47 }
48 let a_jj = a[j * n + j];
49 let beta = if a_jj >= 0.0 { -s } else { s };
51 tau[j] = (beta - a_jj) / beta;
52 let denom = a_jj - beta;
53 for i in (j + 1)..m {
55 a[i * n + j] /= denom;
56 }
57 a[j * n + j] = beta;
58 for col in (j + 1)..n {
60 let mut w = a[j * n + col]; for i in (j + 1)..m {
62 w += a[i * n + j] * a[i * n + col];
63 }
64 w *= tau[j];
65 a[j * n + col] -= w; for i in (j + 1)..m {
67 a[i * n + col] -= w * a[i * n + j];
68 }
69 }
70 }
71 Ok(())
72}
73
74pub fn qr_form_q(
79 m: usize,
80 n: usize,
81 a: &[f64],
82 tau: &[f64],
83 q: &mut [f64],
84) -> Result<(), SolversError> {
85 if m < n || a.len() != m * n || tau.len() != n || q.len() != m * n {
86 return Err(SolversError::InvalidDimension);
87 }
88 for x in q.iter_mut() {
90 *x = 0.0;
91 }
92 for j in 0..n {
93 q[j * n + j] = 1.0;
94 }
95 for j in (0..n).rev() {
97 if tau[j] == 0.0 {
98 continue;
99 }
100 for col in 0..n {
101 let mut w = q[j * n + col]; for i in (j + 1)..m {
103 w += a[i * n + j] * q[i * n + col];
104 }
105 w *= tau[j];
106 q[j * n + col] -= w;
107 for i in (j + 1)..m {
108 q[i * n + col] -= w * a[i * n + j];
109 }
110 }
111 }
112 Ok(())
113}
114
115pub fn qr_solve_least_squares(
123 m: usize,
124 n: usize,
125 a: &[f64],
126 tau: &[f64],
127 b: &mut [f64],
128 x: &mut [f64],
129) -> Result<(), SolversError> {
130 if m < n || a.len() != m * n || tau.len() != n || b.len() != m || x.len() != n {
131 return Err(SolversError::InvalidDimension);
132 }
133 for j in 0..n {
135 if tau[j] == 0.0 {
136 continue;
137 }
138 let mut w = b[j]; for i in (j + 1)..m {
140 w += a[i * n + j] * b[i];
141 }
142 w *= tau[j];
143 b[j] -= w;
144 for i in (j + 1)..m {
145 b[i] -= w * a[i * n + j];
146 }
147 }
148 let mut scale = 0.0_f64;
151 for i in 0..n {
152 let d = a[i * n + i].abs();
153 if d > scale {
154 scale = d;
155 }
156 }
157 let tol = 1e-12 * scale * (n as f64);
158 for i in (0..n).rev() {
160 let pivot = a[i * n + i];
161 if !pivot.is_finite() || pivot.abs() <= tol {
162 return Err(SolversError::SingularMatrix);
163 }
164 let mut s = b[i];
165 for k in (i + 1)..n {
166 s -= a[i * n + k] * x[k];
167 }
168 x[i] = s / pivot;
169 }
170 Ok(())
171}
172
173#[cfg(test)]
174mod tests {
175 use super::*;
176 use crate::solvers::linear_algebra::gemm::{gemm, matmul, Transpose};
177
178 fn approx(a: &[f64], b: &[f64], tol: f64) {
179 assert_eq!(a.len(), b.len());
180 for i in 0..a.len() {
181 assert!(
182 (a[i] - b[i]).abs() < tol,
183 "idx {i}: {} != {} (tol {tol})",
184 a[i],
185 b[i]
186 );
187 }
188 }
189
190 fn extract_r(m: usize, n: usize, a: &[f64]) -> Vec<f64> {
192 let mut r = vec![0.0; n * n];
193 for i in 0..n {
194 for j in i..n {
195 r[i * n + j] = a[i * n + j];
196 }
197 }
198 let _ = m;
199 r
200 }
201
202 #[test]
203 fn factor_reconstructs_a_square() {
204 let a0 = [12.0, -51.0, 4.0, 6.0, 167.0, -68.0, -4.0, 24.0, -41.0];
206 let (m, n) = (3, 3);
207 let mut a = a0;
208 let mut tau = [0.0; 3];
209 qr_factor(m, n, &mut a, &mut tau).unwrap();
210
211 let mut q = [0.0; 9];
212 qr_form_q(m, n, &a, &tau, &mut q).unwrap();
213 let r = extract_r(m, n, &a);
214
215 let mut recon = [0.0; 9];
217 matmul(m, n, n, &q, &r, &mut recon).unwrap();
218 approx(&recon, &a0, 1e-9);
219
220 assert!((a[2 * 3 + 2].abs() - 35.0).abs() < 1e-9);
222 }
223
224 #[test]
225 fn q_has_orthonormal_columns() {
226 let a0 = [12.0, -51.0, 4.0, 6.0, 167.0, -68.0, -4.0, 24.0, -41.0];
227 let mut a = a0;
228 let mut tau = [0.0; 3];
229 qr_factor(3, 3, &mut a, &mut tau).unwrap();
230 let mut q = [0.0; 9];
231 qr_form_q(3, 3, &a, &tau, &mut q).unwrap();
232 let mut qtq = [0.0; 9];
234 gemm(
235 Transpose::Yes,
236 Transpose::No,
237 3,
238 3,
239 3,
240 1.0,
241 &q,
242 &q,
243 0.0,
244 &mut qtq,
245 )
246 .unwrap();
247 approx(&qtq, &[1.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 1.0], 1e-9);
248 }
249
250 #[test]
251 fn square_solve_matches_known() {
252 let a0 = [2.0, 1.0, 1.0, 3.0];
254 let mut a = a0;
255 let mut tau = [0.0; 2];
256 qr_factor(2, 2, &mut a, &mut tau).unwrap();
257 let mut b = [3.0, 5.0];
258 let mut x = [0.0; 2];
259 qr_solve_least_squares(2, 2, &a, &tau, &mut b, &mut x).unwrap();
260 approx(&x, &[0.8, 1.4], 1e-9);
261 }
262
263 #[test]
264 fn least_squares_overdetermined_line_fit() {
265 let a0 = [1.0, 0.0, 1.0, 1.0, 1.0, 2.0, 1.0, 3.0];
268 let mut a = a0;
269 let mut tau = [0.0; 2];
270 qr_factor(4, 2, &mut a, &mut tau).unwrap();
271 let mut b = [1.0, 2.0, 3.0, 4.0];
272 let mut x = [0.0; 2];
273 qr_solve_least_squares(4, 2, &a, &tau, &mut b, &mut x).unwrap();
274 approx(&x, &[1.0, 1.0], 1e-9);
275 }
276
277 #[test]
278 fn least_squares_overdetermined_noisy() {
279 let a0 = [1.0, 0.0, 1.0, 1.0, 1.0, 2.0, 1.0, 3.0];
283 let mut a = a0;
284 let mut tau = [0.0; 2];
285 qr_factor(4, 2, &mut a, &mut tau).unwrap();
286 let mut b = [1.0, 3.0, 4.0, 6.0];
287 let mut x = [0.0; 2];
288 qr_solve_least_squares(4, 2, &a, &tau, &mut b, &mut x).unwrap();
289 approx(&x, &[1.1, 1.6], 1e-9);
291 }
292
293 #[test]
294 fn reconstructs_tall_matrix() {
295 let a0 = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
297 let (m, n) = (4, 2);
298 let mut a = a0;
299 let mut tau = [0.0; 2];
300 qr_factor(m, n, &mut a, &mut tau).unwrap();
301 let mut q = [0.0; 8];
302 qr_form_q(m, n, &a, &tau, &mut q).unwrap();
303 let r = extract_r(m, n, &a);
304 let mut recon = [0.0; 8];
305 matmul(m, n, n, &q, &r, &mut recon).unwrap();
306 approx(&recon, &a0, 1e-9);
307 }
308
309 #[test]
310 fn rank_deficient_fails_closed() {
311 let a0 = [1.0, 1.0, 2.0, 2.0, 3.0, 3.0]; let mut a = a0;
314 let mut tau = [0.0; 2];
315 qr_factor(3, 2, &mut a, &mut tau).unwrap();
316 let mut b = [1.0, 2.0, 3.0];
317 let mut x = [0.0; 2];
318 assert_eq!(
319 qr_solve_least_squares(3, 2, &a, &tau, &mut b, &mut x),
320 Err(SolversError::SingularMatrix)
321 );
322 }
323
324 #[test]
325 fn rejects_bad_dims() {
326 let mut a = [1.0, 2.0, 3.0, 4.0];
327 let mut tau = [0.0; 3]; assert_eq!(
329 qr_factor(2, 2, &mut a, &mut tau),
330 Err(SolversError::InvalidDimension)
331 );
332 let mut a2 = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
334 let mut tau2 = [0.0; 3];
335 assert_eq!(
336 qr_factor(2, 3, &mut a2, &mut tau2),
337 Err(SolversError::InvalidDimension)
338 );
339 }
340}