use num_traits::Zero;
use oxiblas_core::scalar::{ComplexScalar, Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum UnitaryQrError {
EmptyMatrix,
}
impl core::fmt::Display for UnitaryQrError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
}
}
}
impl std::error::Error for UnitaryQrError {}
#[derive(Debug, Clone)]
pub struct UnitaryQr<T: Scalar> {
qr: Mat<T>,
tau: Vec<T>,
m: usize,
n: usize,
}
impl<T: Field + ComplexScalar + bytemuck::Zeroable> UnitaryQr<T>
where
T::Real: Real,
{
pub fn compute(a: MatRef<'_, T>) -> Result<Self, UnitaryQrError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(UnitaryQrError::EmptyMatrix);
}
let mut qr: Mat<T> = Mat::zeros(m, n);
for j in 0..n {
for i in 0..m {
qr[(i, j)] = a[(i, j)];
}
}
let k = m.min(n);
let mut tau = vec![T::zero(); k];
for j in 0..k {
let (tau_j, beta) = complex_householder_vector(&mut qr, j, m);
tau[j] = tau_j;
qr[(j, j)] = beta;
if j < n - 1 {
apply_complex_householder_left(&mut qr, j, m, n, tau_j);
}
}
Ok(Self { qr, tau, m, n })
}
pub fn nrows(&self) -> usize {
self.m
}
pub fn ncols(&self) -> usize {
self.n
}
pub fn r(&self) -> Mat<T> {
let k = self.m.min(self.n);
let mut r: Mat<T> = Mat::zeros(self.m, self.n);
for j in 0..self.n {
for i in 0..=j.min(self.m - 1) {
r[(i, j)] = self.qr[(i, j)];
}
}
for j in 0..k {
for i in (j + 1)..self.m {
r[(i, j)] = T::zero();
}
}
r
}
pub fn r_thin(&self) -> Mat<T> {
let k = self.m.min(self.n);
let mut r: Mat<T> = Mat::zeros(k, self.n);
for j in 0..self.n {
for i in 0..=j.min(k - 1) {
r[(i, j)] = self.qr[(i, j)];
}
}
r
}
pub fn q(&self) -> Mat<T> {
let k = self.m.min(self.n);
let mut q: Mat<T> = Mat::zeros(self.m, self.m);
for i in 0..self.m {
q[(i, i)] = T::one();
}
for j in (0..k).rev() {
apply_complex_householder_to_q(&mut q, &self.qr, j, self.m, self.tau[j]);
}
q
}
pub fn q_thin(&self) -> Mat<T> {
let k = self.m.min(self.n);
let mut q: Mat<T> = Mat::zeros(self.m, k);
for i in 0..k {
q[(i, i)] = T::one();
}
for j in (0..k).rev() {
apply_complex_householder_to_q_thin(&mut q, &self.qr, j, self.m, k, self.tau[j]);
}
q
}
pub fn tau(&self) -> &[T] {
&self.tau
}
}
fn complex_householder_vector<T: Field + ComplexScalar>(
qr: &mut Mat<T>,
j: usize,
m: usize,
) -> (T, T)
where
T::Real: Real,
{
let mut norm_sq = T::Real::zero();
for i in j..m {
norm_sq = norm_sq + qr[(i, j)].abs_sq();
}
let norm = <T::Real as Real>::sqrt(norm_sq);
if norm == T::Real::zero() {
return (T::zero(), T::zero());
}
let x_j = qr[(j, j)];
let x_j_abs = x_j.abs();
let beta = if x_j_abs > T::Real::zero() {
let sign = T::from_real_imag(x_j.real() / x_j_abs, x_j.imag() / x_j_abs);
T::zero() - sign * T::from_real(norm)
} else {
T::from_real(-norm)
};
let diff = beta - x_j;
let tau = diff / beta;
let scale_denom = x_j - beta;
if scale_denom.abs() > T::Real::zero() {
let scale = T::one() / scale_denom;
for i in (j + 1)..m {
qr[(i, j)] = qr[(i, j)] * scale;
}
}
(tau, beta)
}
fn apply_complex_householder_left<T: Field + ComplexScalar>(
qr: &mut Mat<T>,
j: usize,
m: usize,
n: usize,
tau: T,
) where
T::Real: Real,
{
if tau.abs() < T::Real::epsilon() {
return;
}
for k in (j + 1)..n {
let mut w = qr[(j, k)]; for i in (j + 1)..m {
w = w + qr[(i, j)].conj() * qr[(i, k)];
}
let tw = tau * w;
qr[(j, k)] = qr[(j, k)] - tw; for i in (j + 1)..m {
qr[(i, k)] = qr[(i, k)] - tw * qr[(i, j)];
}
}
}
fn apply_complex_householder_to_q<T: Field + ComplexScalar>(
q: &mut Mat<T>,
qr: &Mat<T>,
j: usize,
m: usize,
tau: T,
) where
T::Real: Real,
{
if tau.abs() < T::Real::epsilon() {
return;
}
for k in 0..m {
let mut w = q[(j, k)]; for i in (j + 1)..m {
w = w + qr[(i, j)].conj() * q[(i, k)];
}
let tw = tau * w;
q[(j, k)] = q[(j, k)] - tw;
for i in (j + 1)..m {
q[(i, k)] = q[(i, k)] - tw * qr[(i, j)];
}
}
}
fn apply_complex_householder_to_q_thin<T: Field + ComplexScalar>(
q: &mut Mat<T>,
qr: &Mat<T>,
j: usize,
m: usize,
ncols: usize,
tau: T,
) where
T::Real: Real,
{
if tau.abs() < T::Real::epsilon() {
return;
}
for k in 0..ncols {
let mut w = q[(j, k)];
for i in (j + 1)..m {
w = w + qr[(i, j)].conj() * q[(i, k)];
}
let tw = tau * w;
q[(j, k)] = q[(j, k)] - tw;
for i in (j + 1)..m {
q[(i, k)] = q[(i, k)] - tw * qr[(i, j)];
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use num_complex::{Complex32, Complex64};
#[test]
fn test_unitary_qr_simple() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 0.0), Complex64::new(2.0, 1.0)],
&[Complex64::new(3.0, -1.0), Complex64::new(4.0, 0.0)],
]);
let qr = UnitaryQr::compute(a.as_ref()).expect("Should compute");
let q = qr.q();
let r = qr.r();
let n = q.nrows();
for i in 0..n {
for j in 0..n {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..n {
sum = 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)
};
let diff = (sum - expected).norm();
assert!(
diff < 1e-10,
"Q^H*Q[{},{}] = {:?}, expected {:?}",
i,
j,
sum,
expected
);
}
}
assert!(r[(1, 0)].norm() < 1e-10, "R[1,0] = {:?}", r[(1, 0)]);
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 = sum + q[(i, k)] * r[(k, j)];
}
let diff = (sum - a[(i, j)]).norm();
assert!(
diff < 1e-10,
"QR[{},{}] = {:?}, A = {:?}",
i,
j,
sum,
a[(i, j)]
);
}
}
}
#[test]
fn test_unitary_qr_tall() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 1.0), Complex64::new(2.0, 0.0)],
&[Complex64::new(3.0, 0.0), Complex64::new(4.0, -1.0)],
&[Complex64::new(5.0, -1.0), Complex64::new(6.0, 1.0)],
]);
let qr = UnitaryQr::compute(a.as_ref()).expect("Should compute");
let q = qr.q();
let r = qr.r();
assert_eq!(q.nrows(), 3);
assert_eq!(q.ncols(), 3);
let n = q.nrows();
for i in 0..n {
for j in 0..n {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..n {
sum = 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)
};
let diff = (sum - expected).norm();
assert!(diff < 1e-10);
}
}
for i in 0..3 {
for j in 0..2 {
let mut sum = Complex64::new(0.0, 0.0);
for k in 0..3 {
sum = sum + q[(i, k)] * r[(k, j)];
}
let diff = (sum - a[(i, j)]).norm();
assert!(diff < 1e-10);
}
}
}
#[test]
fn test_unitary_qr_hermitian() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(4.0, 0.0), Complex64::new(1.0, -1.0)],
&[Complex64::new(1.0, 1.0), Complex64::new(3.0, 0.0)],
]);
let qr = UnitaryQr::compute(a.as_ref()).expect("Should compute");
let q = qr.q();
let r = qr.r();
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 = sum + q[(i, k)] * r[(k, j)];
}
let diff = (sum - a[(i, j)]).norm();
assert!(diff < 1e-10);
}
}
}
#[test]
fn test_unitary_qr_complex32() {
let a: Mat<Complex32> = Mat::from_rows(&[
&[Complex32::new(1.0, 0.0), Complex32::new(2.0, 1.0)],
&[Complex32::new(3.0, -1.0), Complex32::new(4.0, 0.0)],
]);
let qr = UnitaryQr::compute(a.as_ref()).expect("Should compute");
let q = qr.q();
let r = qr.r();
for i in 0..2 {
for j in 0..2 {
let mut sum = Complex32::new(0.0, 0.0);
for k in 0..2 {
sum = sum + q[(i, k)] * r[(k, j)];
}
let diff = (sum - a[(i, j)]).norm();
assert!(
diff < 1e-5,
"QR[{},{}] = {:?}, A = {:?}",
i,
j,
sum,
a[(i, j)]
);
}
}
}
#[test]
fn test_unitary_qr_identity() {
let eye: Mat<Complex64> = Mat::eye(3);
let qr = UnitaryQr::compute(eye.as_ref()).expect("Should compute");
let q = qr.q();
let r = qr.r();
for i in 0..3 {
for j in 0..3 {
if i == j {
assert!(q[(i, j)].norm() > 0.99);
assert!(r[(i, j)].norm() > 0.99);
} else {
assert!(q[(i, j)].norm() < 1e-10);
assert!(r[(i, j)].norm() < 1e-10);
}
}
}
}
#[test]
fn test_unitary_qr_thin() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 1.0), Complex64::new(2.0, 0.0)],
&[Complex64::new(3.0, 0.0), Complex64::new(4.0, -1.0)],
&[Complex64::new(5.0, -1.0), Complex64::new(6.0, 1.0)],
]);
let qr = UnitaryQr::compute(a.as_ref()).expect("Should compute");
let q_thin = qr.q_thin();
let r_thin = qr.r_thin();
assert_eq!(q_thin.nrows(), 3);
assert_eq!(q_thin.ncols(), 2);
assert_eq!(r_thin.nrows(), 2);
assert_eq!(r_thin.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 = sum + q_thin[(i, k)] * r_thin[(k, j)];
}
let diff = (sum - a[(i, j)]).norm();
assert!(diff < 1e-10);
}
}
}
}