use num_traits::{One, Zero};
use oxiblas_core::scalar::{ComplexScalar, Field, Real};
use oxiblas_matrix::Mat;
pub(crate) fn tridiagonalize_hermitian<T: Field + ComplexScalar>(
a: &mut Mat<T>,
u: &mut Mat<T>,
n: usize,
) -> (Vec<T::Real>, Vec<T::Real>)
where
T::Real: Real,
{
let mut diag = vec![T::Real::zero(); n];
let mut off_diag = vec![T::Real::zero(); n.saturating_sub(1)];
let mut d_phase: Vec<T> = vec![T::one(); n];
let mut v: Vec<T> = vec![T::zero(); n];
let two = T::Real::one() + T::Real::one();
for k in 0..n.saturating_sub(2) {
let mut norm_sq = T::Real::zero();
for i in (k + 1)..n {
norm_sq = norm_sq + a[(i, k)].abs_sq();
}
let xnorm = <T::Real as Real>::sqrt(norm_sq);
if xnorm <= T::Real::zero() {
d_phase[k + 1] = d_phase[k];
continue;
}
let x1 = a[(k + 1, k)];
let x1_abs = x1.abs();
let beta = if x1_abs > T::Real::zero() {
T::from_real(-xnorm) * (x1 / T::from_real(x1_abs))
} else {
T::from_real(-xnorm)
};
let vdenom = x1 - beta;
v[k + 1] = T::one();
for i in (k + 2)..n {
v[i] = a[(i, k)] / vdenom;
}
let mut v_norm_sq = T::Real::one();
for i in (k + 2)..n {
v_norm_sq = v_norm_sq + v[i].abs_sq();
}
let tau = T::from_real(two / v_norm_sq);
let mut p: Vec<T> = vec![T::zero(); n];
for i in (k + 1)..n {
let mut sum = T::zero();
for j in (k + 1)..n {
sum = sum + a[(i, j)] * v[j];
}
p[i] = tau * sum;
}
let mut vh_p = T::zero();
for i in (k + 1)..n {
vh_p = vh_p + v[i].conj() * p[i];
}
let half_tau_vhp = tau * vh_p / T::from_real(two);
let mut w: Vec<T> = vec![T::zero(); n];
for i in (k + 1)..n {
w[i] = p[i] - half_tau_vhp * v[i];
}
for i in (k + 1)..n {
for j in (k + 1)..n {
a[(i, j)] = a[(i, j)] - v[i] * w[j].conj() - w[i] * v[j].conj();
}
}
for i in 0..n {
let mut uv = T::zero();
for j in (k + 1)..n {
uv = uv + u[(i, j)] * v[j];
}
let tau_uv = tau * uv;
for j in (k + 1)..n {
u[(i, j)] = u[(i, j)] - tau_uv * v[j].conj();
}
}
off_diag[k] = xnorm; d_phase[k + 1] = d_phase[k] * (beta / T::from_real(xnorm));
}
for i in 0..n {
diag[i] = a[(i, i)].real();
}
if n >= 2 {
let e_last = a[(n - 1, n - 2)];
let e_last_abs = e_last.abs();
off_diag[n - 2] = e_last_abs;
if e_last_abs > T::Real::zero() {
d_phase[n - 1] = d_phase[n - 2] * (e_last / T::from_real(e_last_abs));
} else {
d_phase[n - 1] = d_phase[n - 2];
}
}
for j in 0..n {
let dj = d_phase[j];
for i in 0..n {
u[(i, j)] = u[(i, j)] * dj;
}
}
(diag, off_diag)
}