use crate::error::LapackError;
use num_traits::Zero;
use oxiblas_core::scalar::{Field, Real};
use oxiblas_matrix::{Mat, MatRef};
fn lq_reflector<T: Field>(row: &[T]) -> (Vec<T>, T, T)
where
T::Real: Real,
{
let p = row.len();
let mut v = vec![T::zero(); p];
if p == 0 {
return (v, T::zero(), T::zero());
}
v[0] = T::one();
let mut tail_norm_sq = T::Real::zero();
for &entry in &row[1..] {
tail_norm_sq = tail_norm_sq + entry.abs_sq();
}
let c0 = row[0].conj();
if tail_norm_sq == T::Real::zero() && c0.imag() == T::Real::zero() {
return (v, T::zero(), c0);
}
let norm = <T::Real as Real>::sqrt(c0.abs_sq() + tail_norm_sq);
let beta_real = if c0.real() >= T::Real::zero() {
-norm
} else {
norm
};
let beta = T::from_real(beta_real);
let tau = (beta - c0) / beta;
let scale = T::one() / (c0 - beta);
for j in 1..p {
v[j] = row[j].conj() * scale;
}
(v, tau, beta)
}
#[derive(Debug, Clone)]
pub struct Lq<T: Field> {
pub(crate) factors: Mat<T>,
pub(crate) tau: Vec<T>,
}
impl<T: Field + bytemuck::Zeroable> Lq<T>
where
T::Real: Real,
{
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 i in 0..k {
let row_len = n - i;
let mut row = vec![T::zero(); row_len];
for j in 0..row_len {
row[j] = factors[(i, i + j)];
}
let (v, tau_val, beta) = lq_reflector(&row);
tau[i] = tau_val;
factors[(i, i)] = beta;
for j in 1..row_len {
factors[(i, i + j)] = v[j];
}
if tau_val != T::zero() {
for r in (i + 1)..m {
let mut s = T::zero();
for j in 0..row_len {
s = s + factors[(r, i + j)] * v[j];
}
let ts = tau_val * s;
for j in 0..row_len {
factors[(r, i + j)] = factors[(r, i + j)] - ts * v[j].conj();
}
}
}
}
Ok(Self { factors, tau })
}
#[must_use]
pub fn l_factor(&self) -> Mat<T> {
let m = self.factors.nrows();
let n = self.factors.ncols();
let mut l = Mat::zeros(m, n);
if m == 0 || n == 0 {
return l;
}
for i in 0..m {
for j in 0..=i.min(n - 1) {
l[(i, j)] = self.factors[(i, j)];
}
}
l
}
#[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 i in (0..k).rev() {
let tau_i = self.tau[i];
if tau_i == T::zero() {
continue;
}
let ctau = tau_i.conj();
let row_len = n - i;
for c in 0..n {
let mut s = q[(c, i)];
for j in 1..row_len {
s = s + q[(c, i + j)] * self.factors[(i, i + j)];
}
let ts = ctau * s;
q[(c, i)] = q[(c, i)] - ts;
for j in 1..row_len {
q[(c, i + j)] = q[(c, i + j)] - ts * self.factors[(i, i + j)].conj();
}
}
}
q
}
}
#[cfg(test)]
mod tests {
use super::*;
use num_complex::Complex64;
#[test]
fn test_lq_square() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0]]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
assert!(l[(0, 1)].abs() < 1e-10);
}
#[test]
fn test_lq_wide() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
let q = lq.q_factor();
assert_eq!(l.nrows(), 2);
assert_eq!(l.ncols(), 3);
assert_eq!(q.nrows(), 3);
assert_eq!(q.ncols(), 3);
assert!(l[(0, 1)].abs() < 1e-10);
assert!(l[(0, 2)].abs() < 1e-10);
assert!(l[(1, 2)].abs() < 1e-10);
}
#[test]
fn test_lq_reconstruction() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
let q = lq.q_factor();
let mut reconstructed = Mat::zeros(2, 3);
for i in 0..2 {
for j in 0..3 {
let mut sum = 0.0;
for k in 0..3 {
sum += l[(i, k)] * q[(k, j)];
}
reconstructed[(i, j)] = sum;
}
}
for i in 0..2 {
for j in 0..3 {
assert!((reconstructed[(i, j)] - a[(i, j)]).abs() < 1e-10);
}
}
}
#[test]
fn test_lq_q_orthogonal() {
let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0]]);
let lq = Lq::compute(a.as_ref()).unwrap();
let q = lq.q_factor();
let mut qtq = Mat::zeros(3, 3);
for i in 0..3 {
for j in 0..3 {
let mut sum = 0.0;
for k in 0..3 {
sum += q[(k, i)] * q[(k, j)];
}
qtq[(i, j)] = sum;
}
}
for i in 0..3 {
for j in 0..3 {
let expected = if i == j { 1.0 } else { 0.0 };
assert!((qtq[(i, j)] - expected).abs() < 1e-10);
}
}
}
#[test]
fn test_lq_tall_reconstruction() {
let a = Mat::from_rows(&[&[1.0f64, 2.0], &[3.0, 4.0], &[5.0, 6.0], &[7.0, 9.0]]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
let q = lq.q_factor();
assert_eq!(l.nrows(), 4);
assert_eq!(l.ncols(), 2);
assert_eq!(q.nrows(), 2);
assert_eq!(q.ncols(), 2);
for i in 0..4 {
for j in 0..2 {
let mut sum = 0.0;
for k in 0..2 {
sum += l[(i, k)] * q[(k, j)];
}
assert!(
(sum - a[(i, j)]).abs() < 1e-10,
"reconstruction mismatch at ({i}, {j})"
);
}
}
}
#[test]
fn test_lq_empty_l_factor_does_not_panic() {
let a_cols: Mat<f64> = Mat::zeros(3, 0);
let lq_cols = Lq::compute(a_cols.as_ref()).unwrap();
let l_cols = lq_cols.l_factor();
assert_eq!(l_cols.nrows(), 3);
assert_eq!(l_cols.ncols(), 0);
assert_eq!(lq_cols.dims(), (3, 0));
let a_rows: Mat<f64> = Mat::zeros(0, 4);
let lq_rows = Lq::compute(a_rows.as_ref()).unwrap();
let l_rows = lq_rows.l_factor();
assert_eq!(l_rows.nrows(), 0);
assert_eq!(l_rows.ncols(), 4);
let a_empty: Mat<f64> = Mat::zeros(0, 0);
let lq_empty = Lq::compute(a_empty.as_ref()).unwrap();
let l_empty = lq_empty.l_factor();
assert_eq!(l_empty.nrows(), 0);
assert_eq!(l_empty.ncols(), 0);
}
#[test]
fn test_lq_complex_square_reconstruction() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(1.0, 1.0),
Complex64::new(2.0, -1.0),
Complex64::new(0.0, 2.0),
],
&[
Complex64::new(3.0, 0.5),
Complex64::new(1.0, 1.0),
Complex64::new(2.0, 2.0),
],
&[
Complex64::new(0.0, -1.0),
Complex64::new(4.0, 1.0),
Complex64::new(1.0, -3.0),
],
]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
let q = lq.q_factor();
for i in 0..3 {
for j in 0..3 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..3 {
sum += l[(i, k)] * q[(k, j)];
}
let diff = (sum - a[(i, j)]).norm();
assert!(
diff < 1e-10,
"LQ[{i},{j}] = {sum:?}, A = {:?}, diff = {diff}",
a[(i, j)]
);
}
}
for i in 0..3 {
for j in (i + 1)..3 {
assert!(
l[(i, j)].norm() < 1e-10,
"L not lower triangular at ({i}, {j}): {:?}",
l[(i, j)]
);
}
}
for i in 0..3 {
for j in 0..3 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..3 {
sum += q[(i, k)] * q[(j, k)].conj();
}
let expected = if i == j {
Complex64::new(1.0, 0.0)
} else {
Complex64::new(0.0, 0.0)
};
assert!(
(sum - expected).norm() < 1e-10,
"Q not unitary at ({i}, {j}): {sum:?}"
);
}
}
}
#[test]
fn test_lq_complex_wide_reconstruction() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(2.0, 1.0),
Complex64::new(-1.0, 3.0),
Complex64::new(0.0, -2.0),
],
&[
Complex64::new(1.0, -1.0),
Complex64::new(4.0, 0.5),
Complex64::new(-2.0, 1.0),
],
]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
let q = lq.q_factor();
assert_eq!(l.nrows(), 2);
assert_eq!(l.ncols(), 3);
assert_eq!(q.nrows(), 3);
assert_eq!(q.ncols(), 3);
for i in 0..2 {
for j in 0..3 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..3 {
sum += l[(i, k)] * q[(k, j)];
}
assert!(
(sum - a[(i, j)]).norm() < 1e-10,
"wide LQ reconstruction mismatch at ({i}, {j})"
);
}
}
assert!(l[(0, 1)].norm() < 1e-10);
assert!(l[(0, 2)].norm() < 1e-10);
assert!(l[(1, 2)].norm() < 1e-10);
for i in 0..3 {
for j in 0..3 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..3 {
sum += q[(k, i)].conj() * q[(k, j)];
}
let expected = if i == j {
Complex64::new(1.0, 0.0)
} else {
Complex64::new(0.0, 0.0)
};
assert!(
(sum - expected).norm() < 1e-10,
"Q^H Q not identity at ({i}, {j}): {sum:?}"
);
}
}
}
#[test]
fn test_lq_complex_tall_reconstruction() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 2.0), Complex64::new(3.0, -1.0)],
&[Complex64::new(-2.0, 1.0), Complex64::new(0.0, 4.0)],
&[Complex64::new(5.0, -3.0), Complex64::new(1.0, 1.0)],
]);
let lq = Lq::compute(a.as_ref()).unwrap();
let l = lq.l_factor();
let q = lq.q_factor();
assert_eq!(l.nrows(), 3);
assert_eq!(l.ncols(), 2);
assert_eq!(q.nrows(), 2);
assert_eq!(q.ncols(), 2);
for i in 0..3 {
for j in 0..2 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..2 {
sum += l[(i, k)] * q[(k, j)];
}
assert!(
(sum - a[(i, j)]).norm() < 1e-10,
"tall complex LQ reconstruction mismatch at ({i}, {j})"
);
}
}
for i in 0..2 {
for j in 0..2 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..2 {
sum += q[(i, k)] * q[(j, k)].conj();
}
let expected = if i == j {
Complex64::new(1.0, 0.0)
} else {
Complex64::new(0.0, 0.0)
};
assert!((sum - expected).norm() < 1e-10);
}
}
}
}