Skip to main content

gam_linalg/
sparse_exact.rs

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    // Direct column-major CSC construction.  Three-pass: count nnz per
50    // column in parallel, perform the prefix sum serially, then fill each
51    // deterministic column slice in parallel.  Columns are still traversed
52    // in order and rows are written in ascending order within each column,
53    // preserving the same canonical CSC ordering as the previous serial
54    // implementation without requiring a triplet sort/dedup pass.
55    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
76/// Convert a dense symmetric matrix to sparse CSC storing only the upper triangle.
77///
78/// This encoding is required by sparse SPD routines in this module that interpret
79/// entries as symmetric-upper storage and mirror off-diagonals when reconstructing
80/// dense diagnostics.
81pub 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    // Direct CSC build over the upper triangle.  Counts and fills are
88    // parallelized by column, with a serial prefix sum between them so every
89    // column writes to a deterministic, non-overlapping slice.  Iterating rows
90    // from low to high within each column keeps CSC row indices sorted exactly
91    // as in the previous serial implementation.
92    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    // Canonicalize to symmetric-upper storage before factorization.
275    //
276    // Math contract:
277    // - If callers pass upper-only storage, values are preserved.
278    // - If callers pass full symmetric storage, paired (i,j)/(j,i) entries are averaged.
279    // - If callers pass lower-only storage, it is mirrored into upper.
280    //
281    // This prevents off-diagonal double counting in paths that interpret input as
282    // symmetric-upper and makes the sparse factor path robust to caller encoding.
283    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    // Keep an explicit simplicial LLᵀ factor in addition to faer's solver
292    // object. The raw L is needed by callers that must reconstruct H in a
293    // changed basis, such as active-constraint tangent projection.
294    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
312/// Strict SPD factorization of canonical symmetric-upper CSC storage.
313///
314/// Unlike [`factorize_sparse_spd`], this covariance/inference entrypoint does
315/// not average triangles, combine duplicates, or drop small stored entries.
316/// Callers must provide exactly one sorted upper-triangle entry per stored
317/// coordinate. The matrix passed to sparse Cholesky is therefore bit-for-bit
318/// the matrix whose residual downstream inference certifies.
319pub 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
501/// In-place variant of [`solve_sparse_spd`]. Writes the solution directly into
502/// `out`, avoiding the intermediate `Array1` allocation on the hot PIRLS path.
503/// `out` must already be sized to match `factor.n` (typically the reused
504/// Newton-direction buffer).
505pub 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
642// ---------------------------------------------------------------------------
643// Takahashi selected inversion via simplicial Cholesky
644// ---------------------------------------------------------------------------
645
646use 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
651/// A simplicial Cholesky factorization with raw access to L's CSC pattern and
652/// values, plus the AMD permutation.  Built using faer's low-level simplicial
653/// API so that L's sparse structure is directly available for Takahashi
654/// selected inversion.
655pub struct SimplicialFactor {
656    /// Column pointers of L (lower triangular, CSC), length n+1
657    l_col_ptr: Vec<usize>,
658    /// Row indices of L (lower triangular, CSC), length nnz(L)
659    l_row_idx: Vec<usize>,
660    /// Numeric values of L, length nnz(L)
661    l_values: Vec<f64>,
662    /// Inverse permutation returned by faer, used to map original coordinates
663    /// into the permuted simplicial factor basis.
664    perm_inv: Vec<usize>,
665    /// Dimension
666    n: usize,
667    /// log|H| = 2 * sum(log(L_ii))
668    pub logdet: f64,
669}
670
671/// Build a [`SimplicialFactor`] from a symmetric CSC matrix (upper, lower, or
672/// full storage – it is canonicalized to symmetric-upper internally).
673///
674/// The factorization uses AMD fill-reducing ordering and faer's simplicial
675/// LLᵀ numeric factorization.
676pub 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    // 1. AMD ordering
699    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    // perm_fwd and perm_inv have length n and were just populated by
716    // amd::order above for a valid symmetric n×n CSC matrix. On success,
717    // amd::order writes a valid permutation of 0..n into perm_fwd and its
718    // exact inverse into perm_inv.
719    // SAFETY: those are exactly the invariants required by PermRef::new_unchecked.
720    let perm = unsafe { faer::perm::PermRef::new_unchecked(&perm_fwd, &perm_inv, n) };
721
722    // 2. Permute to P A Pᵀ (upper-triangular, unsorted)
723    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            // col_ptrs and row_indices were just produced into preallocated
742            // buffers by permute_self_adjoint_to_unsorted from a valid n×n
743            // symbolic CSC and a valid permutation. That routine writes an
744            // unsorted CSC with col_ptrs length n + 1, monotone column ranges
745            // within row_indices, and every row index in 0..n.
746            // SAFETY: those are the hard SymbolicSparseColMat invariants; the
747            // following faer symbolic Cholesky routines accept this unsorted
748            // self-adjoint permutation.
749            unsafe { SymbolicSparseColMat::new_unchecked(n, n, col_ptrs, None, row_indices) },
750            values,
751        )
752    };
753
754    // 3. Symbolic analysis
755    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    // 4. Numeric LLᵀ factorization
781    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    // 5. Extract col_ptr, row_idx from the symbolic structure
800    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    // 6. Compute logdet from L diagonal: L[j,j] = l_values[l_col_ptr[j]]
804    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    /// Reconstruct the original-order dense SPD matrix represented by this
828    /// permuted sparse Cholesky factor.
829    ///
830    /// The simplicial factor stores `L` for `P H Pᵀ = L Lᵀ`, with
831    /// `perm_inv[original] = permuted`. We first assemble the dense permuted
832    /// product and then map rows/columns back to the caller's coordinate order.
833    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
889/// Result of the Takahashi selected inversion.
890///
891/// Z stores entries of H⁻¹ at positions corresponding to the filled sparsity
892/// pattern of the Cholesky factor L. Off-pattern entries are recovered exactly
893/// on demand by cached column solves against the same simplicial factor.
894pub struct TakahashiInverse {
895    /// Z values stored in the same CSC pattern as L (lower triangular)
896    z_values: Vec<f64>,
897    /// Column pointers (owned copy from L)
898    col_ptr: Vec<usize>,
899    /// Row indices (owned copy from L)
900    row_idx: Vec<usize>,
901    /// Numeric values of the simplicial Cholesky factor L.
902    l_values: Vec<f64>,
903    /// Row-oriented access to L for forward solves in the permuted basis.
904    rows_lower: Arc<Vec<Vec<(usize, f64)>>>,
905    /// Exact inverse columns solved on demand for entries outside the selected
906    /// inverse pattern. Keys are permuted-basis column indices.
907    exact_columns: Mutex<BTreeMap<usize, Arc<Vec<f64>>>>,
908    /// Inverse permutation returned by faer.
909    perm_inv: Vec<usize>,
910    /// Dimension
911    n: usize,
912}
913
914impl TakahashiInverse {
915    /// Binary search for entry (row, col) in lower-triangular CSC.
916    /// Returns the value-array index if the entry exists.
917    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    /// Compute the selected inverse from a simplicial Cholesky factor.
1011    ///
1012    /// Given H = LLᵀ in the permuted basis, this applies the Takahashi
1013    /// recurrence on the filled Cholesky pattern. Off-pattern exact entries are
1014    /// recovered later by cached column solves from the same simplicial factor.
1015    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        // Build row access for forward solves in the permuted basis.
1023        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    /// Get H⁻¹[i,j] in ORIGINAL (unpermuted) coordinates.
1083    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    /// Get Z[pi,pj] in permuted coordinates.
1090    fn get_permuted(&self, pi: usize, pj: usize) -> f64 {
1091        // Z is symmetric and stored as lower-triangular CSC.
1092        // Ensure row >= col for lookup.
1093        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    /// Diagonal of H⁻¹ in original ordering.
1102    pub fn diagonal(&self) -> Array1<f64> {
1103        Array1::from_iter((0..self.n).map(|i| self.get(i, i)))
1104    }
1105
1106    /// H⁻¹[start..end, start..end] block in original ordering.
1107    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    /// tr(H⁻¹ S) where S is given as sparse CSC, symmetric in either upper-
1121    /// triangle-only or full (both triangles stored) format.
1122    ///
1123    /// The algorithm iterates over the upper triangle of S (entries with
1124    /// row ≤ col), doubles off-diagonals, and skips lower-triangle entries.
1125    /// This is correct for both storage conventions:
1126    ///
1127    /// - **Upper-triangle-only** (for example, solver-owned sparse penalty blocks):
1128    ///   every off-diagonal pair has exactly one stored entry with row < col,
1129    ///   which we double.
1130    ///
1131    /// - **Full symmetric** (from `dense_to_sparse`): each off-diagonal pair
1132    ///   has entries at both (i,j) and (j,i).  We process only the row < col
1133    ///   entry and double it; the row > col mirror is skipped.  The diagonal
1134    ///   is stored once and counted once.
1135    ///
1136    /// In both cases: tr(Z S) = Σ_diag Z[i,i] S[i,i] + 2 Σ_{i<j} Z[i,j] S[i,j].
1137    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        // tr(Z S) = Σ_diag Z[i,i] S[i,i] + 2 Σ_{i<j} Z[i,j] S[i,j]. Each column's
1142        // contribution is independent of every other column's, so the expensive
1143        // per-column work (`self.get` — on-demand exact-inverse-column solves,
1144        // `Mutex`-guarded via `exact_columns`, so concurrent lookups including
1145        // cache misses are sound) fans across rayon. But the FINAL reduction must
1146        // NOT be a rayon `.sum()`: that folds partials in a work-stealing tree
1147        // order, so the low-order bits of `tr(ZS)` would vary with thread count /
1148        // scheduling — and this value feeds the REML gradient / EDF, so that drift
1149        // makes fits non-reproducible across machines with different core counts.
1150        // Collect the per-column partials in column order (`collect()` on an
1151        // indexed parallel iterator preserves order) and sum them SERIALLY, so
1152        // the result is bit-identical regardless of the thread pool — matching the
1153        // index-ordered-serial-reduction idiom the Firth outer-Hessian paths use.
1154        // (Parallelization first landed for #759 in b7879667b, lost in the
1155        // gam-linalg crate-extraction refactor a80fe6943, restored here.)
1156        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; // skip lower triangle (handled via its mirror)
1166                    }
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    // ── dense_to_sparse ───────────────────────────────────────────────────
1197
1198    #[test]
1199    fn dense_to_sparse_preserves_all_nonzero_entries() {
1200        // 3x3 matrix with a zero at (1,0) and all others nonzero.
1201        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        // 8 entries should be stored (one zero excluded).
1206        assert_eq!(s.compute_nnz(), 8);
1207    }
1208
1209    #[test]
1210    fn dense_to_sparse_round_trips_via_matvec_identity() {
1211        // Verify that (sparse A) * e_j == column j of A for each column.
1212        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            // Multiply via the raw faer sparse multiply.
1218            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        // Only the two diagonal entries exceed tol.
1243        assert_eq!(
1244            s.compute_nnz(),
1245            2,
1246            "off-diagonal entries below tol must be dropped"
1247        );
1248    }
1249
1250    // ── dense_to_sparse_symmetric_upper ───────────────────────────────────
1251
1252    #[test]
1253    fn dense_to_sparse_symmetric_upper_stores_only_upper_triangle() {
1254        // Full symmetric 3x3 matrix — only upper triangle (i<=j) should be stored.
1255        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        // Upper triangle has 3 diagonal + 3 off-diagonal = 6 entries.
1258        assert_eq!(s.compute_nnz(), 6);
1259    }
1260
1261    // ── sparse_symmetric_upper_matvec_public ──────────────────────────────
1262
1263    #[test]
1264    fn sparse_symmetric_upper_matvec_matches_dense_matvec() {
1265        // Symmetric matrix A; upper-sparse encodes only the upper triangle.
1266        // A * v must equal the result of the symmetric matvec.
1267        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); // dense reference
1270        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        // Pure diagonal matrix: matvec should scale each component.
1280        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    // ── solve_sparse_spd / logdet_from_factor ─────────────────────────────
1290
1291    #[test]
1292    fn solve_sparse_spd_recovers_known_solution() {
1293        // A = [[4,2],[2,5]]; A^{-1} b = [0.5, 2.0] for b = [6, 11].
1294        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        // A^-1 = (1/16)*[[5,-2],[-2,4]]; x = (1/16)*[5*6-2*11, -2*6+4*11] = [0.5, 2.0]
1300        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            // A * x should equal ej.
1335            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        // Diagonal matrix diag(4,9,16): log-det = log(4)+log(9)+log(16)
1345        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        // A = [[4,2],[2,5]]; det(A) = 20-4 = 16; log-det = log(16)
1356        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]; // wrong length
1369        assert!(solve_sparse_spd(&factor, &rhs).is_err());
1370    }
1371
1372    #[test]
1373    fn takahashi_diagonal_matches_dense_inverse() {
1374        // 4x4 SPD matrix
1375        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        // Dense inverse for reference via column solves
1384        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        // Diagonal of selected inverse should match dense inverse diagonal
1400        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        // Dense logdet via existing factor
1416        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    // ── trace_product_sparse (rayon parallel reduction, #759) ─────────────────
1424
1425    /// Dense reference tr(H^{-1} S) computed by a completely different code path:
1426    /// an explicit dense Cholesky-based inverse and an elementwise double sum.
1427    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        // tr(H^{-1} S) = sum_ij (H^{-1})_ij S_ij  (S symmetric here).
1440        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        // Small SPD H with off-diagonal structure; S has off-diagonal and
1452        // off-H-pattern entries, forcing exact-inverse column solves.
1453        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        // S includes entry (0,3)/(3,0) which is OUTSIDE H's sparsity pattern,
1460        // so trace_product_sparse must recover it via a cached column solve.
1461        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        // Full symmetric storage for S (both triangles).
1473        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        // ~40 columns so the rayon reduction fans out across threads.
1487        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            // Strongly diagonally dominant => SPD.
1492            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        // Tridiagonal-ish off-diagonals for H (its sparsity pattern).
1496        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        // S off-diagonals deliberately reach OUTSIDE H's tridiagonal pattern
1502        // (stride-3 and a couple of long-range entries) so that many distinct
1503        // columns require concurrent cache-miss exact-inverse solves.
1504        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    // ── solve_sparse_spdmulti / solve_sparse_spdmulti_rows ───────────────────
1530
1531    #[test]
1532    fn solve_sparse_spdmulti_recovers_identity_inverse() {
1533        // A = diag(4,9): A^{-1} * A = I
1534        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        // Solve against the identity matrix
1538        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        // multi-rhs: two distinct RHS vectors as columns
1552        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        // Each column of x should satisfy A*x_j = rhs_j
1555        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        // A = [[4,2],[2,5]]; A^{-1} has known entries.
1567        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        // Solve against a 2x2 RHS, requesting only row 0..1 (first row only).
1571        let rhs = Array2::<f64>::eye(2);
1572        let row0 = solve_sparse_spdmulti_rows(&factor, &rhs, 0, 1).unwrap();
1573        // Should be a 1x2 matrix: first row of A^{-1}.
1574        // A^{-1} = (1/16)*[[5,-2],[-2,4]]
1575        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        // A = [[4,2],[2,5]]; A^{-1} diagonal = [5/16, 4/16] = [0.3125, 0.25].
1583        // diagonal_sum from row_start=0, rhs=I_2 sums diag(A^{-1})[0..2] = trace.
1584        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        // trace(A^{-1}) = 5/16 + 4/16 = 9/16
1590        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    /// Build the dense inverse of an SPD matrix via per-column Cholesky solves.
1629    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    /// Reference tr(Z·S) computed densely from full matrices.
1645    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        // tr(H⁻¹ S) on a small banded SPD H, with S an arbitrary symmetric
1659        // sparse matrix supplied as upper-triangle-only CSC.
1660        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    /// #759 regression: `trace_product_sparse` is a rayon reduction over columns.
1684    /// On a larger system (many columns, off-pattern S entries that force
1685    /// cache-miss exact-column solves under the `Mutex`-guarded cache) the
1686    /// parallel sum must agree with the dense reference to rounding. This pins
1687    /// the parallelization so a future refactor can't silently revert it to a
1688    /// serial scan again (as happened between b7879667b and a80fe6943).
1689    #[test]
1690    fn trace_product_sparse_parallel_matches_dense_reference_large() {
1691        // n=40 tridiagonal SPD H (diagonally dominant -> SPD).
1692        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        // S is symmetric with both near-diagonal and far off-diagonal entries,
1702        // so tr(Z·S) touches inverse entries OUTSIDE the Cholesky pattern,
1703        // exercising the on-demand exact-column cache concurrently.
1704        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        // Determinism / order-invariance: repeated calls (cache now warm) agree
1733        // bit-for-bit with the first parallel reduction.
1734        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        // Bit-identical across thread-pool sizes. The per-column partials are
1741        // summed in fixed column order, so the reduction cannot depend on how
1742        // many rayon workers fold them. A plain parallel `.sum()` would fold in
1743        // a scheduling-dependent tree order and drift in the low bits, perturbing
1744        // the REML gradient / EDF that consumes tr(H⁻¹S) across machines with
1745        // different core counts. Pin 1-worker == 8-worker exactly.
1746        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}