1use super::*;
2
3pub struct SparseCholeskyOperator {
12 pub(crate) factor: std::sync::Arc<gam_linalg::sparse_exact::SparseExactFactor>,
14 pub(crate) takahashi: Option<std::sync::Arc<gam_linalg::sparse_exact::TakahashiInverse>>,
17 pub(crate) cached_logdet: f64,
19 pub(crate) n_dim: usize,
21}
22
23impl SparseCholeskyOperator {
24 pub fn new(
26 factor: std::sync::Arc<gam_linalg::sparse_exact::SparseExactFactor>,
27 logdet_h: f64,
28 dim: usize,
29 ) -> Self {
30 Self {
31 factor,
32 takahashi: None,
33 cached_logdet: logdet_h,
34 n_dim: dim,
35 }
36 }
37
38 pub fn with_takahashi(
39 mut self,
40 taka: std::sync::Arc<gam_linalg::sparse_exact::TakahashiInverse>,
41 ) -> Self {
42 self.takahashi = Some(taka);
43 self
44 }
45
46 pub(crate) const OPERATOR_SOLVE_CHUNK: usize = 64;
47
48 pub(crate) fn takahashi_block_trace(
49 taka: &gam_linalg::sparse_exact::TakahashiInverse,
50 block: &Array2<f64>,
51 start: usize,
52 ) -> f64 {
53 assert_eq!(block.nrows(), block.ncols());
54 let mut trace = 0.0;
55 for i in 0..block.nrows() {
56 let diag = block[[i, i]];
57 if diag.abs() > 1e-30 {
58 trace += taka.get(start + i, start + i) * diag;
59 }
60 for j in (i + 1)..block.ncols() {
61 let pair = block[[i, j]] + block[[j, i]];
62 if pair.abs() > 1e-30 {
63 trace += taka.get(start + i, start + j) * pair;
64 }
65 }
66 }
67 trace
68 }
69
70 pub(crate) fn takahashi_left_multiply_block(
71 taka: &gam_linalg::sparse_exact::TakahashiInverse,
72 block: &Array2<f64>,
73 start: usize,
74 ) -> Array2<f64> {
75 let dim = block.nrows();
76 let mut out = Array2::<f64>::zeros((dim, dim));
77 for i in 0..dim {
78 let z_diag = taka.get(start + i, start + i);
79 if z_diag.abs() > 1e-30 {
80 for k in 0..dim {
81 out[[i, k]] += z_diag * block[[i, k]];
82 }
83 }
84 for j in (i + 1)..dim {
85 let z = taka.get(start + i, start + j);
86 if z.abs() <= 1e-30 {
87 continue;
88 }
89 for k in 0..dim {
90 out[[i, k]] += z * block[[j, k]];
91 out[[j, k]] += z * block[[i, k]];
92 }
93 }
94 }
95 out
96 }
97
98 pub(crate) fn trace_hinv_operator_exact(&self, op: &dyn HyperOperator) -> f64 {
99 let (range_start, range_end) = op
100 .block_local_data()
101 .map(|(_, start, end)| (start, end))
102 .unwrap_or((0, self.n_dim));
103 let chunk = Self::OPERATOR_SOLVE_CHUNK.min(self.n_dim.max(1));
104 let mut trace = 0.0_f64;
105 let mut rhs_block = Array2::<f64>::zeros((self.n_dim, chunk));
106 let mut start = range_start;
107
108 while start < range_end {
109 let end = (start + chunk).min(range_end);
110 let cols = end - start;
111 op.mul_basis_columns_into(start, rhs_block.slice_mut(ndarray::s![.., ..cols]));
112
113 let diagonal_sum = if cols == chunk {
114 gam_linalg::sparse_exact::solve_sparse_spdmulti_diagonal_sum(
115 &self.factor,
116 &rhs_block,
117 start,
118 )
119 } else {
120 let rhs_view = rhs_block.slice(ndarray::s![.., ..cols]);
121 gam_linalg::sparse_exact::solve_sparse_spdmulti_diagonal_sum(
122 &self.factor,
123 &rhs_view,
124 start,
125 )
126 };
127 trace += diagonal_sum.unwrap_or_else(|e| {
128 reml_contract_panic(format!(
137 "SparseCholeskyOperator exact trace_hinv_operator solve failed: {e}"
138 ))
139 });
140 start = end;
141 }
142
143 trace
144 }
145
146 pub(crate) fn solve_operator_column_range_rows_exact(
147 &self,
148 op: &dyn HyperOperator,
149 col_start: usize,
150 col_end: usize,
151 row_start: usize,
152 row_end: usize,
153 ) -> Result<Array2<f64>, String> {
154 let chunk = Self::OPERATOR_SOLVE_CHUNK.min(self.n_dim.max(1));
155 let cols_total = col_end - col_start;
156 let rows_total = row_end - row_start;
157 let mut solved = Array2::<f64>::zeros((rows_total, cols_total));
158 let mut rhs_block = Array2::<f64>::zeros((self.n_dim, chunk));
159 let mut start = col_start;
160
161 while start < col_end {
162 let end = (start + chunk).min(col_end);
163 let cols = end - start;
164 op.mul_basis_columns_into(start, rhs_block.slice_mut(ndarray::s![.., ..cols]));
165
166 let solved_block = if cols == chunk {
167 gam_linalg::sparse_exact::solve_sparse_spdmulti_rows(
168 &self.factor,
169 &rhs_block,
170 row_start,
171 row_end,
172 )
173 } else {
174 let rhs_view = rhs_block.slice(ndarray::s![.., ..cols]);
175 gam_linalg::sparse_exact::solve_sparse_spdmulti_rows(
176 &self.factor,
177 &rhs_view,
178 row_start,
179 row_end,
180 )
181 }
182 .map_err(|e| {
183 format!(
184 "SparseCholeskyOperator::solve_operator_column_range_rows_exact multi-solve failed: {e}"
185 )
186 })?;
187 solved
188 .slice_mut(ndarray::s![.., start - col_start..end - col_start])
189 .assign(&solved_block);
190 start = end;
191 }
192
193 Ok(solved)
194 }
195
196 pub(crate) fn trace_hinv_matrix_operator_cross_exact(
197 &self,
198 matrix: &Array2<f64>,
199 op: &dyn HyperOperator,
200 ) -> f64 {
201 if let Some((_, range_start, range_end)) = op.block_local_data()
202 && range_end - range_start < self.n_dim
203 {
204 return self.trace_hinv_matrix_block_operator_cross_exact(
205 matrix,
206 op,
207 range_start,
208 range_end,
209 );
210 }
211
212 let solved_matrix = self.solve_multi(matrix);
213 let chunk = Self::OPERATOR_SOLVE_CHUNK.min(self.n_dim.max(1));
214 let mut rhs_block = Array2::<f64>::zeros((self.n_dim, chunk));
215 let mut trace = 0.0_f64;
216 let (range_start, range_end) = op
217 .block_local_data()
218 .map(|(_, start, end)| (start, end))
219 .unwrap_or((0, self.n_dim));
220 let mut start = range_start;
221
222 while start < range_end {
223 let end = (start + chunk).min(range_end);
224 let cols = end - start;
225 op.mul_basis_columns_into(start, rhs_block.slice_mut(ndarray::s![.., ..cols]));
226
227 let solved_op = if cols == chunk {
228 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, &rhs_block)
229 } else {
230 let rhs_view = rhs_block.slice(ndarray::s![.., ..cols]);
231 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, &rhs_view)
232 };
233
234 let solved_op = solved_op.unwrap_or_else(|e| {
235 panic!("SparseCholeskyOperator exact matrix/operator cross solve failed: {e}")
242 });
243
244 for local_col in 0..cols {
245 let matrix_row = start + local_col;
246 for row in 0..self.n_dim {
247 trace += solved_matrix[[matrix_row, row]] * solved_op[[row, local_col]];
248 }
249 }
250 start = end;
251 }
252
253 trace
254 }
255
256 pub(crate) fn trace_hinv_matrix_block_operator_cross_exact(
257 &self,
258 matrix: &Array2<f64>,
259 op: &dyn HyperOperator,
260 range_start: usize,
261 range_end: usize,
262 ) -> f64 {
263 let t_start = std::time::Instant::now();
264 let chunk = Self::OPERATOR_SOLVE_CHUNK.min(self.n_dim.max(1));
265 let mut op_rhs_block = Array2::<f64>::zeros((self.n_dim, chunk));
266 let mut eye_rhs_block = Array2::<f64>::zeros((self.n_dim, chunk));
267 let mut trace = 0.0_f64;
268 let mut start = range_start;
269
270 while start < range_end {
271 let end = (start + chunk).min(range_end);
272 let cols = end - start;
273 op.mul_basis_columns_into(start, op_rhs_block.slice_mut(ndarray::s![.., ..cols]));
274
275 eye_rhs_block.fill(0.0);
276 for local_col in 0..cols {
277 eye_rhs_block[[start + local_col, local_col]] = 1.0;
278 }
279
280 let solved_op = if cols == chunk {
281 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, &op_rhs_block)
282 } else {
283 let rhs_view = op_rhs_block.slice(ndarray::s![.., ..cols]);
284 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, &rhs_view)
285 };
286 let solved_op = solved_op.unwrap_or_else(|e| {
287 panic!(
293 "SparseCholeskyOperator exact matrix/block-operator cross operator solve failed: {e}"
294 )
295 });
296
297 let solved_eye = if cols == chunk {
298 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, &eye_rhs_block)
299 } else {
300 let rhs_view = eye_rhs_block.slice(ndarray::s![.., ..cols]);
301 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, &rhs_view)
302 };
303 let solved_eye = solved_eye.unwrap_or_else(|e| {
304 panic!(
310 "SparseCholeskyOperator exact matrix/block-operator cross identity solve failed: {e}"
311 )
312 });
313
314 let selected_rows_t = matrix.t().dot(&solved_eye);
315 for local_col in 0..cols {
316 for row in 0..self.n_dim {
317 trace += selected_rows_t[[row, local_col]] * solved_op[[row, local_col]];
318 }
319 }
320 start = end;
321 }
322
323 let elapsed_ms = t_start.elapsed().as_secs_f64() * 1000.0;
324 if elapsed_ms > REML_TRACE_SLOW_LOG_MS {
325 log::info!(
326 "[REML-trace] matrix_block_op_cross_exact | n_dim={} | block={} | {:.1}ms",
327 self.n_dim,
328 range_end - range_start,
329 elapsed_ms
330 );
331 }
332 trace
333 }
334
335 pub(crate) fn trace_hinv_operator_cross_exact(
336 &self,
337 left: &dyn HyperOperator,
338 right: &dyn HyperOperator,
339 ) -> f64 {
340 let (left_start, left_end) = left
341 .block_local_data()
342 .map(|(_, start, end)| (start, end))
343 .unwrap_or((0, self.n_dim));
344 let (right_start, right_end) = right
345 .block_local_data()
346 .map(|(_, start, end)| (start, end))
347 .unwrap_or((0, self.n_dim));
348
349 let solved_left = self
350 .solve_operator_column_range_rows_exact(
351 left,
352 left_start,
353 left_end,
354 right_start,
355 right_end,
356 )
357 .unwrap_or_else(|e| {
358 panic!("SparseCholeskyOperator exact operator cross left solve failed: {e}")
365 });
366 let same_operator =
367 std::ptr::addr_eq(left, right) && left_start == right_start && left_end == right_end;
368 let solved_right = if same_operator {
369 None
370 } else {
371 Some(
372 self.solve_operator_column_range_rows_exact(
373 right,
374 right_start,
375 right_end,
376 left_start,
377 left_end,
378 )
379 .unwrap_or_else(|e| {
380 panic!("SparseCholeskyOperator exact operator cross right solve failed: {e}")
386 }),
387 )
388 };
389
390 let right_cols = right_end - right_start;
391 let mut trace = 0.0;
392 for left_col in 0..(left_end - left_start) {
393 for right_col in 0..right_cols {
394 let right_value = match solved_right.as_ref() {
395 Some(solved) => solved[[left_col, right_col]],
396 None => solved_left[[left_col, right_col]],
397 };
398 trace += solved_left[[right_col, left_col]] * right_value;
399 }
400 }
401 trace
402 }
403}
404
405impl HessianFactorization for SparseCholeskyOperator {
406 fn logdet(&self) -> f64 {
407 self.cached_logdet
408 }
409
410 fn assemble_h_dense_for_tangent_projection(&self) -> Result<Array2<f64>, String> {
411 let h = gam_linalg::sparse_exact::assemble_sparse_factor_h_dense(&self.factor)
412 .map_err(|e| e.to_string())?;
413 if h.nrows() != self.n_dim || h.ncols() != self.n_dim {
414 return Err(format!(
415 "sparse Cholesky tangent projection dense H has shape {}x{}, expected {}x{}",
416 h.nrows(),
417 h.ncols(),
418 self.n_dim,
419 self.n_dim
420 ));
421 }
422 Ok(h)
423 }
424
425 fn trace_hinv_product(&self, a: &Array2<f64>) -> f64 {
426 if let Some(ref taka) = self.takahashi {
429 let mut trace = 0.0;
430 for i in 0..a.nrows() {
431 let a_ii = a[[i, i]];
432 if a_ii.abs() > 1e-30 {
433 trace += taka.get(i, i) * a_ii;
434 }
435 for j in (i + 1)..a.ncols() {
436 let pair = a[[i, j]] + a[[j, i]];
437 if pair.abs() > 1e-30 {
438 trace += taka.get(i, j) * pair;
439 }
440 }
441 }
442 return trace;
443 }
444 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, a)
445 .unwrap_or_else(|e| {
446 panic!("SparseCholeskyOperator exact trace_hinv_product solve failed: {e}")
453 })
454 .diag()
455 .sum()
456 }
457
458 fn trace_hinv_operator(&self, op: &dyn HyperOperator) -> f64 {
459 if let Some(ref taka) = self.takahashi {
460 if let Some((local, start, end)) = op.block_local_data() {
461 assert_eq!(local.nrows(), end - start);
462 return Self::takahashi_block_trace(taka, local, start);
463 }
464 if !op.is_implicit() {
466 let dense = op.to_dense();
467 return self.trace_hinv_product(&dense);
468 }
469 }
470 self.trace_hinv_operator_exact(op)
471 }
472
473 fn trace_logdet_operator(&self, op: &dyn HyperOperator) -> f64 {
474 self.trace_hinv_operator(op)
475 }
476
477 fn solve(&self, rhs: &Array1<f64>) -> Array1<f64> {
478 gam_linalg::sparse_exact::solve_sparse_spd(&self.factor, rhs)
483 .unwrap_or_else(|e| panic!("SparseCholeskyOperator exact solve failed: {e}"))
485 }
486
487 fn solve_multi(&self, rhs: &Array2<f64>) -> Array2<f64> {
488 gam_linalg::sparse_exact::solve_sparse_spdmulti(&self.factor, rhs)
492 .unwrap_or_else(|e| panic!("SparseCholeskyOperator exact multi-solve failed: {e}"))
494 }
495
496 fn trace_hinv_product_cross(&self, a: &Array2<f64>, b: &Array2<f64>) -> f64 {
497 let solved_a = self.solve_multi(a);
501 if std::ptr::eq(a, b) {
502 return dense::trace_product(&solved_a, &solved_a);
503 }
504 let solved_b = self.solve_multi(b);
505 dense::trace_product(&solved_a, &solved_b)
506 }
507
508 fn trace_hinv_matrix_operator_cross(
509 &self,
510 matrix: &Array2<f64>,
511 op: &dyn HyperOperator,
512 ) -> f64 {
513 self.trace_hinv_matrix_operator_cross_exact(matrix, op)
517 }
518
519 fn trace_hinv_operator_cross(
520 &self,
521 left: &dyn HyperOperator,
522 right: &dyn HyperOperator,
523 ) -> f64 {
524 if let Some(ref taka) = self.takahashi
527 && let (Some((a_local, a_start, a_end)), Some((b_local, b_start, b_end))) =
528 (left.block_local_data(), right.block_local_data())
529 && a_start == b_start
530 && a_end == b_end
531 {
532 let za = Self::takahashi_left_multiply_block(taka, a_local, a_start);
534 if std::ptr::addr_eq(left, right) {
535 return dense::trace_product(&za, &za);
536 }
537 let zb = Self::takahashi_left_multiply_block(taka, b_local, b_start);
538 return (&za * &zb.t()).sum();
540 }
541 self.trace_hinv_operator_cross_exact(left, right)
544 }
545
546 fn trace_logdet_hessian_cross_matrix_operator(
547 &self,
548 h_i: &Array2<f64>,
549 h_j: &dyn HyperOperator,
550 ) -> f64 {
551 -self.trace_hinv_matrix_operator_cross(h_i, h_j)
552 }
553
554 fn trace_logdet_hessian_cross_operator(
555 &self,
556 h_i: &dyn HyperOperator,
557 h_j: &dyn HyperOperator,
558 ) -> f64 {
559 -self.trace_hinv_operator_cross(h_i, h_j)
560 }
561
562 fn active_rank(&self) -> usize {
563 self.n_dim
564 }
565
566 fn dim(&self) -> usize {
567 self.n_dim
568 }
569}
570
571pub struct DenseCholeskyValueOnlyOperator {
600 pub(crate) chol: gam_linalg::faer_ndarray::FaerCholeskyFactor,
602 pub(crate) cached_logdet: f64,
604 pub(crate) n_dim: usize,
606}
607
608impl DenseCholeskyValueOnlyOperator {
609 pub fn from_spd(h: &Array2<f64>) -> Result<Self, String> {
617 use faer::Side;
618 use gam_linalg::faer_ndarray::FaerCholesky;
619
620 let n = h.nrows();
621 if n != h.ncols() {
622 return Err(format!(
623 "DenseCholeskyValueOnlyOperator: expected square matrix, got {}×{}",
624 n,
625 h.ncols()
626 ));
627 }
628 let chol = h
629 .cholesky(Side::Lower)
630 .map_err(|e| format!("DenseCholeskyValueOnlyOperator LLT failed: {e}"))?;
631 let diag = chol.diag();
632 let cached_logdet = 2.0 * diag.iter().map(|&d| d.ln()).sum::<f64>();
633
634 let epsilon = spectral_epsilon_for_dim(n);
669 let h_inverse = chol.solve_mat(&Array2::<f64>::eye(n));
670 let floor_gap_bound =
671 epsilon * epsilon * h_inverse.iter().map(|entry| entry * entry).sum::<f64>();
672 let agreement_envelope =
673 crate::rho_optimizer::outer_value_agreement_bound(cached_logdet, cached_logdet);
674 if !(floor_gap_bound <= agreement_envelope) {
675 return Err(format!(
676 "DenseCholeskyValueOnlyOperator declines a {n}-dimensional Hessian: its exact \
677 log-determinant can differ from the smooth-floored log|H| the derivative lanes \
678 price by up to {floor_gap_bound:.3e}, above the {agreement_envelope:.3e} \
679 value-agreement envelope (spectral floor eps={epsilon:.3e})"
680 ));
681 }
682
683 Ok(Self {
684 chol,
685 cached_logdet,
686 n_dim: n,
687 })
688 }
689}
690
691impl HessianFactorization for DenseCholeskyValueOnlyOperator {
692 fn logdet(&self) -> f64 {
693 self.cached_logdet
694 }
695
696 fn trace_hinv_product(&self, a: &Array2<f64>) -> f64 {
697 let hinv_a = self.chol.solve_mat(a);
700 hinv_a.diag().iter().sum()
701 }
702
703 fn solve(&self, rhs: &Array1<f64>) -> Array1<f64> {
704 self.chol.solvevec(rhs)
705 }
706
707 fn solve_multi(&self, rhs: &Array2<f64>) -> Array2<f64> {
708 self.chol.solve_mat(rhs)
709 }
710
711 fn active_rank(&self) -> usize {
712 self.n_dim
714 }
715
716 fn dim(&self) -> usize {
717 self.n_dim
718 }
719}
720
721pub struct BlockCoupledOperator {
742 pub(crate) inner: DenseSpectralOperator,
744}
745
746impl BlockCoupledOperator {
747 pub fn from_joint_hessian_with_mode(
751 joint_hessian: &Array2<f64>,
752 mode: PseudoLogdetMode,
753 ) -> Result<Self, String> {
754 let inner = DenseSpectralOperator::from_symmetric_with_mode(joint_hessian, mode)
755 .map_err(|e| format!("BlockCoupledOperator eigendecomposition: {e}"))?;
756
757 Ok(Self { inner })
758 }
759}
760
761impl HessianFactorization for BlockCoupledOperator {
762 fn logdet(&self) -> f64 {
763 self.inner.logdet()
764 }
765
766 fn as_exact_dense_spectral(&self) -> Option<&DenseSpectralOperator> {
767 self.inner.as_exact_dense_spectral()
768 }
769
770 fn assemble_h_dense_for_tangent_projection(&self) -> Result<Array2<f64>, String> {
771 self.inner.assemble_h_dense_for_tangent_projection()
772 }
773
774 fn trace_hinv_product(&self, a: &Array2<f64>) -> f64 {
775 self.inner.trace_hinv_product(a)
776 }
777
778 fn trace_logdet_gradient(&self, a: &Array2<f64>) -> f64 {
779 self.inner.trace_logdet_gradient(a)
780 }
781
782 fn xt_logdet_kernel_x_diagonal(&self, x: &DesignMatrix) -> Array1<f64> {
783 self.inner.xt_logdet_kernel_x_diagonal(x)
784 }
785
786 fn trace_logdet_h_k(
787 &self,
788 a_k: &Array2<f64>,
789 third_deriv_correction: Option<&Array2<f64>>,
790 ) -> f64 {
791 self.inner.trace_logdet_h_k(a_k, third_deriv_correction)
792 }
793
794 fn trace_logdet_operator(&self, op: &dyn HyperOperator) -> f64 {
795 self.inner.trace_logdet_operator(op)
796 }
797
798 fn trace_logdet_hessian_cross(&self, h_i: &Array2<f64>, h_j: &Array2<f64>) -> f64 {
799 self.inner.trace_logdet_hessian_cross(h_i, h_j)
800 }
801
802 fn solve(&self, rhs: &Array1<f64>) -> Array1<f64> {
803 self.inner.solve(rhs)
804 }
805
806 fn solve_multi(&self, rhs: &Array2<f64>) -> Array2<f64> {
807 self.inner.solve_multi(rhs)
808 }
809
810 fn trace_hinv_product_cross(&self, a: &Array2<f64>, b: &Array2<f64>) -> f64 {
811 self.inner.trace_hinv_product_cross(a, b)
812 }
813
814 fn trace_hinv_matrix_operator_cross(
815 &self,
816 matrix: &Array2<f64>,
817 op: &dyn HyperOperator,
818 ) -> f64 {
819 self.inner.trace_hinv_matrix_operator_cross(matrix, op)
820 }
821
822 fn trace_hinv_operator_cross(
823 &self,
824 left: &dyn HyperOperator,
825 right: &dyn HyperOperator,
826 ) -> f64 {
827 self.inner.trace_hinv_operator_cross(left, right)
828 }
829
830 fn active_rank(&self) -> usize {
831 self.inner.active_rank()
832 }
833
834 fn dim(&self) -> usize {
835 self.inner.dim()
836 }
837
838 fn is_dense(&self) -> bool {
839 true
840 }
841
842 fn prefers_stochastic_trace_estimation(&self) -> bool {
843 false
844 }
845
846 fn logdet_traces_match_hinv_kernel(&self) -> bool {
847 false
848 }
849
850 fn as_dense_spectral(&self) -> Option<&DenseSpectralOperator> {
851 Some(&self.inner)
852 }
853}
854
855pub struct MatrixFreeSpdOperator {
868 pub(crate) apply: Arc<dyn Fn(&Array1<f64>) -> Array1<f64> + Send + Sync>,
869 pub(crate) dense_assemble: Option<Arc<dyn Fn() -> Option<Array2<f64>> + Send + Sync>>,
881 pub(crate) cached_logdet: gam_runtime::resource::RayonSafeOnce<f64>,
882 pub(crate) n_dim: usize,
883 pub(crate) dense_spectral: gam_runtime::resource::RayonSafeOnce<Option<DenseSpectralOperator>>,
893 pub(crate) mode: PseudoLogdetMode,
902}
903
904impl MatrixFreeSpdOperator {
905 pub(crate) const EXACT_DENSE_SPECTRAL_MAX_BYTES: usize = 512 * 1024 * 1024;
906 pub(crate) const EXACT_DENSE_SPECTRAL_ARRAYS: usize = 6;
907
908 pub fn new_with_mode<F>(dim: usize, apply: F, mode: PseudoLogdetMode) -> Self
909 where
910 F: Fn(&Array1<f64>) -> Array1<f64> + Send + Sync + 'static,
911 {
912 Self::new_with_mode_and_dense_assemble(dim, apply, mode, None)
913 }
914
915 pub fn new_with_mode_and_dense_assemble<F>(
920 dim: usize,
921 apply: F,
922 mode: PseudoLogdetMode,
923 dense_assemble: Option<Arc<dyn Fn() -> Option<Array2<f64>> + Send + Sync>>,
924 ) -> Self
925 where
926 F: Fn(&Array1<f64>) -> Array1<f64> + Send + Sync + 'static,
927 {
928 let apply = Arc::new(apply);
929
930 Self {
931 apply,
932 dense_assemble,
933 cached_logdet: gam_runtime::resource::RayonSafeOnce::new(),
934 n_dim: dim,
935 dense_spectral: gam_runtime::resource::RayonSafeOnce::new(),
936 mode,
937 }
938 }
939
940 pub(crate) fn exact_dense_spectral_bytes(&self) -> Option<usize> {
941 self.n_dim
942 .checked_mul(self.n_dim)?
943 .checked_mul(std::mem::size_of::<f64>())?
944 .checked_mul(Self::EXACT_DENSE_SPECTRAL_ARRAYS)
945 }
946
947 pub(crate) fn exact_dense_spectral_budget_ok(&self) -> bool {
948 match self.exact_dense_spectral_bytes() {
949 Some(bytes) if bytes <= Self::EXACT_DENSE_SPECTRAL_MAX_BYTES => true,
950 Some(bytes) => {
951 log::error!(
952 "MatrixFreeSpdOperator exact dense spectral materialization requires {:.2} GiB \
953 for dim={}, exceeding the {:.2} GiB cap",
954 bytes as f64 / (1024.0 * 1024.0 * 1024.0),
955 self.n_dim,
956 Self::EXACT_DENSE_SPECTRAL_MAX_BYTES as f64 / (1024.0 * 1024.0 * 1024.0),
957 );
958 false
959 }
960 None => {
961 log::error!(
962 "MatrixFreeSpdOperator exact dense spectral byte count overflow for dim={}",
963 self.n_dim
964 );
965 false
966 }
967 }
968 }
969
970 pub(crate) fn materialize_dense_operator(&self) -> Option<DenseSpectralOperator> {
971 if !self.exact_dense_spectral_budget_ok() {
972 return None;
973 }
974 let materialize_start = std::time::Instant::now();
975 let (matrix, matvec_count) =
981 match self.dense_assemble.as_ref().and_then(|assemble| assemble()) {
982 Some(mut direct)
983 if direct.nrows() == self.n_dim
984 && direct.ncols() == self.n_dim
985 && direct.iter().all(|v| v.is_finite()) =>
986 {
987 for i in 0..self.n_dim {
991 for j in (i + 1)..self.n_dim {
992 let avg = 0.5 * (direct[[i, j]] + direct[[j, i]]);
993 direct[[i, j]] = avg;
994 direct[[j, i]] = avg;
995 }
996 }
997 (direct, 0usize)
998 }
999 _ => {
1000 let mut matrix = Array2::<f64>::zeros((self.n_dim, self.n_dim));
1001 let mut basis = Array1::<f64>::zeros(self.n_dim);
1002 for j in 0..self.n_dim {
1003 basis[j] = 1.0;
1004 let col = (self.apply)(&basis);
1005 basis[j] = 0.0;
1006 if col.len() != self.n_dim || !col.iter().all(|v| v.is_finite()) {
1007 return None;
1008 }
1009 matrix.column_mut(j).assign(&col);
1010 }
1011 for i in 0..self.n_dim {
1012 for j in (i + 1)..self.n_dim {
1013 let avg = 0.5 * (matrix[[i, j]] + matrix[[j, i]]);
1014 matrix[[i, j]] = avg;
1015 matrix[[j, i]] = avg;
1016 }
1017 }
1018 (matrix, self.n_dim)
1019 }
1020 };
1021 let result = DenseSpectralOperator::from_symmetric_with_mode(&matrix, self.mode).ok();
1022 log::info!(
1023 "[STAGE] matrix_free_spd materialize n_dim={} matvec_count={} elapsed={:.3}s",
1024 self.n_dim,
1025 matvec_count,
1026 materialize_start.elapsed().as_secs_f64(),
1027 );
1028 result
1029 }
1030
1031 pub(crate) fn dense_spectral(&self) -> Option<&DenseSpectralOperator> {
1032 self.dense_spectral
1033 .get_or_compute(|| self.materialize_dense_operator())
1034 .as_ref()
1035 }
1036
1037 pub(crate) fn exact_dense_spectral(&self) -> &DenseSpectralOperator {
1038 self.dense_spectral().expect(
1039 "MatrixFreeSpdOperator exact REML algebra requires dense spectral materialization within the configured budget",
1040 )
1041 }
1042
1043 pub(crate) fn use_trace_cg(&self, rel_tol: f64) -> bool {
1044 rel_tol.is_finite()
1045 && rel_tol > 0.0
1046 && self.prefers_stochastic_trace_estimation()
1047 && self.has_matrix_free_trace_cg_operator()
1048 }
1049
1050 pub(crate) fn cg_trace_solve(
1051 &self,
1052 rhs: &Array1<f64>,
1053 rel_tol: f64,
1054 probe_id: Option<u64>,
1055 trace_state: Option<&Arc<Mutex<StochasticTraceState>>>,
1056 ) -> Array1<f64> {
1057 let dim = rhs.len();
1058 if dim != self.n_dim {
1059 return self.solve(rhs);
1060 }
1061
1062 let (initial, warm_start_used) = match (probe_id, trace_state) {
1063 (Some(id), Some(state)) => {
1064 let cached = match state.lock() {
1065 Ok(guard) => guard.cg_warm_starts.get(&id).cloned(),
1066 Err(poisoned) => poisoned.into_inner().cg_warm_starts.get(&id).cloned(),
1067 };
1068 match cached {
1069 Some(x) if x.len() == dim => (x, true),
1070 _ => (Array1::<f64>::zeros(dim), false),
1071 }
1072 }
1073 _ => (Array1::<f64>::zeros(dim), false),
1074 };
1075
1076 let Some((solution, iters, residual_norm)) =
1077 conjugate_gradient_trace_solve(rhs, rel_tol, initial, |v| (self.apply)(v))
1078 else {
1079 return self.solve(rhs);
1080 };
1081
1082 if let Some(state) = trace_state {
1083 let mut guard = match state.lock() {
1084 Ok(guard) => guard,
1085 Err(poisoned) => poisoned.into_inner(),
1086 };
1087 guard.last_linear_residual_norm = Some(
1088 guard
1089 .last_linear_residual_norm
1090 .unwrap_or(0.0)
1091 .max(residual_norm),
1092 );
1093 if let Some(id) = probe_id {
1094 guard.cg_warm_starts.insert(id, solution.clone());
1095 }
1096 }
1097
1098 let probe_label = probe_id
1099 .map(|id| id.to_string())
1100 .unwrap_or_else(|| "untracked".to_string());
1101 log::info!(
1102 "[CG-TRACE] probe_id={} iters={} rel_tol={} warm_start_used={}",
1103 probe_label,
1104 iters,
1105 rel_tol,
1106 warm_start_used
1107 );
1108
1109 solution
1110 }
1111}
1112
1113pub(crate) fn conjugate_gradient_trace_solve<F>(
1114 rhs: &Array1<f64>,
1115 rel_tol: f64,
1116 mut x: Array1<f64>,
1117 apply: F,
1118) -> Option<(Array1<f64>, usize, f64)>
1119where
1120 F: Fn(&Array1<f64>) -> Array1<f64>,
1121{
1122 let dim = rhs.len();
1123 if x.len() != dim {
1124 return None;
1125 }
1126
1127 let rhs_norm_sq = rhs.dot(rhs);
1128 if !rhs_norm_sq.is_finite() {
1129 return None;
1130 }
1131 if rhs_norm_sq <= f64::MIN_POSITIVE {
1132 return Some((Array1::<f64>::zeros(dim), 0, 0.0));
1133 }
1134
1135 let target_sq = (rel_tol * rel_tol * rhs_norm_sq).max(f64::MIN_POSITIVE);
1136 let mut r = rhs.clone();
1137 if x.iter().any(|value| *value != 0.0) {
1138 let ax = apply(&x);
1139 if ax.len() != dim || !ax.iter().all(|value| value.is_finite()) {
1140 return None;
1141 }
1142 r.scaled_add(-1.0, &ax);
1143 }
1144
1145 let mut rs_old = r.dot(&r);
1146 if !rs_old.is_finite() {
1147 return None;
1148 }
1149 if rs_old <= target_sq {
1150 return Some((x, 0, rs_old.max(0.0).sqrt()));
1151 }
1152
1153 let mut p = r.clone();
1154 let mut iters = 0usize;
1155 let mut residual_norm = rs_old.max(0.0).sqrt();
1156 for k in 0..dim.max(1) {
1157 let ap = apply(&p);
1158 if ap.len() != dim || !ap.iter().all(|value| value.is_finite()) {
1159 return None;
1160 }
1161 let denom = p.dot(&ap);
1162 if !denom.is_finite() || denom <= 0.0 {
1163 log::warn!(
1164 "[CG-TRACE] non-positive curvature in trace CG at iter={} denom={}",
1165 k + 1,
1166 denom
1167 );
1168 break;
1169 }
1170 let alpha = rs_old / denom;
1171 if !alpha.is_finite() {
1172 return None;
1173 }
1174 x.scaled_add(alpha, &p);
1175 r.scaled_add(-alpha, &ap);
1176 let rs_new = r.dot(&r);
1177 if !rs_new.is_finite() {
1178 return None;
1179 }
1180 iters = k + 1;
1181 residual_norm = rs_new.max(0.0).sqrt();
1182 if rs_new <= target_sq {
1183 break;
1184 }
1185 let beta = rs_new / rs_old;
1186 if !beta.is_finite() {
1187 return None;
1188 }
1189 p.mapv_inplace(|value| beta * value);
1190 p += &r;
1191 rs_old = rs_new;
1192 }
1193
1194 Some((x, iters, residual_norm))
1195}
1196
1197impl HessianFactorization for MatrixFreeSpdOperator {
1198 fn logdet(&self) -> f64 {
1199 *self
1200 .cached_logdet
1201 .get_or_compute(|| self.exact_dense_spectral().logdet())
1202 }
1203
1204 fn as_exact_dense_spectral(&self) -> Option<&DenseSpectralOperator> {
1205 Some(self.exact_dense_spectral())
1206 }
1207
1208 fn trace_hinv_product(&self, a: &Array2<f64>) -> f64 {
1209 self.exact_dense_spectral().trace_hinv_product(a)
1210 }
1211
1212 fn trace_hinv_operator(&self, op: &dyn HyperOperator) -> f64 {
1213 self.exact_dense_spectral().trace_hinv_operator(op)
1214 }
1215
1216 fn trace_hinv_product_cross(&self, a: &Array2<f64>, b: &Array2<f64>) -> f64 {
1217 self.exact_dense_spectral().trace_hinv_product_cross(a, b)
1218 }
1219
1220 fn trace_hinv_matrix_operator_cross(
1221 &self,
1222 matrix: &Array2<f64>,
1223 op: &dyn HyperOperator,
1224 ) -> f64 {
1225 self.exact_dense_spectral()
1226 .trace_hinv_matrix_operator_cross(matrix, op)
1227 }
1228
1229 fn trace_hinv_operator_cross(
1230 &self,
1231 left: &dyn HyperOperator,
1232 right: &dyn HyperOperator,
1233 ) -> f64 {
1234 self.exact_dense_spectral()
1235 .trace_hinv_operator_cross(left, right)
1236 }
1237
1238 fn trace_logdet_operator(&self, op: &dyn HyperOperator) -> f64 {
1239 let trace_start = std::time::Instant::now();
1240 let result = self.exact_dense_spectral().trace_logdet_operator(op);
1241 log::info!(
1242 "[STAGE] matrix_free_spd trace_logdet_operator implicit={} dim={} elapsed={:.3}s",
1243 op.is_implicit(),
1244 op.dim(),
1245 trace_start.elapsed().as_secs_f64(),
1246 );
1247 result
1248 }
1249
1250 fn solve(&self, rhs: &Array1<f64>) -> Array1<f64> {
1251 self.exact_dense_spectral().solve(rhs)
1252 }
1253
1254 fn solve_multi(&self, rhs: &Array2<f64>) -> Array2<f64> {
1255 self.exact_dense_spectral().solve_multi(rhs)
1256 }
1257
1258 fn stochastic_trace_solve(&self, rhs: &Array1<f64>, rel_tol: f64) -> Array1<f64> {
1259 if self.use_trace_cg(rel_tol) {
1260 return self.cg_trace_solve(rhs, rel_tol, None, None);
1261 }
1262 self.solve(rhs)
1263 }
1264
1265 fn stochastic_trace_solve_for_probe(
1266 &self,
1267 rhs: &Array1<f64>,
1268 rel_tol: f64,
1269 probe_id: u64,
1270 trace_state: Option<&Arc<Mutex<StochasticTraceState>>>,
1271 ) -> Array1<f64> {
1272 if self.use_trace_cg(rel_tol) {
1273 return self.cg_trace_solve(rhs, rel_tol, Some(probe_id), trace_state);
1274 }
1275 self.solve(rhs)
1276 }
1277
1278 fn stochastic_trace_solve_multi(&self, rhs: &Array2<f64>, rel_tol: f64) -> Array2<f64> {
1279 if self.use_trace_cg(rel_tol) {
1280 let mut out = Array2::<f64>::zeros(rhs.raw_dim());
1281 for j in 0..rhs.ncols() {
1282 let solved = self.cg_trace_solve(&rhs.column(j).to_owned(), rel_tol, None, None);
1283 out.column_mut(j).assign(&solved);
1284 }
1285 return out;
1286 }
1287 self.solve_multi(rhs)
1288 }
1289
1290 fn trace_logdet_hessian_cross(&self, h_i: &Array2<f64>, h_j: &Array2<f64>) -> f64 {
1291 self.exact_dense_spectral()
1292 .trace_logdet_hessian_cross(h_i, h_j)
1293 }
1294
1295 fn trace_logdet_hessian_cross_matrix_operator(
1296 &self,
1297 h_i: &Array2<f64>,
1298 h_j: &dyn HyperOperator,
1299 ) -> f64 {
1300 self.exact_dense_spectral()
1301 .trace_logdet_hessian_cross_matrix_operator(h_i, h_j)
1302 }
1303
1304 fn trace_logdet_hessian_cross_operator(
1305 &self,
1306 h_i: &dyn HyperOperator,
1307 h_j: &dyn HyperOperator,
1308 ) -> f64 {
1309 self.exact_dense_spectral()
1310 .trace_logdet_hessian_cross_operator(h_i, h_j)
1311 }
1312
1313 fn active_rank(&self) -> usize {
1314 self.n_dim
1315 }
1316
1317 fn dim(&self) -> usize {
1318 self.n_dim
1319 }
1320
1321 fn is_dense(&self) -> bool {
1322 true
1323 }
1324
1325 fn prefers_stochastic_trace_estimation(&self) -> bool {
1338 !self.exact_dense_spectral_budget_ok()
1339 }
1340
1341 fn logdet_traces_match_hinv_kernel(&self) -> bool {
1352 !self.exact_dense_spectral_budget_ok()
1353 }
1354
1355 fn as_dense_spectral(&self) -> Option<&DenseSpectralOperator> {
1356 self.dense_spectral()
1357 }
1358
1359 fn has_matrix_free_trace_cg_operator(&self) -> bool {
1360 true
1361 }
1362}
1363
1364pub fn penalty_matrix_root(s: &Array2<f64>) -> Result<Array2<f64>, String> {
1373 use faer::Side;
1374 let n = s.nrows();
1375 if n != s.ncols() {
1376 return Err(RemlError::DimensionMismatch {
1377 reason: format!(
1378 "penalty_matrix_root: expected square matrix, got {}×{}",
1379 n,
1380 s.ncols()
1381 ),
1382 }
1383 .into());
1384 }
1385 if n == 0 {
1386 return Ok(Array2::zeros((0, 0)));
1387 }
1388
1389 let (eigenvalues, eigenvectors) = s
1390 .eigh(Side::Lower)
1391 .map_err(|e| format!("penalty_matrix_root eigendecomposition failed: {e}"))?;
1392
1393 let max_ev = eigenvalues.iter().copied().fold(0.0_f64, f64::max);
1394 let tol = (n.max(1) as f64) * f64::EPSILON * max_ev.max(1e-12);
1395
1396 let active: Vec<usize> = eigenvalues
1397 .iter()
1398 .enumerate()
1399 .filter(|(_, v)| **v > tol)
1400 .map(|(i, _)| i)
1401 .collect();
1402 let rank = active.len();
1403
1404 let mut r = Array2::zeros((rank, n));
1405 for (out_row, &idx) in active.iter().enumerate() {
1406 let scale = eigenvalues[idx].sqrt();
1407 for col in 0..n {
1408 r[[out_row, col]] = scale * eigenvectors[[col, idx]];
1409 }
1410 }
1411 Ok(r)
1412}
1413
1414pub fn compute_block_penalty_logdet_derivs(
1433 per_block_rho: &[Array1<f64>],
1434 per_block_penalties: &[&[Array2<f64>]],
1435 ridge: f64,
1436) -> Result<PenaltyLogdetDerivs, String> {
1437 compute_block_penalty_logdet_derivs_with_prior_factors(
1438 per_block_rho,
1439 per_block_penalties,
1440 None,
1441 ridge,
1442 )
1443}
1444
1445pub fn compute_block_penalty_logdet_derivs_with_prior_factors(
1468 per_block_rho: &[Array1<f64>],
1469 per_block_penalties: &[&[Array2<f64>]],
1470 prior_factor_mask: Option<&[Vec<bool>]>,
1471 ridge: f64,
1472) -> Result<PenaltyLogdetDerivs, String> {
1473 use super::super::penalty_logdet::PenaltyPseudologdet;
1474
1475 let total_k: usize = per_block_rho.iter().map(|r| r.len()).sum();
1476 let block_offsets: Vec<usize> = per_block_rho
1477 .iter()
1478 .scan(0usize, |at, rho| {
1479 let current = *at;
1480 *at += rho.len();
1481 Some(current)
1482 })
1483 .collect();
1484
1485 struct BlockPenaltyLogdetResult {
1486 pub(crate) offset: usize,
1487 pub(crate) value: f64,
1488 pub(crate) first: Array1<f64>,
1489 pub(crate) second: Array2<f64>,
1490 }
1491
1492 let compute_block = |(b, block_rho): (usize, &Array1<f64>)| {
1493 let penalties = per_block_penalties[b];
1494 let kb = block_rho.len();
1495 if penalties.is_empty() || kb == 0 {
1496 return Ok(BlockPenaltyLogdetResult {
1497 offset: block_offsets[b],
1498 value: 0.0,
1499 first: Array1::zeros(kb),
1500 second: Array2::zeros((kb, kb)),
1501 });
1502 }
1503 let lambdas = gam_problem::checked_exp_log_strengths(block_rho.iter().copied())
1504 .map_err(|error| format!("penalty-logdet block {b}: {error}"))?;
1505 let mask = prior_factor_mask.map(|m| m[b].as_slice());
1506 let factor_indices: Vec<usize> = (0..kb)
1507 .filter(|&k| mask.is_some_and(|m| m.get(k).copied().unwrap_or(false)))
1508 .collect();
1509
1510 if factor_indices.is_empty() {
1511 let pld = PenaltyPseudologdet::from_components(penalties, &lambdas, ridge)
1517 .map_err(|e| format!("penalty logdet failed for block {b}: {e}"))?;
1518
1519 let value = pld.value();
1520 let (first, second) = pld.rho_derivatives(penalties, &lambdas);
1521 return Ok(BlockPenaltyLogdetResult {
1522 offset: block_offsets[b],
1523 value,
1524 first,
1525 second,
1526 });
1527 }
1528
1529 let mut value = 0.0;
1535 let mut first = Array1::<f64>::zeros(kb);
1536 let mut second = Array2::<f64>::zeros((kb, kb));
1537
1538 let coalesced_indices: Vec<usize> =
1539 (0..kb).filter(|k| !factor_indices.contains(k)).collect();
1540 if !coalesced_indices.is_empty() {
1541 let sub_pens: Vec<Array2<f64>> = coalesced_indices
1542 .iter()
1543 .map(|&k| penalties[k].clone())
1544 .collect();
1545 let sub_lambdas: Vec<f64> = coalesced_indices.iter().map(|&k| lambdas[k]).collect();
1546 let pld = PenaltyPseudologdet::from_components(&sub_pens, &sub_lambdas, ridge)
1547 .map_err(|e| format!("penalty logdet failed for block {b}: {e}"))?;
1548 value += pld.value();
1549 let (sub_first, sub_second) = pld.rho_derivatives(&sub_pens, &sub_lambdas);
1550 for (i, &k) in coalesced_indices.iter().enumerate() {
1551 first[k] = sub_first[i];
1552 for (j, &l) in coalesced_indices.iter().enumerate() {
1553 second[[k, l]] = sub_second[[i, j]];
1554 }
1555 }
1556 }
1557 for &k in &factor_indices {
1558 let factor_pen = std::slice::from_ref(&penalties[k]);
1559 let factor_lambda = [lambdas[k]];
1560 let pld = PenaltyPseudologdet::from_components(factor_pen, &factor_lambda, ridge)
1561 .map_err(|e| {
1562 format!("penalty logdet failed for block {b} prior factor {k}: {e}")
1563 })?;
1564 value += pld.value();
1568 let (sub_first, sub_second) = pld.rho_derivatives(factor_pen, &factor_lambda);
1569 first[k] = sub_first[0];
1570 second[[k, k]] = sub_second[[0, 0]];
1571 }
1572 Ok(BlockPenaltyLogdetResult {
1573 offset: block_offsets[b],
1574 value,
1575 first,
1576 second,
1577 })
1578 };
1579
1580 let block_results: Vec<BlockPenaltyLogdetResult> = if rayon::current_thread_index().is_some() {
1581 per_block_rho
1582 .iter()
1583 .enumerate()
1584 .map(compute_block)
1585 .collect::<Result<Vec<_>, String>>()?
1586 } else {
1587 per_block_rho
1588 .par_iter()
1589 .enumerate()
1590 .map(compute_block)
1591 .collect::<Result<Vec<_>, String>>()?
1592 };
1593
1594 let mut log_det_total = 0.0;
1595 let mut first = Array1::zeros(total_k);
1596 let mut second = Array2::zeros((total_k, total_k));
1597 for block in block_results {
1598 log_det_total += block.value;
1599 let kb = block.first.len();
1600 for k in 0..kb {
1601 first[block.offset + k] = block.first[k];
1602 }
1603 for k in 0..kb {
1604 for l in 0..kb {
1605 second[[block.offset + k, block.offset + l]] = block.second[[k, l]];
1606 }
1607 }
1608 }
1609
1610 Ok(PenaltyLogdetDerivs {
1611 value: log_det_total,
1612 first,
1613 second: Some(second),
1614 })
1615}
1616
1617