Skip to main content

mdarray_linalg_faer/
eig.rs

1// Eigenvalue Decomposition:
2//     A * V = V * Λ  (right eigenvectors)
3//     W^H * A = Λ * W^H  (left eigenvectors)
4// where:
5//     - A is n × n         (input square matrix)
6//     - V is n × n         (right eigenvectors as columns)
7//     - W is n × n         (left eigenvectors as columns)
8//     - Λ is n × n         (diagonal matrix with eigenvalues)
9//
10// For Hermitian/Symmetric matrices:
11//     A = Q * Λ * Q^H
12// where:
13//     - Q is n × n         (orthogonal/unitary eigenvectors)
14//     - Λ is n × n         (diagonal matrix with real eigenvalues)
15//
16// Schur Decomposition:
17//     A = Z * T * Z^H
18// where:
19//     - Z is n × n         (unitary Schur vectors)
20//     - T is n × n         (upper triangular for complex, quasi-upper triangular for real)
21
22use dyn_stack::{MemBuffer, MemStack};
23use faer_traits::ComplexField;
24use mdarray::{Array, Dense, Dim, Layout, Shape, Slice};
25use mdarray_linalg::eig::{Eig, EigDecomp, EigError, EighDecomp, SchurDecomp, SchurError};
26use num_complex::{Complex, ComplexFloat};
27
28use crate::{Faer, into_faer, into_faer_diag_mut, into_faer_mut};
29
30macro_rules! complex_from_faer {
31    ($val:expr, $t:ty) => {{
32        // SAFETY: for the scalar types supported by this backend, faer and num_complex use the
33        // same real component type, so this is a bitwise-preserving reinterpretation.
34        let re: <$t as ComplexFloat>::Real = unsafe { std::mem::transmute_copy(&($val.re)) };
35        // SAFETY: same rationale as above for the imaginary component.
36        let im: <$t as ComplexFloat>::Real = unsafe { std::mem::transmute_copy(&($val.im)) };
37        Complex::new(re, im)
38    }};
39}
40
41// Faer exposes the Hessenberg reduction publicly, but not the final Schur QR step.
42// We keep the same A = Z * T * Z^H interface with the reduced Hessenberg form.
43fn schur_faer_in_place<T, D0: Dim, D1: Dim, L: Layout, Lz: Layout>(
44    t: &mut Slice<T, (D0, D1), L>,
45    z: &mut Slice<T, (D0, D1), Lz>,
46) -> Result<(), SchurError>
47where
48    T: ComplexFloat
49        + ComplexField
50        + Default
51        + std::convert::From<<T as num_complex::ComplexFloat>::Real>,
52{
53    let ash = *t.shape();
54    let (m, n) = (ash.dim(0), ash.dim(1));
55
56    if m != n {
57        return Err(SchurError::NotSquareMatrix);
58    }
59
60    for i in 0..n {
61        for j in 0..n {
62            z[[i, j]] = if i == j { T::one() } else { T::zero() };
63        }
64    }
65
66    if n <= 1 {
67        return Ok(());
68    }
69
70    let par = faer::get_global_parallelism();
71    let bs = faer::linalg::qr::no_pivoting::factor::recommended_block_size::<T>(n - 1, n - 1);
72    let mut householder = faer::Mat::<T>::zeros(bs, n - 1);
73
74    {
75        let mut t_faer = into_faer_mut(t);
76        faer::linalg::evd::hessenberg::hessenberg_in_place(
77            t_faer.as_mut(),
78            householder.as_mut(),
79            par,
80            MemStack::new(&mut MemBuffer::new(
81                faer::linalg::evd::hessenberg::hessenberg_in_place_scratch::<T>(
82                    n,
83                    bs,
84                    par,
85                    faer::prelude::default(),
86                ),
87            )),
88            faer::prelude::default(),
89        );
90    }
91
92    {
93        let t_faer = into_faer(t);
94        let mut z_faer = into_faer_mut(z);
95        faer::linalg::householder::apply_block_householder_sequence_on_the_right_in_place_with_conj(
96            t_faer.submatrix(1, 0, n - 1, n - 1),
97            householder.as_ref(),
98            faer::Conj::No,
99            z_faer.as_mut().submatrix_mut(1, 1, n - 1, n - 1),
100            par,
101            MemStack::new(&mut MemBuffer::new(
102                faer::linalg::householder::apply_block_householder_sequence_on_the_right_in_place_scratch::<T>(
103                    n - 1,
104                    bs,
105                    n - 1,
106                ),
107            )),
108        );
109    }
110
111    for j in 0..n {
112        for i in j + 2..n {
113            t[[i, j]] = T::zero();
114        }
115    }
116
117    Ok(())
118}
119
120fn swap_matrices<T, D0: Dim, D1: Dim, L0: Layout, L1: Layout>(
121    a: &mut Slice<T, (D0, D1), L0>,
122    b: &mut Slice<T, (D0, D1), L1>,
123) {
124    let ash = *a.shape();
125    let (m, n) = (ash.dim(0), ash.dim(1));
126
127    for i in 0..m {
128        for j in 0..n {
129            std::mem::swap(&mut a[[i, j]], &mut b[[i, j]]);
130        }
131    }
132}
133
134impl<T, D0: Dim, D1: Dim> Eig<T, D0, D1> for Faer
135where
136    T: ComplexFloat
137        + ComplexField
138        + Default
139        + std::convert::From<<T as num_complex::ComplexFloat>::Real>,
140    Complex<<T as ComplexFloat>::Real>: ComplexFloat
141        + ComplexField
142        + Default
143        + std::convert::From<<T as ComplexFloat>::Real>
144        + std::convert::From<<Complex<<T as ComplexFloat>::Real> as ComplexFloat>::Real>,
145{
146    type SpectralScalar = Complex<<T as ComplexFloat>::Real>;
147    type RealScalar = <T as ComplexFloat>::Real;
148
149    /// Compute eigenvalues and right eigenvectors with new allocated matrices
150    /// The matrix `A` satisfies: `A * v = λ * v` where v are the right eigenvectors
151    fn eig<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<EigDecomp<Self::SpectralScalar, D0, D1>, EigError> {
152        let ash = *a.shape();
153        let (m, n) = (ash.dim(0), ash.dim(1));
154
155        if m != n {
156            return Err(EigError::NotSquareMatrix);
157        }
158
159        let a_faer = into_faer(a);
160        let eig_result = a_faer.eigen();
161
162        match eig_result {
163            Ok(eig) => {
164                let eigenvalues = eig.S();
165                let right_vecs = eig.U();
166
167                let x = T::default();
168                let ash1 = <(D0,) as Shape>::from_dims(&[n]);
169                let mut eigenvalues_mda = Array::from_elem(ash1, Complex::new(x.re(), x.re()));
170                let mut right_vecs_mda = Array::from_elem(ash, Complex::new(x.re(), x.re()));
171
172                for i in 0..n {
173                    eigenvalues_mda[i] = complex_from_faer!(&eigenvalues[i], T);
174                }
175
176                for i in 0..n {
177                    for j in 0..n {
178                        right_vecs_mda[[i, j]] = complex_from_faer!(&right_vecs[(i, j)], T);
179                    }
180                }
181
182                Ok(EigDecomp {
183                    eigenvalues: eigenvalues_mda,
184                    left_eigenvectors: None,
185                    right_eigenvectors: Some(right_vecs_mda),
186                })
187            }
188            Err(_) => Err(EigError::BackendDidNotConverge { iterations: 0 }),
189        }
190    }
191
192    fn eig_full<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<EigDecomp<Self::SpectralScalar, D0, D1>, EigError> {
193        let ash = *a.shape();
194        let (m, n) = (ash.dim(0), ash.dim(1));
195
196        if m != n {
197            return Err(EigError::NotSquareMatrix);
198        }
199
200        let par = faer::get_global_parallelism();
201        let x = T::default();
202        let xr = x.re();
203        // SAFETY: for the scalar types supported by this backend, both traits expose the same
204        // concrete real scalar, so this preserves the bit pattern without changing layout.
205        let xr_faer: <T as faer_traits::ComplexField>::Real = unsafe {
206            std::mem::transmute_copy(&xr)
207        };
208        let ash1 = <(D0,) as Shape>::from_dims(&[n]);
209        let mut eigenvalues_mda = Array::from_elem(ash1, Complex::new(xr, xr));
210        let a_faer = into_faer(a);
211
212        if T::IS_REAL {
213            let mut s_re_mda = Array::<<T as faer_traits::ComplexField>::Real, (D0,)>::from_elem(
214                ash1,
215                xr_faer.clone(),
216            );
217            let mut s_im_mda = Array::<<T as faer_traits::ComplexField>::Real, (D0,)>::from_elem(
218                ash1,
219                xr_faer.clone(),
220            );
221            let mut left_vecs_tmp = Array::<
222                <T as faer_traits::ComplexField>::Real,
223                (D0, D1),
224            >::from_elem(ash, xr_faer.clone());
225            let mut right_vecs_tmp = Array::<
226                <T as faer_traits::ComplexField>::Real,
227                (D0, D1),
228            >::from_elem(ash, xr_faer.clone());
229            let mut left_vecs_mda = Array::from_elem(ash, Complex::new(xr, xr));
230            let mut right_vecs_mda = Array::from_elem(ash, Complex::new(xr, xr));
231
232            let a_faer_real: faer::MatRef<'_, <T as faer_traits::ComplexField>::Real> = unsafe {
233                // SAFETY: this branch is only taken for real scalar types, so `T` has the same
234                // in-memory representation as its real component and the matrix layout is unchanged.
235                faer::hacks::coerce::<_, faer::MatRef<'_, <T as faer_traits::ComplexField>::Real>>(
236                    a_faer,
237                )
238            };
239
240            let params = <faer::linalg::evd::EvdParams as faer::Auto<
241                <T as faer_traits::ComplexField>::Real,
242            >>::auto();
243
244            let result = faer::linalg::evd::evd_real::<<T as faer_traits::ComplexField>::Real>(
245                a_faer_real,
246                into_faer_diag_mut(&mut s_re_mda),
247                into_faer_diag_mut(&mut s_im_mda),
248                Some(into_faer_mut(&mut left_vecs_tmp)),
249                Some(into_faer_mut(&mut right_vecs_tmp)),
250                par,
251                MemStack::new(&mut MemBuffer::new(faer::linalg::evd::evd_scratch::<
252                    <T as faer_traits::ComplexField>::Real,
253                >(
254                    n,
255                    faer::linalg::evd::ComputeEigenvectors::Yes,
256                    faer::linalg::evd::ComputeEigenvectors::Yes,
257                    par,
258                    params.into(),
259                ))),
260                params.into(),
261            );
262
263            match result {
264                Ok(_) => {
265                    for i in 0..n {
266                        // SAFETY: the temporary storage uses faer's real scalar, which matches
267                        // num_complex's real scalar for the supported backend types.
268                        let re: <T as ComplexFloat>::Real = unsafe {
269                            std::mem::transmute_copy(&s_re_mda[i])
270                        };
271                        // SAFETY: same rationale as above for the imaginary part.
272                        let im: <T as ComplexFloat>::Real = unsafe {
273                            std::mem::transmute_copy(&s_im_mda[i])
274                        };
275                        eigenvalues_mda[i] = Complex::new(re, im);
276                    }
277
278                    let mut j = 0_usize;
279                    while j < n {
280                        let imag_is_zero = s_im_mda[j] == xr_faer;
281                        if imag_is_zero {
282                            for i in 0..n {
283                                // SAFETY: the temporary eigenvector matrices are real in this
284                                // branch, and their scalar matches the public real scalar type.
285                                let vr: <T as ComplexFloat>::Real = unsafe {
286                                    std::mem::transmute_copy(&right_vecs_tmp[[i, j]])
287                                };
288                                // SAFETY: same rationale as above for the left eigenvector entry.
289                                let vl: <T as ComplexFloat>::Real = unsafe {
290                                    std::mem::transmute_copy(&left_vecs_tmp[[i, j]])
291                                };
292                                right_vecs_mda[[i, j]] = Complex::new(vr, xr);
293                                left_vecs_mda[[i, j]] = Complex::new(vl, xr);
294                            }
295                            j += 1;
296                        } else {
297                            for i in 0..n {
298                                // SAFETY: LAPACK-like real output convention from faer stores the
299                                // real and imaginary parts in adjacent real matrices. The scalar
300                                // type matches the public real scalar for supported backends.
301                                let re_right: <T as ComplexFloat>::Real = unsafe {
302                                    std::mem::transmute_copy(&right_vecs_tmp[[i, j]])
303                                };
304                                // SAFETY: same rationale as above.
305                                let im_right: <T as ComplexFloat>::Real = unsafe {
306                                    std::mem::transmute_copy(&right_vecs_tmp[[i, j + 1]])
307                                };
308                                // SAFETY: same rationale as above.
309                                let re_left: <T as ComplexFloat>::Real = unsafe {
310                                    std::mem::transmute_copy(&left_vecs_tmp[[i, j]])
311                                };
312                                // SAFETY: same rationale as above.
313                                let im_left: <T as ComplexFloat>::Real = unsafe {
314                                    std::mem::transmute_copy(&left_vecs_tmp[[i, j + 1]])
315                                };
316
317                                right_vecs_mda[[i, j]] = Complex::new(re_right, im_right);
318                                right_vecs_mda[[i, j + 1]] = Complex::new(re_right, -im_right);
319                                left_vecs_mda[[i, j]] = Complex::new(re_left, im_left);
320                                left_vecs_mda[[i, j + 1]] = Complex::new(re_left, -im_left);
321                            }
322                            j += 2;
323                        }
324                    }
325
326                    Ok(EigDecomp {
327                        eigenvalues: eigenvalues_mda,
328                        left_eigenvectors: Some(left_vecs_mda),
329                        right_eigenvectors: Some(right_vecs_mda),
330                    })
331                }
332                Err(_) => Err(EigError::BackendDidNotConverge { iterations: 0 }),
333            }
334        } else {
335            let mut eigenvalues_tmp = Array::<
336                Complex<<T as faer_traits::ComplexField>::Real>,
337                (D0,),
338            >::from_elem(ash1, Complex::new(xr_faer.clone(), xr_faer.clone()));
339            let mut left_vecs_tmp = Array::<
340                Complex<<T as faer_traits::ComplexField>::Real>,
341                (D0, D1),
342            >::from_elem(ash, Complex::new(xr_faer.clone(), xr_faer.clone()));
343            let mut right_vecs_tmp = Array::<
344                Complex<<T as faer_traits::ComplexField>::Real>,
345                (D0, D1),
346            >::from_elem(ash, Complex::new(xr_faer.clone(), xr_faer.clone()));
347            let mut left_vecs_mda = Array::from_elem(ash, Complex::new(xr, xr));
348            let mut right_vecs_mda = Array::from_elem(ash, Complex::new(xr, xr));
349
350            let a_faer_cplx: faer::MatRef<'_, Complex<<T as faer_traits::ComplexField>::Real>> = unsafe {
351                // SAFETY: this branch is only taken for complex scalar types, so `T` has the same
352                // in-memory representation as `num_complex::Complex<Real>` expected by faer here.
353                faer::hacks::coerce::<_, faer::MatRef<'_, Complex<<T as faer_traits::ComplexField>::Real>>>(
354                    a_faer,
355                )
356            };
357
358            let params = <faer::linalg::evd::EvdParams as faer::Auto<
359                Complex<<T as faer_traits::ComplexField>::Real>,
360            >>::auto();
361
362            let result = faer::linalg::evd::evd_cplx::<<T as faer_traits::ComplexField>::Real>(
363                a_faer_cplx,
364                into_faer_diag_mut(&mut eigenvalues_tmp),
365                Some(into_faer_mut(&mut left_vecs_tmp)),
366                Some(into_faer_mut(&mut right_vecs_tmp)),
367                par,
368                MemStack::new(&mut MemBuffer::new(faer::linalg::evd::evd_scratch::<
369                    Complex<<T as faer_traits::ComplexField>::Real>,
370                >(
371                    n,
372                    faer::linalg::evd::ComputeEigenvectors::Yes,
373                    faer::linalg::evd::ComputeEigenvectors::Yes,
374                    par,
375                    params.into(),
376                ))),
377                params.into(),
378            );
379
380            match result {
381                Ok(_) => {
382                    for i in 0..n {
383                        eigenvalues_mda[i] = complex_from_faer!(&eigenvalues_tmp[i], T);
384                        for j in 0..n {
385                            left_vecs_mda[[i, j]] = complex_from_faer!(&left_vecs_tmp[[i, j]], T);
386                            right_vecs_mda[[i, j]] = complex_from_faer!(&right_vecs_tmp[[i, j]], T);
387                        }
388                    }
389
390                    Ok(EigDecomp {
391                        eigenvalues: eigenvalues_mda,
392                        left_eigenvectors: Some(left_vecs_mda),
393                        right_eigenvectors: Some(right_vecs_mda),
394                    })
395                }
396                Err(_) => Err(EigError::BackendDidNotConverge { iterations: 0 }),
397            }
398        }
399    }
400
401    /// Compute only eigenvalues with new allocated vectors
402    fn eig_values<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<Array<Self::SpectralScalar, (D0,)>, EigError> {
403        let ash = *a.shape();
404        let (m, n) = (ash.dim(0), ash.dim(1));
405
406        if m != n {
407            return Err(EigError::NotSquareMatrix);
408        }
409
410        let a_faer = into_faer(a);
411
412        let eigenvalues_result = a_faer.eigenvalues();
413
414        match eigenvalues_result {
415            Ok(eigenvalues) => {
416                let x = T::default();
417                let ash1 = <(D0,) as Shape>::from_dims(&[n]);
418                let mut eigenvalues_mda = Array::from_elem(ash1, Complex::new(x.re(), x.re()));
419
420                for i in 0..n {
421                    eigenvalues_mda[i] = complex_from_faer!(&eigenvalues[i], T);
422                }
423
424                Ok(eigenvalues_mda)
425            }
426            Err(_) => Err(EigError::BackendDidNotConverge { iterations: 0 }),
427        }
428    }
429
430    /// Compute eigenvalues and eigenvectors of a Hermitian matrix (input should be complex)
431    fn eigh<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<EighDecomp<T, Self::RealScalar, D0, D1>, EigError> {
432        let ash = *a.shape();
433        let (m, n) = (ash.dim(0), ash.dim(1));
434
435        if m != n {
436            return Err(EigError::NotSquareMatrix);
437        }
438
439        let a_faer = into_faer(a);
440        let eig_result = a_faer.self_adjoint_eigen(faer::Side::Lower);
441
442        match eig_result {
443            Ok(eig) => {
444                let eigenvalues = eig.S();
445                let eigenvectors = eig.U();
446
447                let x = T::default();
448                let ash1 = <(D0,) as Shape>::from_dims(&[n]);
449                let mut eigenvalues_mda = Array::from_elem(ash1, x.re());
450                let mut eigenvectors_mda = Array::from_elem(ash, T::default());
451
452                for i in 0..n {
453                    eigenvalues_mda[i] = eigenvalues[i].re();
454                }
455
456                let mut eigenvectors_faer = into_faer_mut(&mut eigenvectors_mda);
457                for i in 0..n {
458                    for j in 0..n {
459                        eigenvectors_faer[(i, j)] = eigenvectors[(i, j)];
460                    }
461                }
462
463                Ok(EighDecomp {
464                    eigenvalues: eigenvalues_mda,
465                    eigenvectors: eigenvectors_mda,
466                })
467            }
468            Err(_) => Err(EigError::BackendDidNotConverge { iterations: 0 }),
469        }
470    }
471
472    /// Compute Schur decomposition with new allocated matrices
473    fn schur<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<SchurDecomp<T, D0, D1>, SchurError> {
474        let ash = *a.shape();
475        let (m, n) = (ash.dim(0), ash.dim(1));
476
477        if m != n {
478            return Err(SchurError::NotSquareMatrix);
479        }
480
481        let mut t = a.to_tensor();
482        let mut z = Array::from_elem(ash, T::zero());
483        schur_faer_in_place(&mut t, &mut z)?;
484
485        Ok(SchurDecomp { t, z })
486    }
487
488    /// Compute Schur decomposition overwriting existing matrices
489    fn schur_write<L: Layout>(
490        &self,
491        a: &mut Slice<T, (D0, D1), L>,
492        t: &mut Slice<T, (D0, D1), Dense>,
493        z: &mut Slice<T, (D0, D1), Dense>,
494    ) -> Result<(), SchurError> {
495        schur_faer_in_place(a, z)?;
496        swap_matrices(a, t);
497        Ok(())
498    }
499
500    /// Compute Schur (complex) decomposition with new allocated matrices
501    fn schur_complex<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<SchurDecomp<Self::SpectralScalar, D0, D1>, SchurError> {
502        let ash = *a.shape();
503        let (m, n) = (ash.dim(0), ash.dim(1));
504
505        if m != n {
506            return Err(SchurError::NotSquareMatrix);
507        }
508
509        let zero = T::default().re();
510        let shape = <(D0, D1) as Shape>::from_dims(&[m, n]);
511        let mut t = Array::from_fn(shape, |idx| {
512            let x = a[idx];
513            Complex::new(x.re(), x.im())
514        });
515        let mut z = Array::from_elem(shape, Complex::new(zero, zero));
516        schur_faer_in_place(&mut t, &mut z)?;
517
518        Ok(SchurDecomp { t, z })
519    }
520
521    /// Compute Schur (complex) decomposition overwriting existing matrices
522    fn schur_complex_write<L: Layout>(
523        &self,
524        a: &mut Slice<T, (D0, D1), L>,
525        t: &mut Slice<Self::SpectralScalar, (D0, D1), Dense>,
526        z: &mut Slice<Self::SpectralScalar, (D0, D1), Dense>,
527    ) -> Result<(), SchurError> {
528        let SchurDecomp { t: t_result, z: z_result } = self.schur_complex(a)?;
529        for (dst, src) in t.iter_mut().zip(t_result.iter()) {
530            *dst = *src;
531        }
532        for (dst, src) in z.iter_mut().zip(z_result.iter()) {
533            *dst = *src;
534        }
535        Ok(())
536    }
537}