Skip to main content

qualia_core_db/solvers/linear_algebra/
qr.rs

1//! Householder QR decomposition `A = Q·R` and least-squares solve.
2//!
3//! Functionality parity with nalgebra's `linalg::qr`, in the qualia idiom:
4//! **zero allocation**, operating on caller-owned **row-major** slices, fail-closed.
5//! No `DMatrix`, no heap, no dependency.
6//!
7//! QR is the stable workhorse the engine was missing entirely (the prior LA had
8//! only fixed-size `Matrix4x4` LU and the silo's heap routines). It gives:
9//! - orthogonal factorisation `A(m×n) = Q(m×m)·R(m×n)`, `m ≥ n`;
10//! - least-squares `min‖A·x − b‖` for overdetermined full-rank systems
11//!   (the normal-equations route proved in [`super::gemm`], but numerically
12//!   stable — no `AᵀA` conditioning blow-up);
13//! - a square linear solve as the `m == n` case (alternative to LU).
14//!
15//! Storage convention (LAPACK `geqrf` style): [`qr_factor`] overwrites `a` in
16//! place — the upper triangle (incl. diagonal) becomes `R`; below the diagonal
17//! holds the essential Householder vectors `v` (with implicit `v[j] = 1`); the
18//! per-column scalings go in `tau`.
19
20use crate::solvers::SolversError;
21
22/// Compute the Householder QR factorisation of the `m×n` row-major matrix `a`
23/// (`m ≥ n`) **in place**. On return:
24/// - the upper triangle of `a` (including the diagonal) holds `R`;
25/// - below the diagonal holds the essential Householder reflector vectors;
26/// - `tau[j]` holds the reflector scaling for column `j`.
27///
28/// `a` must be length `m*n`; `tau` must be length `n`. Returns
29/// [`SolversError::InvalidDimension`] on a length/shape mismatch (including `m < n`).
30/// A zero sub-column yields `tau[j] = 0` (identity reflector) — rank deficiency is
31/// not an error here; it surfaces at solve time as a zero `R` pivot.
32pub 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        // Norm of the sub-column a[j..m, j].
38        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        // beta = new R[j][j]; choose sign to avoid cancellation.
50        let beta = if a_jj >= 0.0 { -s } else { s };
51        tau[j] = (beta - a_jj) / beta;
52        let denom = a_jj - beta;
53        // v[j] = 1 (implicit); v[i>j] = a[i][j]/denom, stored below the diagonal.
54        for i in (j + 1)..m {
55            a[i * n + j] /= denom;
56        }
57        a[j * n + j] = beta;
58        // Apply (I − tau·v·vᵀ) to the trailing columns j+1..n.
59        for col in (j + 1)..n {
60            let mut w = a[j * n + col]; // v[j]·A[j][col], v[j]=1
61            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; // v[j]=1
66            for i in (j + 1)..m {
67                a[i * n + col] -= w * a[i * n + j];
68            }
69        }
70    }
71    Ok(())
72}
73
74/// Materialise the **thin** orthogonal factor `Q` (`m×n`, row-major) from a
75/// factored `a`/`tau` (from [`qr_factor`]). `q` must be length `m*n`. The thin
76/// `Q` satisfies `Q·R_n = A` (with `R_n` the `n×n` upper triangle of `a`) and
77/// has orthonormal columns (`Qᵀ·Q = I_n`).
78pub 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    // Start Q = first n columns of I_m.
89    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    // Q = H_0 · H_1 · … · H_{n-1}; apply in reverse so the product accumulates.
96    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]; // v[j]=1
102            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
115/// Least-squares solve `min‖A·x − b‖` for an `m×n` (`m ≥ n`) full-rank system,
116/// given the factored `a`/`tau` (from [`qr_factor`]).
117///
118/// `b` (length `m`) is overwritten with `Qᵀ·b`; the first `n` entries are then
119/// back-substituted through `R` into `x` (length `n`). For `m == n` this is an
120/// exact square solve. Returns [`SolversError::SingularMatrix`] if an `R` pivot
121/// is ~0 (rank deficient) — fail closed, never a bogus solution.
122pub 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    // Apply Qᵀ to b: b ← H_{n-1} … H_0 · b.
134    for j in 0..n {
135        if tau[j] == 0.0 {
136            continue;
137        }
138        let mut w = b[j]; // v[j]=1
139        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    // Rank check: a pivot small relative to the largest R diagonal means the
149    // columns are (near-)dependent — fail closed rather than divide by ~0.
150    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    // Back-substitute R·x = (Qᵀb)[0..n].
159    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    // Extract the n×n upper-triangular R from a factored buffer.
191    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        // A = [[12,-51,4],[6,167,-68],[-4,24,-41]] — the classic Householder example.
205        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        // Q·R == A
216        let mut recon = [0.0; 9];
217        matmul(m, n, n, &q, &r, &mut recon).unwrap();
218        approx(&recon, &a0, 1e-9);
219
220        // R[2][2] of this matrix is known to be 35 (up to sign).
221        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        // QᵀQ == I
233        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        // [[2,1],[1,3]] x = [3,5] → x = [0.8, 1.4]
253        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        // Fit y = c0 + c1·t to (0,1),(1,2),(2,3),(3,4): exact line y=1+t → c=[1,1].
266        // Design matrix A (4×2): columns [1, t].
267        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        // Slightly perturbed: residual minimised, slope/intercept near the trend.
280        // (0,1),(1,3),(2,4),(3,6). Normal equations: Σt=6, Σt²=14, Σy=14, Σty=29
281        // → 4c0+6c1=14, 6c0+14c1=29 → c1=1.6, c0=3.5−1.5·1.6=1.1.
282        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        // Closed-form normal-equations solution: intercept 1.1, slope 1.6.
290        approx(&x, &[1.1, 1.6], 1e-9);
291    }
292
293    #[test]
294    fn reconstructs_tall_matrix() {
295        // Tall A (4×2): Q_thin(4×2)·R(2×2) == A.
296        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        // Columns identical → rank 1 → zero R pivot → SingularMatrix on solve.
312        let a0 = [1.0, 1.0, 2.0, 2.0, 3.0, 3.0]; // 3×2, col1==col2
313        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]; // wrong length
328        assert_eq!(
329            qr_factor(2, 2, &mut a, &mut tau),
330            Err(SolversError::InvalidDimension)
331        );
332        // m < n rejected
333        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}