use crate::{FloatDType, NdArray, Result};
pub struct QrResult<T: FloatDType> {
pub q: NdArray<T>,
pub r: NdArray<T>,
}
impl<T: FloatDType> QrResult<T> {
pub fn reconstruct(&self) -> Result<NdArray<T>> {
self.q.matmul(&self.r)
}
}
pub fn qr<T: FloatDType>(arr: &NdArray<T>) -> Result<QrResult<T>> {
let a = arr.matrix_view_unsafe()?;
let (m, n) = a.shape();
unsafe {
let r_arr = a.copy(); let q_arr = NdArray::<T>::eye(m)?;
let mut r = r_arr.matrix_view_unsafe().unwrap();
let mut q = q_arr.matrix_view_unsafe().unwrap();
for k in 0..n {
let x_arr = NdArray::<T>::zeros(m - k)?;
let mut x = x_arr.vector_view_unsafe().unwrap();
for i in 0..(m - k) {
x.s(i, r.g(k + i, k));
}
let v_arr = x.copy();
let mut v = v_arr.vector_view_unsafe().unwrap();
let sign = if x.g(0) >= T::zero() { T::one() } else { -T::one() };
let norm_x = x.norm();
v.s(0, v.g(0) + sign * norm_x);
let beta = (T::one() + T::one()) / v.dot(&v)?;
for j in k..n {
let mut proj = T::zero();
for i in 0..v.len() {
proj += v.g(i) * r.g(k + i, j);
}
proj *= beta;
for i in 0..v.len() {
r.s(k + i, j, r.g(k + i, j) - proj * v.g(i));
}
}
for i in 0..m {
let mut proj = T::zero();
for j in 0..v.len() {
proj += q.g(i, k + j) * v.g(j);
}
proj *= beta;
for j in 0..v.len() {
q.s(i, k + j, q.g(i, k + j) - proj * v.g(j));
}
}
}
Ok( QrResult { q: q_arr, r: r_arr })
}
}
#[cfg(test)]
mod test {
use crate::{NdArray, linalg};
#[test]
fn test_qr_simple() {
let a = NdArray::new(&[
[12., -51., 4.],
[6., 167., -68.],
[-4., 24., -41.],
]).unwrap();
let result = linalg::qr(&a).unwrap();
let (q, r) = (result.q, result.r);
let a_rec = q.matmul(&r).unwrap();
println!("{}", a_rec);
assert!(a_rec.allclose(&a, 1e-6, 1e-6));
let qtq = q.transpose_last().unwrap().matmul(&q).unwrap();
let i = NdArray::<f64>::eye(q.dims()[1]).unwrap();
assert!(qtq.allclose(&i, 1e-6, 1e-6));
}
#[test]
fn test_qr_identity() {
let a = NdArray::<f64>::eye(4).unwrap();
let _ = linalg::qr(&a).unwrap();
}
#[test]
fn test_qr_rectangular() {
let a = NdArray::new(&[
[1., 2., 3.],
[4., 5., 6.],
[7., 8., 10.],
[1., 0., 0.],
]).unwrap();
let result = linalg::qr(&a).unwrap();
let (q, r) = (result.q, result.r);
let a_rec = q.matmul(&r).unwrap();
assert!(a_rec.allclose(&a, 1e-6, 1e-6));
let qtq = q.transpose_last().unwrap().matmul(&q).unwrap();
let m = q.dims()[1];
let i = NdArray::<f64>::eye(m).unwrap();
assert!(qtq.allclose(&i, 1e-6, 1e-6));
}
}