use assert2::{assert as fancy_assert, debug_assert as fancy_debug_assert};
use dyn_stack::{DynStack, SizeOverflow, StackReq};
use faer_core::{mul::triangular::BlockStructure, solve, ComplexField, Conj, MatMut, Parallelism};
use reborrow::*;
fn cholesky_in_place_left_looking_impl<T: ComplexField>(
matrix: MatMut<'_, T>,
block_size: usize,
parallelism: Parallelism,
) {
let mut matrix = matrix;
fancy_debug_assert!(
matrix.ncols() == matrix.nrows(),
"only square matrices can be decomposed into cholesky factors",
);
let n = matrix.nrows();
match n {
0 | 1 => return,
_ => (),
};
let mut idx = 0;
loop {
let block_size = (n - idx).min(block_size);
let (top_left, top_right, bottom_left, bottom_right) = matrix.rb_mut().split_at(idx, idx);
let l00 = top_left.into_const();
let d0 = l00.diagonal();
let (_, l10, _, l20) = bottom_left.into_const().split_at(block_size, 0);
let (mut a11, _, mut a21, _) = bottom_right.split_at(block_size, block_size);
let mut l10xd0 = top_right.submatrix(0, 0, idx, block_size).transpose();
for ((l10xd0_col, l10_col), &d_factor) in l10xd0
.rb_mut()
.into_col_iter()
.zip(l10.rb().into_col_iter())
.zip(d0.into_iter())
{
for (l10xd0_elem, l) in l10xd0_col.into_iter().zip(l10_col) {
*l10xd0_elem = *l * d_factor;
}
}
let l10xd0 = l10xd0.into_const();
faer_core::mul::triangular::matmul(
a11.rb_mut(),
BlockStructure::TriangularLower,
Conj::No,
l10xd0,
BlockStructure::Rectangular,
Conj::No,
l10.transpose(),
BlockStructure::Rectangular,
Conj::Yes,
Some(T::one()),
-T::one(),
parallelism,
);
cholesky_in_place_left_looking_impl(a11.rb_mut(), block_size / 2, parallelism);
if idx + block_size == n {
break;
}
let ld11 = a11.into_const();
let l11 = ld11;
let d1 = ld11.diagonal();
faer_core::mul::matmul(
a21.rb_mut(),
Conj::No,
l20,
Conj::No,
l10xd0.transpose(),
Conj::Yes,
Some(T::one()),
-T::one(),
parallelism,
);
solve::solve_unit_lower_triangular_in_place(
l11,
Conj::Yes,
a21.rb_mut().transpose(),
Conj::No,
parallelism,
);
let l21xd1 = a21;
for (l21xd1_col, &d1_elem) in l21xd1.into_col_iter().zip(d1) {
let d1_elem_inv = d1_elem.inv();
for l21xd1_elem in l21xd1_col {
*l21xd1_elem = *l21xd1_elem * d1_elem_inv;
}
}
idx += block_size;
}
}
#[derive(Default, Copy, Clone)]
#[non_exhaustive]
pub struct LdltDiagParams {}
pub fn raw_cholesky_in_place_req<T: 'static>(
dim: usize,
parallelism: Parallelism,
params: LdltDiagParams,
) -> Result<StackReq, SizeOverflow> {
let _ = dim;
let _ = parallelism;
let _ = params;
Ok(StackReq::default())
}
fn cholesky_in_place_impl<T: ComplexField>(
matrix: MatMut<'_, T>,
parallelism: Parallelism,
stack: DynStack<'_>,
) {
fancy_debug_assert!(matrix.nrows() == matrix.ncols());
let mut matrix = matrix;
let mut stack = stack;
let n = matrix.nrows();
if n < 32 {
cholesky_in_place_left_looking_impl(matrix, 16, parallelism);
} else {
let block_size = (n / 2).min(128);
let rem = n - block_size;
let (mut l00, top_right, mut a10, mut a11) =
matrix.rb_mut().split_at(block_size, block_size);
cholesky_in_place_impl(l00.rb_mut(), parallelism, stack.rb_mut());
let l00 = l00.into_const();
let d0 = l00.diagonal();
solve::solve_unit_lower_triangular_in_place(
l00,
Conj::Yes,
a10.rb_mut().transpose(),
Conj::No,
parallelism,
);
{
let mut l10xd0 = top_right.submatrix(0, 0, block_size, rem).transpose();
for ((l10xd0_col, a10_col), &d0_elem) in l10xd0
.rb_mut()
.into_col_iter()
.zip(a10.rb_mut().into_col_iter())
.zip(d0)
{
let d0_elem_inv = d0_elem.inv();
for (l10xd0_elem, a10_elem) in l10xd0_col.into_iter().zip(a10_col) {
*l10xd0_elem = a10_elem.clone();
*a10_elem = *a10_elem * d0_elem_inv;
}
}
faer_core::mul::triangular::matmul(
a11.rb_mut(),
BlockStructure::TriangularLower,
Conj::No,
a10.into_const(),
BlockStructure::Rectangular,
Conj::No,
l10xd0.transpose().into_const(),
BlockStructure::Rectangular,
Conj::Yes,
Some(T::one()),
-T::one(),
parallelism,
);
}
cholesky_in_place_impl(a11, parallelism, stack);
}
}
#[track_caller]
#[inline]
pub fn raw_cholesky_in_place<T: ComplexField>(
matrix: MatMut<'_, T>,
parallelism: Parallelism,
stack: DynStack<'_>,
params: LdltDiagParams,
) {
let _ = params;
fancy_assert!(
matrix.ncols() == matrix.nrows(),
"only square matrices can be decomposed into cholesky factors",
);
cholesky_in_place_impl(matrix, parallelism, stack)
}