1use crate::LinalgError;
2use crate::faer_ndarray::{FaerArrayView, FaerColView};
3use faer::Side;
4use faer::linalg::solvers::Solve;
5use faer::sparse::linalg::solvers::Llt as SparseLlt;
6use faer::sparse::{SparseColMat, SymbolicSparseColMat, Triplet};
7use ndarray::{Array1, Array2, ArrayBase, ArrayView2, Data, Ix1, Ix2};
8use rayon::prelude::*;
9use std::collections::BTreeMap;
10use std::sync::{Arc, Mutex};
11
12const ZERO_TOL: f64 = 1e-12;
13const PARALLEL_SPARSE_FILL_COLUMN_THRESHOLD: usize = 64;
14
15macro_rules! bail_invalid_linalg {
16 ($($arg:tt)*) => {
17 return Err(LinalgError::InvalidInput(format!($($arg)*)))
18 };
19}
20
21#[derive(Clone)]
22pub struct SparseExactFactor {
23 factor: SparseLlt<usize, f64>,
24 simplicial: Arc<SimplicialFactor>,
25 n: usize,
26 logdet: f64,
27}
28
29impl crate::matrix::FactorizedSystem for SparseExactFactor {
30 fn solve(&self, rhs: &Array1<f64>) -> Result<Array1<f64>, String> {
31 solve_sparse_spd(self, rhs).map_err(|e| e.to_string())
32 }
33
34 fn solvemulti(&self, rhs: &Array2<f64>) -> Result<Array2<f64>, String> {
35 solve_sparse_spdmulti(self, rhs).map_err(|e| e.to_string())
36 }
37
38 fn logdet(&self) -> f64 {
39 self.logdet
40 }
41}
42
43pub fn dense_to_sparse(
44 matrix: &Array2<f64>,
45 tol: f64,
46) -> Result<SparseColMat<usize, f64>, LinalgError> {
47 let nrows = matrix.nrows();
48 let ncols = matrix.ncols();
49 let counts: Vec<usize> = (0..ncols)
56 .into_par_iter()
57 .map(|col| {
58 let mut count = 0usize;
59 for row in 0..nrows {
60 if matrix[[row, col]].abs() > tol {
61 count += 1;
62 }
63 }
64 count
65 })
66 .collect();
67 let col_ptr = prefix_sum_counts(&counts);
68 let nnz = col_ptr[ncols];
69 let mut row_idx = vec![0usize; nnz];
70 let mut values = vec![0.0; nnz];
71 fill_dense_to_sparse_columns(matrix, tol, 0, ncols, &col_ptr, &mut row_idx, &mut values);
72 let symbolic = SymbolicSparseColMat::<usize>::new_checked(nrows, ncols, col_ptr, None, row_idx);
73 Ok(SparseColMat::<usize, f64>::new(symbolic, values))
74}
75
76pub fn dense_to_sparse_symmetric_upper(
82 matrix: &Array2<f64>,
83 tol: f64,
84) -> Result<SparseColMat<usize, f64>, LinalgError> {
85 let nrows = matrix.nrows();
86 let ncols = matrix.ncols();
87 let row_limit = nrows.min(ncols);
93 let counts: Vec<usize> = (0..ncols)
94 .into_par_iter()
95 .map(|col| {
96 let mut count = 0usize;
97 let row_end = (col + 1).min(row_limit);
98 for row in 0..row_end {
99 if matrix[[row, col]].abs() > tol {
100 count += 1;
101 }
102 }
103 count
104 })
105 .collect();
106 let col_ptr = prefix_sum_counts(&counts);
107 let nnz = col_ptr[ncols];
108 let mut row_idx = vec![0usize; nnz];
109 let mut values = vec![0.0; nnz];
110 fill_dense_symmetric_upper_columns(
111 matrix,
112 tol,
113 row_limit,
114 0,
115 ncols,
116 &col_ptr,
117 &mut row_idx,
118 &mut values,
119 );
120 let symbolic = SymbolicSparseColMat::<usize>::new_checked(nrows, ncols, col_ptr, None, row_idx);
121 Ok(SparseColMat::<usize, f64>::new(symbolic, values))
122}
123
124fn prefix_sum_counts(counts: &[usize]) -> Vec<usize> {
125 let mut col_ptr = Vec::with_capacity(counts.len() + 1);
126 col_ptr.push(0);
127 let mut running = 0usize;
128 for &count in counts {
129 running += count;
130 col_ptr.push(running);
131 }
132 col_ptr
133}
134
135fn fill_dense_to_sparse_columns(
136 matrix: &Array2<f64>,
137 tol: f64,
138 col_start: usize,
139 col_end: usize,
140 col_ptr: &[usize],
141 row_idx: &mut [usize],
142 values: &mut [f64],
143) {
144 if col_end - col_start <= PARALLEL_SPARSE_FILL_COLUMN_THRESHOLD {
145 let base = col_ptr[col_start];
146 for col in col_start..col_end {
147 let mut write = col_ptr[col] - base;
148 for row in 0..matrix.nrows() {
149 let value = matrix[[row, col]];
150 if value.abs() > tol {
151 row_idx[write] = row;
152 values[write] = value;
153 write += 1;
154 }
155 }
156 }
157 return;
158 }
159
160 let mid = col_start + (col_end - col_start) / 2;
161 let split = col_ptr[mid] - col_ptr[col_start];
162 let (left_rows, right_rows) = row_idx.split_at_mut(split);
163 let (left_values, right_values) = values.split_at_mut(split);
164 rayon::join(
165 || {
166 fill_dense_to_sparse_columns(
167 matrix,
168 tol,
169 col_start,
170 mid,
171 col_ptr,
172 left_rows,
173 left_values,
174 );
175 },
176 || {
177 fill_dense_to_sparse_columns(
178 matrix,
179 tol,
180 mid,
181 col_end,
182 col_ptr,
183 right_rows,
184 right_values,
185 );
186 },
187 );
188}
189
190fn fill_dense_symmetric_upper_columns(
191 matrix: &Array2<f64>,
192 tol: f64,
193 row_limit: usize,
194 col_start: usize,
195 col_end: usize,
196 col_ptr: &[usize],
197 row_idx: &mut [usize],
198 values: &mut [f64],
199) {
200 if col_end - col_start <= PARALLEL_SPARSE_FILL_COLUMN_THRESHOLD {
201 let base = col_ptr[col_start];
202 for col in col_start..col_end {
203 let row_end = (col + 1).min(row_limit);
204 let mut write = col_ptr[col] - base;
205 for row in 0..row_end {
206 let value = matrix[[row, col]];
207 if value.abs() > tol {
208 row_idx[write] = row;
209 values[write] = value;
210 write += 1;
211 }
212 }
213 }
214 return;
215 }
216
217 let mid = col_start + (col_end - col_start) / 2;
218 let split = col_ptr[mid] - col_ptr[col_start];
219 let (left_rows, right_rows) = row_idx.split_at_mut(split);
220 let (left_values, right_values) = values.split_at_mut(split);
221 rayon::join(
222 || {
223 fill_dense_symmetric_upper_columns(
224 matrix,
225 tol,
226 row_limit,
227 col_start,
228 mid,
229 col_ptr,
230 left_rows,
231 left_values,
232 );
233 },
234 || {
235 fill_dense_symmetric_upper_columns(
236 matrix,
237 tol,
238 row_limit,
239 mid,
240 col_end,
241 col_ptr,
242 right_rows,
243 right_values,
244 );
245 },
246 );
247}
248
249pub fn sparse_symmetric_upper_matvec_public<S: Data<Elem = f64>>(
250 matrix: &SparseColMat<usize, f64>,
251 vector: &ArrayBase<S, Ix1>,
252) -> Array1<f64> {
253 let mut out = Array1::<f64>::zeros(matrix.nrows());
254 let (symbolic, values) = matrix.parts();
255 let col_ptr = symbolic.col_ptr();
256 let row_idx = symbolic.row_idx();
257 for col in 0..matrix.ncols() {
258 let x_col = vector[col];
259 for idx in col_ptr[col]..col_ptr[col + 1] {
260 let row = row_idx[idx];
261 let value = values[idx];
262 out[row] += value * x_col;
263 if row != col {
264 out[col] += value * vector[row];
265 }
266 }
267 }
268 out
269}
270
271pub fn factorize_sparse_spd(
272 h: &SparseColMat<usize, f64>,
273) -> Result<SparseExactFactor, LinalgError> {
274 let t_start = std::time::Instant::now();
284 let n_input = h.ncols();
285 let h_upper = canonicalize_sparse_symmetric_upper(h, ZERO_TOL)?;
286 let factor = h_upper.as_ref().sp_cholesky(Side::Upper).map_err(|_| {
287 LinalgError::ModelIsIllConditioned {
288 condition_number: f64::INFINITY,
289 }
290 })?;
291 let simplicial = factorize_simplicial_canonical_upper(&h_upper)?;
295 let logdet = simplicial.logdet;
296 let elapsed_ms = t_start.elapsed().as_secs_f64() * 1000.0;
297 if elapsed_ms > 100.0 {
298 log::info!(
299 "[sparse-chol] factorize_sparse_spd | n={} | {:.1}ms",
300 n_input,
301 elapsed_ms
302 );
303 }
304 Ok(SparseExactFactor {
305 factor,
306 simplicial: Arc::new(simplicial),
307 n: h_upper.ncols(),
308 logdet,
309 })
310}
311
312pub fn factorize_sparse_spd_strict(
320 h_upper: &SparseColMat<usize, f64>,
321) -> Result<SparseExactFactor, LinalgError> {
322 if h_upper.nrows() == 0 || h_upper.nrows() != h_upper.ncols() {
323 bail_invalid_linalg!(
324 "strict sparse SPD factorization requires a non-empty square matrix, got {}x{}",
325 h_upper.nrows(),
326 h_upper.ncols()
327 );
328 }
329 let (symbolic, values) = h_upper.parts();
330 let col_ptr = symbolic.col_ptr();
331 let row_idx = symbolic.row_idx();
332 for col in 0..h_upper.ncols() {
333 let mut previous_row = None;
334 for idx in col_ptr[col]..col_ptr[col + 1] {
335 let row = row_idx[idx];
336 let value = values[idx];
337 if row > col {
338 bail_invalid_linalg!(
339 "strict sparse SPD input must use upper-triangle storage; found lower entry ({row}, {col})"
340 );
341 }
342 if !value.is_finite() {
343 bail_invalid_linalg!(
344 "strict sparse SPD input contains non-finite entry ({row}, {col}) = {value:?}"
345 );
346 }
347 if previous_row.is_some_and(|previous| row <= previous) {
348 bail_invalid_linalg!(
349 "strict sparse SPD input column {col} has duplicate or unsorted row {row}"
350 );
351 }
352 previous_row = Some(row);
353 }
354 }
355
356 let factor = h_upper.as_ref().sp_cholesky(Side::Upper).map_err(|_| {
357 LinalgError::ModelIsIllConditioned {
358 condition_number: f64::INFINITY,
359 }
360 })?;
361 let simplicial = factorize_simplicial_canonical_upper(h_upper)?;
362 let logdet = simplicial.logdet;
363 Ok(SparseExactFactor {
364 factor,
365 simplicial: Arc::new(simplicial),
366 n: h_upper.ncols(),
367 logdet,
368 })
369}
370
371fn canonicalize_sparse_symmetric_upper(
372 matrix: &SparseColMat<usize, f64>,
373 tol: f64,
374) -> Result<SparseColMat<usize, f64>, LinalgError> {
375 if matrix.nrows() != matrix.ncols() {
376 bail_invalid_linalg!(
377 "sparse SPD factorization requires square matrix, got {}x{}",
378 matrix.nrows(),
379 matrix.ncols()
380 );
381 }
382
383 #[derive(Default, Clone, Copy)]
384 struct PairAccum {
385 upper_sum: f64,
386 upper_count: usize,
387 lower_sum: f64,
388 lower_count: usize,
389 }
390
391 let mut accum: BTreeMap<(usize, usize), PairAccum> = BTreeMap::new();
392 let (symbolic, values) = matrix.parts();
393 let col_ptr = symbolic.col_ptr();
394 let row_idx = symbolic.row_idx();
395
396 for col in 0..matrix.ncols() {
397 let start = col_ptr[col];
398 let end = col_ptr[col + 1];
399 for idx in start..end {
400 let row = row_idx[idx];
401 let value = values[idx];
402 let (r, c, is_upper) = if row <= col {
403 (row, col, true)
404 } else {
405 (col, row, false)
406 };
407 let slot = accum.entry((r, c)).or_default();
408 if is_upper {
409 slot.upper_sum += value;
410 slot.upper_count += 1;
411 } else {
412 slot.lower_sum += value;
413 slot.lower_count += 1;
414 }
415 }
416 }
417
418 let mut triplets = Vec::<Triplet<usize, usize, f64>>::new();
419 for ((row, col), slot) in accum {
420 let value = if row == col {
421 let count = slot.upper_count + slot.lower_count;
422 if count == 0 {
423 0.0
424 } else {
425 (slot.upper_sum + slot.lower_sum) / (count as f64)
426 }
427 } else {
428 let upper_avg = if slot.upper_count > 0 {
429 Some(slot.upper_sum / (slot.upper_count as f64))
430 } else {
431 None
432 };
433 let lower_avg = if slot.lower_count > 0 {
434 Some(slot.lower_sum / (slot.lower_count as f64))
435 } else {
436 None
437 };
438 match (upper_avg, lower_avg) {
439 (Some(u), Some(l)) => 0.5 * (u + l),
440 (Some(u), None) => u,
441 (None, Some(l)) => l,
442 (None, None) => 0.0,
443 }
444 };
445
446 if value.abs() > tol {
447 triplets.push(Triplet::new(row, col, value));
448 }
449 }
450
451 SparseColMat::try_new_from_triplets(matrix.nrows(), matrix.ncols(), &triplets).map_err(|_| {
452 LinalgError::InvalidInput(
453 "failed to canonicalize sparse matrix to symmetric-upper CSC".to_string(),
454 )
455 })
456}
457
458fn solve_view<R, I, F>(
459 factor: &SparseExactFactor,
460 rhs: ArrayView2<'_, f64>,
461 indices: I,
462 mut result: R,
463 non_finite_message: &'static str,
464 mut consume: F,
465) -> Result<R, LinalgError>
466where
467 I: IntoIterator<Item = (usize, usize)>,
468 F: FnMut(&mut R, usize, usize, f64),
469{
470 let rhsview = FaerArrayView::new(&rhs);
471 let solved = factor.factor.solve(rhsview.as_ref());
472 for (row, col) in indices {
473 let value = solved[(row, col)];
474 if !value.is_finite() {
475 bail_invalid_linalg!("{}", non_finite_message.to_string());
476 }
477 consume(&mut result, row, col, value);
478 }
479 Ok(result)
480}
481
482pub fn solve_sparse_spd<S>(
483 factor: &SparseExactFactor,
484 rhs: &ArrayBase<S, Ix1>,
485) -> Result<Array1<f64>, LinalgError>
486where
487 S: Data<Elem = f64>,
488{
489 if rhs.len() != factor.n {
490 bail_invalid_linalg!(
491 "sparse SPD solve dimension mismatch: rhs has {}, factor has {}",
492 rhs.len(),
493 factor.n
494 );
495 }
496 let mut result = Array1::<f64>::zeros(rhs.len());
497 solve_sparse_spd_into(factor, rhs, &mut result)?;
498 Ok(result)
499}
500
501pub fn solve_sparse_spd_into<S>(
506 factor: &SparseExactFactor,
507 rhs: &ArrayBase<S, Ix1>,
508 out: &mut Array1<f64>,
509) -> Result<(), LinalgError>
510where
511 S: Data<Elem = f64>,
512{
513 if rhs.len() != factor.n {
514 bail_invalid_linalg!(
515 "sparse SPD solve dimension mismatch: rhs has {}, factor has {}",
516 rhs.len(),
517 factor.n
518 );
519 }
520 if out.len() != factor.n {
521 bail_invalid_linalg!(
522 "sparse SPD solve output dimension mismatch: out has {}, factor has {}",
523 out.len(),
524 factor.n
525 );
526 }
527 let rhsview = FaerColView::new(rhs);
528 let solved = factor.factor.solve(rhsview.as_ref());
529 for i in 0..factor.n {
530 let value = solved[(i, 0)];
531 if !value.is_finite() {
532 bail_invalid_linalg!("sparse SPD solve produced non-finite values");
533 }
534 out[i] = value;
535 }
536 Ok(())
537}
538
539pub fn solve_sparse_spdmulti<S>(
540 factor: &SparseExactFactor,
541 rhs: &ArrayBase<S, Ix2>,
542) -> Result<Array2<f64>, LinalgError>
543where
544 S: Data<Elem = f64>,
545{
546 if rhs.nrows() != factor.n {
547 bail_invalid_linalg!(
548 "sparse SPD multi-solve row mismatch: rhs has {}, factor has {}",
549 rhs.nrows(),
550 factor.n
551 );
552 }
553 let indices = (0..rhs.nrows()).flat_map(|i| (0..rhs.ncols()).map(move |j| (i, j)));
554 solve_view(
555 factor,
556 rhs.view(),
557 indices,
558 Array2::<f64>::zeros(rhs.raw_dim()),
559 "sparse SPD multi-solve produced non-finite values",
560 |result, row, col, value| {
561 result[[row, col]] = value;
562 },
563 )
564}
565
566pub fn solve_sparse_spdmulti_rows<S>(
567 factor: &SparseExactFactor,
568 rhs: &ArrayBase<S, Ix2>,
569 row_start: usize,
570 row_end: usize,
571) -> Result<Array2<f64>, LinalgError>
572where
573 S: Data<Elem = f64>,
574{
575 if rhs.nrows() != factor.n {
576 bail_invalid_linalg!(
577 "sparse SPD multi-solve row mismatch: rhs has {}, factor has {}",
578 rhs.nrows(),
579 factor.n
580 );
581 }
582 if row_start > row_end || row_end > factor.n {
583 bail_invalid_linalg!(
584 "sparse SPD selected rows out of bounds: row_start={}, row_end={}, factor={}",
585 row_start,
586 row_end,
587 factor.n
588 );
589 }
590 let indices = (row_start..row_end).flat_map(|i| (0..rhs.ncols()).map(move |j| (i, j)));
591 solve_view(
592 factor,
593 rhs.view(),
594 indices,
595 Array2::<f64>::zeros((row_end - row_start, rhs.ncols())),
596 "sparse SPD selected-row solve produced non-finite values",
597 |result, row, col, value| {
598 result[[row - row_start, col]] = value;
599 },
600 )
601}
602
603pub fn solve_sparse_spdmulti_diagonal_sum<S>(
604 factor: &SparseExactFactor,
605 rhs: &ArrayBase<S, Ix2>,
606 row_start: usize,
607) -> Result<f64, LinalgError>
608where
609 S: Data<Elem = f64>,
610{
611 if row_start.saturating_add(rhs.ncols()) > rhs.nrows() {
612 bail_invalid_linalg!(
613 "sparse SPD selected diagonal out of bounds: row_start={}, rows={}, cols={}",
614 row_start,
615 rhs.nrows(),
616 rhs.ncols()
617 );
618 }
619 let indices = (0..rhs.ncols()).map(|col| (row_start + col, col));
620 solve_view(
621 factor,
622 rhs.view(),
623 indices,
624 0.0,
625 "sparse SPD selected diagonal solve produced non-finite values",
626 |sum, _, _, value| {
627 *sum += value;
628 },
629 )
630}
631
632pub fn logdet_from_factor(factor: &SparseExactFactor) -> Result<f64, LinalgError> {
633 Ok(factor.logdet)
634}
635
636pub fn assemble_sparse_factor_h_dense(
637 factor: &SparseExactFactor,
638) -> Result<Array2<f64>, LinalgError> {
639 factor.simplicial.assemble_h_dense_original_order()
640}
641
642use faer::dyn_stack::{MemBuffer, MemStack, StackReq};
647use faer::linalg::cholesky::llt::factor::LltRegularization;
648use faer::sparse::linalg::amd;
649use faer::sparse::linalg::cholesky::simplicial;
650
651pub struct SimplicialFactor {
656 l_col_ptr: Vec<usize>,
658 l_row_idx: Vec<usize>,
660 l_values: Vec<f64>,
662 perm_inv: Vec<usize>,
665 n: usize,
667 pub logdet: f64,
669}
670
671pub fn factorize_simplicial(h: &SparseColMat<usize, f64>) -> Result<SimplicialFactor, LinalgError> {
677 let h_upper = canonicalize_sparse_symmetric_upper(h, ZERO_TOL)?;
678 factorize_simplicial_canonical_upper(&h_upper)
679}
680
681fn factorize_simplicial_canonical_upper(
682 h_upper: &SparseColMat<usize, f64>,
683) -> Result<SimplicialFactor, LinalgError> {
684 let n = h_upper.ncols();
685 if n == 0 {
686 return Ok(SimplicialFactor {
687 l_col_ptr: vec![0],
688 l_row_idx: Vec::new(),
689 l_values: Vec::new(),
690 perm_inv: Vec::new(),
691 n: 0,
692 logdet: 0.0,
693 });
694 }
695
696 let a_nnz = h_upper.compute_nnz();
697
698 let mut perm_fwd = vec![0usize; n];
700 let mut perm_inv = vec![0usize; n];
701 {
702 let mut mem = MemBuffer::new(amd::order_scratch::<usize>(n, a_nnz));
703 amd::order(
704 &mut perm_fwd,
705 &mut perm_inv,
706 h_upper.symbolic(),
707 amd::Control::default(),
708 MemStack::new(&mut mem),
709 )
710 .map_err(|_| LinalgError::ModelIsIllConditioned {
711 condition_number: f64::INFINITY,
712 })?;
713 }
714
715 let perm = unsafe { faer::perm::PermRef::new_unchecked(&perm_fwd, &perm_inv, n) };
721
722 let a_perm_upper = {
724 let mut col_ptrs = vec![0usize; n + 1];
725 let mut row_indices = vec![0usize; a_nnz];
726 let mut values = vec![0.0f64; a_nnz];
727 let mut mem = MemBuffer::new(faer::sparse::utils::permute_self_adjoint_scratch::<usize>(
728 n,
729 ));
730 faer::sparse::utils::permute_self_adjoint_to_unsorted(
731 &mut values,
732 &mut col_ptrs,
733 &mut row_indices,
734 h_upper.as_ref(),
735 perm,
736 Side::Upper,
737 Side::Upper,
738 MemStack::new(&mut mem),
739 );
740 SparseColMat::<usize, f64>::new(
741 unsafe { SymbolicSparseColMat::new_unchecked(n, n, col_ptrs, None, row_indices) },
750 values,
751 )
752 };
753
754 let symbolic = {
756 let mut mem = MemBuffer::new(StackReq::any_of(&[
757 simplicial::prefactorize_symbolic_cholesky_scratch::<usize>(n, a_nnz),
758 simplicial::factorize_simplicial_symbolic_cholesky_scratch::<usize>(n),
759 ]));
760 let stack = MemStack::new(&mut mem);
761 let mut etree = vec![0isize; n];
762 let mut col_counts = vec![0usize; n];
763 let etree_ref = simplicial::prefactorize_symbolic_cholesky(
764 &mut etree,
765 &mut col_counts,
766 a_perm_upper.symbolic(),
767 stack,
768 );
769 simplicial::factorize_simplicial_symbolic_cholesky(
770 a_perm_upper.symbolic(),
771 etree_ref,
772 &col_counts,
773 stack,
774 )
775 .map_err(|_| LinalgError::ModelIsIllConditioned {
776 condition_number: f64::INFINITY,
777 })?
778 };
779
780 let mut l_values = vec![0.0f64; symbolic.len_val()];
782 {
783 let mut mem = MemBuffer::new(simplicial::factorize_simplicial_numeric_llt_scratch::<
784 usize,
785 f64,
786 >(n));
787 simplicial::factorize_simplicial_numeric_llt::<usize, f64>(
788 &mut l_values,
789 a_perm_upper.as_ref(),
790 LltRegularization::default(),
791 &symbolic,
792 MemStack::new(&mut mem),
793 )
794 .map_err(|_| LinalgError::HessianNotPositiveDefinite {
795 min_eigenvalue: f64::NAN,
796 })?;
797 }
798
799 let l_col_ptr: Vec<usize> = symbolic.col_ptr().to_vec();
801 let l_row_idx: Vec<usize> = symbolic.row_idx().to_vec();
802
803 let mut logdet = 0.0f64;
805 for j in 0..n {
806 let diag = l_values[l_col_ptr[j]];
807 if diag <= 0.0 {
808 return Err(LinalgError::HessianNotPositiveDefinite {
809 min_eigenvalue: f64::NAN,
810 });
811 }
812 logdet += diag.ln();
813 }
814 logdet *= 2.0;
815
816 Ok(SimplicialFactor {
817 l_col_ptr,
818 l_row_idx,
819 l_values,
820 perm_inv,
821 n,
822 logdet,
823 })
824}
825
826impl SimplicialFactor {
827 fn assemble_h_dense_original_order(&self) -> Result<Array2<f64>, LinalgError> {
834 if self.perm_inv.len() != self.n {
835 bail_invalid_linalg!(
836 "simplicial factor permutation length {} does not match dimension {}",
837 self.perm_inv.len(),
838 self.n
839 );
840 }
841 let mut h_permuted = Array2::<f64>::zeros((self.n, self.n));
842 for col in 0..self.n {
843 let start = self.l_col_ptr[col];
844 let end = self.l_col_ptr[col + 1];
845 for left_idx in start..end {
846 let left_row = self.l_row_idx[left_idx];
847 let left_value = self.l_values[left_idx];
848 if !left_value.is_finite() {
849 bail_invalid_linalg!(
850 "simplicial factor has non-finite L entry at value index {left_idx}"
851 );
852 }
853 for right_idx in start..end {
854 let right_row = self.l_row_idx[right_idx];
855 let right_value = self.l_values[right_idx];
856 h_permuted[[left_row, right_row]] += left_value * right_value;
857 }
858 }
859 }
860
861 let mut h_original = Array2::<f64>::zeros((self.n, self.n));
862 for i in 0..self.n {
863 let pi = self.perm_inv[i];
864 if pi >= self.n {
865 bail_invalid_linalg!(
866 "simplicial factor permutation maps row {i} to out-of-bounds index {pi}"
867 );
868 }
869 for j in 0..self.n {
870 let pj = self.perm_inv[j];
871 if pj >= self.n {
872 bail_invalid_linalg!(
873 "simplicial factor permutation maps column {j} to out-of-bounds index {pj}"
874 );
875 }
876 let value = h_permuted[[pi, pj]];
877 if !value.is_finite() {
878 bail_invalid_linalg!(
879 "dense reconstruction from sparse Cholesky produced non-finite values"
880 );
881 }
882 h_original[[i, j]] = value;
883 }
884 }
885 Ok(h_original)
886 }
887}
888
889pub struct TakahashiInverse {
895 z_values: Vec<f64>,
897 col_ptr: Vec<usize>,
899 row_idx: Vec<usize>,
901 l_values: Vec<f64>,
903 rows_lower: Arc<Vec<Vec<(usize, f64)>>>,
905 exact_columns: Mutex<BTreeMap<usize, Arc<Vec<f64>>>>,
908 perm_inv: Vec<usize>,
910 n: usize,
912}
913
914impl TakahashiInverse {
915 fn find_entry(col_ptr: &[usize], row_idx: &[usize], row: usize, col: usize) -> Option<usize> {
918 let start = col_ptr[col];
919 let end = col_ptr[col + 1];
920 let slice = &row_idx[start..end];
921 slice.binary_search(&row).ok().map(|pos| start + pos)
922 }
923
924 fn solve_permuted_column_from_cholesky(
925 n: usize,
926 col_ptr: &[usize],
927 row_idx: &[usize],
928 l_values: &[f64],
929 rows_lower: &[Vec<(usize, f64)>],
930 rhs_col: usize,
931 ) -> Vec<f64> {
932 let mut rhs = vec![0.0f64; n];
933 rhs[rhs_col] = 1.0;
934 let mut forward = vec![0.0f64; n];
935 let mut solution = vec![0.0f64; n];
936
937 for row in 0..n {
938 let mut sum = rhs[row];
939 let mut diag = None;
940 for &(col, value) in &rows_lower[row] {
941 if col < row {
942 sum -= value * forward[col];
943 } else if col == row {
944 diag = Some(value);
945 }
946 }
947 let l_rr = diag.expect("simplicial factor row should contain its diagonal");
948 forward[row] = sum / l_rr;
949 }
950
951 for row in (0..n).rev() {
952 let col_start = col_ptr[row];
953 let col_end = col_ptr[row + 1];
954 let mut sum = forward[row];
955 let l_rr = l_values[col_start];
956 for idx in (col_start + 1)..col_end {
957 let lower_row = row_idx[idx];
958 sum -= l_values[idx] * solution[lower_row];
959 }
960 solution[row] = sum / l_rr;
961 }
962
963 solution
964 }
965
966 fn exact_permuted_column(&self, col: usize) -> Arc<Vec<f64>> {
967 {
968 let cache = self
969 .exact_columns
970 .lock()
971 .expect("exact Takahashi column cache mutex poisoned");
972 if let Some(solution) = cache.get(&col) {
973 return solution.clone();
974 }
975 }
976
977 let solution = Arc::new(Self::solve_permuted_column_from_cholesky(
978 self.n,
979 &self.col_ptr,
980 &self.row_idx,
981 &self.l_values,
982 self.rows_lower.as_ref(),
983 col,
984 ));
985
986 let mut cache = self
987 .exact_columns
988 .lock()
989 .expect("exact Takahashi column cache mutex poisoned");
990 cache.entry(col).or_insert_with(|| solution.clone()).clone()
991 }
992
993 fn selected_value(
994 z_values: &[f64],
995 col_ptr: &[usize],
996 row_idx: &[usize],
997 row: usize,
998 col: usize,
999 ) -> Result<f64, LinalgError> {
1000 let (lower_row, lower_col) = if row >= col { (row, col) } else { (col, row) };
1001 Self::find_entry(col_ptr, row_idx, lower_row, lower_col)
1002 .map(|idx| z_values[idx])
1003 .ok_or_else(|| {
1004 LinalgError::InvalidInput(format!(
1005 "simplicial selected-inverse pattern is missing entry ({lower_row},{lower_col})"
1006 ))
1007 })
1008 }
1009
1010 pub fn compute(factor: &SimplicialFactor) -> Result<Self, LinalgError> {
1016 let n = factor.n;
1017 let col_ptr = factor.l_col_ptr.clone();
1018 let row_idx = factor.l_row_idx.clone();
1019 let nnz = factor.l_values.len();
1020 let mut z_values = vec![0.0f64; nnz];
1021
1022 let mut rows_lower: Vec<Vec<(usize, f64)>> = vec![Vec::new(); n];
1024 for col in 0..n {
1025 for idx in col_ptr[col]..col_ptr[col + 1] {
1026 let row = row_idx[idx];
1027 rows_lower[row].push((col, factor.l_values[idx]));
1028 }
1029 }
1030
1031 for j in (0..n).rev() {
1032 let diag_idx = col_ptr[j];
1033 let col_end = col_ptr[j + 1];
1034 let diag = factor.l_values[diag_idx];
1035 if !(diag.is_finite() && diag > 0.0) {
1036 return Err(LinalgError::HessianNotPositiveDefinite {
1037 min_eigenvalue: f64::NAN,
1038 });
1039 }
1040 for idx in (diag_idx + 1)..col_end {
1041 let i = row_idx[idx];
1042 let mut correction = 0.0;
1043 for off_idx in (diag_idx + 1)..col_end {
1044 let k = row_idx[off_idx];
1045 let l_kj = factor.l_values[off_idx];
1046 let z_ik = Self::selected_value(&z_values, &col_ptr, &row_idx, i, k)?;
1047 correction += l_kj * z_ik;
1048 }
1049 let value = -correction / diag;
1050 if !value.is_finite() {
1051 bail_invalid_linalg!(
1052 "Takahashi selected inverse produced non-finite entry ({i},{j})"
1053 );
1054 }
1055 z_values[idx] = value;
1056 }
1057 let mut correction = 0.0;
1058 for off_idx in (diag_idx + 1)..col_end {
1059 correction += factor.l_values[off_idx] * z_values[off_idx];
1060 }
1061 let value = (1.0 / diag - correction) / diag;
1062 if !value.is_finite() {
1063 bail_invalid_linalg!(
1064 "Takahashi selected inverse produced non-finite diagonal entry ({j},{j})"
1065 );
1066 }
1067 z_values[diag_idx] = value;
1068 }
1069
1070 Ok(TakahashiInverse {
1071 z_values,
1072 col_ptr,
1073 row_idx,
1074 l_values: factor.l_values.clone(),
1075 rows_lower: Arc::new(rows_lower),
1076 exact_columns: Mutex::new(BTreeMap::new()),
1077 perm_inv: factor.perm_inv.clone(),
1078 n,
1079 })
1080 }
1081
1082 pub fn get(&self, i: usize, j: usize) -> f64 {
1084 let pi = self.perm_inv[i];
1085 let pj = self.perm_inv[j];
1086 self.get_permuted(pi, pj)
1087 }
1088
1089 fn get_permuted(&self, pi: usize, pj: usize) -> f64 {
1091 let (row, col) = if pi >= pj { (pi, pj) } else { (pj, pi) };
1094 if let Some(pos) = Self::find_entry(&self.col_ptr, &self.row_idx, row, col) {
1095 self.z_values[pos]
1096 } else {
1097 self.exact_permuted_column(col)[row]
1098 }
1099 }
1100
1101 pub fn diagonal(&self) -> Array1<f64> {
1103 Array1::from_iter((0..self.n).map(|i| self.get(i, i)))
1104 }
1105
1106 pub fn block(&self, start: usize, end: usize) -> Array2<f64> {
1108 let dim = end - start;
1109 let mut out = Array2::zeros((dim, dim));
1110 for j_local in 0..dim {
1111 let j = start + j_local;
1112 for i_local in 0..dim {
1113 let i = start + i_local;
1114 out[[i_local, j_local]] = self.get(i, j);
1115 }
1116 }
1117 out
1118 }
1119
1120 pub fn trace_product_sparse(&self, s: &SparseColMat<usize, f64>) -> f64 {
1138 let (symbolic, values) = s.parts();
1139 let s_col_ptr = symbolic.col_ptr();
1140 let s_row_idx = symbolic.row_idx();
1141 let per_column: Vec<f64> = (0..s.ncols())
1157 .into_par_iter()
1158 .map(|col| {
1159 let col_start = s_col_ptr[col];
1160 let col_end = s_col_ptr[col + 1];
1161 let mut partial = 0.0;
1162 for idx in col_start..col_end {
1163 let row = s_row_idx[idx];
1164 if row > col {
1165 continue; }
1167 let val = values[idx];
1168 let z_ij = self.get(row, col);
1169 if row == col {
1170 partial += z_ij * val;
1171 } else {
1172 partial += 2.0 * z_ij * val;
1173 }
1174 }
1175 partial
1176 })
1177 .collect();
1178 per_column.iter().sum()
1179 }
1180}
1181
1182#[cfg(test)]
1183mod tests {
1184 use super::*;
1185 use crate::faer_ndarray::FaerCholesky;
1186 use ndarray::{Array1, Array2, array};
1187
1188 fn approx_eq(a: f64, b: f64, tol: f64) {
1189 assert!(
1190 (a - b).abs() <= tol,
1191 "values differ: left={a:.12e}, right={b:.12e}, |diff|={:.12e}, tol={tol:.12e}",
1192 (a - b).abs()
1193 );
1194 }
1195
1196 #[test]
1199 fn dense_to_sparse_preserves_all_nonzero_entries() {
1200 let m = array![[1.0, 2.0, 3.0], [0.0, 5.0, 6.0], [7.0, 8.0, 9.0]];
1202 let s = dense_to_sparse(&m, ZERO_TOL).unwrap();
1203 assert_eq!(s.nrows(), 3);
1204 assert_eq!(s.ncols(), 3);
1205 assert_eq!(s.compute_nnz(), 8);
1207 }
1208
1209 #[test]
1210 fn dense_to_sparse_round_trips_via_matvec_identity() {
1211 let m = array![[4.0, 1.0, 0.5], [1.0, 3.0, 2.0], [0.5, 2.0, 6.0]];
1213 let s = dense_to_sparse(&m, ZERO_TOL).unwrap();
1214 for j in 0..3 {
1215 let mut ej = Array1::<f64>::zeros(3);
1216 ej[j] = 1.0;
1217 let result = {
1219 let mut out = Array1::<f64>::zeros(3);
1220 let (sym, vals) = s.parts();
1221 let col_ptr = sym.col_ptr();
1222 let row_idx = sym.row_idx();
1223 for col in 0..3 {
1224 for idx in col_ptr[col]..col_ptr[col + 1] {
1225 let row = row_idx[idx];
1226 out[row] += vals[idx] * ej[col];
1227 }
1228 }
1229 out
1230 };
1231 for i in 0..3 {
1232 approx_eq(result[i], m[[i, j]], 1e-14);
1233 }
1234 }
1235 }
1236
1237 #[test]
1238 fn dense_to_sparse_filters_entries_below_tolerance() {
1239 let tol = 0.1;
1240 let m = array![[1.0, 0.05], [0.05, 2.0]];
1241 let s = dense_to_sparse(&m, tol).unwrap();
1242 assert_eq!(
1244 s.compute_nnz(),
1245 2,
1246 "off-diagonal entries below tol must be dropped"
1247 );
1248 }
1249
1250 #[test]
1253 fn dense_to_sparse_symmetric_upper_stores_only_upper_triangle() {
1254 let m = array![[4.0, 1.0, 2.0], [1.0, 5.0, 3.0], [2.0, 3.0, 6.0]];
1256 let s = dense_to_sparse_symmetric_upper(&m, ZERO_TOL).unwrap();
1257 assert_eq!(s.compute_nnz(), 6);
1259 }
1260
1261 #[test]
1264 fn sparse_symmetric_upper_matvec_matches_dense_matvec() {
1265 let a = array![[4.0, 2.0, 0.0], [2.0, 5.0, 3.0], [0.0, 3.0, 6.0]];
1268 let v = array![1.0, 2.0, 3.0];
1269 let expected = a.dot(&v); let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1271 let got = sparse_symmetric_upper_matvec_public(&a_sparse, &v);
1272 for i in 0..3 {
1273 approx_eq(got[i], expected[i], 1e-13);
1274 }
1275 }
1276
1277 #[test]
1278 fn sparse_symmetric_upper_matvec_diagonal_only() {
1279 let a = array![[3.0, 0.0, 0.0], [0.0, 5.0, 0.0], [0.0, 0.0, 7.0]];
1281 let v = array![2.0, 4.0, 6.0];
1282 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1283 let got = sparse_symmetric_upper_matvec_public(&a_sparse, &v);
1284 approx_eq(got[0], 6.0, 1e-14);
1285 approx_eq(got[1], 20.0, 1e-14);
1286 approx_eq(got[2], 42.0, 1e-14);
1287 }
1288
1289 #[test]
1292 fn solve_sparse_spd_recovers_known_solution() {
1293 let a = array![[4.0, 2.0], [2.0, 5.0]];
1295 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1296 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1297 let rhs = array![6.0, 11.0];
1298 let x = solve_sparse_spd(&factor, &rhs).unwrap();
1299 approx_eq(x[0], 0.5, 1e-12);
1301 approx_eq(x[1], 2.0, 1e-12);
1302 }
1303
1304 #[test]
1305 fn strict_sparse_spd_preserves_sub_threshold_stored_entries() {
1306 let tiny = 5.0e-13;
1307 let matrix = array![[1.0, tiny], [tiny, 1.0]];
1308 let sparse = dense_to_sparse_symmetric_upper(&matrix, 0.0).unwrap();
1309 let factor = factorize_sparse_spd_strict(&sparse).unwrap();
1310 let solution = solve_sparse_spd(&factor, &array![0.0, 1.0]).unwrap();
1311 assert!(solution[0] < 0.0, "tiny coupling must not be dropped");
1312 approx_eq(solution[0], -tiny / (1.0 - tiny * tiny), 1.0e-27);
1313 }
1314
1315 #[test]
1316 fn strict_sparse_spd_rejects_full_or_lower_triangle_storage() {
1317 let matrix = array![[2.0, 0.5], [0.5, 3.0]];
1318 let full = dense_to_sparse(&matrix, 0.0).unwrap();
1319 let error = factorize_sparse_spd_strict(&full)
1320 .err()
1321 .expect("full symmetric storage must be rejected");
1322 assert!(error.to_string().contains("upper-triangle storage"));
1323 }
1324
1325 #[test]
1326 fn solve_sparse_spd_3x3_round_trip() {
1327 let a: Array2<f64> = array![[9.0, 3.0, 1.0], [3.0, 8.0, 2.0], [1.0, 2.0, 7.0]];
1328 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1329 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1330 for j in 0..3 {
1331 let mut ej = Array1::<f64>::zeros(3);
1332 ej[j] = 1.0;
1333 let col_j = solve_sparse_spd(&factor, &ej).unwrap();
1334 let ax = a.dot(&col_j);
1336 for i in 0..3 {
1337 approx_eq(ax[i], ej[i], 1e-12);
1338 }
1339 }
1340 }
1341
1342 #[test]
1343 fn logdet_from_factor_matches_dense_logdet_diagonal() {
1344 let a: Array2<f64> = array![[4.0, 0.0, 0.0], [0.0, 9.0, 0.0], [0.0, 0.0, 16.0]];
1346 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1347 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1348 let logdet = logdet_from_factor(&factor).unwrap();
1349 let expected = 4.0_f64.ln() + 9.0_f64.ln() + 16.0_f64.ln();
1350 approx_eq(logdet, expected, 1e-12);
1351 }
1352
1353 #[test]
1354 fn logdet_from_factor_matches_2x2_formula() {
1355 let a = array![[4.0, 2.0], [2.0, 5.0]];
1357 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1358 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1359 let logdet = logdet_from_factor(&factor).unwrap();
1360 approx_eq(logdet, 16.0_f64.ln(), 1e-12);
1361 }
1362
1363 #[test]
1364 fn solve_sparse_spd_dimension_mismatch_returns_error() {
1365 let a = array![[4.0, 2.0], [2.0, 5.0]];
1366 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1367 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1368 let rhs = array![1.0, 2.0, 3.0]; assert!(solve_sparse_spd(&factor, &rhs).is_err());
1370 }
1371
1372 #[test]
1373 fn takahashi_diagonal_matches_dense_inverse() {
1374 let h = array![
1376 [4.0, 0.2, 0.0, 0.0],
1377 [0.2, 3.0, 0.1, 0.0],
1378 [0.0, 0.1, 2.5, 0.3],
1379 [0.0, 0.0, 0.3, 2.0]
1380 ];
1381 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1382
1383 let chol = h.cholesky(Side::Lower).unwrap();
1385 let mut h_inv = Array2::<f64>::zeros((4, 4));
1386 for j in 0..4 {
1387 let mut rhs = Array1::<f64>::zeros(4);
1388 rhs[j] = 1.0;
1389 let col = chol.solvevec(&rhs);
1390 for i in 0..4 {
1391 h_inv[[i, j]] = col[i];
1392 }
1393 }
1394
1395 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1396 let taka = TakahashiInverse::compute(&sfactor).unwrap();
1397 let diag = taka.diagonal();
1398
1399 for i in 0..4 {
1401 approx_eq(diag[i], h_inv[[i, i]], 1e-10);
1402 }
1403 }
1404
1405 #[test]
1406 fn takahashi_logdet_matches_dense() {
1407 let h = array![
1408 [4.0, 0.2, 0.0, 0.0],
1409 [0.2, 3.0, 0.1, 0.0],
1410 [0.0, 0.1, 2.5, 0.3],
1411 [0.0, 0.0, 0.3, 2.0]
1412 ];
1413 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1414
1415 let existing = factorize_sparse_spd(&h_sparse).unwrap();
1417 let logdet_dense = existing.logdet;
1418
1419 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1420 approx_eq(sfactor.logdet, logdet_dense, 1e-10);
1421 }
1422
1423 fn dense_trace_ref(h: &Array2<f64>, s: &Array2<f64>) -> f64 {
1428 let n = h.nrows();
1429 let chol = h.cholesky(Side::Lower).unwrap();
1430 let mut h_inv = Array2::<f64>::zeros((n, n));
1431 for j in 0..n {
1432 let mut rhs = Array1::<f64>::zeros(n);
1433 rhs[j] = 1.0;
1434 let col = chol.solvevec(&rhs);
1435 for i in 0..n {
1436 h_inv[[i, j]] = col[i];
1437 }
1438 }
1439 let mut trace = 0.0;
1441 for i in 0..n {
1442 for j in 0..n {
1443 trace += h_inv[[i, j]] * s[[i, j]];
1444 }
1445 }
1446 trace
1447 }
1448
1449 #[test]
1450 fn trace_product_sparse_matches_dense_small() {
1451 let h = array![
1454 [4.0, 0.2, 0.0, 0.0],
1455 [0.2, 3.0, 0.1, 0.0],
1456 [0.0, 0.1, 2.5, 0.3],
1457 [0.0, 0.0, 0.3, 2.0]
1458 ];
1459 let s = array![
1462 [1.0, 0.5, 0.0, 0.7],
1463 [0.5, 2.0, 0.3, 0.0],
1464 [0.0, 0.3, 1.5, 0.4],
1465 [0.7, 0.0, 0.4, 3.0]
1466 ];
1467
1468 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1469 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1470 let taka = TakahashiInverse::compute(&sfactor).unwrap();
1471
1472 let s_sparse = dense_to_sparse(&s, ZERO_TOL).unwrap();
1474 let got = taka.trace_product_sparse(&s_sparse);
1475 let expected = dense_trace_ref(&h, &s);
1476
1477 let rel = (got - expected).abs() / expected.abs().max(1.0);
1478 assert!(
1479 rel <= 1e-9,
1480 "trace mismatch: got={got:.15e}, expected={expected:.15e}, rel={rel:.3e}"
1481 );
1482 }
1483
1484 #[test]
1485 fn trace_product_sparse_matches_dense_large_parallel() {
1486 let n = 40usize;
1488 let mut h = Array2::<f64>::zeros((n, n));
1489 let mut s = Array2::<f64>::zeros((n, n));
1490 for i in 0..n {
1491 h[[i, i]] = (n as f64) + 5.0 + (i as f64) * 0.1;
1493 s[[i, i]] = 1.0 + (i as f64) * 0.05;
1494 }
1495 for i in 0..n - 1 {
1497 let v = 0.3 + 0.01 * (i as f64);
1498 h[[i, i + 1]] = v;
1499 h[[i + 1, i]] = v;
1500 }
1501 for i in 0..n - 3 {
1505 let v = 0.2 + 0.005 * (i as f64);
1506 s[[i, i + 3]] = v;
1507 s[[i + 3, i]] = v;
1508 }
1509 s[[0, n - 1]] = 0.4;
1510 s[[n - 1, 0]] = 0.4;
1511 s[[2, n - 5]] = 0.25;
1512 s[[n - 5, 2]] = 0.25;
1513
1514 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1515 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1516 let taka = TakahashiInverse::compute(&sfactor).unwrap();
1517
1518 let s_sparse = dense_to_sparse(&s, ZERO_TOL).unwrap();
1519 let got = taka.trace_product_sparse(&s_sparse);
1520 let expected = dense_trace_ref(&h, &s);
1521
1522 let rel = (got - expected).abs() / expected.abs().max(1.0);
1523 assert!(
1524 rel <= 1e-9,
1525 "trace mismatch (n={n}): got={got:.15e}, expected={expected:.15e}, rel={rel:.3e}"
1526 );
1527 }
1528
1529 #[test]
1532 fn solve_sparse_spdmulti_recovers_identity_inverse() {
1533 let a = array![[4.0, 0.0], [0.0, 9.0]];
1535 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1536 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1537 let rhs = Array2::<f64>::eye(2);
1539 let inv = solve_sparse_spdmulti(&factor, &rhs).unwrap();
1540 approx_eq(inv[[0, 0]], 0.25, 1e-12);
1541 approx_eq(inv[[0, 1]], 0.0, 1e-12);
1542 approx_eq(inv[[1, 0]], 0.0, 1e-12);
1543 approx_eq(inv[[1, 1]], 1.0 / 9.0, 1e-12);
1544 }
1545
1546 #[test]
1547 fn solve_sparse_spdmulti_3x3_matches_column_wise_solve() {
1548 let a: Array2<f64> = array![[9.0, 3.0, 1.0], [3.0, 8.0, 2.0], [1.0, 2.0, 7.0]];
1549 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1550 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1551 let rhs = array![[1.0, 0.0], [0.0, 1.0], [0.0, 0.0]];
1553 let x = solve_sparse_spdmulti(&factor, &rhs).unwrap();
1554 for j in 0..2 {
1556 let xj = x.column(j);
1557 let ax = a.dot(&xj);
1558 for i in 0..3 {
1559 approx_eq(ax[i], rhs[[i, j]], 1e-11);
1560 }
1561 }
1562 }
1563
1564 #[test]
1565 fn solve_sparse_spdmulti_rows_selects_subset_of_rows() {
1566 let a = array![[4.0, 2.0], [2.0, 5.0]];
1568 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1569 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1570 let rhs = Array2::<f64>::eye(2);
1572 let row0 = solve_sparse_spdmulti_rows(&factor, &rhs, 0, 1).unwrap();
1573 assert_eq!(row0.dim(), (1, 2));
1576 approx_eq(row0[[0, 0]], 5.0 / 16.0, 1e-12);
1577 approx_eq(row0[[0, 1]], -2.0 / 16.0, 1e-12);
1578 }
1579
1580 #[test]
1581 fn solve_sparse_spdmulti_diagonal_sum_matches_trace_of_partial_inverse() {
1582 let a = array![[4.0, 2.0], [2.0, 5.0]];
1585 let a_sparse = dense_to_sparse_symmetric_upper(&a, ZERO_TOL).unwrap();
1586 let factor = factorize_sparse_spd(&a_sparse).unwrap();
1587 let rhs = Array2::<f64>::eye(2);
1588 let diag_sum = solve_sparse_spdmulti_diagonal_sum(&factor, &rhs, 0).unwrap();
1589 approx_eq(diag_sum, 9.0 / 16.0, 1e-12);
1591 }
1592
1593 #[test]
1594 fn takahashi_get_and_block_recover_off_pattern_inverse_entries() {
1595 let h = array![
1596 [4.0, 1.0, 0.0, 0.0],
1597 [1.0, 3.0, 1.0, 0.0],
1598 [0.0, 1.0, 2.5, 1.0],
1599 [0.0, 0.0, 1.0, 2.0]
1600 ];
1601 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1602
1603 let chol = h.cholesky(Side::Lower).unwrap();
1604 let mut h_inv = Array2::<f64>::zeros((4, 4));
1605 for j in 0..4 {
1606 let mut rhs = Array1::<f64>::zeros(4);
1607 rhs[j] = 1.0;
1608 let col = chol.solvevec(&rhs);
1609 for i in 0..4 {
1610 h_inv[[i, j]] = col[i];
1611 }
1612 }
1613
1614 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1615 let taka = TakahashiInverse::compute(&sfactor).unwrap();
1616
1617 assert!(
1618 h_inv[[0, 2]].abs() > 1e-8,
1619 "reference off-pattern inverse entry should be nonzero"
1620 );
1621 approx_eq(taka.get(0, 2), h_inv[[0, 2]], 1e-10);
1622
1623 let block = taka.block(0, 3);
1624 approx_eq(block[[0, 2]], h_inv[[0, 2]], 1e-10);
1625 approx_eq(block[[2, 0]], h_inv[[2, 0]], 1e-10);
1626 }
1627
1628 fn dense_inverse_spd(h: &Array2<f64>) -> Array2<f64> {
1630 let n = h.nrows();
1631 let chol = h.cholesky(Side::Lower).unwrap();
1632 let mut inv = Array2::<f64>::zeros((n, n));
1633 for j in 0..n {
1634 let mut rhs = Array1::<f64>::zeros(n);
1635 rhs[j] = 1.0;
1636 let col = chol.solvevec(&rhs);
1637 for i in 0..n {
1638 inv[[i, j]] = col[i];
1639 }
1640 }
1641 inv
1642 }
1643
1644 fn dense_trace_product(z: &Array2<f64>, s_dense: &Array2<f64>) -> f64 {
1646 let n = z.nrows();
1647 let mut acc = 0.0;
1648 for i in 0..n {
1649 for j in 0..n {
1650 acc += z[[i, j]] * s_dense[[j, i]];
1651 }
1652 }
1653 acc
1654 }
1655
1656 #[test]
1657 fn trace_product_sparse_matches_dense_reference_small() {
1658 let h = array![
1661 [4.0, 1.0, 0.0, 0.0],
1662 [1.0, 3.0, 1.0, 0.0],
1663 [0.0, 1.0, 2.5, 1.0],
1664 [0.0, 0.0, 1.0, 2.0]
1665 ];
1666 let s = array![
1667 [2.0, 0.5, 0.0, 0.1],
1668 [0.5, 1.5, 0.3, 0.0],
1669 [0.0, 0.3, 1.0, 0.4],
1670 [0.1, 0.0, 0.4, 3.0]
1671 ];
1672
1673 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1674 let s_sparse = dense_to_sparse_symmetric_upper(&s, ZERO_TOL).unwrap();
1675 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1676 let taka = TakahashiInverse::compute(&sfactor).unwrap();
1677
1678 let h_inv = dense_inverse_spd(&h);
1679 let expected = dense_trace_product(&h_inv, &s);
1680 approx_eq(taka.trace_product_sparse(&s_sparse), expected, 1e-9);
1681 }
1682
1683 #[test]
1690 fn trace_product_sparse_parallel_matches_dense_reference_large() {
1691 let n = 40usize;
1693 let mut h = Array2::<f64>::zeros((n, n));
1694 for i in 0..n {
1695 h[[i, i]] = 4.0 + (i as f64) * 0.01;
1696 if i + 1 < n {
1697 h[[i, i + 1]] = 1.0;
1698 h[[i + 1, i]] = 1.0;
1699 }
1700 }
1701 let mut s = Array2::<f64>::zeros((n, n));
1705 for i in 0..n {
1706 s[[i, i]] = 2.0 + (i as f64) * 0.05;
1707 }
1708 for &(i, j, v) in &[
1709 (0usize, 3usize, 0.7f64),
1710 (1, 5, -0.4),
1711 (2, 9, 0.3),
1712 (4, 20, 0.25),
1713 (7, 30, -0.15),
1714 (10, 39, 0.2),
1715 (15, 22, 0.35),
1716 ] {
1717 s[[i, j]] = v;
1718 s[[j, i]] = v;
1719 }
1720
1721 let h_sparse = dense_to_sparse_symmetric_upper(&h, ZERO_TOL).unwrap();
1722 let s_sparse = dense_to_sparse_symmetric_upper(&s, ZERO_TOL).unwrap();
1723 let sfactor = factorize_simplicial(&h_sparse).unwrap();
1724 let taka = TakahashiInverse::compute(&sfactor).unwrap();
1725
1726 let h_inv = dense_inverse_spd(&h);
1727 let expected = dense_trace_product(&h_inv, &s);
1728
1729 let got = taka.trace_product_sparse(&s_sparse);
1730 approx_eq(got, expected, 1e-8);
1731
1732 let got_again = taka.trace_product_sparse(&s_sparse);
1735 assert_eq!(
1736 got, got_again,
1737 "parallel trace_product_sparse must be deterministic across calls"
1738 );
1739
1740 let pool1 = rayon::ThreadPoolBuilder::new()
1747 .num_threads(1)
1748 .build()
1749 .unwrap();
1750 let pool8 = rayon::ThreadPoolBuilder::new()
1751 .num_threads(8)
1752 .build()
1753 .unwrap();
1754 let got_1t = pool1.install(|| taka.trace_product_sparse(&s_sparse));
1755 let got_8t = pool8.install(|| taka.trace_product_sparse(&s_sparse));
1756 assert_eq!(
1757 got_1t, got_8t,
1758 "trace_product_sparse must be bit-identical across 1 vs 8 rayon workers"
1759 );
1760 assert_eq!(
1761 got, got_1t,
1762 "default-pool result must match the single-worker reduction"
1763 );
1764 }
1765}