1use faer::Side;
43use ndarray::{Array1, Array2, ArrayView1};
44
45use crate::analytic_penalties::{AnalyticPenalty, PenaltyTier};
46use gam_linalg::faer_ndarray::FaerEigh;
47use gam_linalg::lanczos::{SymmetricLanczosOptions, symmetric_lanczos_eigenpairs};
48
49const DENSE_EIGH_DIM_THRESHOLD: usize = 4096;
54
55#[derive(Debug, Clone)]
61pub struct EdgeRestriction {
62 pub r_uv: Array2<f64>,
63 pub r_vu: Option<Array2<f64>>,
64}
65
66impl EdgeRestriction {
67 #[must_use]
69 pub fn paired(r_uv: Array2<f64>, r_vu: Array2<f64>) -> Self {
70 Self {
71 r_uv,
72 r_vu: Some(r_vu),
73 }
74 }
75
76 #[must_use]
78 pub fn single(r_uv: Array2<f64>) -> Self {
79 Self { r_uv, r_vu: None }
80 }
81
82 pub fn edge_dim(&self) -> usize {
84 self.r_uv.nrows()
85 }
86}
87
88#[derive(Debug, Clone)]
92pub struct SheafConsistencyPenalty {
93 edges: Vec<(usize, usize)>,
94 restrictions: Vec<EdgeRestriction>,
95 weight: f64,
96 stalk_offsets: Vec<usize>,
97 stalk_dims: Vec<usize>,
98}
99
100impl SheafConsistencyPenalty {
101 #[must_use = "build error must be handled"]
111 pub fn new(
112 edges: Vec<(usize, usize)>,
113 restrictions: Vec<EdgeRestriction>,
114 weight: f64,
115 stalk_dims: Vec<usize>,
116 ) -> Result<Self, String> {
117 if !(weight.is_finite() && weight > 0.0) {
118 return Err(format!(
119 "SheafConsistencyPenalty::new requires finite weight > 0, got {weight}"
120 ));
121 }
122 if edges.len() != restrictions.len() {
123 return Err(format!(
124 "SheafConsistencyPenalty::new edge count {} != restriction count {}",
125 edges.len(),
126 restrictions.len()
127 ));
128 }
129 if stalk_dims.is_empty() {
130 return Err("SheafConsistencyPenalty::new requires at least one vertex".into());
131 }
132 for (v, &d) in stalk_dims.iter().enumerate() {
133 if d == 0 {
134 return Err(format!(
135 "SheafConsistencyPenalty::new stalk dim at vertex {v} is zero"
136 ));
137 }
138 }
139 for (e, ((u, v), restriction)) in edges.iter().zip(restrictions.iter()).enumerate() {
140 if *u >= stalk_dims.len() || *v >= stalk_dims.len() {
141 return Err(format!(
142 "SheafConsistencyPenalty::new edge {e} = ({u}, {v}) references vertex \
143 out of range (K = {})",
144 stalk_dims.len()
145 ));
146 }
147 let d_u = stalk_dims[*u];
148 let d_v = stalk_dims[*v];
149 let d_e = restriction.r_uv.nrows();
150 if restriction.r_uv.ncols() != d_u {
151 return Err(format!(
152 "SheafConsistencyPenalty::new edge {e}: r_uv has {} cols, expected d_u = {d_u}",
153 restriction.r_uv.ncols()
154 ));
155 }
156 match &restriction.r_vu {
157 Some(r_vu) => {
158 if r_vu.ncols() != d_v {
159 return Err(format!(
160 "SheafConsistencyPenalty::new edge {e}: r_vu has {} cols, \
161 expected d_v = {d_v}",
162 r_vu.ncols()
163 ));
164 }
165 if r_vu.nrows() != d_e {
166 return Err(format!(
167 "SheafConsistencyPenalty::new edge {e}: r_vu has {} rows, \
168 expected d_e = {d_e}",
169 r_vu.nrows()
170 ));
171 }
172 }
173 None => {
174 if d_e != d_v {
175 return Err(format!(
176 "SheafConsistencyPenalty::new edge {e}: r_vu is identity but \
177 d_e ({d_e}) != d_v ({d_v})"
178 ));
179 }
180 }
181 }
182 if !restriction.r_uv.iter().all(|x| x.is_finite()) {
183 return Err(format!(
184 "SheafConsistencyPenalty::new edge {e}: r_uv contains non-finite entries"
185 ));
186 }
187 if let Some(r_vu) = &restriction.r_vu
188 && !r_vu.iter().all(|x| x.is_finite())
189 {
190 return Err(format!(
191 "SheafConsistencyPenalty::new edge {e}: r_vu contains non-finite entries"
192 ));
193 }
194 }
195 let mut stalk_offsets = Vec::with_capacity(stalk_dims.len() + 1);
196 let mut acc = 0usize;
197 for &d in &stalk_dims {
198 stalk_offsets.push(acc);
199 acc = acc.checked_add(d).ok_or_else(|| {
200 "SheafConsistencyPenalty::new stalk offsets overflow usize".to_string()
201 })?;
202 }
203 stalk_offsets.push(acc);
204 Ok(Self {
205 edges,
206 restrictions,
207 weight,
208 stalk_offsets,
209 stalk_dims,
210 })
211 }
212
213 pub fn total_dim(&self) -> usize {
215 *self.stalk_offsets.last().expect("offsets non-empty")
216 }
217
218 pub fn num_edges(&self) -> usize {
220 self.edges.len()
221 }
222
223 pub fn num_vertices(&self) -> usize {
225 self.stalk_dims.len()
226 }
227
228 pub fn stalk_dims(&self) -> &[usize] {
230 &self.stalk_dims
231 }
232
233 pub fn weight(&self) -> f64 {
235 self.weight
236 }
237
238 fn vertex_slice<'a>(&self, s: ArrayView1<'a, f64>, v: usize) -> ArrayView1<'a, f64> {
239 let start = self.stalk_offsets[v];
240 let end = self.stalk_offsets[v + 1];
241 s.slice_move(ndarray::s![start..end])
242 }
243
244 fn delta(&self, s: ArrayView1<'_, f64>) -> Vec<Array1<f64>> {
247 assert_eq!(
248 s.len(),
249 self.total_dim(),
250 "stacked stalk vector has wrong length",
251 );
252 let mut out = Vec::with_capacity(self.edges.len());
253 for (e, &(u, v)) in self.edges.iter().enumerate() {
254 let s_u = self.vertex_slice(s, u);
255 let s_v = self.vertex_slice(s, v);
256 let restriction = &self.restrictions[e];
257 let mut delta_e = restriction.r_uv.dot(&s_u);
259 match &restriction.r_vu {
261 Some(r_vu) => {
262 let r_vu_s_v = r_vu.dot(&s_v);
263 delta_e.scaled_add(-1.0, &r_vu_s_v);
264 }
265 None => {
266 delta_e.scaled_add(-1.0, &s_v);
267 }
268 }
269 out.push(delta_e);
270 }
271 out
272 }
273
274 fn delta_transpose(&self, y: &[Array1<f64>]) -> Array1<f64> {
277 assert_eq!(
278 y.len(),
279 self.edges.len(),
280 "delta_transpose edge count mismatch"
281 );
282 let mut out = Array1::<f64>::zeros(self.total_dim());
283 for (e, &(u, v)) in self.edges.iter().enumerate() {
284 let restriction = &self.restrictions[e];
285 let y_e = &y[e];
286 assert_eq!(y_e.len(), restriction.edge_dim(), "edge dim mismatch");
287 let contrib_u = restriction.r_uv.t().dot(y_e);
289 let u_start = self.stalk_offsets[u];
290 let u_end = self.stalk_offsets[u + 1];
291 {
292 let mut out_u = out.slice_mut(ndarray::s![u_start..u_end]);
293 out_u.scaled_add(1.0, &contrib_u);
294 }
295 let v_start = self.stalk_offsets[v];
297 let v_end = self.stalk_offsets[v + 1];
298 match &restriction.r_vu {
299 Some(r_vu) => {
300 let contrib_v = r_vu.t().dot(y_e);
301 let mut out_v = out.slice_mut(ndarray::s![v_start..v_end]);
302 out_v.scaled_add(-1.0, &contrib_v);
303 }
304 None => {
305 let mut out_v = out.slice_mut(ndarray::s![v_start..v_end]);
306 out_v.scaled_add(-1.0, y_e);
307 }
308 }
309 }
310 out
311 }
312
313 pub fn laplacian_apply(&self, s: ArrayView1<'_, f64>) -> Array1<f64> {
316 let ds = self.delta(s);
317 self.delta_transpose(&ds)
318 }
319
320 pub fn value(&self, s: ArrayView1<'_, f64>) -> f64 {
322 let ds = self.delta(s);
323 let mut sq = 0.0;
324 for de in &ds {
325 for &x in de.iter() {
326 sq += x * x;
327 }
328 }
329 0.5 * self.weight * sq
330 }
331
332 pub fn gradient(&self, s: ArrayView1<'_, f64>) -> Array1<f64> {
334 let mut g = self.laplacian_apply(s);
335 g *= self.weight;
336 g
337 }
338
339 pub fn hessian_diag(&self, s: ArrayView1<'_, f64>) -> Array1<f64> {
349 assert_eq!(
350 s.len(),
351 self.total_dim(),
352 "stacked stalk vector has wrong length",
353 );
354 let mut diag = Array1::<f64>::zeros(self.total_dim());
358 for (e, &(u, v)) in self.edges.iter().enumerate() {
359 let restriction = &self.restrictions[e];
360 let u_start = self.stalk_offsets[u];
361 let v_start = self.stalk_offsets[v];
362 let r_uv = &restriction.r_uv;
363
364 if u == v {
365 match &restriction.r_vu {
371 Some(r_vu) => {
372 for col in 0..r_uv.ncols() {
373 let mut s2 = 0.0;
374 for row in 0..r_uv.nrows() {
375 let diff = r_uv[[row, col]] - r_vu[[row, col]];
376 s2 += diff * diff;
377 }
378 diag[u_start + col] += s2;
379 }
380 }
381 None => {
382 let d = self.stalk_dims[u];
384 for col in 0..d {
385 let mut s2 = 0.0;
386 for row in 0..r_uv.nrows() {
387 let identity_entry = if row == col { 1.0 } else { 0.0 };
388 let diff = r_uv[[row, col]] - identity_entry;
389 s2 += diff * diff;
390 }
391 diag[u_start + col] += s2;
392 }
393 }
394 }
395 } else {
396 for col in 0..r_uv.ncols() {
400 let mut s2 = 0.0;
401 for row in 0..r_uv.nrows() {
402 let a = r_uv[[row, col]];
403 s2 += a * a;
404 }
405 diag[u_start + col] += s2;
406 }
407 match &restriction.r_vu {
408 Some(r_vu) => {
409 for col in 0..r_vu.ncols() {
410 let mut s2 = 0.0;
411 for row in 0..r_vu.nrows() {
412 let a = r_vu[[row, col]];
413 s2 += a * a;
414 }
415 diag[v_start + col] += s2;
416 }
417 }
418 None => {
419 let d_v = self.stalk_dims[v];
420 for col in 0..d_v {
421 diag[v_start + col] += 1.0;
422 }
423 }
424 }
425 }
426 }
427 diag *= self.weight;
428 diag
429 }
430
431 pub fn hvp(&self, s: ArrayView1<'_, f64>, v: ArrayView1<'_, f64>) -> Array1<f64> {
435 assert_eq!(
436 s.len(),
437 self.total_dim(),
438 "stacked stalk vector has wrong length",
439 );
440 assert_eq!(v.len(), self.total_dim(), "hvp direction has wrong length");
441 let mut hv = self.laplacian_apply(v);
442 hv *= self.weight;
443 hv
444 }
445
446 fn dense_laplacian(&self) -> Array2<f64> {
453 let n = self.total_dim();
454 let mut l = Array2::<f64>::zeros((n, n));
455 let mut e = Array1::<f64>::zeros(n);
456 for j in 0..n {
457 e[j] = 1.0;
458 let col = self.laplacian_apply(e.view());
459 for i in 0..n {
460 l[[i, j]] = col[i];
461 }
462 e[j] = 0.0;
463 }
464 l
465 }
466
467 pub fn harmonic_modes(&self, tol: f64) -> usize {
475 assert!(
476 tol.is_finite() && tol >= 0.0,
477 "harmonic_modes requires finite non-negative tol, got {tol}",
478 );
479 let n = self.total_dim();
480 if n == 0 {
481 return 0;
482 }
483 if n <= DENSE_EIGH_DIM_THRESHOLD {
484 let l = self.dense_laplacian();
485 match l.eigh(Side::Lower) {
486 Ok((evals, _)) => evals.iter().filter(|&&e| e < tol).count(),
487 Err(err) => {
491 panic!("SheafConsistencyPenalty::harmonic_modes faer eigh failed: {err:?}")
492 }
493 }
494 } else {
495 self.harmonic_modes_lanczos(tol)
496 }
497 }
498
499 fn harmonic_modes_lanczos(&self, tol: f64) -> usize {
508 let n = self.total_dim();
509 let k = n.min(64).max(1);
510 let mut q0 = vec![0.0_f64; n];
512 for i in 0..n {
513 let mut state = (i as u64)
517 .wrapping_mul(0x9E37_79B9_7F4A_7C15)
518 .wrapping_sub(0x9E37_79B9_7F4A_7C15);
519 let z = gam_linalg::utils::splitmix64(&mut state);
520 q0[i] = (z as f64 / u64::MAX as f64) - 0.5;
521 }
522 match symmetric_lanczos_eigenpairs(
523 n,
524 &q0,
525 SymmetricLanczosOptions {
526 max_steps: k,
527 residual_tol: 1e-12,
528 local_reorthogonalize: true,
529 full_reorthogonalize: false,
530 },
531 |q, out| {
532 let w = self.laplacian_apply(ArrayView1::from(q));
533 out.copy_from_slice(w.as_slice().ok_or_else(|| {
534 "SheafConsistencyPenalty::harmonic_modes Lanczos matvec produced non-contiguous output"
535 .to_string()
536 })?);
537 Ok(())
538 },
539 ) {
540 Ok(eigen) => eigen.eigenvalues.iter().filter(|&&e| e < tol).count(),
541 Err(err) => {
542 panic!("SheafConsistencyPenalty::harmonic_modes Lanczos failed: {err}")
547 }
548 }
549 }
550}
551
552impl AnalyticPenalty for SheafConsistencyPenalty {
566 fn tier(&self) -> PenaltyTier {
567 PenaltyTier::Psi
568 }
569
570 fn value(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> f64 {
571 assert!(
572 rho.iter().all(|x| x.is_finite()),
573 "SheafConsistencyPenalty: rho must be finite (got {rho:?})",
574 );
575 SheafConsistencyPenalty::value(self, target)
576 }
577
578 fn grad_target(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
579 assert!(
580 rho.iter().all(|x| x.is_finite()),
581 "SheafConsistencyPenalty: rho must be finite (got {rho:?})",
582 );
583 SheafConsistencyPenalty::gradient(self, target)
584 }
585
586 fn hessian_diag(
587 &self,
588 target: ArrayView1<'_, f64>,
589 rho: ArrayView1<'_, f64>,
590 ) -> Option<Array1<f64>> {
591 assert!(
592 rho.iter().all(|x| x.is_finite()),
593 "SheafConsistencyPenalty: rho must be finite (got {rho:?})",
594 );
595 Some(SheafConsistencyPenalty::hessian_diag(self, target))
596 }
597
598 fn hvp(
599 &self,
600 target: ArrayView1<'_, f64>,
601 rho: ArrayView1<'_, f64>,
602 v: ArrayView1<'_, f64>,
603 ) -> Array1<f64> {
604 assert!(
605 rho.iter().all(|x| x.is_finite()),
606 "SheafConsistencyPenalty: rho must be finite (got {rho:?})",
607 );
608 SheafConsistencyPenalty::hvp(self, target, v)
609 }
610
611 fn grad_rho(&self, target: ArrayView1<'_, f64>, rho: ArrayView1<'_, f64>) -> Array1<f64> {
612 assert_eq!(
614 rho.len(),
615 0,
616 "SheafConsistencyPenalty: rho_count is 0 but rho has length {}",
617 rho.len(),
618 );
619 assert_eq!(
620 target.len(),
621 self.total_dim(),
622 "SheafConsistencyPenalty: target length {} != total stalk dim {}",
623 target.len(),
624 self.total_dim(),
625 );
626 Array1::<f64>::zeros(0)
627 }
628
629 fn rho_count(&self) -> usize {
630 0
631 }
632
633 fn name(&self) -> &str {
634 "SheafConsistencyPenalty"
635 }
636}
637
638#[cfg(test)]
639mod tests {
640 use super::*;
641 use approx::assert_abs_diff_eq;
642 use ndarray::array;
643
644 fn identity(d: usize) -> Array2<f64> {
645 let mut m = Array2::<f64>::zeros((d, d));
646 for i in 0..d {
647 m[[i, i]] = 1.0;
648 }
649 m
650 }
651
652 #[test]
653 fn single_edge_identity_restriction_value() {
654 let edges = vec![(0usize, 1usize)];
657 let restrictions = vec![EdgeRestriction::paired(identity(3), identity(3))];
658 let pen =
659 SheafConsistencyPenalty::new(edges, restrictions, 1.0, vec![3, 3]).expect("build");
660 let s = array![1.0_f64, 0.0, 0.0, 0.0, 1.0, 0.0];
661 let v = pen.value(s.view());
662 assert_abs_diff_eq!(v, 1.0, epsilon = 1e-12);
663 }
664
665 #[test]
666 fn gradient_matches_finite_difference_k2_random() {
667 let r_uv = array![[0.7_f64, -0.1, 0.3], [0.2, 0.9, -0.4]];
669 let r_vu = array![[1.0_f64, 0.5], [-0.3, 0.8]];
670 let edges = vec![(0usize, 1usize)];
671 let restrictions = vec![EdgeRestriction::paired(r_uv, r_vu)];
672 let pen =
673 SheafConsistencyPenalty::new(edges, restrictions, 0.3, vec![3, 2]).expect("build");
674 let s = array![0.4_f64, -1.1, 0.2, 0.6, -0.7];
675 let g = pen.gradient(s.view());
676 let eps = 1e-7;
677 for i in 0..s.len() {
678 let mut sp = s.clone();
679 let mut sm = s.clone();
680 sp[i] += eps;
681 sm[i] -= eps;
682 let fd = (pen.value(sp.view()) - pen.value(sm.view())) / (2.0 * eps);
683 assert_abs_diff_eq!(g[i], fd, epsilon = 1e-6);
684 }
685 }
686
687 #[test]
688 fn hvp_matches_reconstructed_laplacian_chain_k3() {
689 let r01_uv = array![[0.9_f64, 0.1], [-0.2, 0.7]];
692 let r01_vu = array![[1.0_f64, 0.0], [0.0, 1.0]];
693 let r12_uv = array![[0.5_f64, -0.3], [0.4, 0.8]];
694 let r12_vu = array![[0.6_f64, 0.0], [0.1, 1.1]];
695 let edges = vec![(0usize, 1usize), (1usize, 2usize)];
696 let restrictions = vec![
697 EdgeRestriction::paired(r01_uv, r01_vu),
698 EdgeRestriction::paired(r12_uv, r12_vu),
699 ];
700 let pen =
701 SheafConsistencyPenalty::new(edges, restrictions, 1.0, vec![2, 2, 2]).expect("build");
702 let l_dense = pen.dense_laplacian();
704 let n = pen.total_dim();
705 let s = array![0.1_f64, -0.2, 0.3, 0.4, -0.5, 0.6];
706 let v = array![0.7_f64, 0.2, -0.1, 0.5, 0.3, -0.4];
707 let hv = pen.hvp(s.view(), v.view());
708 let mut lv = Array1::<f64>::zeros(n);
710 for i in 0..n {
711 let mut acc = 0.0;
712 for j in 0..n {
713 acc += l_dense[[i, j]] * v[j];
714 }
715 lv[i] = acc;
716 }
717 for i in 0..n {
718 assert_abs_diff_eq!(hv[i], lv[i], epsilon = 1e-10);
719 }
720 }
721
722 #[test]
723 fn harmonic_modes_two_components_identity_restrictions() {
724 let pen = SheafConsistencyPenalty::new(vec![], vec![], 1.0, vec![3, 3]).expect("build");
726 let h = pen.harmonic_modes(1e-10);
727 assert_eq!(h, 6);
728
729 let edges = vec![(0usize, 1usize), (2usize, 3usize)];
733 let restrictions = vec![
734 EdgeRestriction::paired(identity(2), identity(2)),
735 EdgeRestriction::paired(identity(2), identity(2)),
736 ];
737 let pen2 = SheafConsistencyPenalty::new(edges, restrictions, 1.0, vec![2, 2, 2, 2])
738 .expect("build");
739 let h2 = pen2.harmonic_modes(1e-10);
740 assert_eq!(h2, 4);
741 }
742
743 #[test]
744 fn value_psd_and_zero_iff_kernel() {
745 let r01_uv = array![[0.9_f64, 0.1], [-0.2, 0.7]];
747 let r01_vu = array![[1.0_f64, 0.0], [0.0, 1.0]];
748 let edges = vec![(0usize, 1usize)];
749 let restrictions = vec![EdgeRestriction::paired(r01_uv.clone(), r01_vu.clone())];
750 let pen =
751 SheafConsistencyPenalty::new(edges, restrictions, 0.5, vec![2, 2]).expect("build");
752
753 let samples = [
755 array![0.0_f64, 0.0, 0.0, 0.0],
756 array![1.0_f64, 2.0, -0.5, 0.3],
757 array![-1.3_f64, 0.7, 0.2, -0.9],
758 ];
759 for s in &samples {
760 let v = pen.value(s.view());
761 assert!(v >= 0.0, "value must be non-negative, got {v}");
762 }
763 let z = Array1::<f64>::zeros(4);
765 assert_abs_diff_eq!(pen.value(z.view()), 0.0, epsilon = 1e-15);
766 let s0 = array![0.3_f64, -1.1];
769 let s1 = r01_uv.dot(&s0);
770 let mut s = Array1::<f64>::zeros(4);
771 s[0] = s0[0];
772 s[1] = s0[1];
773 s[2] = s1[0];
774 s[3] = s1[1];
775 assert_abs_diff_eq!(pen.value(s.view()), 0.0, epsilon = 1e-12);
776 }
777
778 #[test]
779 fn hessian_diag_matches_diag_of_dense_laplacian() {
780 let r_uv = array![[0.7_f64, -0.1, 0.3], [0.2, 0.9, -0.4]];
781 let r_vu = array![[1.0_f64, 0.5], [-0.3, 0.8]];
782 let edges = vec![(0usize, 1usize)];
783 let restrictions = vec![EdgeRestriction::paired(r_uv, r_vu)];
784 let pen =
785 SheafConsistencyPenalty::new(edges, restrictions, 0.3, vec![3, 2]).expect("build");
786 let n = pen.total_dim();
787 let s = Array1::<f64>::zeros(n);
788 let diag = pen.hessian_diag(s.view());
789 let l = pen.dense_laplacian();
790 for i in 0..n {
791 assert_abs_diff_eq!(diag[i], 0.3 * l[[i, i]], epsilon = 1e-12);
792 }
793 }
794
795 #[test]
796 fn hessian_diag_matches_dense_laplacian_on_self_loop_paired() {
797 let r_uv = array![[0.9_f64, 0.1], [-0.2, 0.7]];
801 let r_vu = array![[1.0_f64, 0.5], [-0.3, 0.8]];
802 let edges = vec![(0usize, 0usize)];
803 let restrictions = vec![EdgeRestriction::paired(r_uv, r_vu)];
804 let pen = SheafConsistencyPenalty::new(edges, restrictions, 0.7, vec![2]).expect("build");
805 let n = pen.total_dim();
806 let s = Array1::<f64>::zeros(n);
807 let diag = pen.hessian_diag(s.view());
808 let l = pen.dense_laplacian();
809 for i in 0..n {
810 assert_abs_diff_eq!(diag[i], 0.7 * l[[i, i]], epsilon = 1e-12);
811 }
812 }
813
814 #[test]
815 fn hessian_diag_matches_dense_laplacian_on_self_loop_single() {
816 let r_uv = array![[1.0_f64, 2.0], [3.0, 4.0]];
823 let edges = vec![(0usize, 0usize)];
824 let restrictions = vec![EdgeRestriction::single(r_uv)];
825 let pen = SheafConsistencyPenalty::new(edges, restrictions, 1.3, vec![2]).expect("build");
826 let n = pen.total_dim();
827 let s = Array1::<f64>::zeros(n);
828 let diag = pen.hessian_diag(s.view());
829 let l = pen.dense_laplacian();
830 for i in 0..n {
831 assert_abs_diff_eq!(diag[i], 1.3 * l[[i, i]], epsilon = 1e-12);
832 }
833 assert_abs_diff_eq!(diag[0], 1.3 * 9.0, epsilon = 1e-12);
836 assert_abs_diff_eq!(diag[1], 1.3 * 13.0, epsilon = 1e-12);
837 }
838
839 #[test]
840 fn hessian_diag_matches_dense_laplacian_mixed_self_loop_and_cross_edge() {
841 let r0_uv = array![[0.5_f64, -0.4], [0.3, 0.9]];
847 let r0_vu = array![[0.2_f64, 0.1], [-0.6, 0.7]];
848 let r1_uv = array![[1.1_f64, 0.2], [0.0, -0.5]];
849 let r1_vu = array![[0.8_f64, -0.1], [0.4, 1.0]];
850 let edges = vec![(0usize, 0usize), (0usize, 1usize)];
851 let restrictions = vec![
852 EdgeRestriction::paired(r0_uv, r0_vu),
853 EdgeRestriction::paired(r1_uv, r1_vu),
854 ];
855 let pen =
856 SheafConsistencyPenalty::new(edges, restrictions, 0.5, vec![2, 2]).expect("build");
857 let n = pen.total_dim();
858 let s = Array1::<f64>::zeros(n);
859 let diag = pen.hessian_diag(s.view());
860 let l = pen.dense_laplacian();
861 for i in 0..n {
862 assert_abs_diff_eq!(diag[i], 0.5 * l[[i, i]], epsilon = 1e-12);
863 }
864 }
865
866 #[test]
867 fn single_restriction_edge_form() {
868 let r = array![[1.0_f64, 2.0], [3.0, 4.0]];
870 let edges = vec![(0usize, 1usize)];
871 let restrictions = vec![EdgeRestriction::single(r.clone())];
872 let pen =
873 SheafConsistencyPenalty::new(edges, restrictions, 2.0, vec![2, 2]).expect("build");
874 let s = array![1.0_f64, 0.0, 1.0, 3.0];
876 assert_abs_diff_eq!(pen.value(s.view()), 0.0, epsilon = 1e-12);
877 let s2 = array![1.0_f64, 0.0, 0.0, 0.0];
879 assert_abs_diff_eq!(pen.value(s2.view()), 10.0, epsilon = 1e-12);
880 }
881}