use num_traits::{FromPrimitive, One, Zero};
use oxiblas_core::scalar::{ComplexScalar, Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use super::hessenberg::{Side, Trans};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ComplexHessenbergError {
EmptyMatrix,
NotSquare,
DimensionMismatch,
}
impl core::fmt::Display for ComplexHessenbergError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::EmptyMatrix => write!(f, "Matrix is empty"),
Self::NotSquare => write!(f, "Matrix must be square"),
Self::DimensionMismatch => write!(f, "Dimension mismatch"),
}
}
}
impl std::error::Error for ComplexHessenbergError {}
#[derive(Debug, Clone)]
pub struct ComplexHessenberg<T: Scalar> {
q: Mat<T>,
h: Mat<T>,
n: usize,
}
impl<T: Field + ComplexScalar + bytemuck::Zeroable> ComplexHessenberg<T>
where
T::Real: Real,
{
pub fn compute(a: MatRef<'_, T>) -> Result<Self, ComplexHessenbergError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(ComplexHessenbergError::EmptyMatrix);
}
if m != n {
return Err(ComplexHessenbergError::NotSquare);
}
if n == 1 {
let mut h: Mat<T> = Mat::zeros(1, 1);
h[(0, 0)] = a[(0, 0)];
let mut q: Mat<T> = Mat::zeros(1, 1);
q[(0, 0)] = T::one();
return Ok(Self { q, h, n });
}
if n == 2 {
let mut h: Mat<T> = Mat::zeros(2, 2);
for i in 0..2 {
for j in 0..2 {
h[(i, j)] = a[(i, j)];
}
}
let mut q: Mat<T> = Mat::zeros(2, 2);
q[(0, 0)] = T::one();
q[(1, 1)] = T::one();
return Ok(Self { q, h, n });
}
let mut h: Mat<T> = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
h[(i, j)] = a[(i, j)];
}
}
let mut q: Mat<T> = Mat::zeros(n, n);
for i in 0..n {
q[(i, i)] = T::one();
}
for k in 0..(n - 2) {
let m_size = n - k - 1;
let mut x: Vec<T> = vec![T::zero(); m_size];
for i in 0..m_size {
x[i] = h[(k + 1 + i, k)];
}
let (v, tau) = complex_householder_vector(&x);
if tau.abs() > T::Real::epsilon() {
for j in k..n {
let mut dot = T::zero();
for i in 0..v.len() {
dot = dot + v[i].conj() * h[(k + 1 + i, j)];
}
let scaled = tau * dot;
for i in 0..v.len() {
h[(k + 1 + i, j)] = h[(k + 1 + i, j)] - scaled * v[i];
}
}
for i in 0..n {
let mut dot = T::zero();
for j in 0..v.len() {
dot = dot + h[(i, k + 1 + j)] * v[j];
}
let scaled = tau * dot;
for j in 0..v.len() {
h[(i, k + 1 + j)] = h[(i, k + 1 + j)] - scaled * v[j].conj();
}
}
for i in 0..n {
let mut dot = T::zero();
for j in 0..v.len() {
dot = dot + q[(i, k + 1 + j)] * v[j];
}
let scaled = tau * dot;
for j in 0..v.len() {
q[(i, k + 1 + j)] = q[(i, k + 1 + j)] - scaled * v[j].conj();
}
}
}
}
let hundred: T::Real =
<T::Real as FromPrimitive>::from_f64(100.0).unwrap_or(<T::Real as One>::one());
let eps = <T::Real as Scalar>::epsilon() * hundred;
for j in 0..(n - 2) {
for i in (j + 2)..n {
if h[(i, j)].abs() < eps {
h[(i, j)] = T::zero();
}
}
}
Ok(Self { q, h, n })
}
pub fn q(&self) -> MatRef<'_, T> {
self.q.as_ref()
}
pub fn h(&self) -> MatRef<'_, T> {
self.h.as_ref()
}
pub fn dim(&self) -> usize {
self.n
}
pub fn reconstruct(&self) -> Mat<T> {
let n = self.n;
let mut a: Mat<T> = Mat::zeros(n, n);
let mut hqh: Mat<T> = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
let mut sum = T::zero();
for k in 0..n {
sum = sum + self.h[(i, k)] * self.q[(j, k)].conj();
}
hqh[(i, j)] = sum;
}
}
for i in 0..n {
for j in 0..n {
let mut sum = T::zero();
for k in 0..n {
sum = sum + self.q[(i, k)] * hqh[(k, j)];
}
a[(i, j)] = sum;
}
}
a
}
}
#[derive(Debug, Clone)]
pub struct ComplexHessenbergFactors<T: Scalar> {
qr: Mat<T>,
tau: Vec<T>,
n: usize,
ilo: usize,
ihi: usize,
}
impl<T: Field + ComplexScalar + bytemuck::Zeroable> ComplexHessenbergFactors<T>
where
T::Real: Real,
{
pub fn dim(&self) -> usize {
self.n
}
pub fn ilo(&self) -> usize {
self.ilo
}
pub fn ihi(&self) -> usize {
self.ihi
}
pub fn tau(&self) -> &[T] {
&self.tau
}
pub fn h(&self) -> Mat<T> {
let n = self.n;
let mut h: Mat<T> = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
if i <= j + 1 {
h[(i, j)] = self.qr[(i, j)];
}
}
}
h
}
pub fn q(&self) -> Mat<T> {
let n = self.n;
let mut q: Mat<T> = Mat::zeros(n, n);
for i in 0..n {
q[(i, i)] = T::one();
}
if n <= 2 {
return q;
}
for k in (self.ilo..(self.ihi.saturating_sub(1))).rev() {
let tau = self.tau[k - self.ilo];
if tau.abs() < T::Real::epsilon() {
continue;
}
let m_size = self.ihi - k - 1;
if m_size == 0 {
continue;
}
let mut v: Vec<T> = vec![T::zero(); m_size];
v[0] = T::one();
for i in 1..m_size {
v[i] = self.qr[(k + 1 + i, k)];
}
for i in 0..n {
let mut dot = T::zero();
for j in 0..v.len() {
dot = dot + q[(i, k + 1 + j)] * v[j];
}
let scaled = tau * dot;
for j in 0..v.len() {
q[(i, k + 1 + j)] = q[(i, k + 1 + j)] - scaled * v[j].conj();
}
}
}
q
}
pub fn apply(
&self,
side: Side,
trans: Trans,
c: MatRef<'_, T>,
) -> Result<Mat<T>, ComplexHessenbergError> {
let n = self.n;
let m = c.nrows();
let k = c.ncols();
match side {
Side::Left => {
if m != n {
return Err(ComplexHessenbergError::DimensionMismatch);
}
}
Side::Right => {
if k != n {
return Err(ComplexHessenbergError::DimensionMismatch);
}
}
}
let mut result: Mat<T> = Mat::zeros(m, k);
for i in 0..m {
for j in 0..k {
result[(i, j)] = c[(i, j)];
}
}
if n <= 2 {
return Ok(result);
}
let apply_conj = matches!(trans, Trans::ConjTrans);
let forward = matches!(
(side, trans),
(Side::Left, Trans::ConjTrans) | (Side::Right, Trans::NoTrans)
);
let range: Vec<usize> = if forward {
(self.ilo..(self.ihi.saturating_sub(1))).collect()
} else {
(self.ilo..(self.ihi.saturating_sub(1))).rev().collect()
};
for &idx in &range {
let tau = if apply_conj {
self.tau[idx].conj()
} else {
self.tau[idx]
};
if tau.abs() < T::Real::epsilon() {
continue;
}
let m_size = self.ihi - idx - 1;
let mut v: Vec<T> = vec![T::zero(); m_size + 1];
v[0] = T::one();
for i in 1..=m_size {
v[i] = self.qr[(idx + 1 + i, idx)];
}
match side {
Side::Left => {
for j in 0..k {
let mut dot = T::zero();
for i in 0..v.len() {
dot = dot + v[i].conj() * result[(idx + 1 + i, j)];
}
let scaled = tau * dot;
for i in 0..v.len() {
result[(idx + 1 + i, j)] = result[(idx + 1 + i, j)] - scaled * v[i];
}
}
}
Side::Right => {
for i in 0..m {
let mut dot = T::zero();
for j in 0..v.len() {
dot = dot + result[(i, idx + 1 + j)] * v[j];
}
let scaled = tau * dot;
for j in 0..v.len() {
result[(i, idx + 1 + j)] =
result[(i, idx + 1 + j)] - scaled * v[j].conj();
}
}
}
}
}
Ok(result)
}
}
pub fn zgehrd<T: Field + ComplexScalar + bytemuck::Zeroable>(
a: MatRef<'_, T>,
) -> Result<ComplexHessenbergFactors<T>, ComplexHessenbergError>
where
T::Real: Real,
{
zgehrd_range(a, 0, a.nrows())
}
pub fn zgehrd_range<T: Field + ComplexScalar + bytemuck::Zeroable>(
a: MatRef<'_, T>,
ilo: usize,
ihi: usize,
) -> Result<ComplexHessenbergFactors<T>, ComplexHessenbergError>
where
T::Real: Real,
{
let n = a.nrows();
if n == 0 {
return Err(ComplexHessenbergError::EmptyMatrix);
}
if n != a.ncols() {
return Err(ComplexHessenbergError::NotSquare);
}
let mut qr: Mat<T> = Mat::zeros(n, n);
for i in 0..n {
for j in 0..n {
qr[(i, j)] = a[(i, j)];
}
}
let tau_len = if ihi > ilo + 1 { ihi - ilo - 1 } else { 0 };
let mut tau: Vec<T> = vec![T::zero(); tau_len];
for k in ilo..(ihi.saturating_sub(1)) {
let m_size = ihi - k - 1;
let mut x: Vec<T> = vec![T::zero(); m_size];
for i in 0..m_size {
x[i] = qr[(k + 1 + i, k)];
}
let (v, tau_k) = complex_householder_vector(&x);
tau[k - ilo] = tau_k;
if tau_k.abs() > T::Real::epsilon() {
for i in 1..v.len() {
qr[(k + 1 + i, k)] = v[i];
}
for j in (k + 1)..n {
let mut dot = T::zero();
dot = dot + qr[(k + 1, j)]; for i in 1..v.len() {
dot = dot + v[i].conj() * qr[(k + 1 + i, j)];
}
let scaled = tau_k * dot;
qr[(k + 1, j)] = qr[(k + 1, j)] - scaled;
for i in 1..v.len() {
qr[(k + 1 + i, j)] = qr[(k + 1 + i, j)] - scaled * v[i];
}
}
for i in 0..ihi {
let mut dot = T::zero();
dot = dot + qr[(i, k + 1)]; for j in 1..v.len() {
dot = dot + qr[(i, k + 1 + j)] * v[j];
}
let scaled = tau_k * dot;
qr[(i, k + 1)] = qr[(i, k + 1)] - scaled;
for j in 1..v.len() {
qr[(i, k + 1 + j)] = qr[(i, k + 1 + j)] - scaled * v[j].conj();
}
}
}
qr[(k + 1, k)] = T::from_real(-compute_householder_beta(&x));
}
Ok(ComplexHessenbergFactors {
qr,
tau,
n,
ilo,
ihi,
})
}
pub fn zunhhr<T: Field + ComplexScalar + bytemuck::Zeroable>(
factors: &ComplexHessenbergFactors<T>,
) -> Result<Mat<T>, ComplexHessenbergError>
where
T::Real: Real,
{
Ok(factors.q())
}
pub fn zunmhr<T: Field + ComplexScalar + bytemuck::Zeroable>(
factors: &ComplexHessenbergFactors<T>,
side: Side,
trans: Trans,
c: MatRef<'_, T>,
) -> Result<Mat<T>, ComplexHessenbergError>
where
T::Real: Real,
{
factors.apply(side, trans, c)
}
fn complex_householder_vector<T: Field + ComplexScalar>(x: &[T]) -> (Vec<T>, T)
where
T::Real: Real,
{
let n = x.len();
if n == 0 {
return (Vec::new(), T::zero());
}
let mut norm_sq = T::Real::zero();
for i in 0..n {
norm_sq = norm_sq + x[i].abs_sq();
}
let norm = <T::Real as Real>::sqrt(norm_sq);
if norm < T::Real::epsilon() {
return (vec![T::zero(); n], T::zero());
}
let x0 = x[0];
let x0_abs = x0.abs();
let alpha = if x0_abs > T::Real::epsilon() {
let sign = T::from_real_imag(x0.real() / x0_abs, x0.imag() / x0_abs);
T::zero() - sign * T::from_real(norm)
} else {
T::from_real(-norm)
};
let v0_unscaled = x0 - alpha;
if v0_unscaled.abs() < T::Real::epsilon() {
return (vec![T::zero(); n], T::zero());
}
let mut v: Vec<T> = vec![T::zero(); n];
v[0] = T::one();
for i in 1..n {
v[i] = x[i] / v0_unscaled;
}
let tau = (T::zero() - v0_unscaled) / alpha;
(v, tau)
}
fn compute_householder_beta<T: Field + ComplexScalar>(x: &[T]) -> T::Real
where
T::Real: Real,
{
let mut norm_sq = T::Real::zero();
for val in x {
norm_sq = norm_sq + val.abs_sq();
}
<T::Real as Real>::sqrt(norm_sq)
}
#[cfg(test)]
mod tests {
use super::*;
use num_complex::{Complex32, Complex64};
#[test]
fn test_complex_hessenberg_simple() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(1.0, 0.0),
Complex64::new(2.0, 1.0),
Complex64::new(3.0, 0.0),
],
&[
Complex64::new(4.0, -1.0),
Complex64::new(5.0, 0.0),
Complex64::new(6.0, 1.0),
],
&[
Complex64::new(7.0, 0.0),
Complex64::new(8.0, -1.0),
Complex64::new(9.0, 0.0),
],
]);
let hess = ComplexHessenberg::compute(a.as_ref()).expect("Should compute");
let h = hess.h();
let q = hess.q();
assert!(
h[(2, 0)].norm() < 1e-10,
"H[2,0] = {:?} should be zero",
h[(2, 0)]
);
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
);
}
}
}
#[test]
fn test_complex_hessenberg_reconstruction() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(4.0, 1.0), Complex64::new(1.0, 0.0)],
&[Complex64::new(1.0, 0.0), Complex64::new(3.0, -1.0)],
]);
let hess = ComplexHessenberg::compute(a.as_ref()).expect("Should compute");
let reconstructed = hess.reconstruct();
for i in 0..2 {
for j in 0..2 {
let diff = (reconstructed[(i, j)] - a[(i, j)]).norm();
assert!(
diff < 1e-10,
"Reconstruction[{},{}] = {:?}, A = {:?}",
i,
j,
reconstructed[(i, j)],
a[(i, j)]
);
}
}
}
#[test]
fn test_complex_hessenberg_4x4() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(4.0, 1.0),
Complex64::new(1.0, 0.0),
Complex64::new(-2.0, 1.0),
Complex64::new(2.0, 0.0),
],
&[
Complex64::new(1.0, 0.0),
Complex64::new(2.0, -1.0),
Complex64::new(0.0, 0.0),
Complex64::new(1.0, 1.0),
],
&[
Complex64::new(-2.0, -1.0),
Complex64::new(0.0, 0.0),
Complex64::new(3.0, 0.0),
Complex64::new(-2.0, 0.0),
],
&[
Complex64::new(2.0, 0.0),
Complex64::new(1.0, -1.0),
Complex64::new(-2.0, 0.0),
Complex64::new(-1.0, 1.0),
],
]);
let hess = ComplexHessenberg::compute(a.as_ref()).expect("Should compute");
let h = hess.h();
for j in 0..2 {
for i in (j + 2)..4 {
assert!(
h[(i, j)].norm() < 1e-10,
"H[{},{}] = {:?} should be zero",
i,
j,
h[(i, j)]
);
}
}
}
#[test]
fn test_complex_hessenberg_hermitian() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(4.0, 0.0),
Complex64::new(1.0, -1.0),
Complex64::new(0.0, 2.0),
],
&[
Complex64::new(1.0, 1.0),
Complex64::new(3.0, 0.0),
Complex64::new(1.0, 0.0),
],
&[
Complex64::new(0.0, -2.0),
Complex64::new(1.0, 0.0),
Complex64::new(2.0, 0.0),
],
]);
let hess = ComplexHessenberg::compute(a.as_ref()).expect("Should compute");
let h = hess.h();
assert!(h[(2, 0)].norm() < 1e-10);
let diff = (h[(1, 0)] - h[(0, 1)].conj()).norm();
assert!(
diff < 1e-10,
"H should be Hermitian tridiagonal: H[1,0]={:?}, H[0,1]={:?}",
h[(1, 0)],
h[(0, 1)]
);
}
#[test]
fn test_complex_hessenberg_f32() {
let a: Mat<Complex32> = Mat::from_rows(&[
&[
Complex32::new(1.0, 0.0),
Complex32::new(2.0, 1.0),
Complex32::new(3.0, 0.0),
],
&[
Complex32::new(4.0, -1.0),
Complex32::new(5.0, 0.0),
Complex32::new(6.0, 1.0),
],
&[
Complex32::new(7.0, 0.0),
Complex32::new(8.0, -1.0),
Complex32::new(9.0, 0.0),
],
]);
let hess = ComplexHessenberg::compute(a.as_ref()).expect("Should compute");
let h = hess.h();
assert!(h[(2, 0)].norm() < 1e-5);
}
#[test]
fn test_zgehrd() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(4.0, 0.0),
Complex64::new(1.0, 1.0),
Complex64::new(-2.0, 0.0),
],
&[
Complex64::new(1.0, -1.0),
Complex64::new(2.0, 0.0),
Complex64::new(0.0, 1.0),
],
&[
Complex64::new(-2.0, 0.0),
Complex64::new(0.0, -1.0),
Complex64::new(3.0, 0.0),
],
]);
let factors = zgehrd(a.as_ref()).expect("Should compute");
let h = factors.h();
assert!(h[(2, 0)].norm() < 1e-10);
}
#[test]
fn test_zunhhr() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[
Complex64::new(4.0, 0.0),
Complex64::new(1.0, 1.0),
Complex64::new(-2.0, 0.0),
],
&[
Complex64::new(1.0, -1.0),
Complex64::new(2.0, 0.0),
Complex64::new(0.0, 1.0),
],
&[
Complex64::new(-2.0, 0.0),
Complex64::new(0.0, -1.0),
Complex64::new(3.0, 0.0),
],
]);
let factors = zgehrd(a.as_ref()).expect("Should compute");
let q = zunhhr(&factors).expect("Should generate Q");
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 { 1.0 } else { 0.0 };
let diff = (sum.re - expected).abs() + sum.im.abs();
assert!(
diff < 1e-10,
"Q^H*Q[{},{}] = {:?}, expected {}",
i,
j,
sum,
expected
);
}
}
}
#[test]
fn test_complex_hessenberg_identity() {
let eye: Mat<Complex64> = Mat::eye(4);
let hess = ComplexHessenberg::compute(eye.as_ref()).expect("Should compute");
let h = hess.h();
let q = hess.q();
for i in 0..4 {
for j in 0..4 {
let expected = if i == j {
Complex64::new(1.0, 0.0)
} else {
Complex64::new(0.0, 0.0)
};
let diff = (h[(i, j)] - expected).norm();
assert!(diff < 1e-10, "H[{},{}] = {:?}", i, j, h[(i, j)]);
}
}
for i in 0..4 {
for j in 0..4 {
if i == j {
assert!(q[(i, j)].norm() > 0.99);
} else {
assert!(q[(i, j)].norm() < 1e-10);
}
}
}
}
#[test]
fn test_complex_hessenberg_single() {
let a: Mat<Complex64> = Mat::from_rows(&[&[Complex64::new(5.0, 2.0)]]);
let hess = ComplexHessenberg::compute(a.as_ref()).expect("Should compute");
let h = hess.h();
let q = hess.q();
assert_eq!(h[(0, 0)], Complex64::new(5.0, 2.0));
assert_eq!(q[(0, 0)], Complex64::new(1.0, 0.0));
}
}