use super::matrix::DistFloat;
use super::DistributedLinalgError;
use scirs2_core::ndarray::{Array2, ArrayView2};
#[derive(Debug, Clone)]
pub struct HouseholderQr<T> {
factors: Array2<T>,
tau: Vec<T>,
}
impl<T: DistFloat> HouseholderQr<T> {
pub fn factor(a: Array2<T>) -> Result<Self, DistributedLinalgError> {
let (m, n) = a.dim();
if m < n {
return Err(DistributedLinalgError::UnsupportedShape(format!(
"thin Householder QR needs at least as many rows as columns, got {m}x{n}"
)));
}
let mut factors = a;
let mut tau = vec![T::zero(); n];
for j in 0..n {
let alpha = factors[[j, j]];
let mut max_abs = T::zero();
for i in (j + 1)..m {
let abs = factors[[i, j]].abs();
if abs > max_abs {
max_abs = abs;
}
}
if max_abs == T::zero() {
tau[j] = T::zero();
continue;
}
let mut scaled_sum = T::zero();
for i in (j + 1)..m {
let scaled = factors[[i, j]] / max_abs;
scaled_sum += scaled * scaled;
}
let xnorm = max_abs * scaled_sum.sqrt();
let norm = alpha.hypot(xnorm);
let beta = if alpha < T::zero() { norm } else { -norm };
let scale = alpha - beta;
if scale == T::zero() {
tau[j] = T::zero();
continue;
}
tau[j] = (beta - alpha) / beta;
for i in (j + 1)..m {
factors[[i, j]] /= scale;
}
factors[[j, j]] = beta;
let tau_j = tau[j];
for c in (j + 1)..n {
let mut dot = factors[[j, c]];
for i in (j + 1)..m {
dot += factors[[i, j]] * factors[[i, c]];
}
let w = tau_j * dot;
factors[[j, c]] -= w;
for i in (j + 1)..m {
let v = factors[[i, j]];
factors[[i, c]] -= w * v;
}
}
}
Ok(Self { factors, tau })
}
pub fn nrows(&self) -> usize {
self.factors.nrows()
}
pub fn ncols(&self) -> usize {
self.factors.ncols()
}
pub fn r(&self) -> Array2<T> {
let n = self.ncols();
let mut r = Array2::<T>::zeros((n, n));
for i in 0..n {
for j in i..n {
r[[i, j]] = self.factors[[i, j]];
}
}
r
}
pub fn packed(&self) -> ArrayView2<'_, T> {
self.factors.view()
}
pub fn tau(&self) -> &[T] {
&self.tau
}
fn apply_reflector(&self, j: usize, b: &mut Array2<T>) {
let tau_j = self.tau[j];
if tau_j == T::zero() {
return;
}
let m = self.nrows();
let k = b.ncols();
for c in 0..k {
let mut dot = b[[j, c]];
for i in (j + 1)..m {
dot += self.factors[[i, j]] * b[[i, c]];
}
let w = tau_j * dot;
b[[j, c]] -= w;
for i in (j + 1)..m {
let v = self.factors[[i, j]];
b[[i, c]] -= w * v;
}
}
}
fn check_rows(&self, b: &Array2<T>) -> Result<(), DistributedLinalgError> {
if b.nrows() != self.nrows() {
return Err(DistributedLinalgError::DimensionMismatch(format!(
"operand has {} rows, factorization has {}",
b.nrows(),
self.nrows()
)));
}
Ok(())
}
pub fn apply_qt_in_place(&self, b: &mut Array2<T>) -> Result<(), DistributedLinalgError> {
self.check_rows(b)?;
for j in 0..self.ncols() {
self.apply_reflector(j, b);
}
Ok(())
}
pub fn apply_q_in_place(&self, b: &mut Array2<T>) -> Result<(), DistributedLinalgError> {
self.check_rows(b)?;
for j in (0..self.ncols()).rev() {
self.apply_reflector(j, b);
}
Ok(())
}
pub fn thin_q(&self) -> Result<Array2<T>, DistributedLinalgError> {
let (m, n) = (self.nrows(), self.ncols());
let mut q = Array2::<T>::zeros((m, n));
for i in 0..n {
q[[i, i]] = T::one();
}
self.apply_q_in_place(&mut q)?;
Ok(q)
}
}
#[cfg(test)]
mod tests {
use super::super::matrix::testutil::{deterministic_matrix, frobenius};
use super::*;
use scirs2_core::ndarray::{s, Array2};
fn normalize_row_signs(r: &mut Array2<f64>) {
let n = r.nrows().min(r.ncols());
for i in 0..n {
if r[[i, i]] < 0.0 {
for j in 0..r.ncols() {
r[[i, j]] = -r[[i, j]];
}
}
}
}
fn identity(n: usize) -> Array2<f64> {
let mut eye = Array2::<f64>::zeros((n, n));
for i in 0..n {
eye[[i, i]] = 1.0;
}
eye
}
#[test]
fn square_r_matches_scirs2_qr_up_to_row_signs() {
for n in [1usize, 2, 3, 5, 8] {
let a = deterministic_matrix(n, n, 7 + n as u64);
let qr = HouseholderQr::factor(a.clone()).expect("factors");
let mut mine = qr.r();
let (_, reference_r) = scirs2_linalg::qr(&a.view(), None).expect("scirs2 qr");
let mut reference = reference_r.slice(s![..n, ..n]).to_owned();
normalize_row_signs(&mut mine);
normalize_row_signs(&mut reference);
let diff = &mine - &reference;
assert!(
frobenius(&diff.view()) < 1e-10,
"n={n}: ||R_mine - R_scirs2||_F = {}",
frobenius(&diff.view())
);
}
}
#[test]
fn tall_factorization_reconstructs_the_input() {
for (m, n) in [(8usize, 3usize), (32, 4), (17, 17), (5, 1)] {
let a = deterministic_matrix(m, n, 100 + m as u64);
let qr = HouseholderQr::factor(a.clone()).expect("factors");
let q = qr.thin_q().expect("thin q");
let r = qr.r();
let reconstructed = q.dot(&r);
let diff = &reconstructed - &a;
assert!(
frobenius(&diff.view()) < 1e-10,
"{m}x{n}: ||QR - A||_F = {}",
frobenius(&diff.view())
);
let qtq = q.t().dot(&q);
let ortho = &qtq - &identity(n);
assert!(
frobenius(&ortho.view()) < 1e-10,
"{m}x{n}: ||Q^T Q - I||_F = {}",
frobenius(&ortho.view())
);
for i in 0..n {
for j in 0..i {
assert!(r[[i, j]].abs() < 1e-14, "R[{i},{j}] = {}", r[[i, j]]);
}
}
}
}
#[test]
fn apply_qt_then_apply_q_is_the_identity() {
let (m, n, k) = (12usize, 4usize, 3usize);
let a = deterministic_matrix(m, n, 55);
let qr = HouseholderQr::factor(a).expect("factors");
let b = deterministic_matrix(m, k, 66);
let mut round_trip = b.clone();
qr.apply_qt_in_place(&mut round_trip).expect("Q^T b");
qr.apply_q_in_place(&mut round_trip).expect("Q Q^T b");
let diff = &round_trip - &b;
assert!(frobenius(&diff.view()) < 1e-12);
}
#[test]
fn implicit_and_explicit_thin_q_agree() {
let (m, n, k) = (10usize, 3usize, 2usize);
let a = deterministic_matrix(m, n, 77);
let qr = HouseholderQr::factor(a).expect("factors");
let q = qr.thin_q().expect("thin q");
let d = deterministic_matrix(n, k, 88);
let mut padded = Array2::<f64>::zeros((m, k));
padded.slice_mut(s![..n, ..]).assign(&d);
qr.apply_q_in_place(&mut padded).expect("Q [D; 0]");
let explicit = q.dot(&d);
let diff = &padded - &explicit;
assert!(
frobenius(&diff.view()) < 1e-12,
"||Q[D;0] - Q_thin D||_F = {}",
frobenius(&diff.view())
);
let b = deterministic_matrix(m, k, 99);
let mut qt_b = b.clone();
qr.apply_qt_in_place(&mut qt_b).expect("Q^T B");
let explicit_t = q.t().dot(&b);
let diff_t = &qt_b.slice(s![..n, ..]).to_owned() - &explicit_t;
assert!(frobenius(&diff_t.view()) < 1e-12);
}
#[test]
fn zero_column_leaves_a_zero_reflector() {
let mut a = Array2::<f64>::zeros((4, 2));
a[[0, 0]] = 0.0;
a[[0, 1]] = 1.0;
let qr = HouseholderQr::factor(a.clone()).expect("factors");
assert_eq!(qr.tau()[0], 0.0);
assert_eq!(qr.r()[[0, 0]], 0.0);
let reconstructed = qr.thin_q().expect("thin q").dot(&qr.r());
let diff = &reconstructed - &a;
assert!(frobenius(&diff.view()) < 1e-12);
}
#[test]
fn wide_block_is_refused_not_approximated() {
let a = deterministic_matrix(2, 5, 5);
assert!(matches!(
HouseholderQr::factor(a),
Err(DistributedLinalgError::UnsupportedShape(_))
));
}
}