1use faer_traits::ComplexField;
11use mdarray::{Array, Dense, Dim, Layout, Shape, Slice};
12use mdarray_linalg::svd::{SVD, SVDDecomp, SVDError};
13use num_complex::ComplexFloat;
14
15use super::simple::svd_faer;
16use crate::Faer;
17
18impl<T, D> SVD<T, D> for Faer
19where
20 T: ComplexFloat + ComplexField + Default,
21 D: Dim,
22{
23 type SingularValue = T;
24
25 fn svd<L: Layout>(
27 &self,
28 a: &mut Slice<T, (D, D), L>,
29 ) -> Result<SVDDecomp<T, Self::SingularValue, D>, SVDError> {
30 let ash = *a.shape();
31 let (m, n) = (ash.dim(0), ash.dim(1));
32
33 let min_mn = m.min(n);
34
35 let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
36 let u_shape = <(D, D) as Shape>::from_dims(&[m, m]);
37 let vt_shape = <(D, D) as Shape>::from_dims(&[n, n]);
38
39 let mut s_mda = Array::from_elem(s_shape, T::default());
40 let mut u_mda = Array::from_elem(u_shape, T::default());
41 let mut vt_mda = Array::from_elem(vt_shape, T::default());
42
43 match svd_faer(a, &mut s_mda, Some(&mut u_mda), Some(&mut vt_mda), true) {
58 Err(_) => Err(SVDError::BackendDidNotConverge {
59 superdiagonals: (0),
60 }),
61 Ok(_) => Ok(SVDDecomp {
62 s: s_mda,
63 u: u_mda,
64 vt: vt_mda,
65 }),
66 }
67 }
68
69 fn svd_thin<L: Layout>(
71 &self,
72 a: &mut Slice<T, (D, D), L>,
73 ) -> Result<SVDDecomp<T, Self::SingularValue, D>, SVDError> {
74 let ash = *a.shape();
75 let (m, n) = (ash.dim(0), ash.dim(1));
76
77 let min_mn = m.min(n);
78
79 let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
80 let u_shape = <(D, D) as Shape>::from_dims(&[m, m]);
81 let vt_shape = <(D, D) as Shape>::from_dims(&[n, n]);
82
83 let mut s_mda = Array::from_elem(s_shape, T::default());
84 let mut u_mda = Array::from_elem(u_shape, T::default());
85 let mut vt_mda = Array::from_elem(vt_shape, T::default());
86
87 match svd_faer(a, &mut s_mda, Some(&mut u_mda), Some(&mut vt_mda), false) {
88 Err(_) => Err(SVDError::BackendDidNotConverge {
89 superdiagonals: (0),
90 }),
91 Ok(_) => Ok(SVDDecomp {
92 s: s_mda,
93 u: u_mda,
94 vt: vt_mda,
95 }),
96 }
97 }
98
99 fn svd_s<L: Layout>(
101 &self,
102 a: &mut Slice<T, (D, D), L>,
103 ) -> Result<Array<Self::SingularValue, (D,)>, SVDError> {
104 let ash = *a.shape();
105 let (m, n) = (ash.dim(0), ash.dim(1));
106
107 let min_mn = m.min(n);
108
109 let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
110 let mut s_mda = Array::from_elem(s_shape, T::default());
111
112 match svd_faer::<T, D, L, Dense, Dense, Dense>(a, &mut s_mda, None, None, false) {
117 Err(_) => Err(SVDError::BackendDidNotConverge {
118 superdiagonals: (0),
119 }),
120 Ok(_) => Ok(s_mda),
121 }
122 }
123
124 fn svd_write<L: Layout, Ls: Layout, Lu: Layout, Lvt: Layout>(
126 &self,
127 a: &mut Slice<T, (D, D), L>,
128 s: &mut Slice<Self::SingularValue, (D,), Ls>,
129 u: &mut Slice<T, (D, D), Lu>,
130 vt: &mut Slice<T, (D, D), Lvt>,
131 ) -> Result<(), SVDError> {
132 let compute_svd_full_vectors = u.shape().0 == u.shape().1;
133 svd_faer::<T, D, L, Ls, Lu, Lvt>(a, s, Some(u), Some(vt), compute_svd_full_vectors)
134 }
135
136 fn svd_write_s<L: Layout, Ls: Layout>(
138 &self,
139 a: &mut Slice<T, (D, D), L>,
140 s: &mut Slice<Self::SingularValue, (D,), Ls>,
141 ) -> Result<(), SVDError> {
142 svd_faer::<T, D, L, Ls, Dense, Dense>(a, s, None, None, false)
143 }
144}