mdarray_linalg_lapack/svd/
context.rs1use mdarray::{Array, Dense, Dim, Layout, Shape, Slice};
11use mdarray_linalg::svd::{SVD, SVDDecomp, SVDError};
12use num_complex::ComplexFloat;
13
14use super::{
15 scalar::{LapackScalar, NeedsRwork},
16 simple::gsvd,
17};
18use crate::Lapack;
19
20impl<T, D> SVD<T, D> for Lapack
21where
22 T: ComplexFloat + Default + LapackScalar + NeedsRwork,
23 T::Real: Into<T>,
24 D: Dim,
25{
26 type SingularValue = T;
27
28 fn svd<L: Layout>(
30 &self,
31 a: &mut Slice<T, (D, D), L>,
32 ) -> Result<SVDDecomp<T, Self::SingularValue, D>, SVDError> {
33 let ash = *a.shape();
34 let (m, n) = (ash.dim(0), ash.dim(1));
35 let min_mn = m.min(n);
36
37 let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
38 let u_shape = <(D, D) as Shape>::from_dims(&[m, m]);
39 let vt_shape = <(D, D) as Shape>::from_dims(&[n, n]);
40
41 let mut s = Array::from_elem(s_shape, T::default());
42 let mut u = Array::from_elem(u_shape, T::default());
43 let mut vt = Array::from_elem(vt_shape, T::default());
44
45 match gsvd(
46 a,
47 &mut s,
48 Some(&mut u),
49 Some(&mut vt),
50 self.svd_config,
51 true,
52 ) {
53 Ok(_) => Ok(SVDDecomp { s, u, vt }),
54 Err(e) => Err(e),
55 }
56 }
57
58 fn svd_thin<L: Layout>(
60 &self,
61 a: &mut Slice<T, (D, D), L>,
62 ) -> Result<SVDDecomp<T, Self::SingularValue, D>, SVDError> {
63 let ash = *a.shape();
64 let (m, n) = (ash.dim(0), ash.dim(1));
65 let min_mn = m.min(n);
66
67 let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
68 let u_shape = <(D, D) as Shape>::from_dims(&[m, min_mn]);
69 let vt_shape = <(D, D) as Shape>::from_dims(&[min_mn, n]);
70
71 let mut s = Array::from_elem(s_shape, T::default());
72 let mut u = Array::from_elem(u_shape, T::default());
73 let mut vt = Array::from_elem(vt_shape, T::default());
74
75 match gsvd(
76 a,
77 &mut s,
78 Some(&mut u),
79 Some(&mut vt),
80 self.svd_config,
81 false,
82 ) {
83 Ok(_) => Ok(SVDDecomp { s, u, vt }),
84 Err(e) => Err(e),
85 }
86 }
87
88 fn svd_s<L: Layout>(
90 &self,
91 a: &mut Slice<T, (D, D), L>,
92 ) -> Result<Array<Self::SingularValue, (D,)>, SVDError> {
93 let ash = *a.shape();
94 let (m, n) = (ash.dim(0), ash.dim(1));
95
96 let min_mn = m.min(n);
97
98 let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
100 let mut s = Array::from_elem(s_shape, T::default());
101
102 match gsvd::<T, D, L, Dense, Dense, Dense>(a, &mut s, None, None, self.svd_config, false) {
103 Ok(_) => Ok(s),
104 Err(err) => Err(err),
105 }
106 }
107
108 fn svd_write<L: Layout, Ls: Layout, Lu: Layout, Lvt: Layout>(
110 &self,
111 a: &mut Slice<T, (D, D), L>,
112 s: &mut Slice<Self::SingularValue, (D,), Ls>,
113 u: &mut Slice<T, (D, D), Lu>,
114 vt: &mut Slice<T, (D, D), Lvt>,
115 ) -> Result<(), SVDError> {
116 let compute_full_svd_vectors = u.shape().0 == u.shape().1;
117 gsvd(
118 a,
119 s,
120 Some(u),
121 Some(vt),
122 self.svd_config,
123 compute_full_svd_vectors,
124 )
125 }
126
127 fn svd_write_s<L: Layout, Ls: Layout>(
129 &self,
130 a: &mut Slice<T, (D, D), L>,
131 s: &mut Slice<Self::SingularValue, (D,), Ls>,
132 ) -> Result<(), SVDError> {
133 gsvd::<T, D, L, Ls, Dense, Dense>(a, s, None, None, self.svd_config, false)
134 }
135}