use super::CholeskyError;
use crate::{
ldlt_diagonal::update::{delete_rows_and_cols_triangular, rank_update_indices},
llt::compute::{cholesky_in_place, cholesky_in_place_req},
};
use assert2::{assert as fancy_assert, debug_assert as fancy_debug_assert};
use core::{any::TypeId, mem::size_of};
use dyn_stack::{DynStack, SizeOverflow, StackReq};
use faer_core::{
mul, mul::triangular::BlockStructure, solve, temp_mat_req, temp_mat_uninit, ColMut,
ComplexField, Conj, MatMut, Parallelism, RealField,
};
use num_traits::Zero;
use pulp::Arch;
use reborrow::*;
use seq_macro::seq;
macro_rules! generate {
($name: ident, $r: tt, $ty: ty, $tys: ty, $splat: ident, $mul_add: ident, $mul: ident) => {
#[inline(always)]
unsafe fn $name(
arch: Arch,
n: usize,
l_col: *mut $ty,
w: *mut $ty,
w_col_stride: isize,
neg_wj_over_ljj_array: *const $ty,
alpha_wj_over_nljj_array: *const $ty,
nljj_over_ljj_array: *const $ty,
) {
struct Impl {
n: usize,
l_col: *mut $ty,
w: *mut $ty,
w_col_stride: isize,
neg_wj_over_ljj_array: *const $ty,
alpha_wj_over_nljj_array: *const $ty,
nljj_over_ljj_array: *const $ty,
}
impl pulp::WithSimd for Impl {
type Output = ();
#[inline(always)]
fn with_simd<S: pulp::Simd>(self, simd: S) {
unsafe {
let Self { n, l_col, w, w_col_stride, neg_wj_over_ljj_array, alpha_wj_over_nljj_array, nljj_over_ljj_array } = self;
let l_col = l_col as *mut $ty;
let w = w as *mut $ty;
let neg_wj_over_ljj_array = neg_wj_over_ljj_array as *const $ty;
let nljj_over_ljj_array = nljj_over_ljj_array as *const $ty;
let alpha_wj_over_nljj_array = alpha_wj_over_nljj_array as *const $ty;
let lanes = size_of::<$tys>() / size_of::<$ty>();
let n_vec = n / lanes;
let n_rem = n % lanes;
seq!(I in 0..$r {
let neg_wj_over_ljj~I = *neg_wj_over_ljj_array.add(I);
let nljj_over_ljj~I = *nljj_over_ljj_array.add(I);
let alpha_wj_over_nljj~I = *alpha_wj_over_nljj_array.add(I);
let w_col~I = w.offset(I * w_col_stride);
});
{
let l_col = l_col as *mut $tys;
seq!(I in 0..$r {
let neg_wj_over_ljj~I = simd.$splat(neg_wj_over_ljj~I);
let nljj_over_ljj~I = simd.$splat(nljj_over_ljj~I);
let alpha_wj_over_nljj~I = simd.$splat(alpha_wj_over_nljj~I);
let w_col~I = w_col~I as *mut $tys;
});
for i in 0..n_vec {
let mut l = *l_col.add(i);
seq!(I in 0..$r {
let mut w~I = *w_col~I.add(i);
w~I = simd.$mul_add(neg_wj_over_ljj~I, l, w~I);
l = simd.$mul_add(alpha_wj_over_nljj~I, w~I, simd.$mul(nljj_over_ljj~I, l));
w_col~I.add(i).write(w~I);
});
l_col.add(i).write(l);
}
}
{
for i in n - n_rem..n {
let mut l = *l_col.add(i);
seq!(I in 0..$r {
let mut w~I = *w_col~I.add(i);
w~I = $ty::mul_add(neg_wj_over_ljj~I, l, w~I);
l = $ty::mul_add(alpha_wj_over_nljj~I, w~I, nljj_over_ljj~I * l);
w_col~I.add(i).write(w~I);
});
l_col.add(i).write(l);
}
}
}
}
}
arch.dispatch(Impl {
n,
l_col,
w,
w_col_stride,
neg_wj_over_ljj_array,
alpha_wj_over_nljj_array,
nljj_over_ljj_array,
})
}
};
}
macro_rules! generate_generic {
($name: ident, $r: tt) => {
unsafe fn $name<T: ComplexField>(
n: usize,
l_col: *mut T,
l_row_stride: isize,
w: *mut T,
w_row_stride: isize,
w_col_stride: isize,
neg_wj_over_ljj_array: *const T,
alpha_wj_over_nljj_array: *const T,
nljj_over_ljj_array: *const T,
) {
seq!(I in 0..$r {
let neg_wj_over_ljj~I = *neg_wj_over_ljj_array.add(I);
let nljj_over_ljj~I = *nljj_over_ljj_array.add(I);
let alpha_wj_over_nljj~I = *alpha_wj_over_nljj_array.add(I);
let w_col~I = w.offset(I * w_col_stride);
});
for i in 0..n {
let mut l = (*l_col.offset(i as isize * l_row_stride)).clone();
seq!(I in 0..$r {
let mut w~I = (*w_col~I.offset(i as isize * w_row_stride)).clone();
w~I = (neg_wj_over_ljj~I * l) + w~I;
l = (alpha_wj_over_nljj~I * w~I) + (nljj_over_ljj~I * l);
*w_col~I.offset(i as isize * w_row_stride) = w~I;
});
*l_col.offset(i as isize * l_row_stride) = l;
}
}
};
}
generate_generic!(r1, 1);
generate_generic!(r2, 2);
generate_generic!(r3, 3);
generate_generic!(r4, 4);
#[rustfmt::skip]
generate!(rank_1_f64, 1, f64, S::f64s, f64s_splat, f64s_mul_adde, f64s_mul);
#[rustfmt::skip]
generate!(rank_2_f64, 2, f64, S::f64s, f64s_splat, f64s_mul_adde, f64s_mul);
#[rustfmt::skip]
generate!(rank_3_f64, 3, f64, S::f64s, f64s_splat, f64s_mul_adde, f64s_mul);
#[rustfmt::skip]
generate!(rank_4_f64, 4, f64, S::f64s, f64s_splat, f64s_mul_adde, f64s_mul);
#[rustfmt::skip]
generate!(rank_1_f32, 1, f32, S::f32s, f32s_splat, f32s_mul_adde, f32s_mul);
#[rustfmt::skip]
generate!(rank_2_f32, 2, f32, S::f32s, f32s_splat, f32s_mul_adde, f32s_mul);
#[rustfmt::skip]
generate!(rank_3_f32, 3, f32, S::f32s, f32s_splat, f32s_mul_adde, f32s_mul);
#[rustfmt::skip]
generate!(rank_4_f32, 4, f32, S::f32s, f32s_splat, f32s_mul_adde, f32s_mul);
struct RankRUpdate<'a, T> {
l: MatMut<'a, T>,
w: MatMut<'a, T>,
alpha: ColMut<'a, T>,
r: &'a mut dyn FnMut() -> usize,
}
impl<'a, T: ComplexField> RankRUpdate<'a, T> {
fn run(self) -> Result<(), CholeskyError> {
let RankRUpdate {
mut l,
mut w,
mut alpha,
r,
} = self;
let n = l.nrows();
let k = w.ncols();
fancy_debug_assert!(l.ncols() == n);
fancy_debug_assert!(w.nrows() == n);
fancy_debug_assert!(alpha.nrows() == k);
let l_rs = l.row_stride();
let w_cs = w.col_stride();
let w_rs = w.row_stride();
let arch = Arch::new();
unsafe {
for j in 0..n {
let r = (*r)().min(k);
let mut r_idx = 0;
while r_idx < r {
let r_chunk = (r - r_idx).min(4);
let mut neg_wj_over_ljj_array = [T::zero(), T::zero(), T::zero(), T::zero()];
let mut alpha_wj_over_nljj_array = [T::zero(), T::zero(), T::zero(), T::zero()];
let mut nljj_over_ljj_array = [T::zero(), T::zero(), T::zero(), T::zero()];
let mut ljj = l.rb().get_unchecked(j, j).clone();
for k in 0..r_chunk {
let neg_wj_over_ljj = neg_wj_over_ljj_array.get_unchecked_mut(k);
let alpha_conj_wj_over_nljj = alpha_wj_over_nljj_array.get_unchecked_mut(k);
let nljj_over_ljj = nljj_over_ljj_array.get_unchecked_mut(k);
let alpha = alpha.rb_mut().get_unchecked(r_idx + k);
let wj = *w.rb().get_unchecked(j, r_idx + k);
let alpha_conj_wj = *alpha * wj.conj();
let sqr_nljj = ljj * ljj + alpha_conj_wj * wj;
if !(sqr_nljj.real() > T::Real::zero()) {
return Err(CholeskyError);
}
let nljj = T::from_real(sqr_nljj.real().sqrt());
let inv_ljj = ljj.inv();
let inv_nljj = nljj.inv();
*neg_wj_over_ljj = -(wj * inv_ljj);
*nljj_over_ljj = nljj * inv_ljj;
*alpha_conj_wj_over_nljj = alpha_conj_wj * inv_nljj;
*alpha =
*alpha - *alpha_conj_wj_over_nljj * (*alpha_conj_wj_over_nljj).conj();
ljj = nljj;
}
*l.rb_mut().get_unchecked(j, j) = ljj;
let rem = n - j - 1;
let l_ptr = l.rb_mut().ptr_at(j + 1, j);
let w_ptr = w.rb_mut().ptr_at(j + 1, r_idx);
let neg_wj_over_ljj = neg_wj_over_ljj_array.as_ptr();
let alpha_wj_over_nljj = alpha_wj_over_nljj_array.as_ptr();
let nljj_over_ljj = nljj_over_ljj_array.as_ptr();
if TypeId::of::<T>() == TypeId::of::<f64>() && l_rs == 1 && w_rs == 1 {
match r_chunk {
1 => rank_1_f64(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
2 => rank_2_f64(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
3 => rank_3_f64(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
4 => rank_4_f64(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
_ => unreachable!(),
};
} else if TypeId::of::<T>() == TypeId::of::<f32>() && l_rs == 1 && w_rs == 1 {
match r_chunk {
1 => rank_1_f32(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
2 => rank_2_f32(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
3 => rank_3_f32(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
4 => rank_4_f32(
arch,
rem,
l_ptr as _,
w_ptr as _,
w_cs,
neg_wj_over_ljj as _,
alpha_wj_over_nljj as _,
nljj_over_ljj as _,
),
_ => unreachable!(),
};
} else {
match r_chunk {
1 => r1(
rem,
l_ptr,
l_rs,
w_ptr,
w_rs,
w_cs,
neg_wj_over_ljj,
alpha_wj_over_nljj,
nljj_over_ljj,
),
2 => r2(
rem,
l_ptr,
l_rs,
w_ptr,
w_rs,
w_cs,
neg_wj_over_ljj,
alpha_wj_over_nljj,
nljj_over_ljj,
),
3 => r3(
rem,
l_ptr,
l_rs,
w_ptr,
w_rs,
w_cs,
neg_wj_over_ljj,
alpha_wj_over_nljj,
nljj_over_ljj,
),
4 => r4(
rem,
l_ptr,
l_rs,
w_ptr,
w_rs,
w_cs,
neg_wj_over_ljj,
alpha_wj_over_nljj,
nljj_over_ljj,
),
_ => unreachable!(),
};
}
r_idx += r_chunk;
}
}
}
Ok(())
}
}
#[track_caller]
pub fn rank_r_update_clobber<T: ComplexField>(
cholesky_factor: MatMut<'_, T>,
w: MatMut<'_, T>,
alpha: ColMut<'_, T>,
) -> Result<(), CholeskyError> {
let n = cholesky_factor.nrows();
let k = w.ncols();
fancy_assert!(cholesky_factor.ncols() == n);
fancy_assert!(w.nrows() == n);
fancy_assert!(alpha.nrows() == k);
RankRUpdate {
l: cholesky_factor,
w,
alpha,
r: &mut || k,
}
.run()
}
#[track_caller]
pub fn delete_rows_and_cols_clobber_req<T: 'static>(
dim: usize,
number_of_rows_to_remove: usize,
) -> Result<StackReq, SizeOverflow> {
let r = number_of_rows_to_remove;
StackReq::try_all_of([temp_mat_req::<T>(dim, r)?, temp_mat_req::<T>(r, 1)?])
}
#[track_caller]
pub fn delete_rows_and_cols_clobber<T: ComplexField>(
cholesky_factor: MatMut<'_, T>,
indices: &mut [usize],
stack: DynStack<'_>,
) {
let n = cholesky_factor.nrows();
let r = indices.len();
fancy_assert!(cholesky_factor.ncols() == n);
fancy_assert!(indices.len() < n);
if r == 0 {
return;
}
indices.sort_unstable();
for i in 0..r - 1 {
fancy_assert!(indices[i + 1] > indices[i]);
}
fancy_assert!(indices[r - 1] < n);
let first = indices[0];
temp_mat_uninit! {
let (mut w, stack) = unsafe { temp_mat_uninit::<T>(n - first - r, r, stack) };
let (alpha, _) = unsafe { temp_mat_uninit::<T>(r, 1, stack) };
}
let mut alpha = alpha.col(0);
Arch::new().dispatch(|| {
for k in 0..r {
let j = indices[k];
unsafe {
*alpha.rb_mut().ptr_in_bounds_at(k) = T::one();
}
for chunk_i in k..r {
let chunk_i = chunk_i + 1;
let i_start = indices[chunk_i - 1] + 1;
#[rustfmt::skip]
let i_finish = if chunk_i == r { n } else { indices[chunk_i] };
for i in i_start..i_finish {
unsafe {
*w.rb_mut()
.ptr_in_bounds_at_unchecked(i - chunk_i - first, k) =
cholesky_factor.rb().get_unchecked(i, j).clone();
}
}
}
}
});
let mut cholesky_factor = cholesky_factor;
delete_rows_and_cols_triangular(cholesky_factor.rb_mut(), indices);
RankRUpdate {
l: unsafe {
cholesky_factor.submatrix_unchecked(first, first, n - first - r, n - first - r)
},
w,
alpha,
r: &mut rank_update_indices(first, indices),
}
.run()
.unwrap();
}
pub fn insert_rows_and_cols_clobber_req<T: 'static>(
inserted_matrix_ncols: usize,
parallelism: Parallelism,
) -> Result<StackReq, SizeOverflow> {
cholesky_in_place_req::<T>(inserted_matrix_ncols, parallelism, Default::default())
}
#[track_caller]
pub fn insert_rows_and_cols_clobber<T: ComplexField>(
cholesky_factor_extended: MatMut<'_, T>,
insertion_index: usize,
inserted_matrix: MatMut<'_, T>,
parallelism: Parallelism,
stack: DynStack<'_>,
) -> Result<(), CholeskyError> {
let new_n = cholesky_factor_extended.nrows();
let r = inserted_matrix.ncols();
fancy_assert!(cholesky_factor_extended.nrows() == cholesky_factor_extended.ncols());
fancy_assert!(cholesky_factor_extended.ncols() == new_n);
fancy_assert!(r < new_n);
let old_n = new_n - r;
fancy_assert!(insertion_index <= old_n);
if r == 0 {
return Ok(());
}
let mut current_col = old_n;
let mut ld = cholesky_factor_extended;
while current_col != insertion_index {
current_col -= 1;
unsafe {
for i in (current_col..old_n).rev() {
*ld.rb_mut()
.ptr_in_bounds_at_unchecked(i + r, current_col + r) =
(*ld.rb().ptr_in_bounds_at_unchecked(i, current_col)).clone();
}
}
}
while current_col != 0 {
current_col -= 1;
unsafe {
for i in (insertion_index..old_n).rev() {
*ld.rb_mut().ptr_in_bounds_at_unchecked(i + r, current_col) =
(*ld.rb().ptr_in_bounds_at_unchecked(i, current_col)).clone();
}
}
}
let (l00, _, l_bot_left, ld_bot_right) = ld.split_at(insertion_index, insertion_index);
let l00 = l00.into_const();
let (_, mut l10, _, l20) = l_bot_left.split_at(r, 0);
let (mut l11, _, mut l21, ld22) = ld_bot_right.split_at(r, r);
let (_, mut a01, _, a_bottom) = inserted_matrix.split_at(insertion_index, 0);
let (_, a11, _, a21) = a_bottom.split_at(r, 0);
let mut stack = stack;
solve::solve_lower_triangular_in_place(l00.rb(), Conj::No, a01.rb_mut(), Conj::No, parallelism);
let a10 = a01.rb().transpose();
for j in 0..insertion_index {
for i in 0..r {
unsafe {
*l10.rb_mut().ptr_in_bounds_at_unchecked(i, j) = (*a10.get_unchecked(i, j)).conj();
}
}
}
for j in 0..r {
for i in j..r {
unsafe {
*l11.rb_mut().ptr_in_bounds_at_unchecked(i, j) = *a11.rb().get_unchecked(i, j);
}
}
}
mul::triangular::matmul(
l11.rb_mut(),
BlockStructure::TriangularLower,
Conj::No,
l10.rb(),
BlockStructure::Rectangular,
Conj::No,
a01.rb(),
BlockStructure::Rectangular,
Conj::No,
Some(T::one()),
-T::one(),
parallelism,
);
cholesky_in_place(
l11.rb_mut(),
parallelism,
stack.rb_mut(),
Default::default(),
)?;
let l11 = l11.into_const();
let rem = l21.nrows();
for j in 0..r {
for i in 0..rem {
unsafe {
*l21.rb_mut().ptr_in_bounds_at_unchecked(i, j) =
a21.rb().get_unchecked(i, j).clone();
}
}
}
mul::matmul(
l21.rb_mut(),
Conj::No,
l20.rb(),
Conj::No,
a01.rb(),
Conj::No,
Some(T::one()),
-T::one(),
parallelism,
);
solve::solve_lower_triangular_in_place(
l11,
Conj::Yes,
l21.rb_mut().transpose(),
Conj::No,
parallelism,
);
let mut alpha = a11.col(0);
let mut w = a21;
for j in 0..r {
unsafe {
*alpha.rb_mut().ptr_in_bounds_at_unchecked(j) = -T::one();
for i in 0..rem {
*w.rb_mut().ptr_in_bounds_at_unchecked(i, j) = l21.rb().get(i, j).clone();
}
}
}
rank_r_update_clobber(ld22, w, alpha)
}