use mdarray::{Array, Dense, Dim, Layout, Shape, Slice};
use mdarray_linalg::svd::{SVD, SVDDecomp, SVDError};
use num_complex::ComplexFloat;
use super::{
scalar::{LapackScalar, NeedsRwork},
simple::gsvd,
};
use crate::Lapack;
impl<T, D> SVD<T, D> for Lapack
where
T: ComplexFloat + Default + LapackScalar + NeedsRwork,
T::Real: Into<T>,
D: Dim,
{
type SingularValue = T;
fn svd<L: Layout>(
&self,
a: &mut Slice<T, (D, D), L>,
) -> Result<SVDDecomp<T, Self::SingularValue, D>, SVDError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
let min_mn = m.min(n);
let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
let u_shape = <(D, D) as Shape>::from_dims(&[m, m]);
let vt_shape = <(D, D) as Shape>::from_dims(&[n, n]);
let mut s = Array::from_elem(s_shape, T::default());
let mut u = Array::from_elem(u_shape, T::default());
let mut vt = Array::from_elem(vt_shape, T::default());
match gsvd(
a,
&mut s,
Some(&mut u),
Some(&mut vt),
self.svd_config,
true,
) {
Ok(_) => Ok(SVDDecomp { s, u, vt }),
Err(e) => Err(e),
}
}
fn svd_thin<L: Layout>(
&self,
a: &mut Slice<T, (D, D), L>,
) -> Result<SVDDecomp<T, Self::SingularValue, D>, SVDError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
let min_mn = m.min(n);
let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
let u_shape = <(D, D) as Shape>::from_dims(&[m, min_mn]);
let vt_shape = <(D, D) as Shape>::from_dims(&[min_mn, n]);
let mut s = Array::from_elem(s_shape, T::default());
let mut u = Array::from_elem(u_shape, T::default());
let mut vt = Array::from_elem(vt_shape, T::default());
match gsvd(
a,
&mut s,
Some(&mut u),
Some(&mut vt),
self.svd_config,
false,
) {
Ok(_) => Ok(SVDDecomp { s, u, vt }),
Err(e) => Err(e),
}
}
fn svd_s<L: Layout>(
&self,
a: &mut Slice<T, (D, D), L>,
) -> Result<Array<Self::SingularValue, (D,)>, SVDError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
let min_mn = m.min(n);
let s_shape = <(D,) as Shape>::from_dims(&[min_mn]);
let mut s = Array::from_elem(s_shape, T::default());
match gsvd::<T, D, L, Dense, Dense, Dense>(a, &mut s, None, None, self.svd_config, false) {
Ok(_) => Ok(s),
Err(err) => Err(err),
}
}
fn svd_write<L: Layout, Ls: Layout, Lu: Layout, Lvt: Layout>(
&self,
a: &mut Slice<T, (D, D), L>,
s: &mut Slice<Self::SingularValue, (D,), Ls>,
u: &mut Slice<T, (D, D), Lu>,
vt: &mut Slice<T, (D, D), Lvt>,
) -> Result<(), SVDError> {
let compute_full_svd_vectors = u.shape().0 == u.shape().1;
gsvd(
a,
s,
Some(u),
Some(vt),
self.svd_config,
compute_full_svd_vectors,
)
}
fn svd_write_s<L: Layout, Ls: Layout>(
&self,
a: &mut Slice<T, (D, D), L>,
s: &mut Slice<Self::SingularValue, (D,), Ls>,
) -> Result<(), SVDError> {
gsvd::<T, D, L, Ls, Dense, Dense>(a, s, None, None, self.svd_config, false)
}
}