1use 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 let re: <$t as ComplexFloat>::Real = unsafe { std::mem::transmute_copy(&($val.re)) };
35 let im: <$t as ComplexFloat>::Real = unsafe { std::mem::transmute_copy(&($val.im)) };
37 Complex::new(re, im)
38 }};
39}
40
41fn 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 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 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 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 let re: <T as ComplexFloat>::Real = unsafe {
269 std::mem::transmute_copy(&s_re_mda[i])
270 };
271 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 let vr: <T as ComplexFloat>::Real = unsafe {
286 std::mem::transmute_copy(&right_vecs_tmp[[i, j]])
287 };
288 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 let re_right: <T as ComplexFloat>::Real = unsafe {
302 std::mem::transmute_copy(&right_vecs_tmp[[i, j]])
303 };
304 let im_right: <T as ComplexFloat>::Real = unsafe {
306 std::mem::transmute_copy(&right_vecs_tmp[[i, j + 1]])
307 };
308 let re_left: <T as ComplexFloat>::Real = unsafe {
310 std::mem::transmute_copy(&left_vecs_tmp[[i, j]])
311 };
312 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 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 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 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 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 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 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 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}