#[derive(Debug, Clone)]
pub struct QrAccumulator<const N: usize> {
r: [[f64; N]; N],
qtb: [f64; N],
}
impl<const N: usize> Default for QrAccumulator<N> {
fn default() -> Self {
Self::new()
}
}
impl<const N: usize> QrAccumulator<N> {
pub fn new() -> Self {
Self {
r: [[0.0; N]; N],
qtb: [0.0; N],
}
}
pub fn push_row(&mut self, a: &[f64; N], b: f64) {
let mut row = *a;
let mut rhs = b;
for k in 0..N {
if row[k] == 0.0 {
continue;
}
let (rkk, ak) = (self.r[k][k], row[k]);
let norm = rkk.hypot(ak);
if norm == 0.0 {
continue;
}
let c = rkk / norm;
let s = ak / norm;
self.r[k][k] = norm;
row[k] = 0.0;
#[allow(clippy::needless_range_loop)]
for j in (k + 1)..N {
let (rkj, aj) = (self.r[k][j], row[j]);
self.r[k][j] = c * rkj + s * aj;
row[j] = c * aj - s * rkj;
}
let (zk, rb) = (self.qtb[k], rhs);
self.qtb[k] = c * zk + s * rb;
rhs = c * rb - s * zk;
}
}
pub fn push_damping(&mut self, mu: f64) {
if mu <= 0.0 {
return;
}
let root_mu = mu.sqrt();
for i in 0..N {
let mut row = [0.0_f64; N];
row[i] = root_mu;
self.push_row(&row, 0.0);
}
}
pub fn solve(&self) -> Option<[f64; N]> {
let mut h = [0.0_f64; N];
for i in (0..N).rev() {
let mut acc = self.qtb[i];
#[allow(clippy::needless_range_loop)]
for j in (i + 1)..N {
acc -= self.r[i][j] * h[j];
}
let rii = self.r[i][i];
if rii == 0.0 || !rii.is_finite() {
return None;
}
let v = acc / rii;
if !v.is_finite() {
return None;
}
h[i] = v;
}
Some(h)
}
pub fn r(&self) -> &[[f64; N]; N] {
&self.r
}
}
pub fn solve_least_squares<const N: usize>(rows: &[[f64; N]], rhs: &[f64]) -> Option<[f64; N]> {
if rows.len() != rhs.len() {
return None;
}
let mut qr = QrAccumulator::<N>::new();
for (a, b) in rows.iter().zip(rhs) {
qr.push_row(a, *b);
}
qr.solve()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn matches_normal_equations_when_well_conditioned() {
let rows = [[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]];
let rhs = [1.0, 2.0, 3.05];
let h = solve_least_squares::<2>(&rows, &rhs).expect("solvable");
let want = [(2.0 * 4.05 - 5.05) / 3.0, (2.0 * 5.05 - 4.05) / 3.0];
for i in 0..2 {
assert!(
(h[i] - want[i]).abs() < 1e-12,
"component {i}: QR {} vs normal equations {}",
h[i],
want[i]
);
}
}
#[test]
fn survives_row_weights_that_would_square_out_of_range() {
const BIG: f64 = 1.0e8;
let rows = [[BIG, BIG], [1.0, 0.0], [0.0, 1.0]];
let rhs = [BIG * (2.0 - 3.0), 2.0, -3.0];
let h = solve_least_squares::<2>(&rows, &rhs).expect("solvable");
assert!(
(h[0] - 2.0).abs() < 1e-6 && (h[1] + 3.0).abs() < 1e-6,
"stiff system: got {h:?}, want [2, -3]"
);
}
#[test]
fn damping_stays_finite_at_extreme_mu() {
let rows = [[1.0, 0.0], [0.0, 1.0]];
let rhs = [1.0, 1.0];
let mut prev = f64::INFINITY;
for mu in [0.0_f64, 1.0, 1e10, 1e30, 1e60] {
let mut qr = QrAccumulator::<2>::new();
for (a, b) in rows.iter().zip(&rhs) {
qr.push_row(a, *b);
}
qr.push_damping(mu);
let h = qr.solve().expect("damped system is always full rank");
let norm = (h[0] * h[0] + h[1] * h[1]).sqrt();
assert!(
norm.is_finite(),
"mu {mu:.0e} produced a non-finite step {h:?}"
);
assert!(
norm <= prev + 1e-12,
"step norm grew with damping: mu {mu:.0e} gave {norm:.3e} after {prev:.3e}"
);
prev = norm;
}
assert!(
prev < 1e-20,
"at mu = 1e60 the step should be crushed to nothing, got {prev:.3e}"
);
}
#[test]
fn rank_deficient_system_is_refused() {
let rows = [[1.0, 2.0], [2.0, 4.0]];
let rhs = [1.0, 2.0];
assert!(
solve_least_squares::<2>(&rows, &rhs).is_none(),
"a singular system must be refused rather than answered"
);
}
#[test]
fn r_transpose_r_reproduces_the_normal_matrix() {
let rows = [[1.0, 2.0], [3.0, -1.0], [0.5, 0.25]];
let rhs = [0.0, 0.0, 0.0];
let mut qr = QrAccumulator::<2>::new();
for (a, b) in rows.iter().zip(&rhs) {
qr.push_row(a, *b);
}
let r = qr.r();
for i in 0..2 {
for j in 0..2 {
let ata: f64 = rows.iter().map(|row| row[i] * row[j]).sum();
let rtr: f64 = (0..2).map(|k| r[k][i] * r[k][j]).sum();
assert!(
(ata - rtr).abs() < 1e-12,
"R^T R [{i}][{j}] = {rtr} != A^T A = {ata}"
);
}
}
}
}