use faer_traits::ComplexField;
use mdarray::{Array, Dense, Dim, Layout, Shape, Slice};
use mdarray_linalg::svd::{SVD, SVDDecomp, SVDError};
use num_complex::ComplexFloat;
use super::simple::svd_faer;
use crate::Faer;
impl<T, D> SVD<T, D> for Faer
where
T: ComplexFloat + ComplexField + Default,
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_mda = Array::from_elem(s_shape, T::default());
let mut u_mda = Array::from_elem(u_shape, T::default());
let mut vt_mda = Array::from_elem(vt_shape, T::default());
match svd_faer(a, &mut s_mda, Some(&mut u_mda), Some(&mut vt_mda), true) {
Err(_) => Err(SVDError::BackendDidNotConverge {
superdiagonals: (0),
}),
Ok(_) => Ok(SVDDecomp {
s: s_mda,
u: u_mda,
vt: vt_mda,
}),
}
}
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, m]);
let vt_shape = <(D, D) as Shape>::from_dims(&[n, n]);
let mut s_mda = Array::from_elem(s_shape, T::default());
let mut u_mda = Array::from_elem(u_shape, T::default());
let mut vt_mda = Array::from_elem(vt_shape, T::default());
match svd_faer(a, &mut s_mda, Some(&mut u_mda), Some(&mut vt_mda), false) {
Err(_) => Err(SVDError::BackendDidNotConverge {
superdiagonals: (0),
}),
Ok(_) => Ok(SVDDecomp {
s: s_mda,
u: u_mda,
vt: vt_mda,
}),
}
}
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_mda = Array::from_elem(s_shape, T::default());
match svd_faer::<T, D, L, Dense, Dense, Dense>(a, &mut s_mda, None, None, false) {
Err(_) => Err(SVDError::BackendDidNotConverge {
superdiagonals: (0),
}),
Ok(_) => Ok(s_mda),
}
}
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_svd_full_vectors = u.shape().0 == u.shape().1;
svd_faer::<T, D, L, Ls, Lu, Lvt>(a, s, Some(u), Some(vt), compute_svd_full_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> {
svd_faer::<T, D, L, Ls, Dense, Dense>(a, s, None, None, false)
}
}