use dyn_stack::{MemBuffer, MemStack};
use faer_traits::ComplexField;
use mdarray::{Array, Dim, Layout, Shape, Slice};
use mdarray_linalg::lu::{InvError, LU};
use num_complex::ComplexFloat;
use super::simple::lu_faer;
use crate::{Faer, into_faer_mut};
fn map_cholesky_error(err: faer::linalg::cholesky::llt::factor::LltError) -> InvError {
match err {
faer::linalg::cholesky::llt::factor::LltError::NonPositivePivot { index } => {
InvError::NotPositiveDefinite {
lpm: index as i32 + 1,
}
}
}
}
impl<T, D0: Dim, D1: Dim> LU<T, D0, D1> for Faer
where
T: ComplexFloat
+ ComplexField
+ Default
+ std::convert::From<<T as num_complex::ComplexFloat>::Real>,
{
fn lu<L: Layout>(
&self,
a: &mut Slice<T, (D0, D1), L>,
) -> (Array<T, (D0, D0)>, Array<T, (D0, D1)>, Array<T, (D0, D0)>) {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
let min_mn = m.min(n);
let l_shape = <(D0, D0) as Shape>::from_dims(&[m, min_mn]);
let u_shape = <(D0, D1) as Shape>::from_dims(&[min_mn, n]);
let p_shape = <(D0, D0) as Shape>::from_dims(&[m, m]);
let mut l_mda = Array::from_elem(l_shape, T::default());
let mut u_mda = Array::from_elem(u_shape, T::default());
let mut p_mda = Array::from_elem(p_shape, T::default());
lu_faer(a, &mut l_mda, &mut u_mda, &mut p_mda);
(l_mda, u_mda, p_mda)
}
fn lu_write<L: Layout, Ll: Layout, Lu: Layout, Lp: Layout>(
&self,
a: &mut Slice<T, (D0, D1), L>,
l: &mut Slice<T, (D0, D0), Ll>,
u: &mut Slice<T, (D0, D1), Lu>,
p: &mut Slice<T, (D0, D0), Lp>,
) {
lu_faer::<T, D0, D1, L, Ll, Lu, Lp>(a, l, u, p);
}
fn inv<L: Layout>(
&self,
a: &mut Slice<T, (D0, D1), L>,
) -> Result<Array<T, (D0, D1)>, InvError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
if m != n {
return Err(InvError::NotSquare {
rows: m as i32,
cols: n as i32,
});
}
let par = faer::get_global_parallelism();
let mut a_faer = into_faer_mut(a);
let mut row_perm_fwd = vec![0usize; m];
let mut row_perm_bwd = vec![0usize; m];
faer::linalg::lu::partial_pivoting::factor::lu_in_place(
a_faer.as_mut(),
&mut row_perm_fwd,
&mut row_perm_bwd,
par,
MemStack::new(&mut MemBuffer::new(
faer::linalg::lu::partial_pivoting::factor::lu_in_place_scratch::<usize, T>(
m,
n,
par,
faer::prelude::default(),
),
)),
faer::prelude::default(),
);
let l_mat = a_faer.as_ref();
let u_mat = a_faer.as_ref();
let perm = unsafe {
faer::perm::Perm::new_unchecked(
row_perm_fwd.into_boxed_slice(),
row_perm_bwd.into_boxed_slice(),
)
};
let mut inv_mat = Array::<T, (D0, D1)>::from_elem(ash, T::zero());
let mut inv_mat_faer = into_faer_mut(&mut inv_mat);
faer::linalg::lu::partial_pivoting::inverse::inverse(
inv_mat_faer.as_mut(),
l_mat,
u_mat,
perm.as_ref(),
par,
MemStack::new(&mut MemBuffer::new(
faer::linalg::lu::partial_pivoting::inverse::inverse_scratch::<usize, T>(m, par),
)),
);
Ok(inv_mat)
}
fn inv_write<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<(), InvError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
if m != n {
return Err(InvError::NotSquare {
rows: m as i32,
cols: n as i32,
});
}
let par = faer::get_global_parallelism();
let mut a_faer = into_faer_mut(a);
let mut row_perm_fwd = vec![0usize; m];
let mut row_perm_bwd = vec![0usize; m];
faer::linalg::lu::partial_pivoting::factor::lu_in_place(
a_faer.as_mut(),
&mut row_perm_fwd,
&mut row_perm_bwd,
par,
MemStack::new(&mut MemBuffer::new(
faer::linalg::lu::partial_pivoting::factor::lu_in_place_scratch::<usize, T>(
m,
n,
par,
faer::prelude::default(),
),
)),
faer::prelude::default(),
);
let l_mat = a_faer.as_ref();
let u_mat = a_faer.as_ref();
let perm = unsafe {
faer::perm::Perm::new_unchecked(
row_perm_fwd.into_boxed_slice(),
row_perm_bwd.into_boxed_slice(),
)
};
let mut inv_mat = faer::Mat::<T>::zeros(m, n);
faer::linalg::lu::partial_pivoting::inverse::inverse(
inv_mat.as_mut(),
l_mat,
u_mat,
perm.as_ref(),
par,
MemStack::new(&mut MemBuffer::new(
faer::linalg::lu::partial_pivoting::inverse::inverse_scratch::<usize, T>(m, par),
)),
);
for i in 0..m {
for j in 0..n {
a_faer[(i, j)] = inv_mat[(i, j)];
}
}
Ok(())
}
fn det<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> T {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
assert_eq!(m, n, "determinant is only defined for square matrices");
let a_faer = into_faer_mut(a);
a_faer.determinant()
}
fn cholesky<L: Layout>(
&self,
a: &mut Slice<T, (D0, D1), L>,
) -> Result<Array<T, (D0, D1)>, InvError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
if m != n {
return Err(InvError::NotSquare {
rows: m as i32,
cols: n as i32,
});
}
let mut l = a.to_tensor();
self.cholesky_write(&mut l)?;
Ok(l)
}
fn cholesky_write<L: Layout>(&self, a: &mut Slice<T, (D0, D1), L>) -> Result<(), InvError> {
let ash = *a.shape();
let (m, n) = (ash.dim(0), ash.dim(1));
if m != n {
return Err(InvError::NotSquare {
rows: m as i32,
cols: n as i32,
});
}
let par = faer::get_global_parallelism();
let result = {
let mut a_faer = into_faer_mut(a);
faer::linalg::cholesky::llt::factor::cholesky_in_place(
a_faer.as_mut(),
Default::default(),
par,
MemStack::new(&mut MemBuffer::new(
faer::linalg::cholesky::llt::factor::cholesky_in_place_scratch::<T>(
n,
par,
faer::prelude::default(),
),
)),
faer::prelude::default(),
)
};
match result {
Ok(_) => {
for i in 0..n {
for j in i + 1..n {
a[[i, j]] = T::zero();
}
}
Ok(())
}
Err(err) => Err(map_cholesky_error(err)),
}
}
}