Skip to main content

oxiblas_lapack/evd/
schur.rs

1//! Schur decomposition.
2//!
3//! Computes the Schur decomposition A = Q T Q^T where Q is orthogonal
4//! and T is quasi-upper triangular (upper triangular with possible 2×2 blocks
5//! on the diagonal for complex eigenvalue pairs).
6
7use oxiblas_core::scalar::{Field, Real, Scalar};
8use oxiblas_matrix::{Mat, MatRef};
9
10use super::hessenberg::Hessenberg;
11
12/// Error type for Schur decomposition.
13#[derive(Debug, Clone, Copy, PartialEq, Eq)]
14pub enum SchurError {
15    /// Matrix is empty.
16    EmptyMatrix,
17    /// Matrix is not square.
18    NotSquare,
19    /// Algorithm did not converge.
20    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/// Represents a real or complex eigenvalue.
36#[derive(Debug, Clone, Copy, PartialEq)]
37pub struct Eigenvalue<T> {
38    /// Real part of the eigenvalue.
39    pub real: T,
40    /// Imaginary part of the eigenvalue (zero for real eigenvalues).
41    pub imag: T,
42}
43
44impl<T: Scalar> Eigenvalue<T> {
45    /// Creates a real eigenvalue.
46    pub fn real_only(value: T) -> Self {
47        Self {
48            real: value,
49            imag: T::zero(),
50        }
51    }
52
53    /// Creates a complex eigenvalue.
54    pub fn complex(real: T, imag: T) -> Self {
55        Self { real, imag }
56    }
57
58    /// Returns true if this is a real eigenvalue.
59    pub fn is_real(&self) -> bool {
60        self.imag == T::zero()
61    }
62}
63
64/// Schur decomposition of a matrix.
65///
66/// For a matrix A, computes A = Q T Q^T where:
67/// - Q is orthogonal (Q^T Q = I)
68/// - T is quasi-upper triangular (real Schur form)
69#[derive(Debug, Clone)]
70pub struct Schur<T: Scalar> {
71    /// The orthogonal matrix Q (Schur vectors).
72    q: Mat<T>,
73    /// The quasi-upper triangular matrix T.
74    t: Mat<T>,
75    /// Eigenvalues (real and complex pairs).
76    eigenvalues: Vec<Eigenvalue<T>>,
77    /// Matrix dimension.
78    n: usize,
79}
80
81impl<T: Field + Real + bytemuck::Zeroable> Schur<T> {
82    /// Maximum iterations for QR iteration.
83    const MAX_ITERATIONS: usize = 100;
84
85    /// Computes the Schur decomposition of a square matrix.
86    ///
87    /// # Example
88    ///
89    /// ```
90    /// use oxiblas_lapack::evd::Schur;
91    /// use oxiblas_matrix::Mat;
92    ///
93    /// let a = Mat::from_rows(&[
94    ///     &[1.0f64, 2.0],
95    ///     &[0.0, 3.0],
96    /// ]);
97    ///
98    /// let schur = Schur::compute(a.as_ref()).unwrap();
99    /// let eigenvalues = schur.eigenvalues();
100    ///
101    /// // Eigenvalues are 1 and 3
102    /// assert!((eigenvalues[0].real - 1.0).abs() < 1e-10 || (eigenvalues[0].real - 3.0).abs() < 1e-10);
103    /// ```
104    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        // Handle 1×1 case
117        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        // Handle 2×2 case
132        if n == 2 {
133            return Self::compute_2x2(a);
134        }
135
136        // n >= 3: reduce to Hessenberg form and run the Francis double-shift QR
137        // iteration. The iteration budget scales with n — `MAX_ITERATIONS` sweeps
138        // per row is the classic heuristic used by LAPACK's xHSEQR (`saturating_mul`
139        // guards the pathological huge-`n` case against usize overflow).
140        Self::compute_hessenberg_qr(a, Self::MAX_ITERATIONS.saturating_mul(n))
141    }
142
143    /// Reduces `a` (which must be square with `n >= 3`) to real Schur form via
144    /// Hessenberg reduction followed by the Francis double-shift QR iteration.
145    ///
146    /// `max_total_iterations` bounds the *total* number of QR sweeps across all
147    /// deflation blocks. If that budget is exhausted while the active window has
148    /// neither deflated to size `<= 2` nor reached real Schur form, this returns
149    /// [`SchurError::NotConverged`] instead of silently handing back a partially
150    /// reduced `T`/`Q` as if the decomposition had succeeded — matching the
151    /// `INFO > 0` failure contract of LAPACK's DHSEQR.
152    ///
153    /// # Why the post-loop convergence check is correct
154    ///
155    /// Every loop pass either deflates a converged 1×1/2×2 block (shrinking `p`
156    /// by 1 or 2) or performs a bulge-chasing QR sweep on the active window
157    /// `[0, p)`. On a successful run the loop therefore exits with `p <= 2`. The
158    /// *only* way to exit with `p > 2` is to trip the `iter_count` cap, i.e. the
159    /// iteration ran out of budget before finishing. Even then, the residual
160    /// window might coincidentally already be quasi-triangular, so we do not fail
161    /// blindly on `p > 2`: we fail only when the window still contains two
162    /// consecutive non-negligible sub-diagonals (a diagonal block of order `>= 3`),
163    /// which is exactly the LAPACK definition of "not yet in real Schur form".
164    fn compute_hessenberg_qr(
165        a: MatRef<'_, T>,
166        max_total_iterations: usize,
167    ) -> Result<Self, SchurError> {
168        let n = a.ncols();
169
170        // Step 1: Reduce to upper Hessenberg form
171        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        // Step 2: Apply QR iteration with implicit shifts
189        let eps = <T as Scalar>::epsilon();
190        let tol = eps * T::from_f64(100.0).unwrap_or(T::one());
191
192        // Process from bottom to top
193        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            // Find the active block
200            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                // 1×1 block converged
214                p -= 1;
215            } else if q_idx == p - 2 {
216                // Check if 2×2 block has converged (complex eigenvalues)
217                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                    // Complex eigenvalues, keep as 2×2 block
227                    p -= 2;
228                } else {
229                    // Real eigenvalues, continue iteration
230                    Self::francis_qr_step(&mut t, &mut q, q_idx, p);
231                }
232            } else {
233                // Apply Francis QR step
234                Self::francis_qr_step(&mut t, &mut q, q_idx, p);
235            }
236        }
237
238        // Report non-convergence honestly. If the QR sweeps exhausted their budget
239        // while an active window of order `> 2` remained *and* that window is still
240        // not in real Schur form, the eigenvalues extracted from `T` would be
241        // garbage — return NotConverged rather than a plausible-looking wrong answer.
242        if p > 2 && !Self::active_block_converged(&t, p, tol) {
243            return Err(SchurError::NotConverged);
244        }
245
246        // Handle remaining 2×2 block if needed
247        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        // Extract eigenvalues
256        let eigenvalues = Self::extract_eigenvalues(&t);
257
258        Ok(Self {
259            q,
260            t,
261            eigenvalues,
262            n,
263        })
264    }
265
266    /// Returns `true` when the leading `p`×`p` active window of `t` is already in
267    /// real Schur form, i.e. it contains no two *consecutive* non-negligible
268    /// sub-diagonal entries.
269    ///
270    /// A quasi-triangular matrix is built from isolated 1×1 and 2×2 diagonal
271    /// blocks, so a lone 2×2 block's single non-zero sub-diagonal is perfectly
272    /// converged. Only two adjacent non-negligible sub-diagonals — which would
273    /// imply an unreduced diagonal block of order `>= 3` — mean the QR iteration
274    /// still had work to do. This is the criterion used to distinguish a genuine
275    /// convergence failure from a budget exit that nonetheless landed on a valid
276    /// Schur form.
277    fn active_block_converged(t: &Mat<T>, p: usize, tol: T) -> bool {
278        // Sub-diagonal entry `t[k, k-1]` is negligible relative to its neighbouring
279        // diagonal magnitudes (the same deflation test used inside the QR loop).
280        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        // Scan for adjacent non-negligible sub-diagonals within the window.
287        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    /// Test-only hook: runs the Schur decomposition with a custom QR iteration
296    /// budget so the [`SchurError::NotConverged`] path can be exercised
297    /// deterministically (a genuine slow-converging matrix would otherwise need to
298    /// be enormous to defeat the default `MAX_ITERATIONS * n` budget).
299    #[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            // Tiny cases are closed-form and never iterate; the budget is irrelevant.
315            return Self::compute(a);
316        }
317        Self::compute_hessenberg_qr(a, max_total_iterations)
318    }
319
320    /// Computes Schur decomposition for 2×2 matrix.
321    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            // Real eigenvalues - triangularize
343            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            // Find rotation to triangularize
348            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                // T = Q^T * A * Q
363                let mut temp = Mat::zeros(2, 2);
364                // temp = Q^T * A
365                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                // t = temp * Q
375                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                // Already triangular
386                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            // Complex eigenvalues - keep as 2×2 block
396            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    /// Applies one step of Francis double-shift QR iteration.
418    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        // Compute shift from bottom 2×2 block
426        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; // trace
432        let p = h11 * h22 - h12 * h21; // determinant
433
434        // First column of (H - s1*I)(H - s2*I) = H² - s*H + p*I
435        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        // Chase the bulge
448        for k in start..end.saturating_sub(1) {
449            // Compute Householder to zero out y, z
450            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                // Apply from left: T := (I - tau * v * v^T) * T
456                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                // Apply from right: T := T * (I - tau * v * v^T)
471                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                // Accumulate Q
485                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            // Prepare for next iteration
499            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                // 2×2 Householder for the last step
509                x = t[(k + 1, k)];
510                y = t[(k + 2, k)];
511                z = T::zero();
512            }
513        }
514    }
515
516    /// Extracts eigenvalues from the quasi-upper triangular Schur form.
517    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                // Last element is a 1×1 block
526                eigenvalues.push(Eigenvalue::real_only(t[(i, i)]));
527                i += 1;
528            } else {
529                // Check if this is a 2×2 block
530                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                    // 1×1 block
535                    eigenvalues.push(Eigenvalue::real_only(t[(i, i)]));
536                    i += 1;
537                } else {
538                    // 2×2 block - compute eigenvalues
539                    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                        // Real eigenvalues
550                        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                        // Complex conjugate pair
559                        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    /// Returns the orthogonal matrix Q (Schur vectors).
574    pub fn q(&self) -> MatRef<'_, T> {
575        self.q.as_ref()
576    }
577
578    /// Returns the quasi-upper triangular matrix T (Schur form).
579    pub fn t(&self) -> MatRef<'_, T> {
580        self.t.as_ref()
581    }
582
583    /// Returns the eigenvalues.
584    pub fn eigenvalues(&self) -> &[Eigenvalue<T>] {
585        &self.eigenvalues
586    }
587
588    /// Returns only the real parts of eigenvalues.
589    pub fn eigenvalues_real(&self) -> Vec<T> {
590        self.eigenvalues.iter().map(|e| e.real).collect()
591    }
592
593    /// Reconstructs the original matrix: A = Q T Q^T.
594    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        // QT = Q * T
599        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        // A = QT * Q^T
610        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    /// Computes right eigenvectors from the Schur form (LAPACK TREVC).
624    ///
625    /// Returns the right eigenvectors V such that T V = V D where D is the
626    /// diagonal matrix of eigenvalues. The eigenvectors are normalized.
627    ///
628    /// For real eigenvalues, returns real eigenvectors.
629    /// For complex conjugate pairs, returns two columns: the real and imaginary
630    /// parts of the eigenvector.
631    ///
632    /// # Example
633    ///
634    /// ```
635    /// use oxiblas_lapack::evd::Schur;
636    /// use oxiblas_matrix::Mat;
637    ///
638    /// let a = Mat::from_rows(&[
639    ///     &[1.0f64, 2.0, 3.0],
640    ///     &[0.0, 4.0, 5.0],
641    ///     &[0.0, 0.0, 6.0],
642    /// ]);
643    ///
644    /// let schur = Schur::compute(a.as_ref()).unwrap();
645    /// let vr = schur.right_eigenvectors();
646    /// // Each column of vr is a right eigenvector
647    /// ```
648    #[must_use]
649    pub fn right_eigenvectors(&self) -> Mat<T> {
650        trevc_right(&self.t)
651    }
652
653    /// Computes left eigenvectors from the Schur form (LAPACK TREVC).
654    ///
655    /// Returns the left eigenvectors U such that U^T T = D U^T where D is the
656    /// diagonal matrix of eigenvalues. The eigenvectors are normalized.
657    ///
658    /// # Example
659    ///
660    /// ```
661    /// use oxiblas_lapack::evd::Schur;
662    /// use oxiblas_matrix::Mat;
663    ///
664    /// let a = Mat::from_rows(&[
665    ///     &[1.0f64, 2.0, 3.0],
666    ///     &[0.0, 4.0, 5.0],
667    ///     &[0.0, 0.0, 6.0],
668    /// ]);
669    ///
670    /// let schur = Schur::compute(a.as_ref()).unwrap();
671    /// let vl = schur.left_eigenvectors();
672    /// ```
673    #[must_use]
674    pub fn left_eigenvectors(&self) -> Mat<T> {
675        trevc_left(&self.t)
676    }
677
678    /// Computes eigenvectors of the original matrix A = Q T Q^T.
679    ///
680    /// The eigenvectors of A are Q * V where V are the eigenvectors of T.
681    ///
682    /// # Returns
683    ///
684    /// (right_eigenvectors, left_eigenvectors) of A.
685    #[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        // Transform to eigenvectors of A: A_vr = Q * T_vr, A_vl = Q * T_vl
691        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    /// Computes reciprocal condition numbers for eigenvalues (LAPACK DTRSNA).
711    ///
712    /// For each eigenvalue λ, the reciprocal condition number s measures how
713    /// sensitive λ is to perturbations in the matrix. A small value indicates
714    /// a poorly conditioned eigenvalue.
715    ///
716    /// The condition number is computed as:
717    /// - For simple eigenvalues: s = 1 / |y^H x| where x is the right eigenvector
718    ///   and y is the left eigenvector, both normalized to unit length.
719    /// - For complex conjugate pairs: uses the average of the pair.
720    ///
721    /// # Returns
722    ///
723    /// Vector of reciprocal condition numbers, one per eigenvalue.
724    /// Smaller values indicate more sensitive eigenvalues.
725    ///
726    /// # Example
727    ///
728    /// ```
729    /// use oxiblas_lapack::evd::Schur;
730    /// use oxiblas_matrix::Mat;
731    ///
732    /// let a = Mat::from_rows(&[
733    ///     &[1.0f64, 0.0],
734    ///     &[0.0, 1000.0],
735    /// ]);
736    ///
737    /// let schur = Schur::compute(a.as_ref()).unwrap();
738    /// let cond = schur.eigenvalue_condition_numbers();
739    /// // Both eigenvalues of a diagonal matrix are well-conditioned
740    /// assert!(cond[0] > 0.9);
741    /// assert!(cond[1] > 0.9);
742    /// ```
743    #[must_use]
744    pub fn eigenvalue_condition_numbers(&self) -> Vec<T> {
745        trsna_s(&self.t)
746    }
747
748    /// Computes reciprocal condition numbers for eigenvectors (LAPACK DTRSNA).
749    ///
750    /// For each right eigenvector x_j, computes the separation sep_j which
751    /// measures how close the eigenvalue is to the rest of the spectrum.
752    ///
753    /// # Returns
754    ///
755    /// Vector of separation values, one per eigenvalue.
756    /// Smaller values indicate more sensitive eigenvectors.
757    #[must_use]
758    pub fn eigenvector_separation(&self) -> Vec<T> {
759        trsna_sep(&self.t)
760    }
761}
762
763/// Computes reciprocal condition numbers for eigenvalues (s values from LAPACK DTRSNA).
764///
765/// For each simple eigenvalue λ_j, s_j = |y_j^H * x_j| where x_j is the right
766/// eigenvector and y_j is the left eigenvector, both normalized.
767///
768/// # Arguments
769///
770/// * `t` - The quasi-upper triangular Schur matrix
771///
772/// # Returns
773///
774/// Vector of reciprocal condition numbers for each eigenvalue.
775pub 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        // Check for 2×2 block
790        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            // Complex conjugate pair: compute s = |y^H * x| using both columns
800            // For complex eigenvector stored as (real, imag) in consecutive columns:
801            // y^H * x = (yr - i*yi)^T * (xr + i*xi) = (yr^T*xr + yi^T*xi) + i*(yr^T*xi - yi^T*xr)
802            let jp1 = j + 1;
803
804            let mut prod_rr = T::zero(); // yr^T * xr
805            let mut prod_ii = T::zero(); // yi^T * xi
806            let mut prod_ri = T::zero(); // yr^T * xi
807            let mut prod_ir = T::zero(); // yi^T * xr
808
809            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            // Both eigenvalues in the pair have the same condition number
821            s[j] = abs_inner;
822            s[jp1] = abs_inner;
823
824            j += 2;
825        } else {
826            // Simple real eigenvalue: s = |y^T * x|
827            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
839/// Computes separation (sep) for eigenvectors (from LAPACK DTRSNA).
840///
841/// For each eigenvalue λ_j, sep_j = σ_min(T_22 - λ_j * I) where T_22 is the
842/// (n-1)×(n-1) trailing principal submatrix with λ_j removed.
843///
844/// This measures how separated λ_j is from the rest of the spectrum.
845///
846/// # Arguments
847///
848/// * `t` - The quasi-upper triangular Schur matrix
849///
850/// # Returns
851///
852/// Vector of separation values for each eigenvalue.
853pub 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        // Check for 2×2 block
865        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            // Complex conjugate pair
875            let jp1 = j + 1;
876
877            // Compute approximate separation as minimum distance to other eigenvalues
878            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            // Check distance to all other eigenvalues
893            let mut k = 0;
894            while k < n {
895                if k == j || k == jp1 {
896                    k += 1;
897                    continue;
898                }
899
900                // Check if k is part of a 2×2 block
901                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                // Distance between eigenvalues
930                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                // Also check conjugate
938                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            // Simple real eigenvalue
958            let lambda = t[(j, j)];
959            let mut min_sep = T::one() / eps;
960
961            // Check distance to all other eigenvalues
962            let mut k = 0;
963            while k < n {
964                if k == j {
965                    k += 1;
966                    continue;
967                }
968
969                // Check if k is part of a 2×2 block
970                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                // Distance between eigenvalues
999                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
1020/// Computes right eigenvectors of a quasi-upper triangular matrix (LAPACK DTREVC).
1021///
1022/// The input T must be in real Schur form (quasi-upper triangular with 1×1 and
1023/// 2×2 diagonal blocks).
1024///
1025/// # Arguments
1026///
1027/// * `t` - The quasi-upper triangular Schur matrix
1028///
1029/// # Returns
1030///
1031/// Matrix V of right eigenvectors (column-wise). For complex conjugate pairs,
1032/// consecutive columns contain the real and imaginary parts.
1033pub 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    // Initialize to identity for back-substitution starting point
1038    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    // Process eigenvalues from last to first
1045    let mut j = n;
1046    while j > 0 {
1047        j -= 1;
1048
1049        // Check if this is part of a 2×2 block
1050        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            // 2×2 block: complex conjugate eigenvalues
1060            // Process columns j-1 and j together
1061            let jm1 = j - 1;
1062
1063            // Eigenvalues of 2×2 block
1064            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            // Complex eigenvalues: λ = (trace ± i*sqrt(-disc)) / 2
1074            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            // For the 2×2 block, set up eigenvector components
1079            // v[jm1] = real part, v[j] = imag part for first eigenvector
1080            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            // Compute actual eigenvector of 2×2 block
1086            // (a11 - λ) v1 + a12 v2 = 0
1087            // a21 v1 + (a22 - λ) v2 = 0
1088            // For λ = real_part + i*imag_part:
1089            // Let v = vr + i*vi
1090            // Then: (T - λI)(vr + i*vi) = 0
1091            // Real: (T - real_part*I)vr + imag_part*vi = 0
1092            // Imag: (T - real_part*I)vi - imag_part*vr = 0
1093
1094            // For the 2×2 block itself, find normalized eigenvector
1095            // The eigenvector satisfies (T - λI)v = 0 where λ = real_part + i*imag_part
1096            // For complex eigenvector v = vr + i*vi:
1097            // (T - real_part*I)vr + imag_part*vi = 0  (real part)
1098            // (T - real_part*I)vi - imag_part*vr = 0  (imag part)
1099            //
1100            // For the 2×2 case, we can set v2 = 1 + 0i and solve for v1
1101            // From row 2: a21*v1r + (a22-real)*v2r + imag*v2i = 0 (real)
1102            //             a21*v1i + (a22-real)*v2i - imag*v2r = 0 (imag)
1103            // With v2r=1, v2i=0:
1104            //   a21*v1r + (a22-real) = 0  =>  v1r = -(a22-real)/a21 = -d22/a21
1105            //   a21*v1i - imag = 0        =>  v1i = imag/a21
1106
1107            // Use the more stable formulation based on LAPACK
1108            // For standardized eigenvector, use row with larger coefficient
1109            let d11 = a11 - real_part;
1110            let d22 = a22 - real_part;
1111
1112            // Choose the row with larger off-diagonal to avoid division by small number
1113            if Scalar::abs(a21) >= Scalar::abs(a12) && Scalar::abs(a21) > eps {
1114                // Use row 2: a21*v1 + d22*v2 = 0 (with v2=1)
1115                // v1r = -d22/a21, v1i = imag/a21
1116                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                // Use row 1: d11*v1 + a12*v2 = 0 (with v1=1)
1124                // v2r = -d11/a12, v2i = -imag/a12
1125                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                // Fallback to standard basis
1133                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            // Back-substitute for rows above the 2×2 block
1140            for i in (0..jm1).rev() {
1141                // Solve for v[i,jm1] and v[i,j] (real and imag parts)
1142                // (T[i,i] - real_part) * vr[i] + imag_part * vi[i] = -sum of upper terms (real)
1143                // (T[i,i] - real_part) * vi[i] - imag_part * vr[i] = -sum of upper terms (imag)
1144
1145                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                    // Solve 2×2 system:
1157                    // [d, imag] [vr]   [-sum_r]
1158                    // [-imag, d] [vi] = [-sum_i]
1159                    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            // Normalize the two columns
1165            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; // Skip the already-processed column
1180        } else {
1181            // 1×1 block: real eigenvalue
1182            let lambda = t[(j, j)];
1183
1184            // Initialize: v[j,j] = 1, others computed by back-substitution
1185            v[(j, j)] = T::one();
1186
1187            // Back-substitute: (T[i,i] - λ) v[i] = -sum_{k>i} T[i,k] v[k]
1188            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                    // Near-singular: use small perturbation
1199                    v[(i, j)] = -sum / eps;
1200                }
1201            }
1202
1203            // Normalize
1204            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
1220/// Computes left eigenvectors of a quasi-upper triangular matrix (LAPACK DTREVC).
1221///
1222/// The input T must be in real Schur form. Returns left eigenvectors U such that
1223/// U^T T = D U^T where D is diagonal.
1224///
1225/// # Arguments
1226///
1227/// * `t` - The quasi-upper triangular Schur matrix
1228///
1229/// # Returns
1230///
1231/// Matrix U of left eigenvectors (column-wise).
1232pub 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    // Initialize
1237    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    // Process eigenvalues from first to last (forward substitution for left eigenvectors)
1244    let mut j = 0;
1245    while j < n {
1246        // Check if this is part of a 2×2 block
1247        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            // 2×2 block: complex conjugate eigenvalues
1257            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            // Initialize 2×2 block eigenvector
1273            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            // Forward substitute for rows after the 2×2 block
1291            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            // Normalize
1309            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            // 1×1 block: real eigenvalue
1324            let lambda = t[(j, j)];
1325            v[(j, j)] = T::one();
1326
1327            // Forward substitute: (T[i,i] - λ) v[i] = -sum_{k<i} T[k,i] v[k]
1328            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            // Normalize
1343            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
1361/// Computes a Householder vector for a 3-element (or smaller) vector.
1362fn 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    // `!= zero` (not `> zero`): `v_norm_sq` is NaN whenever `x` contains a
1396    // NaN (the un-gated sum-of-squares above propagates it correctly), and
1397    // `>` is always false for NaN — silently downgrading a poisoned norm to
1398    // the identity-reflector fallback (`tau = 0`) instead of honestly
1399    // returning a NaN `tau`.
1400    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        // Already upper triangular
1419        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        // Eigenvalues should be 1 and 3
1425        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        // Rotation matrix - has complex eigenvalues
1448        let theta = core::f64::consts::FRAC_PI_4; // 45 degrees
1449        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        // Eigenvalues should be cos(θ) ± i*sin(θ)
1457        assert_eq!(eigenvalues.len(), 2);
1458        // They should be complex conjugates
1459        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        // Check Q^T * Q = I
1494        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        // Upper triangular matrix - eigenvectors should be standard basis vectors
1573        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        // First eigenvector for λ=1 should be proportional to [1, 0, 0]
1578        assert!(v[(1, 0)].abs() < 1e-10);
1579        assert!(v[(2, 0)].abs() < 1e-10);
1580
1581        // Third eigenvector for λ=6 should be proportional to [*, *, 1]
1582        // (normalized, so third component is non-zero)
1583        assert!(v[(2, 2)].abs() > 0.1);
1584    }
1585
1586    #[test]
1587    fn test_trevc_right_diagonal() {
1588        // Diagonal matrix - eigenvectors are standard basis
1589        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        // Should be identity (or close to it with normalization)
1594        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        // Test that T * v = λ * v for computed eigenvectors
1608        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        // Eigenvalues are 1, 4, 6 (diagonal elements)
1613        let eigenvalues = [1.0, 4.0, 6.0];
1614
1615        for (j, &lambda) in eigenvalues.iter().enumerate() {
1616            // Compute T * v[:,j]
1617            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            // Check T * v = λ * v
1625            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        // Should be identity
1644        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        // Test eigenvectors through the Schur decomposition
1658        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 upper triangular matrix, A = T, so eigenvectors are the same
1664        // Check that A * v = λ * v
1665        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        // Create a 2x2 block with complex eigenvalues: rotation matrix
1690        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        // For complex eigenvalues, columns should contain real and imag parts
1698        // The eigenvector equation (T - λI)v = 0 where λ = c + i*s
1699        // Check that both columns are normalized
1700        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        // Diagonal matrix: left and right eigenvectors are standard basis
1714        // so s = |e_i^T e_i| = 1 for all eigenvalues
1715        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        // Upper triangular: eigenvalues are diagonal elements
1733        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        // All condition numbers should be positive
1739        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        // Diagonal matrix: separation is the minimum distance to other eigenvalues
1747        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        // λ=2: min dist to 3,5 is 1
1753        assert!(
1754            approx_eq(sep[0], 1.0, 1e-10),
1755            "sep[0] = {}, expected 1.0",
1756            sep[0]
1757        );
1758        // λ=5: min dist to 2,3 is 2
1759        assert!(
1760            approx_eq(sep[1], 2.0, 1e-10),
1761            "sep[1] = {}, expected 2.0",
1762            sep[1]
1763        );
1764        // λ=3: min dist to 2,5 is 1
1765        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        // Eigenvalues very close together should have small separation
1775        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        // Both should have separation close to 0.001
1781        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        // Test through Schur decomposition
1796        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        // Diagonal matrix should have well-conditioned eigenvalues
1803        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        // Test through Schur decomposition
1816        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        // Eigenvalues are 2 and 4, so separation should be 2
1823        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        // Rotation matrix with complex eigenvalues
1836        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        // Both eigenvalues in a complex pair should have the same condition
1848        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    /// A 5×5 non-symmetric matrix with tightly clustered eigenvalues. It is
1863    /// perfectly solvable, but a single QR sweep cannot possibly deflate it down
1864    /// to a size-2 window, so an artificially tiny iteration budget must be
1865    /// reported as a convergence failure — not silently returned as garbage.
1866    ///
1867    /// The eigenvalues cluster near 5 (this is `5*I` plus a small nilpotent-ish
1868    /// bidiagonal + coupling perturbation), which is exactly the kind of spectrum
1869    /// that makes the QR iteration work hardest.
1870    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        // A total budget of 1 QR sweep is nowhere near enough to reduce a 5×5
1885        // active window (which needs several sweeps and multiple deflations).
1886        // Before the fix this fell through and returned a partially-reduced T as
1887        // if valid; now it must surface NotConverged.
1888        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        // Control for the non-convergence test: the SAME matrix must converge
1899        // (and reconstruct as A = Q T Q^T) under the real default budget, proving
1900        // the error above is caused purely by the artificial cap, not a defect in
1901        // the matrix or the algorithm.
1902        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        // Trace is preserved by a similarity transform: sum of eigenvalue real
1918        // parts must equal the trace (= 25 here).
1919        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}