use mdarray::{Dim, Layout, Slice};
use num_complex::ComplexFloat;
use num_traits::{MulAdd, One, Zero};
pub fn naive_qr<T, D0: Dim, D1: Dim, D2: Dim, L, Lq, Lr>(
a: &mut Slice<T, (D0, D1), L>,
q: &mut Slice<T, (D0, D2), Lq>,
r: &mut Slice<T, (D2, D1), Lr>,
) where
T: ComplexFloat + Zero + One + MulAdd<Output = T>,
L: Layout,
Lq: Layout,
Lr: Layout,
{
let (m, n) = *a.shape();
let m_size = m.size();
let n_size = n.size();
assert_eq!(q.shape().0.size(), m_size);
assert_eq!(q.shape().1.size(), n_size);
assert_eq!(r.shape().0.size(), n_size);
assert_eq!(r.shape().1.size(), n_size);
for i in 0..n_size {
for j in 0..n_size {
r[[i, j]] = T::zero();
}
}
for j in 0..n_size {
for i in 0..m_size {
q[[i, j]] = a[[i, j]];
}
for i in 0..j {
let mut dot = T::zero();
for k in 0..m_size {
dot = q[[k, i]].conj().mul_add(q[[k, j]], dot);
}
r[[i, j]] = dot;
for k in 0..m_size {
q[[k, j]] = q[[k, j]] - r[[i, j]] * q[[k, i]];
}
}
let mut norm_sq = T::zero();
for k in 0..m_size {
norm_sq = q[[k, j]].conj().mul_add(q[[k, j]], norm_sq);
}
let norm = norm_sq.sqrt();
r[[j, j]] = norm;
let inv_norm = T::one() / norm;
for k in 0..m_size {
q[[k, j]] = q[[k, j]] * inv_norm;
}
}
}