use num_traits::{FromPrimitive, One, Zero};
use oxiblas_core::scalar::{ComplexScalar, Field, Real, Scalar};
use oxiblas_matrix::{Mat, MatRef};
use super::complex_hessenberg::ComplexHessenberg;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ComplexSchurError {
EmptyMatrix,
NotSquare,
NotConverged,
}
impl core::fmt::Display for ComplexSchurError {
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::NotConverged => write!(f, "Schur decomposition did not converge"),
}
}
}
impl std::error::Error for ComplexSchurError {}
#[derive(Debug, Clone)]
pub struct ComplexSchur<T: Scalar> {
q: Mat<T>,
t: Mat<T>,
eigenvalues: Vec<T>,
n: usize,
}
impl<T: Field + ComplexScalar + bytemuck::Zeroable> ComplexSchur<T>
where
T::Real: Real,
{
const MAX_ITERATIONS: usize = 100;
pub fn compute(a: MatRef<'_, T>) -> Result<Self, ComplexSchurError> {
let m = a.nrows();
let n = a.ncols();
if m == 0 || n == 0 {
return Err(ComplexSchurError::EmptyMatrix);
}
if m != n {
return Err(ComplexSchurError::NotSquare);
}
if n == 1 {
let mut t: Mat<T> = Mat::zeros(1, 1);
t[(0, 0)] = a[(0, 0)];
let mut q: Mat<T> = Mat::zeros(1, 1);
q[(0, 0)] = T::one();
let eigenvalues = vec![a[(0, 0)]];
return Ok(Self {
q,
t,
eigenvalues,
n,
});
}
if n == 2 {
return Self::compute_2x2(a);
}
let hess = ComplexHessenberg::compute(a).map_err(|_| ComplexSchurError::NotSquare)?;
let mut t: Mat<T> = Mat::zeros(n, n);
let h = hess.h();
for i in 0..n {
for j in 0..n {
t[(i, j)] = h[(i, j)];
}
}
let mut q: Mat<T> = Mat::zeros(n, n);
let q_hess = hess.q();
for i in 0..n {
for j in 0..n {
q[(i, j)] = q_hess[(i, j)];
}
}
let eps = <T::Real as Scalar>::epsilon();
let tol = eps * T::Real::from_f64(100.0).unwrap_or(T::Real::one());
let mut p = n;
let mut iter_count = 0;
let mut stagnation_count = 0;
while p > 1 && iter_count < Self::MAX_ITERATIONS * n {
iter_count += 1;
let mut q_idx = p - 1;
while q_idx > 0 {
let sub = t[(q_idx, q_idx - 1)].abs();
let diag_sum = t[(q_idx - 1, q_idx - 1)].abs() + t[(q_idx, q_idx)].abs();
if sub <= tol * (diag_sum + T::Real::one()) {
t[(q_idx, q_idx - 1)] = T::zero();
break;
}
q_idx -= 1;
}
if q_idx == p - 1 {
p -= 1;
stagnation_count = 0;
} else {
let shift = compute_wilkinson_shift(&t, p);
stagnation_count += 1;
let actual_shift = if stagnation_count > 10 {
stagnation_count = 0;
let mag = t[(p - 1, p - 2)].abs() + t[(p - 1, p - 1)].abs();
T::from_real(mag)
} else {
shift
};
complex_qr_step(&mut t, &mut q, q_idx, p, actual_shift);
}
}
if iter_count >= Self::MAX_ITERATIONS * n && p > 1 {
return Err(ComplexSchurError::NotConverged);
}
for i in 1..n {
let sub = t[(i, i - 1)].abs();
let diag_sum = t[(i - 1, i - 1)].abs() + t[(i, i)].abs();
if sub
<= tol
* (diag_sum + T::Real::one())
* T::Real::from_f64(10.0).unwrap_or(T::Real::one())
{
t[(i, i - 1)] = T::zero();
}
}
let mut eigenvalues = Vec::with_capacity(n);
for i in 0..n {
eigenvalues.push(t[(i, i)]);
}
Ok(Self {
q,
t,
eigenvalues,
n,
})
}
fn compute_2x2(a: MatRef<'_, T>) -> Result<Self, ComplexSchurError> {
let a00 = a[(0, 0)];
let a01 = a[(0, 1)];
let a10 = a[(1, 0)];
let a11 = a[(1, 1)];
let trace = a00 + a11;
let det = a00 * a11 - a01 * a10;
let two = T::one() + T::one();
let four = two + two;
let disc = trace * trace - four * det;
let disc_sqrt = complex_sqrt(disc);
let lambda1 = (trace + disc_sqrt) / two;
let _lambda2 = (trace - disc_sqrt) / two;
let eps = T::Real::epsilon();
if a10.abs() < eps {
let mut q: Mat<T> = Mat::zeros(2, 2);
q[(0, 0)] = T::one();
q[(1, 1)] = T::one();
let mut t: Mat<T> = Mat::zeros(2, 2);
for i in 0..2 {
for j in 0..2 {
t[(i, j)] = a[(i, j)];
}
}
let eigenvalues = vec![a00, a11];
return Ok(Self {
q,
t,
eigenvalues,
n: 2,
});
}
let d0 = a00 - lambda1;
let d1 = a11 - lambda1;
let row0_mag = d0.abs_sq() + a01.abs_sq();
let row1_mag = a10.abs_sq() + d1.abs_sq();
let (v0, v1) = if row0_mag >= row1_mag && a01.abs() > eps {
(T::one(), T::zero() - d0 / a01)
} else if a10.abs() > eps {
(T::zero() - d1 / a10, T::one())
} else {
(T::one(), T::zero())
};
let norm_sq = v0.abs_sq() + v1.abs_sq();
let norm = <T::Real as Real>::sqrt(norm_sq);
let u0 = if norm > eps {
(v0 / T::from_real(norm), v1 / T::from_real(norm))
} else {
(T::one(), T::zero())
};
let u1 = (T::zero() - u0.1.conj(), u0.0.conj());
let mut q: Mat<T> = Mat::zeros(2, 2);
q[(0, 0)] = u0.0;
q[(1, 0)] = u0.1;
q[(0, 1)] = u1.0;
q[(1, 1)] = u1.1;
let mut t: Mat<T> = Mat::zeros(2, 2);
let mut qha: Mat<T> = Mat::zeros(2, 2);
for i in 0..2 {
for j in 0..2 {
let mut sum = T::zero();
for k in 0..2 {
sum = sum + q[(k, i)].conj() * a[(k, j)];
}
qha[(i, j)] = sum;
}
}
for i in 0..2 {
for j in 0..2 {
let mut sum = T::zero();
for k in 0..2 {
sum = sum + qha[(i, k)] * q[(k, j)];
}
t[(i, j)] = sum;
}
}
let hundred: T::Real =
<T::Real as FromPrimitive>::from_f64(100.0).unwrap_or(<T::Real as One>::one());
if t[(1, 0)].abs() < eps * hundred {
t[(1, 0)] = T::zero();
}
let eigenvalues = vec![t[(0, 0)], t[(1, 1)]];
Ok(Self {
q,
t,
eigenvalues,
n: 2,
})
}
pub fn q(&self) -> MatRef<'_, T> {
self.q.as_ref()
}
pub fn t(&self) -> MatRef<'_, T> {
self.t.as_ref()
}
pub fn eigenvalues(&self) -> &[T] {
&self.eigenvalues
}
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 tqh: 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.t[(i, k)] * self.q[(j, k)].conj();
}
tqh[(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)] * tqh[(k, j)];
}
a[(i, j)] = sum;
}
}
a
}
pub fn residual(&self, a: MatRef<'_, T>) -> T::Real {
let n = self.n;
let reconstructed = self.reconstruct();
let mut diff_norm_sq = T::Real::zero();
let mut a_norm_sq = T::Real::zero();
for i in 0..n {
for j in 0..n {
let diff = a[(i, j)] - reconstructed[(i, j)];
diff_norm_sq = diff_norm_sq + diff.abs_sq();
a_norm_sq = a_norm_sq + a[(i, j)].abs_sq();
}
}
if a_norm_sq < T::Real::epsilon() {
return diff_norm_sq;
}
<T::Real as Real>::sqrt(diff_norm_sq / a_norm_sq)
}
}
fn complex_sqrt<T: Field + ComplexScalar>(z: T) -> T
where
T::Real: Real,
{
let (re, im) = (z.real(), z.imag());
let mag = <T::Real as Real>::sqrt(re * re + im * im);
if mag < T::Real::epsilon() {
return T::zero();
}
let two = T::Real::one() + T::Real::one();
let r = <T::Real as Real>::sqrt((mag + re) / two);
let i = if im >= T::Real::zero() {
<T::Real as Real>::sqrt((mag - re) / two)
} else {
-<T::Real as Real>::sqrt((mag - re) / two)
};
T::from_real_imag(r, i)
}
fn compute_wilkinson_shift<T: Field + ComplexScalar + bytemuck::Zeroable>(t: &Mat<T>, p: usize) -> T
where
T::Real: Real,
{
if p < 2 {
return t[(p - 1, p - 1)];
}
let h11 = t[(p - 2, p - 2)];
let h12 = t[(p - 2, p - 1)];
let h21 = t[(p - 1, p - 2)];
let h22 = t[(p - 1, p - 1)];
let trace = h11 + h22;
let det = h11 * h22 - h12 * h21;
let two = T::one() + T::one();
let four = two + two;
let disc = trace * trace - four * det;
let disc_sqrt = complex_sqrt(disc);
let half = T::one() / two;
let lambda1 = half * (trace + disc_sqrt);
let lambda2 = half * (trace - disc_sqrt);
let diff1 = lambda1 - h22;
let diff2 = lambda2 - h22;
if diff1.abs_sq() < diff2.abs_sq() {
lambda1
} else {
lambda2
}
}
fn complex_qr_step<T: Field + ComplexScalar + bytemuck::Zeroable>(
h: &mut Mat<T>,
q: &mut Mat<T>,
start: usize,
end: usize,
shift: T,
) where
T::Real: Real,
{
let n = h.nrows();
let x0 = h[(start, start)] - shift;
let x1 = h[(start + 1, start)];
let (c, s) = givens_rotation_complex(x0, x1);
apply_givens_left(h, c, s, start, start, n);
apply_givens_right(h, c, s, start, 0, end.min(start + 3));
apply_givens_right_to_q(q, c, s, start, n);
for k in start..(end - 2) {
let a = h[(k + 1, k)];
let b = h[(k + 2, k)];
let (c, s) = givens_rotation_complex(a, b);
apply_givens_left(h, c, s, k + 1, k, n);
let row_end = end.min(k + 5);
apply_givens_right(h, c, s, k + 1, 0, row_end);
apply_givens_right_to_q(q, c, s, k + 1, n);
h[(k + 2, k)] = T::zero();
}
}
#[inline]
fn apply_givens_left<T: Field + ComplexScalar>(
h: &mut Mat<T>,
c: T,
s: T,
row: usize,
col_start: usize,
col_end: usize,
) where
T::Real: Real,
{
for j in col_start..col_end {
let t1 = h[(row, j)];
let t2 = h[(row + 1, j)];
h[(row, j)] = c * t1 + s.conj() * t2;
h[(row + 1, j)] = (T::zero() - s) * t1 + c * t2;
}
}
#[inline]
fn apply_givens_right<T: Field + ComplexScalar>(
h: &mut Mat<T>,
c: T,
s: T,
col: usize,
row_start: usize,
row_end: usize,
) where
T::Real: Real,
{
for i in row_start..row_end {
let t1 = h[(i, col)];
let t2 = h[(i, col + 1)];
h[(i, col)] = t1 * c + t2 * s;
h[(i, col + 1)] = (T::zero() - t1) * s.conj() + t2 * c;
}
}
#[inline]
fn apply_givens_right_to_q<T: Field + ComplexScalar>(
q: &mut Mat<T>,
c: T,
s: T,
col: usize,
n: usize,
) where
T::Real: Real,
{
for i in 0..n {
let t1 = q[(i, col)];
let t2 = q[(i, col + 1)];
q[(i, col)] = t1 * c + t2 * s;
q[(i, col + 1)] = (T::zero() - t1) * s.conj() + t2 * c;
}
}
fn givens_rotation_complex<T: Field + ComplexScalar>(a: T, b: T) -> (T, T)
where
T::Real: Real,
{
let b_norm = b.abs();
if b_norm < T::Real::epsilon() {
return (T::one(), T::zero());
}
let a_norm = a.abs();
let r = <T::Real as Real>::sqrt(a_norm * a_norm + b_norm * b_norm);
if r < T::Real::epsilon() {
return (T::one(), T::zero());
}
let c = T::from_real(a_norm / r);
let s = if a_norm > T::Real::epsilon() {
a.conj() * b / T::from_real(a_norm * r)
} else {
b / T::from_real(r)
};
(c, s)
}
#[cfg(test)]
mod tests {
use super::*;
use num_complex::{Complex32, Complex64};
fn approx_eq_complex(a: Complex64, b: Complex64, tol: f64) -> bool {
(a - b).norm() < tol
}
#[test]
fn test_complex_schur_upper_triangular() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 0.0), Complex64::new(2.0, 1.0)],
&[Complex64::new(0.0, 0.0), Complex64::new(3.0, 0.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let t = schur.t();
let eigenvalues = schur.eigenvalues();
assert!(
t[(1, 0)].norm() < 1e-10,
"T[1,0] = {:?} should be zero",
t[(1, 0)]
);
let mut eigs: Vec<f64> = eigenvalues.iter().map(|e| e.re).collect();
eigs.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert!(
(eigs[0] - 1.0).abs() < 1e-10,
"eig[0] = {}, expected 1",
eigs[0]
);
assert!(
(eigs[1] - 3.0).abs() < 1e-10,
"eig[1] = {}, expected 3",
eigs[1]
);
}
#[test]
fn test_complex_schur_general() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 0.0), Complex64::new(2.0, 0.0)],
&[Complex64::new(3.0, 0.0), Complex64::new(4.0, 0.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let t = schur.t();
let q = schur.q();
assert!(
t[(1, 0)].norm() < 1e-10,
"T[1,0] = {:?} should be zero",
t[(1, 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 { 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_schur_reconstruction() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(4.0, 1.0), Complex64::new(1.0, -1.0)],
&[Complex64::new(1.0, 1.0), Complex64::new(3.0, 0.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let residual = schur.residual(a.as_ref());
assert!(
residual < 1e-10,
"Reconstruction residual = {} is too large",
residual
);
}
#[test]
fn test_complex_schur_3x3() {
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 schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let t = schur.t();
for j in 0..2 {
for i in (j + 1)..3 {
assert!(
t[(i, j)].norm() < 1e-10,
"T[{},{}] = {:?} should be zero",
i,
j,
t[(i, j)]
);
}
}
let residual = schur.residual(a.as_ref());
assert!(residual < 1e-10);
}
#[test]
fn test_complex_schur_trace_determinant() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(1.0, 2.0), Complex64::new(3.0, 1.0)],
&[Complex64::new(-1.0, 1.0), Complex64::new(4.0, -1.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let eigenvalues = schur.eigenvalues();
let sum = eigenvalues[0] + eigenvalues[1];
let trace = a[(0, 0)] + a[(1, 1)];
assert!(
approx_eq_complex(sum, trace, 1e-10),
"sum = {:?}, trace = {:?}",
sum,
trace
);
let prod = eigenvalues[0] * eigenvalues[1];
let det = a[(0, 0)] * a[(1, 1)] - a[(0, 1)] * a[(1, 0)];
assert!(
approx_eq_complex(prod, det, 1e-10),
"prod = {:?}, det = {:?}",
prod,
det
);
}
#[test]
fn test_complex_schur_hermitian() {
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(2.0, 0.0), Complex64::new(1.0, -1.0)],
&[Complex64::new(1.0, 1.0), Complex64::new(3.0, 0.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let eigenvalues = schur.eigenvalues();
for e in eigenvalues {
assert!(e.im.abs() < 1e-10, "Expected real eigenvalue, got {:?}", e);
}
}
#[test]
fn test_complex_schur_single() {
let a: Mat<Complex64> = Mat::from_rows(&[&[Complex64::new(5.0, 2.0)]]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let eigenvalues = schur.eigenvalues();
assert_eq!(eigenvalues.len(), 1);
assert!(approx_eq_complex(
eigenvalues[0],
Complex64::new(5.0, 2.0),
1e-10
));
}
#[test]
fn test_complex_schur_f32() {
let a: Mat<Complex32> = Mat::from_rows(&[
&[Complex32::new(1.0, 0.0), Complex32::new(2.0, 0.0)],
&[Complex32::new(3.0, 0.0), Complex32::new(4.0, 0.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let t = schur.t();
assert!(
t[(1, 0)].norm() < 1e-5,
"T[1,0] = {:?} should be zero",
t[(1, 0)]
);
}
#[test]
fn test_complex_schur_rotation() {
let theta = std::f64::consts::FRAC_PI_4;
let c = theta.cos();
let s = theta.sin();
let a: Mat<Complex64> = Mat::from_rows(&[
&[Complex64::new(c, 0.0), Complex64::new(-s, 0.0)],
&[Complex64::new(s, 0.0), Complex64::new(c, 0.0)],
]);
let schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let eigenvalues = schur.eigenvalues();
for e in eigenvalues {
let mag = e.norm();
assert!(
(mag - 1.0).abs() < 1e-10,
"Expected |eigenvalue| = 1, got {}",
mag
);
}
}
#[test]
fn test_complex_schur_identity() {
let eye: Mat<Complex64> = Mat::eye(3);
let schur = ComplexSchur::compute(eye.as_ref()).expect("Should compute");
let t = schur.t();
for i in 0..3 {
for j in 0..3 {
let expected = if i == j {
Complex64::new(1.0, 0.0)
} else {
Complex64::new(0.0, 0.0)
};
let diff = (t[(i, j)] - expected).norm();
assert!(diff < 1e-10, "T[{},{}] = {:?}", i, j, t[(i, j)]);
}
}
}
#[test]
fn test_complex_schur_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 schur = ComplexSchur::compute(a.as_ref()).expect("Should compute");
let t = schur.t();
for j in 0..3 {
for i in (j + 1)..4 {
assert!(
t[(i, j)].norm() < 1e-9,
"T[{},{}] = {:?} should be zero",
i,
j,
t[(i, j)]
);
}
}
let residual = schur.residual(a.as_ref());
assert!(residual < 1e-9, "Residual {} is too large", residual);
}
}