Skip to main content

mdarray_linalg_lapack/svd/
context.rs

1//! Singular Value Decomposition (SVD):
2//!     A = U * Σ * V^T
3//! where:
4//!     - A is m × n         (input matrix)
5//!     - U is m × m         (left singular vectors, orthogonal)
6//!     - Σ is µ × µ         (diagonal matrix with singular values on the diagonal, µ = min(m,n))
7//!     - V^T is n × n       (transpose of right singular vectors, orthogonal)
8//!     - s (Σ) contains min(m, n) singular values (non-negative, sorted in descending order) in the first row
9
10use 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    // Computes full SVD with new allocated matrices
29    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    // Computes thin SVD with new allocated matrices
59    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    // Computes only singular values with new allocated matrix
89    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        // Only allocate space for singular values
99        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    // Computes full SVD, overwriting existing matrices
109    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    // Computes only singular values, overwriting existing matrix
128    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}