1use oxiblas_core::scalar::{Field, Real, Scalar};
8use oxiblas_matrix::{Mat, MatRef};
9
10use super::hessenberg::Hessenberg;
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum SchurError {
15 EmptyMatrix,
17 NotSquare,
19 NotConverged,
21}
22
23impl core::fmt::Display for SchurError {
24 fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
25 match self {
26 Self::EmptyMatrix => write!(f, "Matrix is empty"),
27 Self::NotSquare => write!(f, "Matrix must be square"),
28 Self::NotConverged => write!(f, "Schur decomposition did not converge"),
29 }
30 }
31}
32
33impl std::error::Error for SchurError {}
34
35#[derive(Debug, Clone, Copy, PartialEq)]
37pub struct Eigenvalue<T> {
38 pub real: T,
40 pub imag: T,
42}
43
44impl<T: Scalar> Eigenvalue<T> {
45 pub fn real_only(value: T) -> Self {
47 Self {
48 real: value,
49 imag: T::zero(),
50 }
51 }
52
53 pub fn complex(real: T, imag: T) -> Self {
55 Self { real, imag }
56 }
57
58 pub fn is_real(&self) -> bool {
60 self.imag == T::zero()
61 }
62}
63
64#[derive(Debug, Clone)]
70pub struct Schur<T: Scalar> {
71 q: Mat<T>,
73 t: Mat<T>,
75 eigenvalues: Vec<Eigenvalue<T>>,
77 n: usize,
79}
80
81impl<T: Field + Real + bytemuck::Zeroable> Schur<T> {
82 const MAX_ITERATIONS: usize = 100;
84
85 pub fn compute(a: MatRef<'_, T>) -> Result<Self, SchurError> {
105 let m = a.nrows();
106 let n = a.ncols();
107
108 if m == 0 || n == 0 {
109 return Err(SchurError::EmptyMatrix);
110 }
111
112 if m != n {
113 return Err(SchurError::NotSquare);
114 }
115
116 if n == 1 {
118 let mut t = Mat::zeros(1, 1);
119 t[(0, 0)] = a[(0, 0)];
120 let mut q = Mat::zeros(1, 1);
121 q[(0, 0)] = T::one();
122 let eigenvalues = vec![Eigenvalue::real_only(a[(0, 0)])];
123 return Ok(Self {
124 q,
125 t,
126 eigenvalues,
127 n,
128 });
129 }
130
131 if n == 2 {
133 return Self::compute_2x2(a);
134 }
135
136 Self::compute_hessenberg_qr(a, Self::MAX_ITERATIONS.saturating_mul(n))
141 }
142
143 fn compute_hessenberg_qr(
165 a: MatRef<'_, T>,
166 max_total_iterations: usize,
167 ) -> Result<Self, SchurError> {
168 let n = a.ncols();
169
170 let hess = Hessenberg::compute(a).map_err(|_| SchurError::NotSquare)?;
172 let mut t = Mat::zeros(n, n);
173 let h = hess.h();
174 for i in 0..n {
175 for j in 0..n {
176 t[(i, j)] = h[(i, j)];
177 }
178 }
179
180 let mut q = Mat::zeros(n, n);
181 let q_hess = hess.q();
182 for i in 0..n {
183 for j in 0..n {
184 q[(i, j)] = q_hess[(i, j)];
185 }
186 }
187
188 let eps = <T as Scalar>::epsilon();
190 let tol = eps * T::from_f64(100.0).unwrap_or(T::one());
191
192 let mut p = n;
194 let mut iter_count = 0;
195
196 while p > 2 && iter_count < max_total_iterations {
197 iter_count += 1;
198
199 let mut q_idx = p - 1;
201 while q_idx > 0 {
202 let sub = Scalar::abs(t[(q_idx, q_idx - 1)]);
203 let diag_sum =
204 Scalar::abs(t[(q_idx - 1, q_idx - 1)]) + Scalar::abs(t[(q_idx, q_idx)]);
205 if sub <= tol * diag_sum {
206 t[(q_idx, q_idx - 1)] = T::zero();
207 break;
208 }
209 q_idx -= 1;
210 }
211
212 if q_idx == p - 1 {
213 p -= 1;
215 } else if q_idx == p - 2 {
216 let a11 = t[(p - 2, p - 2)];
218 let a12 = t[(p - 2, p - 1)];
219 let a21 = t[(p - 1, p - 2)];
220 let a22 = t[(p - 1, p - 1)];
221 let trace = a11 + a22;
222 let det = a11 * a22 - a12 * a21;
223 let disc = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
224
225 if disc < T::zero() {
226 p -= 2;
228 } else {
229 Self::francis_qr_step(&mut t, &mut q, q_idx, p);
231 }
232 } else {
233 Self::francis_qr_step(&mut t, &mut q, q_idx, p);
235 }
236 }
237
238 if p > 2 && !Self::active_block_converged(&t, p, tol) {
243 return Err(SchurError::NotConverged);
244 }
245
246 if p == 2 {
248 let sub = Scalar::abs(t[(1, 0)]);
249 let diag_sum = Scalar::abs(t[(0, 0)]) + Scalar::abs(t[(1, 1)]);
250 if sub <= tol * diag_sum {
251 t[(1, 0)] = T::zero();
252 }
253 }
254
255 let eigenvalues = Self::extract_eigenvalues(&t);
257
258 Ok(Self {
259 q,
260 t,
261 eigenvalues,
262 n,
263 })
264 }
265
266 fn active_block_converged(t: &Mat<T>, p: usize, tol: T) -> bool {
278 let negligible = |k: usize| -> bool {
281 let sub = Scalar::abs(t[(k, k - 1)]);
282 let diag_sum = Scalar::abs(t[(k - 1, k - 1)]) + Scalar::abs(t[(k, k)]);
283 sub <= tol * diag_sum
284 };
285
286 for k in 1..p.saturating_sub(1) {
288 if !negligible(k) && !negligible(k + 1) {
289 return false;
290 }
291 }
292 true
293 }
294
295 #[cfg(test)]
300 pub(crate) fn compute_with_iteration_budget(
301 a: MatRef<'_, T>,
302 max_total_iterations: usize,
303 ) -> Result<Self, SchurError> {
304 let m = a.nrows();
305 let n = a.ncols();
306
307 if m == 0 || n == 0 {
308 return Err(SchurError::EmptyMatrix);
309 }
310 if m != n {
311 return Err(SchurError::NotSquare);
312 }
313 if n <= 2 {
314 return Self::compute(a);
316 }
317 Self::compute_hessenberg_qr(a, max_total_iterations)
318 }
319
320 fn compute_2x2(a: MatRef<'_, T>) -> Result<Self, SchurError> {
322 let mut t = Mat::zeros(2, 2);
323 for i in 0..2 {
324 for j in 0..2 {
325 t[(i, j)] = a[(i, j)];
326 }
327 }
328
329 let a11 = a[(0, 0)];
330 let a12 = a[(0, 1)];
331 let a21 = a[(1, 0)];
332 let a22 = a[(1, 1)];
333
334 let trace = a11 + a22;
335 let det = a11 * a22 - a12 * a21;
336 let disc = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
337
338 let mut q = Mat::zeros(2, 2);
339 let eigenvalues: Vec<Eigenvalue<T>>;
340
341 if disc >= T::zero() {
342 let sqrt_disc = Real::sqrt(disc);
344 let lambda1 = (trace + sqrt_disc) / T::from_f64(2.0).unwrap_or_else(T::zero);
345 let lambda2 = (trace - sqrt_disc) / T::from_f64(2.0).unwrap_or_else(T::zero);
346
347 if Scalar::abs(a21) > <T as Scalar>::epsilon() {
349 let theta = if Scalar::abs(a11 - lambda1) > <T as Scalar>::epsilon() {
350 Real::atan2(a21, a11 - lambda1)
351 } else {
352 T::from_f64(core::f64::consts::FRAC_PI_2).unwrap_or_else(T::zero)
353 };
354 let c = Real::cos(theta);
355 let s = Real::sin(theta);
356
357 q[(0, 0)] = c;
358 q[(0, 1)] = -s;
359 q[(1, 0)] = s;
360 q[(1, 1)] = c;
361
362 let mut temp = Mat::zeros(2, 2);
364 for i in 0..2 {
366 for j in 0..2 {
367 let mut sum = T::zero();
368 for k in 0..2 {
369 sum = sum + q[(k, i)] * a[(k, j)];
370 }
371 temp[(i, j)] = sum;
372 }
373 }
374 for i in 0..2 {
376 for j in 0..2 {
377 let mut sum = T::zero();
378 for k in 0..2 {
379 sum = sum + temp[(i, k)] * q[(k, j)];
380 }
381 t[(i, j)] = sum;
382 }
383 }
384 } else {
385 q[(0, 0)] = T::one();
387 q[(1, 1)] = T::one();
388 }
389
390 eigenvalues = vec![
391 Eigenvalue::real_only(lambda1),
392 Eigenvalue::real_only(lambda2),
393 ];
394 } else {
395 let sqrt_disc = Real::sqrt(-disc);
397 let real_part = trace / T::from_f64(2.0).unwrap_or_else(T::zero);
398 let imag_part = sqrt_disc / T::from_f64(2.0).unwrap_or_else(T::zero);
399
400 q[(0, 0)] = T::one();
401 q[(1, 1)] = T::one();
402
403 eigenvalues = vec![
404 Eigenvalue::complex(real_part, imag_part),
405 Eigenvalue::complex(real_part, -imag_part),
406 ];
407 }
408
409 Ok(Self {
410 q,
411 t,
412 eigenvalues,
413 n: 2,
414 })
415 }
416
417 fn francis_qr_step(t: &mut Mat<T>, q: &mut Mat<T>, start: usize, end: usize) {
419 let n = t.nrows();
420
421 if end - start < 2 {
422 return;
423 }
424
425 let h11 = t[(end - 2, end - 2)];
427 let h12 = t[(end - 2, end - 1)];
428 let h21 = t[(end - 1, end - 2)];
429 let h22 = t[(end - 1, end - 1)];
430
431 let s = h11 + h22; let p = h11 * h22 - h12 * h21; let h_00 = t[(start, start)];
436 let h_01 = t[(start, start + 1)];
437 let h_10 = t[(start + 1, start)];
438
439 let mut x = h_00 * h_00 + h_01 * h_10 - s * h_00 + p;
440 let mut y = h_10 * (h_00 + t[(start + 1, start + 1)] - s);
441 let mut z = if start + 2 < end {
442 h_10 * t[(start + 2, start + 1)]
443 } else {
444 T::zero()
445 };
446
447 for k in start..end.saturating_sub(1) {
449 let (v, tau) = householder_3(&[x, y, z]);
451
452 if tau != T::zero() {
453 let r = if k > start { k - 1 } else { k };
454
455 let col_start = r;
457 let col_end = n;
458 for j in col_start..col_end {
459 let rows = (k..(k + 3).min(end)).collect::<Vec<_>>();
460 let mut dot = T::zero();
461 for (vi, &row) in rows.iter().enumerate() {
462 dot = dot + v[vi] * t[(row, j)];
463 }
464 let scaled = tau * dot;
465 for (vi, &row) in rows.iter().enumerate() {
466 t[(row, j)] = t[(row, j)] - scaled * v[vi];
467 }
468 }
469
470 let row_end = (k + 4).min(end);
472 for i in 0..row_end {
473 let cols = (k..(k + 3).min(end)).collect::<Vec<_>>();
474 let mut dot = T::zero();
475 for (vi, &col) in cols.iter().enumerate() {
476 dot = dot + t[(i, col)] * v[vi];
477 }
478 let scaled = tau * dot;
479 for (vi, &col) in cols.iter().enumerate() {
480 t[(i, col)] = t[(i, col)] - scaled * v[vi];
481 }
482 }
483
484 for i in 0..n {
486 let cols = (k..(k + 3).min(end)).collect::<Vec<_>>();
487 let mut dot = T::zero();
488 for (vi, &col) in cols.iter().enumerate() {
489 dot = dot + q[(i, col)] * v[vi];
490 }
491 let scaled = tau * dot;
492 for (vi, &col) in cols.iter().enumerate() {
493 q[(i, col)] = q[(i, col)] - scaled * v[vi];
494 }
495 }
496 }
497
498 if k + 3 < end {
500 x = t[(k + 1, k)];
501 y = t[(k + 2, k)];
502 z = if k + 3 < end {
503 t[(k + 3, k)]
504 } else {
505 T::zero()
506 };
507 } else if k + 2 < end {
508 x = t[(k + 1, k)];
510 y = t[(k + 2, k)];
511 z = T::zero();
512 }
513 }
514 }
515
516 fn extract_eigenvalues(t: &Mat<T>) -> Vec<Eigenvalue<T>> {
518 let n = t.nrows();
519 let mut eigenvalues = Vec::with_capacity(n);
520 let eps = <T as Scalar>::epsilon() * T::from_f64(100.0).unwrap_or(T::one());
521
522 let mut i = 0;
523 while i < n {
524 if i == n - 1 {
525 eigenvalues.push(Eigenvalue::real_only(t[(i, i)]));
527 i += 1;
528 } else {
529 let sub = Scalar::abs(t[(i + 1, i)]);
531 let diag_sum = Scalar::abs(t[(i, i)]) + Scalar::abs(t[(i + 1, i + 1)]);
532
533 if sub <= eps * diag_sum {
534 eigenvalues.push(Eigenvalue::real_only(t[(i, i)]));
536 i += 1;
537 } else {
538 let a11 = t[(i, i)];
540 let a12 = t[(i, i + 1)];
541 let a21 = t[(i + 1, i)];
542 let a22 = t[(i + 1, i + 1)];
543
544 let trace = a11 + a22;
545 let det = a11 * a22 - a12 * a21;
546 let disc = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
547
548 if disc >= T::zero() {
549 let sqrt_disc = Real::sqrt(disc);
551 let lambda1 =
552 (trace + sqrt_disc) / T::from_f64(2.0).unwrap_or_else(T::zero);
553 let lambda2 =
554 (trace - sqrt_disc) / T::from_f64(2.0).unwrap_or_else(T::zero);
555 eigenvalues.push(Eigenvalue::real_only(lambda1));
556 eigenvalues.push(Eigenvalue::real_only(lambda2));
557 } else {
558 let sqrt_disc = Real::sqrt(-disc);
560 let real_part = trace / T::from_f64(2.0).unwrap_or_else(T::zero);
561 let imag_part = sqrt_disc / T::from_f64(2.0).unwrap_or_else(T::zero);
562 eigenvalues.push(Eigenvalue::complex(real_part, imag_part));
563 eigenvalues.push(Eigenvalue::complex(real_part, -imag_part));
564 }
565 i += 2;
566 }
567 }
568 }
569
570 eigenvalues
571 }
572
573 pub fn q(&self) -> MatRef<'_, T> {
575 self.q.as_ref()
576 }
577
578 pub fn t(&self) -> MatRef<'_, T> {
580 self.t.as_ref()
581 }
582
583 pub fn eigenvalues(&self) -> &[Eigenvalue<T>] {
585 &self.eigenvalues
586 }
587
588 pub fn eigenvalues_real(&self) -> Vec<T> {
590 self.eigenvalues.iter().map(|e| e.real).collect()
591 }
592
593 pub fn reconstruct(&self) -> Mat<T> {
595 let mut qt = Mat::zeros(self.n, self.n);
596 let mut a = Mat::zeros(self.n, self.n);
597
598 for i in 0..self.n {
600 for j in 0..self.n {
601 let mut sum = T::zero();
602 for k in 0..self.n {
603 sum = sum + self.q[(i, k)] * self.t[(k, j)];
604 }
605 qt[(i, j)] = sum;
606 }
607 }
608
609 for i in 0..self.n {
611 for j in 0..self.n {
612 let mut sum = T::zero();
613 for k in 0..self.n {
614 sum = sum + qt[(i, k)] * self.q[(j, k)];
615 }
616 a[(i, j)] = sum;
617 }
618 }
619
620 a
621 }
622
623 #[must_use]
649 pub fn right_eigenvectors(&self) -> Mat<T> {
650 trevc_right(&self.t)
651 }
652
653 #[must_use]
674 pub fn left_eigenvectors(&self) -> Mat<T> {
675 trevc_left(&self.t)
676 }
677
678 #[must_use]
686 pub fn eigenvectors(&self) -> (Mat<T>, Mat<T>) {
687 let vr_t = trevc_right(&self.t);
688 let vl_t = trevc_left(&self.t);
689
690 let mut vr_a = Mat::zeros(self.n, self.n);
692 let mut vl_a = Mat::zeros(self.n, self.n);
693
694 for i in 0..self.n {
695 for j in 0..self.n {
696 let mut sum_r = T::zero();
697 let mut sum_l = T::zero();
698 for k in 0..self.n {
699 sum_r = sum_r + self.q[(i, k)] * vr_t[(k, j)];
700 sum_l = sum_l + self.q[(i, k)] * vl_t[(k, j)];
701 }
702 vr_a[(i, j)] = sum_r;
703 vl_a[(i, j)] = sum_l;
704 }
705 }
706
707 (vr_a, vl_a)
708 }
709
710 #[must_use]
744 pub fn eigenvalue_condition_numbers(&self) -> Vec<T> {
745 trsna_s(&self.t)
746 }
747
748 #[must_use]
758 pub fn eigenvector_separation(&self) -> Vec<T> {
759 trsna_sep(&self.t)
760 }
761}
762
763pub fn trsna_s<T: Field + Real + bytemuck::Zeroable>(t: &Mat<T>) -> Vec<T> {
776 let n = t.nrows();
777 if n == 0 {
778 return Vec::new();
779 }
780
781 let vr = trevc_right(t);
782 let vl = trevc_left(t);
783
784 let mut s = vec![T::zero(); n];
785 let eps = <T as Scalar>::epsilon();
786
787 let mut j = 0;
788 while j < n {
789 let is_2x2 = if j + 1 < n {
791 let sub = Scalar::abs(t[(j + 1, j)]);
792 let diag_sum = Scalar::abs(t[(j, j)]) + Scalar::abs(t[(j + 1, j + 1)]);
793 sub > eps * T::from_f64(100.0).unwrap_or_else(T::zero) * diag_sum
794 } else {
795 false
796 };
797
798 if is_2x2 {
799 let jp1 = j + 1;
803
804 let mut prod_rr = T::zero(); let mut prod_ii = T::zero(); let mut prod_ri = T::zero(); let mut prod_ir = T::zero(); for k in 0..n {
810 prod_rr = prod_rr + vl[(k, j)] * vr[(k, j)];
811 prod_ii = prod_ii + vl[(k, jp1)] * vr[(k, jp1)];
812 prod_ri = prod_ri + vl[(k, j)] * vr[(k, jp1)];
813 prod_ir = prod_ir + vl[(k, jp1)] * vr[(k, j)];
814 }
815
816 let real_part = prod_rr + prod_ii;
817 let imag_part = prod_ri - prod_ir;
818 let abs_inner = Real::sqrt(real_part * real_part + imag_part * imag_part);
819
820 s[j] = abs_inner;
822 s[jp1] = abs_inner;
823
824 j += 2;
825 } else {
826 let mut inner = T::zero();
828 for k in 0..n {
829 inner = inner + vl[(k, j)] * vr[(k, j)];
830 }
831 s[j] = Scalar::abs(inner);
832 j += 1;
833 }
834 }
835
836 s
837}
838
839pub fn trsna_sep<T: Field + Real + bytemuck::Zeroable>(t: &Mat<T>) -> Vec<T> {
854 let n = t.nrows();
855 if n == 0 {
856 return Vec::new();
857 }
858
859 let mut sep = vec![T::zero(); n];
860 let eps = <T as Scalar>::epsilon();
861
862 let mut j = 0;
863 while j < n {
864 let is_2x2 = if j + 1 < n {
866 let sub = Scalar::abs(t[(j + 1, j)]);
867 let diag_sum = Scalar::abs(t[(j, j)]) + Scalar::abs(t[(j + 1, j + 1)]);
868 sub > eps * T::from_f64(100.0).unwrap_or_else(T::zero) * diag_sum
869 } else {
870 false
871 };
872
873 if is_2x2 {
874 let jp1 = j + 1;
876
877 let a11 = t[(j, j)];
879 let a22 = t[(jp1, jp1)];
880 let lambda_real = (a11 + a22) / T::from_f64(2.0).unwrap_or_else(T::zero);
881 let trace = a11 + a22;
882 let det = a11 * a22 - t[(j, jp1)] * t[(jp1, j)];
883 let disc = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
884 let lambda_imag = if disc < T::zero() {
885 Real::sqrt(-disc) / T::from_f64(2.0).unwrap_or_else(T::zero)
886 } else {
887 T::zero()
888 };
889
890 let mut min_sep = T::one() / eps;
891
892 let mut k = 0;
894 while k < n {
895 if k == j || k == jp1 {
896 k += 1;
897 continue;
898 }
899
900 let adjacent = j > 0 && k == j - 1;
902 let k_is_2x2 = if k + 1 < n && !adjacent {
903 let sub = Scalar::abs(t[(k + 1, k)]);
904 let diag_sum = Scalar::abs(t[(k, k)]) + Scalar::abs(t[(k + 1, k + 1)]);
905 sub > eps * T::from_f64(100.0).unwrap_or_else(T::zero) * diag_sum
906 } else {
907 false
908 };
909
910 let (other_real, other_imag) = if k_is_2x2 {
911 let kp1 = k + 1;
912 let b11 = t[(k, k)];
913 let b22 = t[(kp1, kp1)];
914 let other_trace = b11 + b22;
915 let other_det = b11 * b22 - t[(k, kp1)] * t[(kp1, k)];
916 let other_disc = other_trace * other_trace
917 - T::from_f64(4.0).unwrap_or_else(T::zero) * other_det;
918 let r = (b11 + b22) / T::from_f64(2.0).unwrap_or_else(T::zero);
919 let i = if other_disc < T::zero() {
920 Real::sqrt(-other_disc) / T::from_f64(2.0).unwrap_or_else(T::zero)
921 } else {
922 T::zero()
923 };
924 (r, i)
925 } else {
926 (t[(k, k)], T::zero())
927 };
928
929 let dr = lambda_real - other_real;
931 let di = lambda_imag - other_imag;
932 let dist = Real::sqrt(dr * dr + di * di);
933 if dist < min_sep && dist > T::zero() {
934 min_sep = dist;
935 }
936
937 if other_imag != T::zero() {
939 let di_conj = lambda_imag + other_imag;
940 let dist_conj = Real::sqrt(dr * dr + di_conj * di_conj);
941 if dist_conj < min_sep && dist_conj > T::zero() {
942 min_sep = dist_conj;
943 }
944 }
945
946 if k_is_2x2 {
947 k += 2;
948 } else {
949 k += 1;
950 }
951 }
952
953 sep[j] = min_sep;
954 sep[jp1] = min_sep;
955 j += 2;
956 } else {
957 let lambda = t[(j, j)];
959 let mut min_sep = T::one() / eps;
960
961 let mut k = 0;
963 while k < n {
964 if k == j {
965 k += 1;
966 continue;
967 }
968
969 let adjacent_to_j = (j > 0 && k == j - 1) || k == j + 1;
971 let k_is_2x2 = if k + 1 < n && !adjacent_to_j {
972 let sub = Scalar::abs(t[(k + 1, k)]);
973 let diag_sum = Scalar::abs(t[(k, k)]) + Scalar::abs(t[(k + 1, k + 1)]);
974 sub > eps * T::from_f64(100.0).unwrap_or_else(T::zero) * diag_sum
975 } else {
976 false
977 };
978
979 let (other_real, other_imag) = if k_is_2x2 {
980 let kp1 = k + 1;
981 let b11 = t[(k, k)];
982 let b22 = t[(kp1, kp1)];
983 let other_trace = b11 + b22;
984 let other_det = b11 * b22 - t[(k, kp1)] * t[(kp1, k)];
985 let other_disc = other_trace * other_trace
986 - T::from_f64(4.0).unwrap_or_else(T::zero) * other_det;
987 let r = (b11 + b22) / T::from_f64(2.0).unwrap_or_else(T::zero);
988 let i = if other_disc < T::zero() {
989 Real::sqrt(-other_disc) / T::from_f64(2.0).unwrap_or_else(T::zero)
990 } else {
991 T::zero()
992 };
993 (r, i)
994 } else {
995 (t[(k, k)], T::zero())
996 };
997
998 let dr = lambda - other_real;
1000 let dist = Real::sqrt(dr * dr + other_imag * other_imag);
1001 if dist < min_sep && dist > T::zero() {
1002 min_sep = dist;
1003 }
1004
1005 if k_is_2x2 {
1006 k += 2;
1007 } else {
1008 k += 1;
1009 }
1010 }
1011
1012 sep[j] = min_sep;
1013 j += 1;
1014 }
1015 }
1016
1017 sep
1018}
1019
1020pub fn trevc_right<T: Field + Real + bytemuck::Zeroable>(t: &Mat<T>) -> Mat<T> {
1034 let n = t.nrows();
1035 let mut v = Mat::zeros(n, n);
1036
1037 for i in 0..n {
1039 v[(i, i)] = T::one();
1040 }
1041
1042 let eps = <T as Scalar>::epsilon() * T::from_f64(100.0).unwrap_or(T::one());
1043
1044 let mut j = n;
1046 while j > 0 {
1047 j -= 1;
1048
1049 let is_2x2 = if j > 0 {
1051 let sub = Scalar::abs(t[(j, j - 1)]);
1052 let diag_sum = Scalar::abs(t[(j - 1, j - 1)]) + Scalar::abs(t[(j, j)]);
1053 sub > eps * diag_sum
1054 } else {
1055 false
1056 };
1057
1058 if is_2x2 {
1059 let jm1 = j - 1;
1062
1063 let a11 = t[(jm1, jm1)];
1065 let a12 = t[(jm1, j)];
1066 let a21 = t[(j, jm1)];
1067 let a22 = t[(j, j)];
1068
1069 let trace = a11 + a22;
1070 let det = a11 * a22 - a12 * a21;
1071 let disc = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
1072
1073 let two = T::from_f64(2.0).unwrap_or_else(T::zero);
1075 let real_part = trace / two;
1076 let imag_part = Real::sqrt(-disc) / two;
1077
1078 v[(jm1, jm1)] = T::one();
1081 v[(j, jm1)] = T::zero();
1082 v[(jm1, j)] = T::zero();
1083 v[(j, j)] = T::one();
1084
1085 let d11 = a11 - real_part;
1110 let d22 = a22 - real_part;
1111
1112 if Scalar::abs(a21) >= Scalar::abs(a12) && Scalar::abs(a21) > eps {
1114 let v1r = -d22 / a21;
1117 let v1i = imag_part / a21;
1118 v[(jm1, jm1)] = v1r;
1119 v[(j, jm1)] = T::one();
1120 v[(jm1, j)] = v1i;
1121 v[(j, j)] = T::zero();
1122 } else if Scalar::abs(a12) > eps {
1123 let v2r = -d11 / a12;
1126 let v2i = -imag_part / a12;
1127 v[(jm1, jm1)] = T::one();
1128 v[(j, jm1)] = v2r;
1129 v[(jm1, j)] = T::zero();
1130 v[(j, j)] = v2i;
1131 } else {
1132 v[(jm1, jm1)] = T::one();
1134 v[(j, jm1)] = T::zero();
1135 v[(jm1, j)] = T::zero();
1136 v[(j, j)] = T::one();
1137 }
1138
1139 for i in (0..jm1).rev() {
1141 let mut sum_r = T::zero();
1146 let mut sum_i = T::zero();
1147 for k in (i + 1)..=j {
1148 sum_r = sum_r + t[(i, k)] * v[(k, jm1)];
1149 sum_i = sum_i + t[(i, k)] * v[(k, j)];
1150 }
1151
1152 let d = t[(i, i)] - real_part;
1153 let det_2x2 = d * d + imag_part * imag_part;
1154
1155 if Scalar::abs(det_2x2) > eps {
1156 v[(i, jm1)] = (-d * sum_r - imag_part * sum_i) / det_2x2;
1160 v[(i, j)] = (imag_part * sum_r - d * sum_i) / det_2x2;
1161 }
1162 }
1163
1164 let mut norm_r_sq = T::zero();
1166 let mut norm_i_sq = T::zero();
1167 for i in 0..n {
1168 norm_r_sq = norm_r_sq + v[(i, jm1)] * v[(i, jm1)];
1169 norm_i_sq = norm_i_sq + v[(i, j)] * v[(i, j)];
1170 }
1171 let norm = Real::sqrt(norm_r_sq + norm_i_sq);
1172 if norm > T::zero() {
1173 for i in 0..n {
1174 v[(i, jm1)] = v[(i, jm1)] / norm;
1175 v[(i, j)] = v[(i, j)] / norm;
1176 }
1177 }
1178
1179 j = jm1; } else {
1181 let lambda = t[(j, j)];
1183
1184 v[(j, j)] = T::one();
1186
1187 for i in (0..j).rev() {
1189 let mut sum = T::zero();
1190 for k in (i + 1)..=j {
1191 sum = sum + t[(i, k)] * v[(k, j)];
1192 }
1193
1194 let d = t[(i, i)] - lambda;
1195 if Scalar::abs(d) > eps {
1196 v[(i, j)] = -sum / d;
1197 } else {
1198 v[(i, j)] = -sum / eps;
1200 }
1201 }
1202
1203 let mut norm_sq = T::zero();
1205 for i in 0..n {
1206 norm_sq = norm_sq + v[(i, j)] * v[(i, j)];
1207 }
1208 let norm = Real::sqrt(norm_sq);
1209 if norm > T::zero() {
1210 for i in 0..n {
1211 v[(i, j)] = v[(i, j)] / norm;
1212 }
1213 }
1214 }
1215 }
1216
1217 v
1218}
1219
1220pub fn trevc_left<T: Field + Real + bytemuck::Zeroable>(t: &Mat<T>) -> Mat<T> {
1233 let n = t.nrows();
1234 let mut v = Mat::zeros(n, n);
1235
1236 for i in 0..n {
1238 v[(i, i)] = T::one();
1239 }
1240
1241 let eps = <T as Scalar>::epsilon() * T::from_f64(100.0).unwrap_or(T::one());
1242
1243 let mut j = 0;
1245 while j < n {
1246 let is_2x2 = if j + 1 < n {
1248 let sub = Scalar::abs(t[(j + 1, j)]);
1249 let diag_sum = Scalar::abs(t[(j, j)]) + Scalar::abs(t[(j + 1, j + 1)]);
1250 sub > eps * diag_sum
1251 } else {
1252 false
1253 };
1254
1255 if is_2x2 {
1256 let jp1 = j + 1;
1258
1259 let a11 = t[(j, j)];
1260 let a12 = t[(j, jp1)];
1261 let a21 = t[(jp1, j)];
1262 let a22 = t[(jp1, jp1)];
1263
1264 let trace = a11 + a22;
1265 let det = a11 * a22 - a12 * a21;
1266 let disc = trace * trace - T::from_f64(4.0).unwrap_or_else(T::zero) * det;
1267
1268 let two = T::from_f64(2.0).unwrap_or_else(T::zero);
1269 let real_part = trace / two;
1270 let imag_part = Real::sqrt(-disc) / two;
1271
1272 let d11 = a11 - real_part;
1274 let det_factor = d11 * d11 + imag_part * imag_part;
1275
1276 if Scalar::abs(det_factor) > eps {
1277 let vr2 = -a21 * d11 / det_factor;
1278 let vi2 = -imag_part * a21 / det_factor;
1279 v[(j, j)] = T::one();
1280 v[(jp1, j)] = vr2;
1281 v[(j, jp1)] = T::zero();
1282 v[(jp1, jp1)] = vi2;
1283 } else {
1284 v[(j, j)] = T::one();
1285 v[(jp1, j)] = T::zero();
1286 v[(j, jp1)] = T::zero();
1287 v[(jp1, jp1)] = T::one();
1288 }
1289
1290 for i in (jp1 + 1)..n {
1292 let mut sum_r = T::zero();
1293 let mut sum_i = T::zero();
1294 for k in j..i {
1295 sum_r = sum_r + t[(k, i)] * v[(k, j)];
1296 sum_i = sum_i + t[(k, i)] * v[(k, jp1)];
1297 }
1298
1299 let d = t[(i, i)] - real_part;
1300 let det_2x2 = d * d + imag_part * imag_part;
1301
1302 if Scalar::abs(det_2x2) > eps {
1303 v[(i, j)] = (-d * sum_r - imag_part * sum_i) / det_2x2;
1304 v[(i, jp1)] = (imag_part * sum_r - d * sum_i) / det_2x2;
1305 }
1306 }
1307
1308 let mut norm_sq = T::zero();
1310 for i in 0..n {
1311 norm_sq = norm_sq + v[(i, j)] * v[(i, j)] + v[(i, jp1)] * v[(i, jp1)];
1312 }
1313 let norm = Real::sqrt(norm_sq);
1314 if norm > T::zero() {
1315 for i in 0..n {
1316 v[(i, j)] = v[(i, j)] / norm;
1317 v[(i, jp1)] = v[(i, jp1)] / norm;
1318 }
1319 }
1320
1321 j = jp1 + 1;
1322 } else {
1323 let lambda = t[(j, j)];
1325 v[(j, j)] = T::one();
1326
1327 for i in (j + 1)..n {
1329 let mut sum = T::zero();
1330 for k in j..i {
1331 sum = sum + t[(k, i)] * v[(k, j)];
1332 }
1333
1334 let d = t[(i, i)] - lambda;
1335 if Scalar::abs(d) > eps {
1336 v[(i, j)] = -sum / d;
1337 } else {
1338 v[(i, j)] = -sum / eps;
1339 }
1340 }
1341
1342 let mut norm_sq = T::zero();
1344 for i in 0..n {
1345 norm_sq = norm_sq + v[(i, j)] * v[(i, j)];
1346 }
1347 let norm = Real::sqrt(norm_sq);
1348 if norm > T::zero() {
1349 for i in 0..n {
1350 v[(i, j)] = v[(i, j)] / norm;
1351 }
1352 }
1353
1354 j += 1;
1355 }
1356 }
1357
1358 v
1359}
1360
1361fn householder_3<T: Field + Real>(x: &[T]) -> (Vec<T>, T) {
1363 let n = x.len().min(3);
1364 if n == 0 {
1365 return (Vec::new(), T::zero());
1366 }
1367
1368 let mut norm_sq = T::zero();
1369 for i in 0..n {
1370 norm_sq = norm_sq + x[i] * x[i];
1371 }
1372 let norm = Real::sqrt(norm_sq);
1373
1374 if norm == T::zero() {
1375 return (vec![T::zero(); n], T::zero());
1376 }
1377
1378 let mut v = vec![T::zero(); n];
1379 for i in 0..n {
1380 v[i] = x[i];
1381 }
1382
1383 let sign = if x[0] >= T::zero() {
1384 T::one()
1385 } else {
1386 -T::one()
1387 };
1388 v[0] = v[0] + sign * norm;
1389
1390 let mut v_norm_sq = T::zero();
1391 for i in 0..n {
1392 v_norm_sq = v_norm_sq + v[i] * v[i];
1393 }
1394
1395 if v_norm_sq != T::zero() {
1401 let tau = T::from_f64(2.0).unwrap_or_else(T::zero) / v_norm_sq;
1402 (v, tau)
1403 } else {
1404 (v, T::zero())
1405 }
1406}
1407
1408#[cfg(test)]
1409mod tests {
1410 use super::*;
1411
1412 fn approx_eq(a: f64, b: f64, tol: f64) -> bool {
1413 (a - b).abs() < tol
1414 }
1415
1416 #[test]
1417 fn test_schur_upper_triangular() {
1418 let a = Mat::from_rows(&[&[1.0f64, 2.0], &[0.0, 3.0]]);
1420
1421 let schur = Schur::compute(a.as_ref()).unwrap();
1422 let eigenvalues = schur.eigenvalues();
1423
1424 let mut eigs: Vec<f64> = eigenvalues.iter().map(|e| e.real).collect();
1426 eigs.sort_by(|a, b| a.partial_cmp(b).unwrap());
1427 assert!(approx_eq(eigs[0], 1.0, 1e-10));
1428 assert!(approx_eq(eigs[1], 3.0, 1e-10));
1429 }
1430
1431 #[test]
1432 fn test_schur_diagonal() {
1433 let a = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, 5.0, 0.0], &[0.0, 0.0, 3.0]]);
1434
1435 let schur = Schur::compute(a.as_ref()).unwrap();
1436 let eigenvalues = schur.eigenvalues();
1437
1438 let mut eigs: Vec<f64> = eigenvalues.iter().map(|e| e.real).collect();
1439 eigs.sort_by(|a, b| a.partial_cmp(b).unwrap());
1440 assert!(approx_eq(eigs[0], 2.0, 1e-10));
1441 assert!(approx_eq(eigs[1], 3.0, 1e-10));
1442 assert!(approx_eq(eigs[2], 5.0, 1e-10));
1443 }
1444
1445 #[test]
1446 fn test_schur_complex_eigenvalues() {
1447 let theta = core::f64::consts::FRAC_PI_4; let c = theta.cos();
1450 let s = theta.sin();
1451 let a = Mat::from_rows(&[&[c, -s], &[s, c]]);
1452
1453 let schur = Schur::compute(a.as_ref()).unwrap();
1454 let eigenvalues = schur.eigenvalues();
1455
1456 assert_eq!(eigenvalues.len(), 2);
1458 assert!(approx_eq(eigenvalues[0].real, eigenvalues[1].real, 1e-10));
1460 assert!(approx_eq(eigenvalues[0].imag, -eigenvalues[1].imag, 1e-10));
1461 assert!(approx_eq(eigenvalues[0].real, c, 1e-10));
1462 assert!(approx_eq(eigenvalues[0].imag.abs(), s, 1e-10));
1463 }
1464
1465 #[test]
1466 fn test_schur_reconstruction() {
1467 let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
1468
1469 let schur = Schur::compute(a.as_ref()).unwrap();
1470 let reconstructed = schur.reconstruct();
1471
1472 for i in 0..3 {
1473 for j in 0..3 {
1474 assert!(
1475 approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-10),
1476 "reconstructed[{},{}] = {}, a = {}",
1477 i,
1478 j,
1479 reconstructed[(i, j)],
1480 a[(i, j)]
1481 );
1482 }
1483 }
1484 }
1485
1486 #[test]
1487 fn test_schur_q_orthogonal() {
1488 let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[4.0, 5.0, 6.0], &[7.0, 8.0, 9.0]]);
1489
1490 let schur = Schur::compute(a.as_ref()).unwrap();
1491 let q = schur.q();
1492
1493 let n = 3;
1495 for i in 0..n {
1496 for j in 0..n {
1497 let mut dot = 0.0;
1498 for k in 0..n {
1499 dot += q[(k, i)] * q[(k, j)];
1500 }
1501 let expected = if i == j { 1.0 } else { 0.0 };
1502 assert!(
1503 approx_eq(dot, expected, 1e-10),
1504 "Q^T*Q[{},{}] = {}, expected {}",
1505 i,
1506 j,
1507 dot,
1508 expected
1509 );
1510 }
1511 }
1512 }
1513
1514 #[test]
1515 fn test_schur_single() {
1516 let a = Mat::from_rows(&[&[5.0f64]]);
1517 let schur = Schur::compute(a.as_ref()).unwrap();
1518
1519 assert_eq!(schur.eigenvalues().len(), 1);
1520 assert!(approx_eq(schur.eigenvalues()[0].real, 5.0, 1e-10));
1521 }
1522
1523 #[test]
1524 fn test_schur_4x4() {
1525 let a = Mat::from_rows(&[
1526 &[4.0f64, 1.0, -2.0, 2.0],
1527 &[1.0, 2.0, 0.0, 1.0],
1528 &[-2.0, 0.0, 3.0, -2.0],
1529 &[2.0, 1.0, -2.0, -1.0],
1530 ]);
1531
1532 let schur = Schur::compute(a.as_ref()).unwrap();
1533 let reconstructed = schur.reconstruct();
1534
1535 for i in 0..4 {
1536 for j in 0..4 {
1537 assert!(
1538 approx_eq(reconstructed[(i, j)], a[(i, j)], 1e-8),
1539 "reconstructed[{},{}] = {}, a = {}",
1540 i,
1541 j,
1542 reconstructed[(i, j)],
1543 a[(i, j)]
1544 );
1545 }
1546 }
1547 }
1548
1549 #[test]
1550 fn test_schur_f32() {
1551 let a = Mat::from_rows(&[&[1.0f32, 2.0], &[3.0, 4.0]]);
1552
1553 let schur = Schur::compute(a.as_ref()).unwrap();
1554 let reconstructed = schur.reconstruct();
1555
1556 for i in 0..2 {
1557 for j in 0..2 {
1558 assert!(
1559 (reconstructed[(i, j)] - a[(i, j)]).abs() < 1e-4,
1560 "reconstructed[{},{}] = {}, a = {}",
1561 i,
1562 j,
1563 reconstructed[(i, j)],
1564 a[(i, j)]
1565 );
1566 }
1567 }
1568 }
1569
1570 #[test]
1571 fn test_trevc_right_upper_triangular() {
1572 let t = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[0.0, 4.0, 5.0], &[0.0, 0.0, 6.0]]);
1574
1575 let v = trevc_right(&t);
1576
1577 assert!(v[(1, 0)].abs() < 1e-10);
1579 assert!(v[(2, 0)].abs() < 1e-10);
1580
1581 assert!(v[(2, 2)].abs() > 0.1);
1584 }
1585
1586 #[test]
1587 fn test_trevc_right_diagonal() {
1588 let t = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, 5.0, 0.0], &[0.0, 0.0, 3.0]]);
1590
1591 let v = trevc_right(&t);
1592
1593 for i in 0..3 {
1595 assert!(
1596 approx_eq(v[(i, i)].abs(), 1.0, 1e-10),
1597 "v[{},{}] = {}",
1598 i,
1599 i,
1600 v[(i, i)]
1601 );
1602 }
1603 }
1604
1605 #[test]
1606 fn test_trevc_eigenvector_equation() {
1607 let t = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[0.0, 4.0, 5.0], &[0.0, 0.0, 6.0]]);
1609
1610 let v = trevc_right(&t);
1611
1612 let eigenvalues = [1.0, 4.0, 6.0];
1614
1615 for (j, &lambda) in eigenvalues.iter().enumerate() {
1616 let mut tv = [0.0; 3];
1618 for i in 0..3 {
1619 for k in 0..3 {
1620 tv[i] += t[(i, k)] * v[(k, j)];
1621 }
1622 }
1623
1624 for i in 0..3 {
1626 assert!(
1627 approx_eq(tv[i], lambda * v[(i, j)], 1e-10),
1628 "T*v[{}] = {}, λ*v = {}",
1629 i,
1630 tv[i],
1631 lambda * v[(i, j)]
1632 );
1633 }
1634 }
1635 }
1636
1637 #[test]
1638 fn test_trevc_left_diagonal() {
1639 let t = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, 5.0, 0.0], &[0.0, 0.0, 3.0]]);
1640
1641 let u = trevc_left(&t);
1642
1643 for i in 0..3 {
1645 assert!(
1646 approx_eq(u[(i, i)].abs(), 1.0, 1e-10),
1647 "u[{},{}] = {}",
1648 i,
1649 i,
1650 u[(i, i)]
1651 );
1652 }
1653 }
1654
1655 #[test]
1656 fn test_schur_eigenvectors() {
1657 let a = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[0.0, 4.0, 5.0], &[0.0, 0.0, 6.0]]);
1659
1660 let schur = Schur::compute(a.as_ref()).unwrap();
1661 let (vr, _vl) = schur.eigenvectors();
1662
1663 for j in 0..3 {
1666 let lambda = schur.eigenvalues()[j].real;
1667
1668 let mut av = [0.0; 3];
1669 for i in 0..3 {
1670 for k in 0..3 {
1671 av[i] += a[(i, k)] * vr[(k, j)];
1672 }
1673 }
1674
1675 for i in 0..3 {
1676 assert!(
1677 approx_eq(av[i], lambda * vr[(i, j)], 1e-8),
1678 "A*v[{}] = {}, λ*v = {}",
1679 i,
1680 av[i],
1681 lambda * vr[(i, j)]
1682 );
1683 }
1684 }
1685 }
1686
1687 #[test]
1688 fn test_trevc_2x2_block() {
1689 let theta = core::f64::consts::FRAC_PI_4;
1691 let c = theta.cos();
1692 let s = theta.sin();
1693 let t = Mat::from_rows(&[&[c, -s], &[s, c]]);
1694
1695 let v = trevc_right(&t);
1696
1697 let norm0_sq = v[(0, 0)] * v[(0, 0)] + v[(1, 0)] * v[(1, 0)];
1701 let norm1_sq = v[(0, 1)] * v[(0, 1)] + v[(1, 1)] * v[(1, 1)];
1702 let total_norm = (norm0_sq + norm1_sq).sqrt();
1703
1704 assert!(
1705 approx_eq(total_norm, 1.0, 1e-10),
1706 "eigenvector norm = {}",
1707 total_norm
1708 );
1709 }
1710
1711 #[test]
1712 fn test_trsna_s_diagonal() {
1713 let t = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, 5.0, 0.0], &[0.0, 0.0, 3.0]]);
1716
1717 let s = trsna_s(&t);
1718
1719 assert_eq!(s.len(), 3);
1720 for i in 0..3 {
1721 assert!(
1722 approx_eq(s[i], 1.0, 1e-10),
1723 "s[{}] = {}, expected 1.0",
1724 i,
1725 s[i]
1726 );
1727 }
1728 }
1729
1730 #[test]
1731 fn test_trsna_s_upper_triangular() {
1732 let t = Mat::from_rows(&[&[1.0f64, 2.0, 3.0], &[0.0, 4.0, 5.0], &[0.0, 0.0, 6.0]]);
1734
1735 let s = trsna_s(&t);
1736
1737 assert_eq!(s.len(), 3);
1738 for i in 0..3 {
1740 assert!(s[i] > 0.0, "s[{}] = {} should be positive", i, s[i]);
1741 }
1742 }
1743
1744 #[test]
1745 fn test_trsna_sep_diagonal() {
1746 let t = Mat::from_rows(&[&[2.0f64, 0.0, 0.0], &[0.0, 5.0, 0.0], &[0.0, 0.0, 3.0]]);
1748
1749 let sep = trsna_sep(&t);
1750
1751 assert_eq!(sep.len(), 3);
1752 assert!(
1754 approx_eq(sep[0], 1.0, 1e-10),
1755 "sep[0] = {}, expected 1.0",
1756 sep[0]
1757 );
1758 assert!(
1760 approx_eq(sep[1], 2.0, 1e-10),
1761 "sep[1] = {}, expected 2.0",
1762 sep[1]
1763 );
1764 assert!(
1766 approx_eq(sep[2], 1.0, 1e-10),
1767 "sep[2] = {}, expected 1.0",
1768 sep[2]
1769 );
1770 }
1771
1772 #[test]
1773 fn test_trsna_sep_close_eigenvalues() {
1774 let t = Mat::from_rows(&[&[1.0f64, 0.0], &[0.0, 1.001]]);
1776
1777 let sep = trsna_sep(&t);
1778
1779 assert_eq!(sep.len(), 2);
1780 assert!(
1782 approx_eq(sep[0], 0.001, 1e-10),
1783 "sep[0] = {}, expected 0.001",
1784 sep[0]
1785 );
1786 assert!(
1787 approx_eq(sep[1], 0.001, 1e-10),
1788 "sep[1] = {}, expected 0.001",
1789 sep[1]
1790 );
1791 }
1792
1793 #[test]
1794 fn test_schur_eigenvalue_condition_numbers() {
1795 let a = Mat::from_rows(&[&[4.0f64, 0.0], &[0.0, 2.0]]);
1797
1798 let schur = Schur::compute(a.as_ref()).unwrap();
1799 let cond = schur.eigenvalue_condition_numbers();
1800
1801 assert_eq!(cond.len(), 2);
1802 for i in 0..2 {
1804 assert!(
1805 cond[i] > 0.5,
1806 "cond[{}] = {} should be > 0.5 for diagonal matrix",
1807 i,
1808 cond[i]
1809 );
1810 }
1811 }
1812
1813 #[test]
1814 fn test_schur_eigenvector_separation() {
1815 let a = Mat::from_rows(&[&[4.0f64, 0.0], &[0.0, 2.0]]);
1817
1818 let schur = Schur::compute(a.as_ref()).unwrap();
1819 let sep = schur.eigenvector_separation();
1820
1821 assert_eq!(sep.len(), 2);
1822 for i in 0..2 {
1824 assert!(
1825 approx_eq(sep[i], 2.0, 1e-10),
1826 "sep[{}] = {}, expected 2.0",
1827 i,
1828 sep[i]
1829 );
1830 }
1831 }
1832
1833 #[test]
1834 fn test_trsna_complex_eigenvalues() {
1835 let theta = core::f64::consts::FRAC_PI_4;
1837 let c = theta.cos();
1838 let s = theta.sin();
1839 let t = Mat::from_rows(&[&[c, -s], &[s, c]]);
1840
1841 let cond = trsna_s(&t);
1842 let sep = trsna_sep(&t);
1843
1844 assert_eq!(cond.len(), 2);
1845 assert_eq!(sep.len(), 2);
1846
1847 assert!(
1849 approx_eq(cond[0], cond[1], 1e-10),
1850 "cond[0]={}, cond[1]={} should be equal",
1851 cond[0],
1852 cond[1]
1853 );
1854 assert!(
1855 approx_eq(sep[0], sep[1], 1e-10),
1856 "sep[0]={}, sep[1]={} should be equal",
1857 sep[0],
1858 sep[1]
1859 );
1860 }
1861
1862 fn slow_converging_5x5() -> Mat<f64> {
1871 Mat::from_rows(&[
1872 &[5.0, 1.0, 0.2, 0.0, 0.1],
1873 &[0.3, 5.0, 1.0, 0.15, 0.0],
1874 &[0.0, 0.25, 5.0, 1.0, 0.2],
1875 &[0.1, 0.0, 0.3, 5.0, 1.0],
1876 &[0.2, 0.1, 0.0, 0.35, 5.0],
1877 ])
1878 }
1879
1880 #[test]
1881 fn test_schur_reports_non_convergence_when_budget_exhausted() {
1882 let a = slow_converging_5x5();
1883
1884 let result = Schur::compute_with_iteration_budget(a.as_ref(), 1);
1889 assert_eq!(
1890 result.err(),
1891 Some(SchurError::NotConverged),
1892 "tiny iteration budget must report NotConverged, not fabricate a result"
1893 );
1894 }
1895
1896 #[test]
1897 fn test_schur_converges_with_full_budget() {
1898 let a = slow_converging_5x5();
1903
1904 let schur = Schur::compute(a.as_ref()).expect("full budget must converge");
1905 let rec = schur.reconstruct();
1906 for i in 0..5 {
1907 for j in 0..5 {
1908 assert!(
1909 approx_eq(rec[(i, j)], a[(i, j)], 1e-8),
1910 "reconstruction mismatch at ({i},{j}): {} vs {}",
1911 rec[(i, j)],
1912 a[(i, j)]
1913 );
1914 }
1915 }
1916
1917 let trace: f64 = (0..5).map(|i| a[(i, i)]).sum();
1920 let eig_sum: f64 = schur.eigenvalues().iter().map(|e| e.real).sum();
1921 assert!(
1922 approx_eq(eig_sum, trace, 1e-8),
1923 "eigenvalue real-part sum {eig_sum} != trace {trace}"
1924 );
1925 }
1926}