use assert2::assert as fancy_assert;
use dyn_stack::{DynStack, SizeOverflow, StackReq};
use faer_core::{
inverse::invert_lower_triangular,
mul::triangular::{self, BlockStructure},
temp_mat_req, temp_mat_uninit, ComplexField, Conj, MatMut, MatRef, Parallelism,
};
use reborrow::*;
fn invert_lower_impl<T: ComplexField>(
dst: MatMut<'_, T>,
cholesky_factor: Option<MatRef<'_, T>>,
parallelism: Parallelism,
stack: DynStack,
) {
let cholesky_factor = match cholesky_factor {
Some(cholesky_factor) => cholesky_factor,
None => dst.rb(),
};
let n = cholesky_factor.nrows();
temp_mat_uninit! {
let (mut tmp, _) = unsafe { temp_mat_uninit::<T>(n, n, stack) };
}
invert_lower_triangular(tmp.rb_mut(), cholesky_factor, Conj::No, parallelism);
triangular::matmul(
dst,
BlockStructure::TriangularLower,
Conj::No,
tmp.rb().transpose(),
BlockStructure::TriangularUpper,
Conj::Yes,
tmp.rb(),
BlockStructure::TriangularLower,
Conj::No,
None,
T::one(),
parallelism,
);
}
pub fn invert_lower_req<T: 'static>(
dimension: usize,
parallelism: Parallelism,
) -> Result<StackReq, SizeOverflow> {
let _ = parallelism;
temp_mat_req::<T>(dimension, dimension)
}
pub fn invert_lower_in_place_req<T: 'static>(
dimension: usize,
parallelism: Parallelism,
) -> Result<StackReq, SizeOverflow> {
invert_lower_req::<T>(dimension, parallelism)
}
#[track_caller]
pub fn invert_lower_in_place<T: ComplexField>(
cholesky_factor: MatMut<'_, T>,
parallelism: Parallelism,
stack: DynStack,
) {
fancy_assert!(cholesky_factor.nrows() == cholesky_factor.ncols());
invert_lower_impl(cholesky_factor, None, parallelism, stack);
}
#[track_caller]
pub fn invert_lower<T: ComplexField>(
dst: MatMut<'_, T>,
cholesky_factor: MatRef<'_, T>,
parallelism: Parallelism,
stack: DynStack,
) {
fancy_assert!(cholesky_factor.nrows() == cholesky_factor.ncols());
fancy_assert!((dst.nrows(), dst.ncols()) == (cholesky_factor.nrows(), cholesky_factor.ncols()));
invert_lower_impl(dst, Some(cholesky_factor), parallelism, stack);
}