use crate::error::LapackError;
use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone)]
pub struct Rq<T: Field> {
pub(crate) factors: Mat<T>,
pub(crate) tau: Vec<T>,
}
impl<T: Field + Real + bytemuck::Zeroable> Rq<T> {
pub fn compute(a: MatRef<T>) -> Result<Self, LapackError> {
let m = a.nrows();
let n = a.ncols();
let k = m.min(n);
let mut factors = Mat::zeros(m, n);
for i in 0..m {
for j in 0..n {
factors[(i, j)] = a[(i, j)];
}
}
let mut tau = vec![T::zero(); k];
for step in 0..k {
let row = m - 1 - step;
let vec_len = n - step;
if vec_len == 0 {
continue;
}
let pivot_col = vec_len - 1;
let (tau_val, beta) = compute_householder_right(&mut factors, row, vec_len);
tau[step] = tau_val;
factors[(row, pivot_col)] = beta;
if tau_val != T::zero() {
for i in 0..m {
if i == row {
continue; }
let mut w = factors[(i, pivot_col)]; for j in 0..pivot_col {
w = w + factors[(i, j)] * factors[(row, j)];
}
let tw = tau_val * w;
factors[(i, pivot_col)] = factors[(i, pivot_col)] - tw;
for j in 0..pivot_col {
factors[(i, j)] = factors[(i, j)] - tw * factors[(row, j)].conj();
}
}
}
}
Ok(Self { factors, tau })
}
#[must_use]
pub fn r_factor(&self) -> Mat<T> {
let m = self.factors.nrows();
let n = self.factors.ncols();
let mut r = Mat::zeros(m, n);
for i in 0..m {
if m <= n {
let start_col = n - m + i;
for j in start_col..n {
r[(i, j)] = self.factors[(i, j)];
}
} else {
if i < m - n {
for j in 0..n {
r[(i, j)] = self.factors[(i, j)];
}
} else {
let local_i = i - (m - n);
for j in local_i..n {
r[(i, j)] = self.factors[(i, j)];
}
}
}
}
r
}
#[must_use]
pub fn dims(&self) -> (usize, usize) {
(self.factors.nrows(), self.factors.ncols())
}
#[must_use]
pub fn q_factor(&self) -> Mat<T> {
let m = self.factors.nrows();
let n = self.factors.ncols();
let k = m.min(n);
let mut q = Mat::zeros(n, n);
for i in 0..n {
q[(i, i)] = T::one();
}
for step in (0..k).rev() {
let row = m - 1 - step;
let vec_len = n - step;
if vec_len == 0 || self.tau[step] == T::zero() {
continue;
}
let pivot_col = vec_len - 1;
for qcol in 0..n {
let mut w = q[(pivot_col, qcol)]; for j in 0..pivot_col {
w = w + self.factors[(row, j)].conj() * q[(j, qcol)];
}
let tw = self.tau[step] * w;
q[(pivot_col, qcol)] = q[(pivot_col, qcol)] - tw;
for j in 0..pivot_col {
q[(j, qcol)] = q[(j, qcol)] - tw * self.factors[(row, j)];
}
}
}
q
}
}
fn compute_householder_right<T: Field + Real>(
factors: &mut Mat<T>,
row: usize,
len: usize,
) -> (T, T) {
if len == 0 {
return (T::zero(), T::zero());
}
if len == 1 {
return (T::zero(), factors[(row, 0)]);
}
let mut norm_sq = T::zero();
for j in 0..len {
norm_sq = norm_sq + factors[(row, j)] * factors[(row, j)].conj();
}
let norm = Real::sqrt(norm_sq);
if norm == T::zero() {
return (T::zero(), T::zero());
}
let pivot_idx = len - 1;
let x_pivot = factors[(row, pivot_idx)];
let beta = if x_pivot.real() >= T::Real::zero() {
-norm
} else {
norm
};
let tau = (beta - x_pivot) / beta;
let denom = x_pivot - beta;
if Scalar::abs(denom) > <T as Scalar>::epsilon() {
let scale = T::one() / denom;
for j in 0..pivot_idx {
factors[(row, j)] = factors[(row, j)] * scale;
}
}
(tau, beta)
}
#[cfg(test)]
mod tests {
use super::*;
fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn test_rq_square() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let rq = Rq::compute(a.as_ref()).unwrap();
let r = rq.r_factor();
assert!(r[(1, 0)].abs() < 1e-10);
}
#[test]
fn test_rq_wide() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0, 4.0], &[5.0, 6.0, 7.0, 8.0]]);
let rq = Rq::compute(a.as_ref()).unwrap();
let r = rq.r_factor();
let q = rq.q_factor();
assert_eq!(r.nrows(), 2);
assert_eq!(r.ncols(), 4);
assert_eq!(q.nrows(), 4);
assert_eq!(q.ncols(), 4);
}
#[test]
fn test_rq_tall() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0]]);
let rq = Rq::compute(a.as_ref()).unwrap();
let r = rq.r_factor();
assert_eq!(r.nrows(), 3);
assert_eq!(r.ncols(), 2);
}
#[test]
fn test_rq_q_orthogonal() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let rq = Rq::compute(a.as_ref()).unwrap();
let q = rq.q_factor();
for i in 0..3 {
for j in 0..3 {
let mut dot = 0.0;
for k in 0..3 {
dot += q[(k, i)] * q[(k, j)];
}
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
approx_eq(dot, expected, 1e-10),
"Q^T*Q[{},{}] = {}, expected {}",
i,
j,
dot,
expected
);
}
}
}
#[test]
fn test_rq_identity() {
let a: Mat<f64> = Mat::eye(3);
let rq = Rq::compute(a.as_ref()).unwrap();
let r = rq.r_factor();
let q = rq.q_factor();
for i in 0..3 {
assert!(r[(i, i)].abs() > 0.99, "R diagonal should be ~1");
assert!(q[(i, i)].abs() > 0.99, "Q diagonal should be ~1");
}
}
#[test]
fn test_rq_reconstruction() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let rq = Rq::compute(a.as_ref()).unwrap();
let r = rq.r_factor();
let q = rq.q_factor();
let m = a.nrows();
let n = a.ncols();
for i in 0..m {
for j in 0..n {
let mut sum = 0.0;
for k in 0..n {
sum += r[(i, k)] * q[(k, j)];
}
assert!(
approx_eq(sum, a[(i, j)], 1e-10),
"Reconstruction[{},{}] = {}, expected {}",
i,
j,
sum,
a[(i, j)]
);
}
}
}
#[test]
fn test_rq_reconstruction_tall() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0], &[7.0, 9.0]]);
let rq = Rq::compute(a.as_ref()).unwrap();
let r = rq.r_factor();
let q = rq.q_factor();
let m = a.nrows();
let n = a.ncols();
assert_eq!(r.nrows(), m);
assert_eq!(r.ncols(), n);
assert_eq!(q.nrows(), n);
assert_eq!(q.ncols(), n);
let mut top_block_all_zero = true;
for i in 0..(m - n) {
for j in 0..n {
if r[(i, j)].abs() > 1e-10 {
top_block_all_zero = false;
}
}
}
assert!(
!top_block_all_zero,
"top (m-n)x n block of R is all zero -- regression of the r_factor zeroing bug"
);
for i in 0..m {
for j in 0..n {
let mut sum = 0.0;
for k in 0..n {
sum += r[(i, k)] * q[(k, j)];
}
assert!(
approx_eq(sum, a[(i, j)], 1e-10),
"Tall reconstruction[{},{}] = {}, expected {}",
i,
j,
sum,
a[(i, j)]
);
}
}
}
}