1use ndarray::{Array1, Array2, ArrayView1, ArrayView2, Axis};
115use rayon::prelude::*;
116use serde::{Deserialize, Serialize};
117
118use faer::Side;
119
120use gam_linalg::faer_ndarray::{FaerEigh, default_rrqr_rank_alpha, rrqr_nullspace_basis};
121
122use super::{
123 AnisoBasisPsiDerivatives, AnisoPenaltyCrossProvider, BasisBuildResult, BasisError,
124 BasisMetadata, CenterStrategy, ConstructiveQuadratic, PenaltyCandidate, PenaltySource,
125 filter_penalty_candidates, normalize_penalty, normalize_penalty_cross_psi_derivative,
126 normalize_penaltywith_psi_derivatives, select_centers_by_strategy, trace_of_product,
127};
128
129pub(crate) const MEASURE_JET_PROFILE_CUTOFF: f64 = 3.0;
135
136pub(crate) const MEASURE_JET_PSEUDOINVERSE_RTOL: f64 = 64.0 * f64::EPSILON;
140
141pub(crate) const MEASURE_JET_DEFAULT_ORDER_S: f64 = 1.5;
147
148pub(crate) const MEASURE_JET_MIN_AUTO_SCALES: usize = 3;
152pub(crate) const MEASURE_JET_MAX_AUTO_SCALES: usize = 8;
153
154pub(crate) const MEASURE_JET_AUTO_LENGTH_SCALE_FACTOR: f64 = 1.0;
169
170pub(crate) const MEASURE_JET_PARALLEL_FORM_BUDGET_DOUBLES: usize = 1 << 26;
176
177#[derive(Debug, Clone, Serialize, Deserialize, Default)]
183pub enum MeasureJetIdentifiability {
184 #[default]
189 CenterSumToZero,
190 FrozenTransform { transform: Array2<f64> },
193}
194
195#[derive(Debug, Clone, Serialize, Deserialize)]
200pub struct MeasureJetFrozenQuadrature {
201 pub masses: Array1<f64>,
203 pub eps_band: Vec<f64>,
205 pub support_means: Vec<f64>,
208 pub penalty_normalization_scales: Vec<f64>,
211 pub raw_penalty_normalization_scales: Vec<f64>,
214 pub fused_penalty_normalization_scale: Option<f64>,
217 #[serde(default)]
226 pub sigma_coord: Option<f64>,
227}
228
229fn measure_jet_learn_length_scale_default() -> bool {
233 false
234}
235
236#[derive(Debug, Clone, Serialize, Deserialize)]
244pub struct MeasureJetBasisSpec {
245 pub center_strategy: CenterStrategy,
247 pub order_s: f64,
250 pub alpha: f64,
252 pub tau0: f64,
256 pub num_scales: usize,
258 pub length_scale: f64,
261 pub double_penalty: bool,
264 #[serde(default = "measure_jet_learn_length_scale_default")]
272 pub learn_length_scale: bool,
273 #[serde(default)]
281 pub multiscale: bool,
282 #[serde(default)]
284 pub identifiability: MeasureJetIdentifiability,
285 #[serde(default)]
288 pub frozen_quadrature: Option<MeasureJetFrozenQuadrature>,
289}
290
291impl Default for MeasureJetBasisSpec {
292 fn default() -> Self {
293 Self {
294 center_strategy: CenterStrategy::FarthestPoint { num_centers: 50 },
295 order_s: 0.0,
296 alpha: 1.0,
308 tau0: 1e-3,
309 num_scales: 0,
310 length_scale: 0.0,
311 double_penalty: true,
312 learn_length_scale: false,
313 multiscale: false,
314 identifiability: MeasureJetIdentifiability::CenterSumToZero,
315 frozen_quadrature: None,
316 }
317 }
318}
319
320pub struct MeasureJetBand {
323 pub eps: Vec<f64>,
324 pub log_step: f64,
325}
326
327pub struct MeasureJetEnergyJets {
333 pub q: Array2<f64>,
334 pub dq_ds: Array2<f64>,
335 pub d2q_ds2: Array2<f64>,
336 pub dq_dalpha: Array2<f64>,
337 pub d2q_dalpha2: Array2<f64>,
338 pub d2q_ds_dalpha: Array2<f64>,
339 pub dq_dlogtau: Array2<f64>,
340 pub d2q_dlogtau2: Array2<f64>,
341 pub d2q_ds_dlogtau: Array2<f64>,
342 pub d2q_dalpha_dlogtau: Array2<f64>,
343}
344
345pub(crate) fn householder_sum_to_zero_u(m: usize) -> Array1<f64> {
353 let c = 1.0 / (m as f64).sqrt();
354 let mut u = Array1::<f64>::from_elem(m, c);
355 u[0] -= 1.0;
356 let norm = u.dot(&u).sqrt();
357 u.mapv_inplace(|v| v / norm);
358 u
359}
360
361pub(crate) fn householder_sum_to_zero_z(u: &Array1<f64>) -> Array2<f64> {
365 let m = u.len();
366 let mut z = Array2::<f64>::zeros((m, m - 1));
367 for j in 0..(m - 1) {
368 for i in 0..m {
369 let h = if i == j + 1 { 1.0 } else { 0.0 } - 2.0 * u[i] * u[j + 1];
370 z[(i, j)] = h;
371 }
372 }
373 z
374}
375
376pub(crate) fn symmetric_pseudoinverse(
377 a: &Array2<f64>,
378 label: &str,
379) -> Result<Array2<f64>, BasisError> {
380 let n = a.nrows();
381 if a.ncols() != n {
382 crate::bail_dim_basis!(
383 "measure-jet pseudo-inverse `{label}` needs a square matrix, got {:?}",
384 a.dim()
385 );
386 }
387 let (evals, evecs) = a.eigh(Side::Lower).map_err(|e| {
388 BasisError::InvalidInput(format!(
389 "measure-jet pseudo-inverse `{label}` eigendecomposition failed: {e}"
390 ))
391 })?;
392 let lam_max = evals.iter().fold(0.0_f64, |acc, v| acc.max((*v).max(0.0)));
393 let rank_tol = MEASURE_JET_PSEUDOINVERSE_RTOL * (n.max(1) as f64) * lam_max;
394 let mut scaled = evecs.clone();
395 for (k, mut col) in scaled.axis_iter_mut(Axis(1)).enumerate() {
396 let lam = evals[k].max(0.0);
397 let inv = if lam > rank_tol { 1.0 / lam } else { 0.0 };
398 col.mapv_inplace(|v| v * inv);
399 }
400 Ok(scaled.dot(&evecs.t()))
401}
402
403fn measure_jet_affine_value_basis(
410 centers: ArrayView2<'_, f64>,
411 masses: ArrayView1<'_, f64>,
412) -> Array2<f64> {
413 let m = centers.nrows();
414 let head_transform = measure_jet_affine_head_transform(centers, masses);
415 let head_rank = head_transform.ncols();
416 let mut affine = Array2::<f64>::ones((m, head_rank + 1));
417 if head_rank > 0 {
418 affine
419 .slice_mut(ndarray::s![.., 1..])
420 .assign(¢ers.dot(&head_transform));
421 }
422 affine
423}
424
425fn affine_function_nullspace_form(
433 centers: ArrayView2<'_, f64>,
434 masses: ArrayView1<'_, f64>,
435) -> Result<Array2<f64>, BasisError> {
436 let m = centers.nrows();
437 if masses.len() != m {
438 crate::bail_dim_basis!(
439 "measure-jet affine function-space form shape mismatch: centers {:?}, masses {}",
440 centers.dim(),
441 masses.len()
442 );
443 }
444 let affine = measure_jet_affine_value_basis(centers, masses);
445 let mut weighted_affine = affine.clone();
446 for (i, mut row) in weighted_affine.outer_iter_mut().enumerate() {
447 row.mapv_inplace(|v| v * masses[i]);
448 }
449 let affine_gram = affine.t().dot(&weighted_affine);
450 let affine_gram_pinv = symmetric_pseudoinverse(&affine_gram, "affine function-space Gram")?;
451 let form = weighted_affine
452 .dot(&affine_gram_pinv)
453 .dot(&weighted_affine.t());
454 Ok((&form + &form.t()) * 0.5)
455}
456
457fn pullback_center_form(evaluation: &Array2<f64>, form: &Array2<f64>) -> Array2<f64> {
459 let pulled = evaluation.t().dot(form).dot(evaluation);
460 (&pulled + &pulled.t()) * 0.5
461}
462
463fn pullback_center_form_log_length_jets(
466 evaluation: &Array2<f64>,
467 evaluation_first: &Array2<f64>,
468 evaluation_second: &Array2<f64>,
469 form: &Array2<f64>,
470) -> (Array2<f64>, Array2<f64>) {
471 let h_e = form.dot(evaluation);
472 let h_e_first = form.dot(evaluation_first);
473 let h_e_second = form.dot(evaluation_second);
474 let first_raw = evaluation_first.t().dot(&h_e) + evaluation.t().dot(&h_e_first);
475 let second_raw = evaluation_second.t().dot(&h_e)
476 + evaluation.t().dot(&h_e_second)
477 + evaluation_first.t().dot(&h_e_first) * 2.0;
478 (
479 (&first_raw + &first_raw.t()) * 0.5,
480 (&second_raw + &second_raw.t()) * 0.5,
481 )
482}
483
484fn pullback_center_form_log_length_cross(
486 evaluation: &Array2<f64>,
487 evaluation_first: &Array2<f64>,
488 form_first: &Array2<f64>,
489) -> Array2<f64> {
490 let h_e = form_first.dot(evaluation);
491 let h_e_first = form_first.dot(evaluation_first);
492 let cross_raw = evaluation_first.t().dot(&h_e) + evaluation.t().dot(&h_e_first);
493 (&cross_raw + &cross_raw.t()) * 0.5
494}
495
496pub(crate) fn affine_function_nullspace_penalty(
501 evaluation: &Array2<f64>,
502 centers: ArrayView2<'_, f64>,
503 masses: ArrayView1<'_, f64>,
504) -> Result<Array2<f64>, BasisError> {
505 if evaluation.nrows() != centers.nrows() {
506 crate::bail_dim_basis!(
507 "measure-jet affine function-space penalty shape mismatch: evaluation {:?}, centers {:?}",
508 evaluation.dim(),
509 centers.dim()
510 );
511 }
512 let form = affine_function_nullspace_form(centers, masses)?;
513 Ok(pullback_center_form(evaluation, &form))
514}
515
516pub(crate) fn pairwise_sq_dists(a: ArrayView2<'_, f64>, b: ArrayView2<'_, f64>) -> Array2<f64> {
527 let an: Vec<f64> = a.outer_iter().map(|r| r.dot(&r)).collect();
528 let bn: Vec<f64> = b.outer_iter().map(|r| r.dot(&r)).collect();
529 let mut g = a.dot(&b.t());
530 g.axis_iter_mut(Axis(0))
531 .into_par_iter()
532 .enumerate()
533 .for_each(|(i, mut row)| {
534 for (j, v) in row.iter_mut().enumerate() {
535 *v = (an[i] + bn[j] - 2.0 * *v).max(0.0);
536 }
537 });
538 g
539}
540
541pub(crate) const MEASURE_JET_ASSIGN_BLOCK_ROWS: usize = 65_536;
545
546pub(crate) fn validate_finite_points(
547 points: ArrayView2<'_, f64>,
548 what: &str,
549) -> Result<(), BasisError> {
550 for (i, row) in points.outer_iter().enumerate() {
551 if row.iter().any(|v| !v.is_finite()) {
552 crate::bail_invalid_basis!("measure-jet {what} row {i} has a non-finite coordinate");
553 }
554 }
555 Ok(())
556}
557
558pub(crate) fn median_nearest_center_spacing(dist2: &Array2<f64>) -> Result<f64, BasisError> {
561 let m = dist2.nrows();
562 if m < 2 {
563 return Err(BasisError::InsufficientColumnsForConstraint { found: m });
564 }
565 let mut nearest: Vec<f64> = Vec::with_capacity(m);
566 for i in 0..m {
567 let mut best = f64::INFINITY;
568 for j in 0..m {
569 if j != i && dist2[(i, j)] < best {
570 best = dist2[(i, j)];
571 }
572 }
573 nearest.push(best.sqrt());
574 }
575 nearest.sort_by(|a, b| a.partial_cmp(b).expect("finite center spacings"));
576 let median = nearest[nearest.len() / 2];
577 if !(median.is_finite() && median > 0.0) {
578 crate::bail_invalid_basis!(
579 "measure-jet centers are degenerate (median nearest-center spacing = {median}); \
580 duplicate centers cannot carry a scale band"
581 );
582 }
583 Ok(median)
584}
585
586pub fn measure_jet_band(
594 centers: ArrayView2<'_, f64>,
595 num_scales: usize,
596) -> Result<MeasureJetBand, BasisError> {
597 validate_finite_points(centers, "centers")?;
598 let dist2 = pairwise_sq_dists(centers, centers);
599 let eps_min = median_nearest_center_spacing(&dist2)?;
600 let d = centers.ncols();
602 let mut diag2 = 0.0_f64;
603 for k in 0..d {
604 let col = centers.column(k);
605 let mut lo = f64::INFINITY;
606 let mut hi = f64::NEG_INFINITY;
607 for &v in col.iter() {
608 lo = lo.min(v);
609 hi = hi.max(v);
610 }
611 diag2 += (hi - lo) * (hi - lo);
612 }
613 let eps_max = 0.5 * diag2.sqrt();
614 if !(eps_max.is_finite() && eps_max > eps_min) {
615 return Ok(MeasureJetBand {
616 eps: vec![eps_min],
617 log_step: std::f64::consts::LN_2,
618 });
619 }
620 let auto = ((eps_max / eps_min).log2().ceil() as usize + 1)
621 .clamp(MEASURE_JET_MIN_AUTO_SCALES, MEASURE_JET_MAX_AUTO_SCALES);
622 let count = if num_scales == 0 { auto } else { num_scales };
623 if count == 1 {
624 return Ok(MeasureJetBand {
625 eps: vec![eps_min],
626 log_step: std::f64::consts::LN_2,
627 });
628 }
629 let ratio = (eps_max / eps_min).powf(1.0 / (count as f64 - 1.0));
630 let mut eps = Vec::with_capacity(count);
631 let mut e = eps_min;
632 for _ in 0..count {
633 eps.push(e);
634 e *= ratio;
635 }
636 Ok(MeasureJetBand {
637 eps,
638 log_step: ratio.ln(),
639 })
640}
641
642pub fn measure_jet_quadrature_nodes(
649 data: ArrayView2<'_, f64>,
650 centers: ArrayView2<'_, f64>,
651) -> Result<(Array2<f64>, Array1<f64>), BasisError> {
652 if data.ncols() != centers.ncols() {
653 crate::bail_dim_basis!(
654 "measure-jet mass assignment dimension mismatch: data d={} centers d={}",
655 data.ncols(),
656 centers.ncols()
657 );
658 }
659 validate_finite_points(data, "data")?;
660 validate_finite_points(centers, "centers")?;
661 let n = data.nrows();
662 let m = centers.nrows();
663 let d = centers.ncols();
664 if n == 0 || m == 0 {
665 crate::bail_invalid_basis!("measure-jet mass assignment needs nonempty data and centers");
666 }
667 let cn: Vec<f64> = centers.outer_iter().map(|r| r.dot(&r)).collect();
672 let assignments: Vec<usize> = (0..n)
673 .step_by(MEASURE_JET_ASSIGN_BLOCK_ROWS)
674 .flat_map(|start| {
675 let end = (start + MEASURE_JET_ASSIGN_BLOCK_ROWS).min(n);
676 let g = data.slice(ndarray::s![start..end, ..]).dot(¢ers.t());
677 let block: Vec<usize> = g
678 .axis_iter(Axis(0))
679 .into_par_iter()
680 .map(|row| {
681 let mut best_j = 0usize;
682 let mut best = f64::INFINITY;
683 for (j, &gij) in row.iter().enumerate() {
684 let s = cn[j] - 2.0 * gij;
685 if s < best {
686 best = s;
687 best_j = j;
688 }
689 }
690 best_j
691 })
692 .collect();
693 block
694 })
695 .collect();
696 let mut masses = Array1::<f64>::zeros(m);
697 let mut nodes = centers.to_owned();
698 let mut sums = Array2::<f64>::zeros((m, d));
699 let unit = 1.0 / n as f64;
700 for (i, &j) in assignments.iter().enumerate() {
701 masses[j] += unit;
702 for k in 0..d {
703 sums[(j, k)] += data[(i, k)];
704 }
705 }
706 let mut barycenter = sums;
709 for j in 0..m {
710 let count = masses[j] * n as f64;
711 if count > 0.0 {
712 for k in 0..d {
713 barycenter[(j, k)] /= count;
714 nodes[(j, k)] = barycenter[(j, k)];
715 }
716 }
717 }
718 Ok((nodes, masses))
719}
720
721pub fn measure_jet_center_masses(
724 data: ArrayView2<'_, f64>,
725 centers: ArrayView2<'_, f64>,
726) -> Result<Array1<f64>, BasisError> {
727 measure_jet_quadrature_nodes(data, centers).map(|(_, masses)| masses)
728}
729
730pub(crate) fn assemble_weighted_forms<F>(
753 centers: ArrayView2<'_, f64>,
754 masses: ArrayView1<'_, f64>,
755 band: &MeasureJetBand,
756 order_s: f64,
757 alpha: f64,
758 tau0: f64,
759 n_forms: usize,
760 channels: usize,
761 weights: &F,
762) -> Result<Vec<Array2<f64>>, BasisError>
763where
764 F: Fn(usize, f64, f64, f64, &mut [[f64; 3]]) + Sync,
765{
766 let m = centers.nrows();
767 let d = centers.ncols();
768 if n_forms == 0 || !(1..=3).contains(&channels) {
769 crate::bail_invalid_basis!(
770 "measure-jet assembly needs at least one output form and 1..=3 block channels"
771 );
772 }
773 if masses.len() != m {
774 crate::bail_dim_basis!(
775 "measure-jet energy mass/center mismatch: {} masses for {} centers",
776 masses.len(),
777 m
778 );
779 }
780 if band.eps.is_empty() || band.eps.iter().any(|e| !(e.is_finite() && *e > 0.0)) {
781 crate::bail_invalid_basis!("measure-jet energy needs a nonempty positive scale band");
782 }
783 if !(order_s.is_finite() && order_s > 0.0 && order_s < 2.0) {
784 crate::bail_invalid_basis!(
785 "measure-jet order s must lie in (0, 2) for the affine-jet energy; got {order_s}"
786 );
787 }
788 if !(alpha.is_finite() && tau0.is_finite() && tau0 >= 0.0) {
789 crate::bail_invalid_basis!(
790 "measure-jet energy needs finite alpha and finite tau0 >= 0; got alpha={alpha}, tau0={tau0}"
791 );
792 }
793 if masses.iter().any(|v| !(v.is_finite() && *v >= 0.0)) {
794 crate::bail_invalid_basis!("measure-jet energy needs finite nonnegative center masses");
795 }
796 let dist2 = pairwise_sq_dists(centers, centers);
797
798 let assemble_scale = |scale_idx: usize, eps: f64| -> Result<Vec<Array2<f64>>, BasisError> {
803 let mut out: Vec<Array2<f64>> =
804 (0..n_forms).map(|_| Array2::<f64>::zeros((m, m))).collect();
805 let cutoff2 = (MEASURE_JET_PROFILE_CUTOFF * eps) * (MEASURE_JET_PROFILE_CUTOFF * eps);
806 let inv_two_eps2 = 1.0 / (2.0 * eps * eps);
807 let eta = 2.0 * order_s + (d as f64) * (2.0 - 2.0 * alpha);
808 let scale_weight = band.log_step * eps.powf(-eta);
809 let net_radius2 = 0.25 * eps * eps;
813 let mut outer: Vec<usize> = Vec::new();
814 for i in 0..m {
815 if masses[i] <= 0.0 {
816 continue;
817 }
818 let covered = outer.iter().any(|&o| dist2[(i, o)] <= net_radius2);
819 if !covered {
820 outer.push(i);
821 }
822 }
823 let mut net_mass = vec![0.0_f64; m];
824 for i in 0..m {
825 if masses[i] <= 0.0 {
826 continue;
827 }
828 let mut best = f64::INFINITY;
829 let mut best_o = usize::MAX;
830 for &o in &outer {
831 if dist2[(i, o)] < best {
832 best = dist2[(i, o)];
833 best_o = o;
834 }
835 }
836 if best_o != usize::MAX {
837 net_mass[best_o] += masses[i];
838 }
839 }
840 let mut wbuf = vec![[0.0_f64; 3]; n_forms];
841 for &i in &outer {
842 let mut idx: Vec<usize> = Vec::new();
844 for j in 0..m {
845 if dist2[(i, j)] <= cutoff2 {
846 idx.push(j);
847 }
848 }
849 let ml = idx.len();
850 let mut w = Array1::<f64>::zeros(ml);
852 let mut q = 0.0_f64;
853 for (a, &j) in idx.iter().enumerate() {
854 let wj = masses[j] * (-dist2[(i, j)] * inv_two_eps2).exp();
855 w[a] = wj;
856 q += wj;
857 }
858 if !(q > 0.0) {
859 continue;
860 }
861 let mut phi = Array2::<f64>::zeros((ml, d));
863 for (a, &j) in idx.iter().enumerate() {
864 for k in 0..d {
865 phi[(a, k)] = (centers[(j, k)] - centers[(i, k)]) / eps;
866 }
867 }
868 let a_mean = phi.t().dot(&w) / q;
869 let mut wphi = phi.clone();
871 for (a, mut row) in wphi.outer_iter_mut().enumerate() {
872 row.mapv_inplace(|v| v * w[a]);
873 }
874 let mut b = wphi.clone();
875 for (a, mut row) in b.outer_iter_mut().enumerate() {
876 for k in 0..d {
877 row[k] -= w[a] * a_mean[k];
878 }
879 }
880 let mut g = phi.t().dot(&wphi);
881 g.mapv_inplace(|v| v / q);
882 for r in 0..d {
883 for c in 0..d {
884 g[(r, c)] -= a_mean[r] * a_mean[c];
885 }
886 }
887 let g_pinv = symmetric_pseudoinverse(&g, "local affine Gram")?;
888 let bm = b.dot(&g_pinv);
889 let base = scale_weight * net_mass[i] * q.powf(1.0 - 2.0 * alpha);
890 weights(scale_idx, eps, q, base, &mut wbuf);
891 for (a, &ja) in idx.iter().enumerate() {
894 let bma = bm.row(a);
895 for (c, &jc) in idx.iter().enumerate() {
896 let b_c = b.row(c);
897 let mut val_r = -w[a] * w[c] / q - bma.dot(&b_c) / q;
898 if a == c {
899 val_r += w[a];
900 }
901 for (k, out_k) in out.iter_mut().enumerate() {
902 let wk = wbuf[k];
903 out_k[(ja, jc)] += wk[0] * val_r;
904 }
905 }
906 }
907 }
908 Ok(out)
909 };
910
911 let n_scales = band.eps.len();
912 let parallel_ok = m
913 .saturating_mul(m)
914 .saturating_mul(n_scales)
915 .saturating_mul(n_forms)
916 <= MEASURE_JET_PARALLEL_FORM_BUDGET_DOUBLES;
917 let per_scale: Vec<Vec<Array2<f64>>> = if parallel_ok {
918 band.eps
919 .par_iter()
920 .enumerate()
921 .map(|(scale_idx, &eps)| assemble_scale(scale_idx, eps))
922 .collect::<Result<Vec<_>, BasisError>>()?
923 } else {
924 band.eps
925 .iter()
926 .enumerate()
927 .map(|(scale_idx, &eps)| assemble_scale(scale_idx, eps))
928 .collect::<Result<Vec<_>, BasisError>>()?
929 };
930
931 let mut totals: Vec<Array2<f64>> = (0..n_forms).map(|_| Array2::<f64>::zeros((m, m))).collect();
932 for scale_forms in per_scale {
933 for (total, part) in totals.iter_mut().zip(scale_forms) {
934 *total += ∂
935 }
936 }
937 Ok(totals.into_iter().map(|t| (&t + &t.t()) * 0.5).collect())
939}
940
941pub fn measure_jet_energy_form(
956 centers: ArrayView2<'_, f64>,
957 masses: ArrayView1<'_, f64>,
958 band: &MeasureJetBand,
959 order_s: f64,
960 alpha: f64,
961 tau0: f64,
962) -> Result<Array2<f64>, BasisError> {
963 let mut forms = assemble_weighted_forms(
964 centers,
965 masses,
966 band,
967 order_s,
968 alpha,
969 tau0,
970 1,
971 1,
972 &|_, _, _, base, out: &mut [[f64; 3]]| out[0] = [base, 0.0, 0.0],
973 )?;
974 let q = forms.swap_remove(0);
975 project_symmetric_psd(q, "measure-jet energy form")
983}
984
985pub(crate) fn project_symmetric_psd(
991 a: Array2<f64>,
992 label: &str,
993) -> Result<Array2<f64>, BasisError> {
994 let n = a.nrows();
995 if n == 0 {
996 return Ok(a);
997 }
998 let (evals, evecs) = a.eigh(Side::Lower).map_err(|e| {
999 BasisError::InvalidInput(format!(
1000 "measure-jet PSD projection `{label}` eigendecomposition failed: {e}"
1001 ))
1002 })?;
1003 if evals.iter().all(|&lam| lam >= 0.0) {
1004 return Ok(a);
1005 }
1006 let mut scaled = evecs.clone();
1007 for (k, mut col) in scaled.axis_iter_mut(Axis(1)).enumerate() {
1008 let lam = evals[k].max(0.0);
1009 col.mapv_inplace(|v| v * lam);
1010 }
1011 let psd = scaled.dot(&evecs.t());
1012 Ok((&psd + &psd.t()) * 0.5)
1013}
1014
1015pub fn measure_jet_energy_form_with_jets(
1030 centers: ArrayView2<'_, f64>,
1031 masses: ArrayView1<'_, f64>,
1032 band: &MeasureJetBand,
1033 order_s: f64,
1034 alpha: f64,
1035 tau0: f64,
1036) -> Result<MeasureJetEnergyJets, BasisError> {
1037 if !(tau0.is_finite() && tau0 > 0.0) {
1038 crate::bail_invalid_basis!(
1039 "measure-jet jets need tau0 > 0 because the retained τ coordinate is ln τ; got {tau0}"
1040 );
1041 }
1042 let mut forms = assemble_weighted_forms(
1043 centers,
1044 masses,
1045 band,
1046 order_s,
1047 alpha,
1048 tau0,
1049 10,
1050 3,
1051 &|_, eps: f64, q: f64, base: f64, out: &mut [[f64; 3]]| {
1052 let gs = -2.0 * eps.ln();
1053 let intrinsic_dim = centers.ncols() as f64;
1054 let ga = 2.0 * intrinsic_dim * eps.ln() - 2.0 * q.max(f64::MIN_POSITIVE).ln();
1055 out[0] = [base, 0.0, 0.0];
1056 out[1] = [gs * base, 0.0, 0.0];
1057 out[2] = [gs * gs * base, 0.0, 0.0];
1058 out[3] = [ga * base, 0.0, 0.0];
1059 out[4] = [ga * ga * base, 0.0, 0.0];
1060 out[5] = [gs * ga * base, 0.0, 0.0];
1061 out[6] = [0.0, 0.0, 0.0];
1062 out[7] = [0.0, 0.0, 0.0];
1063 out[8] = [0.0, 0.0, 0.0];
1064 out[9] = [0.0, 0.0, 0.0];
1065 },
1066 )?;
1067 let d2q_dalpha_dlogtau = forms.pop().expect("ten assembled forms");
1068 let d2q_ds_dlogtau = forms.pop().expect("ten assembled forms");
1069 let d2q_dlogtau2 = forms.pop().expect("ten assembled forms");
1070 let dq_dlogtau = forms.pop().expect("ten assembled forms");
1071 let d2q_ds_dalpha = forms.pop().expect("ten assembled forms");
1072 let d2q_dalpha2 = forms.pop().expect("ten assembled forms");
1073 let dq_dalpha = forms.pop().expect("ten assembled forms");
1074 let d2q_ds2 = forms.pop().expect("ten assembled forms");
1075 let dq_ds = forms.pop().expect("ten assembled forms");
1076 let q = forms.pop().expect("ten assembled forms");
1077 Ok(MeasureJetEnergyJets {
1078 q,
1079 dq_ds,
1080 d2q_ds2,
1081 dq_dalpha,
1082 d2q_dalpha2,
1083 d2q_ds_dalpha,
1084 dq_dlogtau,
1085 d2q_dlogtau2,
1086 d2q_ds_dlogtau,
1087 d2q_dalpha_dlogtau,
1088 })
1089}
1090
1091pub fn measure_jet_scale_spectrum(
1097 centers: ArrayView2<'_, f64>,
1098 masses: ArrayView1<'_, f64>,
1099 band: &MeasureJetBand,
1100 order_s: f64,
1101 alpha: f64,
1102 tau0: f64,
1103 values: ArrayView1<'_, f64>,
1104) -> Result<Vec<f64>, BasisError> {
1105 if values.len() != centers.nrows() {
1106 crate::bail_dim_basis!(
1107 "measure-jet scale spectrum needs one value per center: {} values for {} centers",
1108 values.len(),
1109 centers.nrows()
1110 );
1111 }
1112 let forms = measure_jet_energy_forms_per_scale(centers, masses, band, order_s, alpha, tau0)?;
1113 Ok(forms
1114 .iter()
1115 .map(|q_l| values.dot(&q_l.dot(&values)))
1116 .collect())
1117}
1118
1119pub fn measure_jet_energy_forms_per_scale(
1125 centers: ArrayView2<'_, f64>,
1126 masses: ArrayView1<'_, f64>,
1127 band: &MeasureJetBand,
1128 order_s: f64,
1129 alpha: f64,
1130 tau0: f64,
1131) -> Result<Vec<Array2<f64>>, BasisError> {
1132 let n_scales = band.eps.len();
1133 assemble_weighted_forms(
1134 centers,
1135 masses,
1136 band,
1137 order_s,
1138 alpha,
1139 tau0,
1140 n_scales,
1141 1,
1142 &|scale_idx, _, _, base, out: &mut [[f64; 3]]| {
1143 for (k, slot) in out.iter_mut().enumerate() {
1144 *slot = if k == scale_idx {
1145 [base, 0.0, 0.0]
1146 } else {
1147 [0.0, 0.0, 0.0]
1148 };
1149 }
1150 },
1151 )
1152}
1153
1154pub fn measure_jet_support_curve(
1162 queries: ArrayView2<'_, f64>,
1163 centers: ArrayView2<'_, f64>,
1164 masses: ArrayView1<'_, f64>,
1165 eps_band: &[f64],
1166) -> Result<Array2<f64>, BasisError> {
1167 if queries.ncols() != centers.ncols() {
1168 crate::bail_dim_basis!(
1169 "measure-jet support curve dimension mismatch: queries d={} centers d={}",
1170 queries.ncols(),
1171 centers.ncols()
1172 );
1173 }
1174 if masses.len() != centers.nrows() {
1175 crate::bail_dim_basis!(
1176 "measure-jet support curve mass/center mismatch: {} masses for {} centers",
1177 masses.len(),
1178 centers.nrows()
1179 );
1180 }
1181 if eps_band.is_empty() || eps_band.iter().any(|e| !(e.is_finite() && *e > 0.0)) {
1182 crate::bail_invalid_basis!("measure-jet support curve needs a nonempty positive band");
1183 }
1184 validate_finite_points(queries, "queries")?;
1185 validate_finite_points(centers, "centers")?;
1186 let nq = queries.nrows();
1187 let nl = eps_band.len();
1188 let d2 = pairwise_sq_dists(queries, centers);
1191 let mut out = Array2::<f64>::zeros((nq, nl));
1192 out.axis_iter_mut(Axis(0))
1193 .into_par_iter()
1194 .enumerate()
1195 .for_each(|(qi, mut row)| {
1196 let d2_row = d2.row(qi);
1197 for (li, &eps) in eps_band.iter().enumerate() {
1198 let inv_two_eps2 = 1.0 / (2.0 * eps * eps);
1199 let mut acc = 0.0_f64;
1200 for (j, &dd) in d2_row.iter().enumerate() {
1201 acc += masses[j] * (-dd * inv_two_eps2).exp();
1202 }
1203 row[li] = acc;
1204 }
1205 });
1206 Ok(out)
1207}
1208
1209pub(crate) fn measure_jet_support_means(
1210 centers: ArrayView2<'_, f64>,
1211 masses: ArrayView1<'_, f64>,
1212 eps_band: &[f64],
1213) -> Result<Vec<f64>, BasisError> {
1214 let total_mass = masses.sum();
1215 if !(total_mass.is_finite() && total_mass > 0.0) {
1216 crate::bail_invalid_basis!(
1217 "measure-jet support means need positive finite total mass; got {total_mass}"
1218 );
1219 }
1220 let support = measure_jet_support_curve(centers, centers, masses, eps_band)?;
1221 let mut means = vec![0.0_f64; eps_band.len()];
1222 for (i, row) in support.rows().into_iter().enumerate() {
1223 let mass = masses[i];
1224 for (mean, &q) in means.iter_mut().zip(row.iter()) {
1225 *mean += mass * q;
1226 }
1227 }
1228 for mean in &mut means {
1229 *mean /= total_mass;
1230 if !(*mean).is_finite() || *mean <= 0.0 {
1231 crate::bail_invalid_basis!(
1232 "measure-jet support mean must be positive and finite; got {mean}"
1233 );
1234 }
1235 }
1236 Ok(means)
1237}
1238
1239pub fn measure_jet_design_matrix(
1241 data: ArrayView2<'_, f64>,
1242 centers: ArrayView2<'_, f64>,
1243 length_scale: f64,
1244) -> Result<Array2<f64>, BasisError> {
1245 if data.ncols() != centers.ncols() {
1246 crate::bail_dim_basis!(
1247 "measure-jet design dimension mismatch: data d={} centers d={}",
1248 data.ncols(),
1249 centers.ncols()
1250 );
1251 }
1252 if !(length_scale.is_finite() && length_scale > 0.0) {
1253 crate::bail_invalid_basis!(
1254 "measure-jet design needs a positive finite length_scale; got {length_scale}"
1255 );
1256 }
1257 validate_finite_points(data, "data")?;
1258 validate_finite_points(centers, "centers")?;
1259 let inv_two_l2 = 1.0 / (2.0 * length_scale * length_scale);
1260 let mut out = pairwise_sq_dists(data, centers);
1263 out.axis_iter_mut(Axis(0))
1264 .into_par_iter()
1265 .for_each(|mut row| {
1266 row.mapv_inplace(|d2| (-d2 * inv_two_l2).exp());
1267 });
1268 Ok(out)
1269}
1270
1271fn measure_jet_design_log_length_jets(
1274 data: ArrayView2<'_, f64>,
1275 centers: ArrayView2<'_, f64>,
1276 length_scale: f64,
1277) -> Result<(Array2<f64>, Array2<f64>), BasisError> {
1278 let kernel = measure_jet_design_matrix(data, centers, length_scale)?;
1279 let squared_distances = pairwise_sq_dists(data, centers);
1280 let inv_l2 = 1.0 / (length_scale * length_scale);
1281 let mut first = kernel.clone();
1282 let mut second = kernel;
1283 for ((first_value, second_value), &distance_squared) in first
1284 .iter_mut()
1285 .zip(second.iter_mut())
1286 .zip(squared_distances.iter())
1287 {
1288 let a = distance_squared * inv_l2;
1289 let kernel_value = *first_value;
1290 *first_value = kernel_value * a;
1291 *second_value = kernel_value * (a * a - 2.0 * a);
1292 }
1293 Ok((first, second))
1294}
1295
1296pub fn measure_jet_affine_head_transform(
1327 centers: ArrayView2<'_, f64>,
1328 masses: ArrayView1<'_, f64>,
1329) -> Array2<f64> {
1330 let m = centers.nrows();
1331 let d = centers.ncols();
1332 let total_mass = masses.sum();
1333 let mdot = |u: &Array1<f64>, v: &Array1<f64>| -> f64 {
1335 let mut acc = 0.0;
1336 for i in 0..m {
1337 acc += masses[i] * u[i] * v[i];
1338 }
1339 acc
1340 };
1341 let cols: Vec<Array1<f64>> = (0..d)
1345 .map(|k| {
1346 let col = centers.column(k).to_owned();
1347 let mean = if total_mass > 0.0 {
1348 mdot(&col, &Array1::ones(m)) / total_mass
1349 } else {
1350 0.0
1351 };
1352 col.mapv(|x| x - mean)
1353 })
1354 .collect();
1355 let max_norm = cols
1357 .iter()
1358 .fold(0.0_f64, |acc, c| acc.max(mdot(c, c).sqrt()));
1359 let drop_below =
1360 (MEASURE_JET_PSEUDOINVERSE_RTOL * (d.max(1) as f64) * max_norm).max(f64::MIN_POSITIVE);
1361 let mut q_cols: Vec<Array1<f64>> = Vec::new();
1365 let mut t_cols: Vec<Array1<f64>> = Vec::new();
1366 for k in 0..d {
1367 let mut v = cols[k].clone();
1368 let mut t = Array1::<f64>::zeros(d);
1369 t[k] = 1.0;
1370 for (q, tq) in q_cols.iter().zip(t_cols.iter()) {
1371 let proj = mdot(q, &v);
1372 v.scaled_add(-proj, q);
1373 t.scaled_add(-proj, tq);
1374 }
1375 let norm = mdot(&v, &v).sqrt();
1376 if norm > drop_below {
1377 v.mapv_inplace(|x| x / norm);
1378 t.mapv_inplace(|x| x / norm);
1379 q_cols.push(v);
1380 t_cols.push(t);
1381 }
1382 }
1383 let head_rank = t_cols.len();
1384 let mut t_mat = Array2::<f64>::zeros((d, head_rank));
1385 for (r, t) in t_cols.into_iter().enumerate() {
1386 t_mat.column_mut(r).assign(&t);
1387 }
1388 t_mat
1389}
1390
1391pub fn realized_measure_jet_length_scale(
1396 centers: ArrayView2<'_, f64>,
1397 spec_length_scale: f64,
1398) -> Result<f64, BasisError> {
1399 if spec_length_scale.is_finite() && spec_length_scale > 0.0 {
1400 return Ok(spec_length_scale);
1401 }
1402 if spec_length_scale != 0.0 {
1403 crate::bail_invalid_basis!(
1404 "measure-jet length_scale must be positive (or 0.0 for auto); got {spec_length_scale}"
1405 );
1406 }
1407 let dist2 = pairwise_sq_dists(centers, centers);
1408 let spacing = median_nearest_center_spacing(&dist2)?;
1409 Ok(MEASURE_JET_AUTO_LENGTH_SCALE_FACTOR * spacing)
1410}
1411
1412pub(crate) struct RealizedMeasureJetGeometry {
1417 pub(crate) centers: Array2<f64>,
1418 pub(crate) masses: Array1<f64>,
1419 pub(crate) eps_band: Vec<f64>,
1420 pub(crate) log_step: f64,
1421 pub(crate) length_scale: f64,
1422 pub(crate) order_s_eval: f64,
1426 pub(crate) per_level: bool,
1428 pub(crate) z: Array2<f64>,
1429 pub(crate) coefficient_gauge: gam_problem::Gauge,
1430 pub(crate) kz: Array2<f64>,
1431 pub(crate) head_transform: Array2<f64>,
1437}
1438
1439pub(crate) fn realize_measure_jet_geometry(
1440 data: ArrayView2<'_, f64>,
1441 spec: &MeasureJetBasisSpec,
1442) -> Result<RealizedMeasureJetGeometry, BasisError> {
1443 if data.ncols() == 0 {
1444 crate::bail_invalid_basis!("measure-jet smooth needs at least one feature column");
1445 }
1446 validate_finite_points(data, "data")?;
1447 let seed_centers = select_centers_by_strategy(data, &spec.center_strategy)?;
1448 let m = seed_centers.nrows();
1449 if m < 3 {
1450 return Err(BasisError::InsufficientColumnsForConstraint { found: m });
1451 }
1452 let order_s = if spec.order_s == 0.0 {
1453 MEASURE_JET_DEFAULT_ORDER_S
1454 } else {
1455 spec.order_s
1456 };
1457 let (centers, masses, eps_band, log_step) = match &spec.frozen_quadrature {
1464 Some(frozen) => {
1465 if frozen.masses.len() != m {
1466 crate::bail_dim_basis!(
1467 "frozen measure-jet quadrature mismatch: {} masses for {} centers",
1468 frozen.masses.len(),
1469 m
1470 );
1471 }
1472 if frozen.eps_band.is_empty() {
1473 crate::bail_invalid_basis!("frozen measure-jet quadrature has an empty band");
1474 }
1475 let log_step = if frozen.eps_band.len() >= 2 {
1476 (frozen.eps_band[1] / frozen.eps_band[0]).ln()
1477 } else {
1478 std::f64::consts::LN_2
1479 };
1480 (
1481 seed_centers,
1482 frozen.masses.clone(),
1483 frozen.eps_band.clone(),
1484 log_step,
1485 )
1486 }
1487 None => {
1488 let (nodes, masses) = measure_jet_quadrature_nodes(data, seed_centers.view())?;
1489 let band = measure_jet_band(nodes.view(), spec.num_scales)?;
1490 (nodes, masses, band.eps, band.log_step)
1491 }
1492 };
1493 let length_scale = realized_measure_jet_length_scale(centers.view(), spec.length_scale)?;
1494 let head_transform = if spec.multiscale {
1505 Array2::<f64>::zeros((centers.ncols(), 0))
1506 } else {
1507 measure_jet_affine_head_transform(centers.view(), masses.view())
1508 };
1509 let head_rank = head_transform.ncols();
1510 let m_aug = m + head_rank;
1511 let k_cc = measure_jet_design_matrix(centers.view(), centers.view(), length_scale)?;
1512 let head_cc = centers.dot(&head_transform);
1513 let (z, coefficient_gauge) = match &spec.identifiability {
1528 MeasureJetIdentifiability::FrozenTransform { transform } => {
1529 if transform.nrows() != m_aug {
1530 crate::bail_dim_basis!(
1531 "frozen measure-jet identifiability transform mismatch: {} representers + {} head columns but transform has {} rows",
1532 m,
1533 head_rank,
1534 transform.nrows()
1535 );
1536 }
1537 (
1538 transform.clone(),
1539 gam_problem::Gauge::from_block_transforms(&[transform.clone()]),
1540 )
1541 }
1542 MeasureJetIdentifiability::CenterSumToZero => {
1543 let z_rbf = if head_rank > 0 {
1544 let affine = measure_jet_affine_value_basis(centers.view(), masses.view());
1545 let mut weighted_affine = affine.clone();
1546 for (i, mut row) in weighted_affine.outer_iter_mut().enumerate() {
1547 row.mapv_inplace(|v| v * masses[i]);
1548 }
1549 let constraint_cross = k_cc.t().dot(&weighted_affine);
1553 rrqr_nullspace_basis(&constraint_cross, default_rrqr_rank_alpha())
1554 .map_err(BasisError::LinalgError)?
1555 .0
1556 } else {
1557 let u = householder_sum_to_zero_u(m);
1558 householder_sum_to_zero_z(&u)
1559 };
1560 let rbf_rank = z_rbf.ncols();
1561 let mut z_block = Array2::<f64>::zeros((m_aug, rbf_rank + head_rank));
1562 z_block
1563 .slice_mut(ndarray::s![..m, ..rbf_rank])
1564 .assign(&z_rbf);
1565 for r in 0..head_rank {
1566 z_block[(m + r, rbf_rank + r)] = 1.0;
1567 }
1568 (
1569 z_block.clone(),
1570 gam_problem::Gauge::from_block_transforms(&[z_block]),
1571 )
1572 }
1573 };
1574 let mut k_aug = Array2::<f64>::zeros((m, m_aug));
1579 k_aug.slice_mut(ndarray::s![.., ..m]).assign(&k_cc);
1580 if head_rank > 0 {
1581 k_aug.slice_mut(ndarray::s![.., m..]).assign(&head_cc);
1582 }
1583 let kz = coefficient_gauge.restrict_design(&k_aug);
1584 Ok(RealizedMeasureJetGeometry {
1585 centers,
1586 masses,
1587 eps_band,
1588 log_step,
1589 length_scale,
1590 order_s_eval: order_s,
1591 per_level: spec.multiscale,
1596 z,
1597 coefficient_gauge,
1598 kz,
1599 head_transform,
1600 })
1601}
1602
1603pub fn measure_jet_input_noise_scale(
1626 data: ArrayView2<'_, f64>,
1627 centers: ArrayView2<'_, f64>,
1628) -> Result<Option<f64>, BasisError> {
1629 let d = data.ncols();
1630 let m = centers.nrows();
1631 if d == 0 || m == 0 || data.nrows() == 0 {
1632 return Ok(None);
1633 }
1634 if centers.ncols() != d {
1635 crate::bail_dim_basis!(
1636 "measure-jet input-noise estimate: data d={d} disagrees with centers d={}",
1637 centers.ncols()
1638 );
1639 }
1640 validate_finite_points(data, "data")?;
1641 validate_finite_points(centers, "centers")?;
1642 let sq = pairwise_sq_dists(data, centers);
1645 let mut members: Vec<Vec<usize>> = vec![Vec::new(); m];
1646 for (j, row) in sq.axis_iter(Axis(0)).enumerate() {
1647 let mut best = 0usize;
1648 let mut best_d = f64::INFINITY;
1649 for (i, &dij) in row.iter().enumerate() {
1650 if dij < best_d {
1651 best_d = dij;
1652 best = i;
1653 }
1654 }
1655 members[best].push(j);
1656 }
1657 let mut weighted_sum = 0.0_f64;
1658 let mut weight = 0.0_f64;
1659 for cell in &members {
1660 let n_i = cell.len();
1661 if n_i < d + 1 {
1664 continue;
1665 }
1666 let mut mean = Array1::<f64>::zeros(d);
1668 for &j in cell {
1669 mean += &data.row(j);
1670 }
1671 mean /= n_i as f64;
1672 let mut cov = Array2::<f64>::zeros((d, d));
1673 for &j in cell {
1674 let mut centered = data.row(j).to_owned();
1675 centered -= &mean;
1676 for a in 0..d {
1677 for b in 0..d {
1678 cov[(a, b)] += centered[a] * centered[b];
1679 }
1680 }
1681 }
1682 cov /= n_i as f64;
1683 let cov_sym = (&cov + &cov.t()) * 0.5;
1686 let (evals, _) = cov_sym.eigh(Side::Lower).map_err(|e| {
1687 BasisError::InvalidInput(format!(
1688 "measure-jet input-noise estimate: local covariance eigendecomposition failed: {e}"
1689 ))
1690 })?;
1691 let smallest = evals
1692 .iter()
1693 .copied()
1694 .fold(f64::INFINITY, |acc, v| acc.min(v))
1695 .max(0.0);
1696 if smallest.is_finite() {
1697 weighted_sum += n_i as f64 * smallest;
1698 weight += n_i as f64;
1699 }
1700 }
1701 if weight <= 0.0 {
1702 return Ok(None);
1703 }
1704 let sigma2 = weighted_sum / weight;
1705 if !(sigma2.is_finite() && sigma2 > 0.0) {
1706 return Ok(None);
1707 }
1708 Ok(Some(sigma2.sqrt()))
1709}
1710
1711pub fn measure_jet_multiscale_mode(spec: &MeasureJetBasisSpec) -> bool {
1719 spec.multiscale
1720}
1721
1722pub fn build_measure_jet_basis(
1730 data: ArrayView2<'_, f64>,
1731 spec: &MeasureJetBasisSpec,
1732) -> Result<BasisBuildResult, BasisError> {
1733 let RealizedMeasureJetGeometry {
1734 centers,
1735 masses,
1736 eps_band,
1737 log_step,
1738 length_scale,
1739 order_s_eval: order_s,
1740 per_level,
1741 z,
1742 coefficient_gauge,
1743 kz,
1744 head_transform,
1745 } = realize_measure_jet_geometry(data, spec)?;
1746 let band = MeasureJetBand {
1747 eps: eps_band.clone(),
1748 log_step,
1749 };
1750 let m = centers.nrows();
1751 let head_rank = head_transform.ncols();
1752 let m_aug = m + head_rank;
1753 let kernel_design = measure_jet_design_matrix(data, centers.view(), length_scale)?;
1758 let mut raw_design = Array2::<f64>::zeros((data.nrows(), m_aug));
1759 raw_design
1760 .slice_mut(ndarray::s![.., ..m])
1761 .assign(&kernel_design);
1762 if head_rank > 0 {
1763 let head_design = data.dot(&head_transform);
1764 raw_design
1765 .slice_mut(ndarray::s![.., m..])
1766 .assign(&head_design);
1767 }
1768 let constrained_design = coefficient_gauge.restrict_design(&raw_design);
1769 let design = gam_linalg::matrix::DesignMatrix::Dense(
1770 gam_linalg::matrix::DenseDesignMatrix::from(constrained_design),
1771 );
1772 let support_means = measure_jet_support_means(centers.view(), masses.view(), &eps_band)?;
1773 let mut candidates = Vec::new();
1786 let mut penalty_normalization_scales = Vec::new();
1787 let mut raw_penalty_normalization_scales = Vec::new();
1788 let mut fused_penalty_normalization_scale = None;
1789 if per_level {
1790 let forms = measure_jet_energy_forms_per_scale(
1791 centers.view(),
1792 masses.view(),
1793 &band,
1794 order_s,
1795 spec.alpha,
1796 spec.tau0,
1797 )?;
1798 for (level, q_l) in forms.into_iter().enumerate() {
1799 let s_l = kz.t().dot(&q_l).dot(&kz);
1800 let (s_norm, c_l) = normalize_penalty(&((&s_l + &s_l.t()) * 0.5));
1801 let intrinsic_dim = centers.ncols() as f64;
1802 let eta = 2.0 * order_s + intrinsic_dim * (2.0 - 2.0 * spec.alpha);
1803 let scale_weight = log_step * eps_band[level].powf(-eta);
1804 penalty_normalization_scales.push(c_l);
1805 raw_penalty_normalization_scales.push(c_l / scale_weight);
1806 candidates.push(PenaltyCandidate {
1807 matrix: ConstructiveQuadratic::try_from_dense_psd(
1808 s_norm,
1809 "measure-jet scale penalty",
1810 )?,
1811 source: PenaltySource::Other(format!("measure_jet_scale_{level}")),
1812 normalization_scale: c_l,
1813 kronecker_factors: None,
1814 op: None,
1815 });
1816 }
1817 } else {
1818 let q_form = measure_jet_energy_form(
1819 centers.view(),
1820 masses.view(),
1821 &band,
1822 order_s,
1823 spec.alpha,
1824 spec.tau0,
1825 )?;
1826 let penalty = pullback_center_form(&kz, &q_form);
1831 let (penalty_norm, c_primary) = normalize_penalty(&penalty);
1832 fused_penalty_normalization_scale = Some(c_primary);
1833 candidates.push(PenaltyCandidate {
1834 matrix: ConstructiveQuadratic::try_from_dense_psd(
1835 penalty_norm,
1836 "measure-jet primary penalty",
1837 )?,
1838 source: PenaltySource::Primary,
1839 normalization_scale: c_primary,
1840 kronecker_factors: None,
1841 op: None,
1842 });
1843 }
1844 if spec.double_penalty {
1850 let null_penalty = affine_function_nullspace_penalty(&kz, centers.view(), masses.view())?;
1851 let (null_penalty_norm, c_null) = normalize_penalty(&null_penalty);
1852 candidates.push(PenaltyCandidate {
1853 matrix: ConstructiveQuadratic::try_from_dense_psd(
1854 null_penalty_norm,
1855 "measure-jet null-function penalty",
1856 )?,
1857 source: PenaltySource::DoublePenaltyNullspace,
1858 normalization_scale: c_null,
1859 kronecker_factors: None,
1860 op: None,
1861 });
1862 }
1863 let filtered = filter_penalty_candidates(candidates)?;
1864 let sigma_coord = measure_jet_input_noise_scale(data, centers.view())?;
1867 Ok(BasisBuildResult {
1868 design,
1869 affine_offset: None,
1870 active_penalties: filtered.active,
1871 dropped_penalties: filtered.dropped,
1872 metadata: BasisMetadata::MeasureJet {
1873 centers,
1874 input_scale: crate::IsotropicScale::ONE,
1875 length_scale,
1876 eps_band,
1877 order_s: spec.order_s,
1882 alpha: spec.alpha,
1883 tau0: spec.tau0,
1884 masses,
1885 support_means,
1886 penalty_normalization_scales,
1887 raw_penalty_normalization_scales,
1888 fused_penalty_normalization_scale,
1889 constraint_transform: Some(z),
1890 sigma_coord,
1893 },
1894 kronecker_factored: None,
1895 joint_null_rotation: None,
1896 })
1897}
1898
1899pub fn build_measure_jet_basis_psi_derivatives(
1922 data: ArrayView2<'_, f64>,
1923 spec: &MeasureJetBasisSpec,
1924) -> Result<AnisoBasisPsiDerivatives, BasisError> {
1925 if !(spec.tau0.is_finite() && spec.tau0 > 0.0) {
1926 crate::bail_invalid_basis!(
1927 "measure-jet ψ derivatives need tau0 > 0 because the retained τ coordinate is ln τ; got {}",
1928 spec.tau0
1929 );
1930 }
1931 let geom = realize_measure_jet_geometry(data, spec)?;
1932 let band = MeasureJetBand {
1933 eps: geom.eps_band.clone(),
1934 log_step: geom.log_step,
1935 };
1936 let n = data.nrows();
1937 let p = geom.kz.ncols();
1938 let m = geom.centers.nrows();
1939 let m_aug = m + geom.head_transform.ncols();
1940
1941 struct LengthScaleJets {
1942 evaluation_first: Array2<f64>,
1943 evaluation_second: Array2<f64>,
1944 design_first: Array2<f64>,
1945 design_second: Array2<f64>,
1946 }
1947
1948 let length_scale_jets = if spec.learn_length_scale {
1955 let (dk_data, d2k_data) =
1956 measure_jet_design_log_length_jets(data, geom.centers.view(), geom.length_scale)?;
1957 let mut dk_data_aug = Array2::<f64>::zeros((n, m_aug));
1958 let mut d2k_data_aug = Array2::<f64>::zeros((n, m_aug));
1959 dk_data_aug.slice_mut(ndarray::s![.., ..m]).assign(&dk_data);
1960 d2k_data_aug
1961 .slice_mut(ndarray::s![.., ..m])
1962 .assign(&d2k_data);
1963
1964 let (dk_centers, d2k_centers) = measure_jet_design_log_length_jets(
1965 geom.centers.view(),
1966 geom.centers.view(),
1967 geom.length_scale,
1968 )?;
1969 let mut dk_centers_aug = Array2::<f64>::zeros((m, m_aug));
1970 let mut d2k_centers_aug = Array2::<f64>::zeros((m, m_aug));
1971 dk_centers_aug
1972 .slice_mut(ndarray::s![.., ..m])
1973 .assign(&dk_centers);
1974 d2k_centers_aug
1975 .slice_mut(ndarray::s![.., ..m])
1976 .assign(&d2k_centers);
1977
1978 Some(LengthScaleJets {
1979 evaluation_first: geom.coefficient_gauge.restrict_design(&dk_centers_aug),
1980 evaluation_second: geom.coefficient_gauge.restrict_design(&d2k_centers_aug),
1981 design_first: geom.coefficient_gauge.restrict_design(&dk_data_aug),
1982 design_second: geom.coefficient_gauge.restrict_design(&d2k_data_aug),
1983 })
1984 } else {
1985 None
1986 };
1987
1988 let coord_offset = usize::from(length_scale_jets.is_some());
1989 let n_coords = coord_offset + if geom.per_level { 2 } else { 0 };
1990 let pairs: Vec<(usize, usize)> = (0..n_coords)
1991 .flat_map(|a| ((a + 1)..n_coords).map(move |b| (a, b)))
1992 .collect();
1993 let zero_p = || Array2::<f64>::zeros((p, p));
1994
1995 struct RawPenaltyJets {
1996 value: Array2<f64>,
1997 first: Vec<Array2<f64>>,
1998 second_diag: Vec<Array2<f64>>,
1999 cross: Vec<Array2<f64>>,
2000 }
2001
2002 let sandwich = |form: &Array2<f64>| pullback_center_form(&geom.kz, form);
2003 let length_diag = |form: &Array2<f64>| {
2004 let jets = length_scale_jets
2005 .as_ref()
2006 .expect("length-scale form jets require an enrolled length coordinate");
2007 pullback_center_form_log_length_jets(
2008 &geom.kz,
2009 &jets.evaluation_first,
2010 &jets.evaluation_second,
2011 form,
2012 )
2013 };
2014 let length_cross = |form_first: &Array2<f64>| {
2015 let jets = length_scale_jets
2016 .as_ref()
2017 .expect("length-scale cross jets require an enrolled length coordinate");
2018 pullback_center_form_log_length_cross(&geom.kz, &jets.evaluation_first, form_first)
2019 };
2020
2021 let mut raw: Vec<RawPenaltyJets> = if geom.per_level {
2027 let l_count = band.eps.len();
2028 let forms = assemble_weighted_forms(
2031 geom.centers.view(),
2032 geom.masses.view(),
2033 &band,
2034 geom.order_s_eval,
2035 spec.alpha,
2036 spec.tau0,
2037 6 * l_count,
2038 3,
2039 &|scale_idx, eps: f64, q: f64, base: f64, out: &mut [[f64; 3]]| {
2040 for slot in out.iter_mut() {
2041 *slot = [0.0, 0.0, 0.0];
2042 }
2043 let intrinsic_dim = geom.centers.ncols() as f64;
2044 let ga = 2.0 * intrinsic_dim * eps.ln() - 2.0 * q.max(f64::MIN_POSITIVE).ln();
2045 let k0 = 6 * scale_idx;
2046 out[k0] = [base, 0.0, 0.0];
2047 out[k0 + 1] = [ga * base, 0.0, 0.0];
2048 out[k0 + 2] = [ga * ga * base, 0.0, 0.0];
2049 out[k0 + 3] = [0.0, 0.0, 0.0];
2050 out[k0 + 4] = [0.0, 0.0, 0.0];
2051 out[k0 + 5] = [0.0, 0.0, 0.0];
2052 },
2053 )?;
2054 let alpha_coord = coord_offset;
2055 let tau_coord = coord_offset + 1;
2056 let mut raw = Vec::with_capacity(l_count + usize::from(spec.double_penalty));
2057 for level in 0..l_count {
2058 let chunk = &forms[6 * level..6 * level + 6];
2059 let mut first: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2060 let mut second_diag: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2061 first[alpha_coord] = sandwich(&chunk[1]);
2062 first[tau_coord] = sandwich(&chunk[3]);
2063 second_diag[alpha_coord] = sandwich(&chunk[2]);
2064 second_diag[tau_coord] = sandwich(&chunk[4]);
2065 if coord_offset == 1 {
2066 let (ell_first, ell_second) = length_diag(&chunk[0]);
2067 first[0] = ell_first;
2068 second_diag[0] = ell_second;
2069 }
2070 let mut cross: Vec<Array2<f64>> = (0..pairs.len()).map(|_| zero_p()).collect();
2071 for (pair_idx, &(a, b)) in pairs.iter().enumerate() {
2072 cross[pair_idx] = if coord_offset == 1 && a == 0 && b == alpha_coord {
2073 length_cross(&chunk[1])
2074 } else if coord_offset == 1 && a == 0 && b == tau_coord {
2075 length_cross(&chunk[3])
2076 } else if a == alpha_coord && b == tau_coord {
2077 sandwich(&chunk[5])
2078 } else {
2079 zero_p()
2080 };
2081 }
2082 raw.push(RawPenaltyJets {
2083 value: sandwich(&chunk[0]),
2084 first,
2085 second_diag,
2086 cross,
2087 });
2088 }
2089 raw
2090 } else {
2091 let q_form = measure_jet_energy_form(
2095 geom.centers.view(),
2096 geom.masses.view(),
2097 &band,
2098 geom.order_s_eval,
2099 spec.alpha,
2100 spec.tau0,
2101 )?;
2102 let mut first: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2103 let mut second_diag: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2104 if coord_offset == 1 {
2105 let (ell_first, ell_second) = length_diag(&q_form);
2106 first[0] = ell_first;
2107 second_diag[0] = ell_second;
2108 }
2109 vec![RawPenaltyJets {
2110 value: sandwich(&q_form),
2111 first,
2112 second_diag,
2113 cross: Vec::new(),
2114 }]
2115 };
2116
2117 if spec.double_penalty {
2118 let null_form = affine_function_nullspace_form(geom.centers.view(), geom.masses.view())?;
2119 let mut first: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2120 let mut second_diag: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2121 if coord_offset == 1 {
2122 let (ell_first, ell_second) = length_diag(&null_form);
2123 first[0] = ell_first;
2124 second_diag[0] = ell_second;
2125 }
2126 raw.push(RawPenaltyJets {
2127 value: sandwich(&null_form),
2128 first,
2129 second_diag,
2130 cross: (0..pairs.len()).map(|_| zero_p()).collect(),
2133 });
2134 }
2135
2136 let n_cands = raw.len();
2137 let mut penalties_first: Vec<Vec<Array2<f64>>> =
2138 (0..n_coords).map(|_| Vec::with_capacity(n_cands)).collect();
2139 let mut penalties_second_diag: Vec<Vec<Array2<f64>>> =
2140 (0..n_coords).map(|_| Vec::with_capacity(n_cands)).collect();
2141 let mut crosses: Vec<Vec<Array2<f64>>> = (0..pairs.len()).map(|_| Vec::new()).collect();
2145 for candidate in &raw {
2146 let s_raw = &candidate.value;
2147 let fro = trace_of_product(s_raw, s_raw).sqrt();
2156 let c = if fro.is_finite() && fro > 1e-12 {
2157 fro
2158 } else {
2159 1.0
2160 };
2161 for coord in 0..n_coords {
2162 let (_, s_first, s_second, _) = normalize_penaltywith_psi_derivatives(
2163 s_raw,
2164 &candidate.first[coord],
2165 &candidate.second_diag[coord],
2166 );
2167 penalties_first[coord].push(s_first);
2168 penalties_second_diag[coord].push(s_second);
2169 }
2170 for (pair_idx, &(a, b)) in pairs.iter().enumerate() {
2171 let cross_raw_mat = normalize_penalty_cross_psi_derivative(
2172 s_raw,
2173 &candidate.first[a],
2174 &candidate.first[b],
2175 &candidate.cross[pair_idx],
2176 c,
2177 );
2178 crosses[pair_idx].push(cross_raw_mat);
2179 }
2180 }
2181
2182 let pair_index: Vec<((usize, usize), Vec<Array2<f64>>)> =
2183 pairs.iter().copied().zip(crosses.into_iter()).collect();
2184 let provider = AnisoPenaltyCrossProvider::new(move |a, b| {
2185 pair_index
2186 .iter()
2187 .find(|((pa, pb), _)| (*pa, *pb) == (a, b) || (*pa, *pb) == (b, a))
2188 .map(|(_, mats)| mats.clone())
2189 .ok_or_else(|| {
2190 BasisError::InvalidInput(format!(
2191 "measure-jet ψ cross derivative requested for unknown pair ({a}, {b})"
2192 ))
2193 })
2194 });
2195 let mut design_first: Vec<Array2<f64>> = (0..n_coords)
2196 .map(|_| Array2::<f64>::zeros((n, p)))
2197 .collect();
2198 let mut design_second_diag: Vec<Array2<f64>> = (0..n_coords)
2199 .map(|_| Array2::<f64>::zeros((n, p)))
2200 .collect();
2201 if let Some(jets) = &length_scale_jets {
2202 design_first[0] = jets.design_first.clone();
2203 design_second_diag[0] = jets.design_second.clone();
2204 }
2205 Ok(AnisoBasisPsiDerivatives {
2206 design_first,
2207 design_second_diag,
2208 design_second_cross: Vec::new(),
2209 design_second_cross_pairs: Vec::new(),
2210 penalties_first,
2211 penalties_second_diag,
2212 penalties_cross_pairs: pairs,
2213 penalties_cross_provider: Some(provider),
2214 implicit_operator: None,
2215 })
2216}
2217
2218#[cfg(test)]
2219mod tests {
2220 use super::*;
2221
2222 fn lcg_normal(state: &mut u64) -> f64 {
2225 let mut next = || {
2226 *state = state
2227 .wrapping_mul(6364136223846793005)
2228 .wrapping_add(1442695040888963407);
2229 (((*state >> 11) as f64) + 0.5) / (1u64 << 53) as f64
2231 };
2232 let u1 = next();
2233 let u2 = next();
2234 (-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
2235 }
2236
2237 #[test]
2243 pub(crate) fn input_noise_scale_recovers_known_perpendicular_sigma() {
2244 let tang = [1.0 / 5f64.sqrt(), 2.0 / 5f64.sqrt()];
2246 let perp = [2.0 / 5f64.sqrt(), -1.0 / 5f64.sqrt()];
2247 let sigma = 0.05_f64;
2248 let n = 600usize;
2249 let mut state = 0x1234_5678_9abc_def0u64;
2250 let mut data = Array2::<f64>::zeros((n, 2));
2251 for j in 0..n {
2252 let t = 3.0 * (j as f64) / (n as f64 - 1.0);
2254 let noise = sigma * lcg_normal(&mut state);
2255 for a in 0..2 {
2256 data[(j, a)] = t * tang[a] + noise * perp[a];
2257 }
2258 }
2259 let n_centers = 8usize;
2262 let mut centers = Array2::<f64>::zeros((n_centers, 2));
2263 for i in 0..n_centers {
2264 let t = 3.0 * (i as f64 + 0.5) / (n_centers as f64);
2265 for a in 0..2 {
2266 centers[(i, a)] = t * tang[a];
2267 }
2268 }
2269 let est = measure_jet_input_noise_scale(data.view(), centers.view())
2270 .expect("estimate ok")
2271 .expect("noise scale present");
2272 assert!(
2275 (est - sigma).abs() <= 0.4 * sigma,
2276 "estimated σ_coord {est} far from true {sigma}"
2277 );
2278 }
2279
2280 #[test]
2283 pub(crate) fn input_noise_scale_none_when_cells_too_small() {
2284 let data = array![[0.0, 0.0], [1.0, 2.0], [2.0, 4.0]];
2285 let centers = array![[0.0, 0.0], [1.0, 2.0], [2.0, 4.0]];
2286 assert!(
2288 measure_jet_input_noise_scale(data.view(), centers.view())
2289 .expect("estimate ok")
2290 .is_none()
2291 );
2292 }
2293
2294 pub(crate) fn two_cluster_centers() -> (ndarray::Array2<f64>, ndarray::Array1<f64>) {
2295 let centers = array![
2296 [0.00, 0.00],
2297 [0.31, 0.05],
2298 [0.58, -0.07],
2299 [0.93, 0.11],
2300 [1.22, 0.02],
2301 [1.49, -0.04],
2302 [3.10, 2.00],
2303 [3.42, 2.13],
2304 [3.71, 1.91],
2305 [4.05, 2.07],
2306 [4.33, 1.96],
2307 [4.61, 2.12],
2308 ];
2309 let m = centers.nrows();
2310 let masses = ndarray::Array1::<f64>::from_elem(m, 1.0 / m as f64);
2311 (centers, masses)
2312 }
2313 use ndarray::array;
2314
2315 pub(crate) fn band_for(centers: &Array2<f64>) -> MeasureJetBand {
2316 measure_jet_band(centers.view(), 0).expect("band")
2317 }
2318
2319 #[test]
2322 pub(crate) fn energy_form_annihilates_constants_exactly() {
2323 let (centers, masses) = two_cluster_centers();
2324 let band = band_for(¢ers);
2325 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2326 .expect("energy form");
2327 let m = q.nrows();
2328 let ones = Array1::<f64>::ones(m);
2329 let qv = q.dot(&ones);
2330 let scale = q.iter().fold(0.0_f64, |acc, v| acc.max(v.abs()));
2331 assert!(scale > 0.0, "energy form is identically zero");
2332 for (i, v) in qv.iter().enumerate() {
2333 assert!(
2334 v.abs() <= 1e-12 * scale,
2335 "Q·1 leak at row {i}: {v:.3e} vs scale {scale:.3e}"
2336 );
2337 }
2338 let vqv = ones.dot(&qv);
2339 assert!(
2340 vqv.abs() <= 1e-12 * scale,
2341 "constant carries energy: 1ᵀQ1 = {vqv:.3e}"
2342 );
2343 }
2344
2345 #[test]
2348 pub(crate) fn energy_form_annihilates_affine_at_default_tau() {
2349 let (centers, masses) = two_cluster_centers();
2350 let band = band_for(¢ers);
2351 let m = centers.nrows();
2352 let mut affine = Array1::<f64>::zeros(m);
2354 let mut rough = Array1::<f64>::zeros(m);
2355 for i in 0..m {
2356 affine[i] = 0.7 + 1.3 * centers[(i, 0)] - 0.4 * centers[(i, 1)];
2357 rough[i] = if i % 2 == 0 { 1.0 } else { -1.0 };
2358 }
2359 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2360 .expect("energy form");
2361 let e_affine = affine.dot(&q.dot(&affine));
2362 let e_rough = rough.dot(&q.dot(&rough));
2363 assert!(e_rough > 0.0, "rough vector must pay energy");
2364 assert!(
2365 e_affine.abs() <= 1e-12 * e_rough,
2366 "default affine energy {e_affine:.3e} vs rough {e_rough:.3e}"
2367 );
2368 }
2369
2370 #[test]
2372 pub(crate) fn energy_form_is_psd() {
2373 let (centers, masses) = two_cluster_centers();
2374 let band = band_for(¢ers);
2375 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2376 .expect("energy form");
2377 let m = q.nrows();
2378 for trial in 0..5usize {
2379 let v = Array1::<f64>::from_shape_fn(m, |i| {
2380 ((i * 7 + trial * 13) % 11) as f64 / 11.0 - 0.5
2381 });
2382 let e = v.dot(&q.dot(&v));
2383 assert!(e >= -1e-10, "vᵀQv = {e:.3e} < 0 on trial {trial}");
2384 }
2385 }
2386
2387 #[test]
2390 pub(crate) fn rough_vector_pays_more_than_smooth() {
2391 let m = 24usize;
2392 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2393 let t = i as f64 / (m as f64 - 1.0);
2394 if k == 0 {
2395 t * 4.0
2396 } else {
2397 0.3 * (t * 4.0).sin()
2398 }
2399 });
2400 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2401 let band = band_for(¢ers);
2402 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2403 .expect("energy form");
2404 let slow = Array1::<f64>::from_shape_fn(m, |i| (i as f64 / (m as f64 - 1.0)).powi(2));
2405 let fast = Array1::<f64>::from_shape_fn(m, |i| if i % 2 == 0 { 0.5 } else { -0.5 });
2406 let e_slow = slow.dot(&q.dot(&slow));
2407 let e_fast = fast.dot(&q.dot(&fast));
2408 assert!(
2409 e_fast > 10.0 * e_slow,
2410 "alternating values must pay >> a slow trend: fast {e_fast:.3e} vs slow {e_slow:.3e}"
2411 );
2412 }
2413
2414 #[test]
2419 pub(crate) fn energy_jets_match_finite_differences() {
2420 let (centers, masses) = two_cluster_centers();
2421 let band = band_for(¢ers);
2422 let (s0, a0, tau) = (1.3, 0.8, 1e-3);
2423 let jets =
2424 measure_jet_energy_form_with_jets(centers.view(), masses.view(), &band, s0, a0, tau)
2425 .expect("jets");
2426 let q_at = |s: f64, a: f64| {
2427 measure_jet_energy_form(centers.view(), masses.view(), &band, s, a, tau)
2428 .expect("energy form")
2429 };
2430 let q_plain = q_at(s0, a0);
2432 for (a, b) in jets.q.iter().zip(q_plain.iter()) {
2433 assert!(
2434 (a - b).abs() <= 1e-14 * (1.0 + b.abs()),
2435 "Q drift {a} vs {b}"
2436 );
2437 }
2438 let lt0 = tau.ln();
2439 let q_at_lt = |lt: f64| {
2440 measure_jet_energy_form(centers.view(), masses.view(), &band, s0, a0, lt.exp())
2441 .expect("energy form")
2442 };
2443 let h = 1e-4;
2449 let checks: [(&str, &Array2<f64>, Array2<f64>); 9] = [
2450 ("dq_ds", &jets.dq_ds, {
2451 let (p, m_) = (q_at(s0 + h, a0), q_at(s0 - h, a0));
2452 (&p - &m_) / (2.0 * h)
2453 }),
2454 ("d2q_ds2", &jets.d2q_ds2, {
2455 let (p, c, m_) = (q_at(s0 + h, a0), q_at(s0, a0), q_at(s0 - h, a0));
2456 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2457 }),
2458 ("dq_dalpha", &jets.dq_dalpha, {
2459 let (p, m_) = (q_at(s0, a0 + h), q_at(s0, a0 - h));
2460 (&p - &m_) / (2.0 * h)
2461 }),
2462 ("d2q_dalpha2", &jets.d2q_dalpha2, {
2463 let (p, c, m_) = (q_at(s0, a0 + h), q_at(s0, a0), q_at(s0, a0 - h));
2464 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2465 }),
2466 ("d2q_ds_dalpha", &jets.d2q_ds_dalpha, {
2467 let pp = q_at(s0 + h, a0 + h);
2468 let pm = q_at(s0 + h, a0 - h);
2469 let mp = q_at(s0 - h, a0 + h);
2470 let mm = q_at(s0 - h, a0 - h);
2471 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2472 }),
2473 ("dq_dlogtau", &jets.dq_dlogtau, {
2474 let (p, m_) = (q_at_lt(lt0 + h), q_at_lt(lt0 - h));
2475 (&p - &m_) / (2.0 * h)
2476 }),
2477 ("d2q_dlogtau2", &jets.d2q_dlogtau2, {
2478 let (p, c, m_) = (q_at_lt(lt0 + h), q_at_lt(lt0), q_at_lt(lt0 - h));
2479 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2480 }),
2481 ("d2q_ds_dlogtau", &jets.d2q_ds_dlogtau, {
2482 let f = |s: f64, lt: f64| {
2483 measure_jet_energy_form(centers.view(), masses.view(), &band, s, a0, lt.exp())
2484 .expect("energy form")
2485 };
2486 let pp = f(s0 + h, lt0 + h);
2487 let pm = f(s0 + h, lt0 - h);
2488 let mp = f(s0 - h, lt0 + h);
2489 let mm = f(s0 - h, lt0 - h);
2490 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2491 }),
2492 ("d2q_dalpha_dlogtau", &jets.d2q_dalpha_dlogtau, {
2493 let f = |a: f64, lt: f64| {
2494 measure_jet_energy_form(centers.view(), masses.view(), &band, s0, a, lt.exp())
2495 .expect("energy form")
2496 };
2497 let pp = f(a0 + h, lt0 + h);
2498 let pm = f(a0 + h, lt0 - h);
2499 let mp = f(a0 - h, lt0 + h);
2500 let mm = f(a0 - h, lt0 - h);
2501 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2502 }),
2503 ];
2504 for (name, analytic, fd) in checks.iter() {
2505 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2506 for (a, b) in analytic.iter().zip(fd.iter()) {
2507 assert!(
2508 (a - b).abs() <= 5e-5 * scale,
2509 "{name} jet mismatch: analytic {a:.6e} vs FD {b:.6e} (scale {scale:.3e})"
2510 );
2511 }
2512 }
2513 }
2514
2515 #[test]
2519 pub(crate) fn scale_spectrum_sums_to_total_and_localizes_roughness() {
2520 let m = 24usize;
2521 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2522 let t = i as f64 / (m as f64 - 1.0);
2523 if k == 0 { t * 4.0 } else { 0.0 }
2524 });
2525 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2526 let band = band_for(¢ers);
2527 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2528 .expect("energy form");
2529 let fast = Array1::<f64>::from_shape_fn(m, |i| if i % 2 == 0 { 0.5 } else { -0.5 });
2530 let spec = measure_jet_scale_spectrum(
2531 centers.view(),
2532 masses.view(),
2533 &band,
2534 1.5,
2535 1.0,
2536 1e-3,
2537 fast.view(),
2538 )
2539 .expect("spectrum");
2540 assert_eq!(spec.len(), band.eps.len());
2541 let total = fast.dot(&q.dot(&fast));
2542 let sum: f64 = spec.iter().sum();
2543 assert!(
2544 (sum - total).abs() <= 1e-10 * total.abs().max(1e-30),
2545 "spectrum must sum to vᵀQv: {sum:.6e} vs {total:.6e}"
2546 );
2547 let finest = spec[0];
2549 let coarsest = *spec.last().expect("nonempty spectrum");
2550 assert!(
2551 finest > coarsest,
2552 "alternating values must charge fine scales hardest: fine {finest:.3e} vs coarse {coarsest:.3e}"
2553 );
2554 }
2555
2556 #[test]
2559 pub(crate) fn support_curve_separates_on_web_from_off_web() {
2560 let m = 24usize;
2561 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2562 let t = i as f64 / (m as f64 - 1.0);
2563 if k == 0 { t * 4.0 } else { 0.0 }
2564 });
2565 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2566 let band = band_for(¢ers);
2567 let queries = array![[2.0, 0.0], [2.0, 1.5]];
2568 let curves =
2569 measure_jet_support_curve(queries.view(), centers.view(), masses.view(), &band.eps)
2570 .expect("support curve");
2571 assert!(
2573 curves[(0, 0)] > 10.0 * curves[(1, 0)],
2574 "fine-scale support must separate web from void: on {:.3e} vs off {:.3e}",
2575 curves[(0, 0)],
2576 curves[(1, 0)]
2577 );
2578 for qi in 0..2 {
2580 for li in 1..band.eps.len() {
2581 assert!(
2582 curves[(qi, li)] >= curves[(qi, li - 1)] - 1e-15,
2583 "support curve must be monotone in scale (query {qi}, level {li})"
2584 );
2585 }
2586 }
2587 }
2588
2589 #[test]
2597 pub(crate) fn default_stays_single_scale_until_multiscale_opt_in() {
2598 let n = 200usize;
2599 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2600 let t = i as f64 / (n as f64 - 1.0);
2601 if k == 0 {
2602 t * 3.0
2603 } else {
2604 0.4 * (t * 3.0).sin()
2605 }
2606 });
2607 let single = MeasureJetBasisSpec {
2611 center_strategy: CenterStrategy::FarthestPoint { num_centers: 80 },
2612 ..MeasureJetBasisSpec::default()
2613 };
2614 assert!(
2615 !measure_jet_multiscale_mode(&single),
2616 "default must resolve to single-scale at any center count"
2617 );
2618 let built_single =
2619 build_measure_jet_basis(data.view(), &single).expect("single-scale build");
2620 assert_eq!(
2621 built_single.active_penalties.len(),
2622 2,
2623 "single-scale double-penalty mode emits Primary + affine/null component"
2624 );
2625 assert!(matches!(
2626 built_single.active_penalties[0].info.source,
2627 PenaltySource::Primary
2628 ));
2629 assert!(matches!(
2630 built_single.active_penalties[1].info.source,
2631 PenaltySource::DoublePenaltyNullspace
2632 ));
2633 let multi = MeasureJetBasisSpec {
2637 center_strategy: CenterStrategy::FarthestPoint { num_centers: 80 },
2638 multiscale: true,
2639 ..MeasureJetBasisSpec::default()
2640 };
2641 assert!(
2642 measure_jet_multiscale_mode(&multi),
2643 "multiscale=true must resolve to multiscale mode"
2644 );
2645 let built_multi = build_measure_jet_basis(data.view(), &multi).expect("multiscale build");
2646 assert!(
2647 built_multi.active_penalties.len() > built_single.active_penalties.len(),
2648 "multiscale mode emits the per-scale spectral split plus null selection, got {} (vs single-scale {})",
2649 built_multi.active_penalties.len(),
2650 built_single.active_penalties.len()
2651 );
2652 }
2653
2654 #[test]
2658 pub(crate) fn fused_mode_without_double_penalty_emits_single_primary_candidate() {
2659 let n = 40usize;
2660 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2661 let t = i as f64 / (n as f64 - 1.0);
2662 if k == 0 {
2663 t * 3.0
2664 } else {
2665 0.4 * (t * 3.0).sin()
2666 }
2667 });
2668 let spec = MeasureJetBasisSpec {
2669 center_strategy: CenterStrategy::FarthestPoint { num_centers: 14 },
2670 order_s: 1.3,
2671 double_penalty: false,
2672 ..MeasureJetBasisSpec::default()
2673 };
2674 let built = build_measure_jet_basis(data.view(), &spec).expect("fused build");
2675 assert_eq!(
2676 built.active_penalties.len(),
2677 1,
2678 "single-scale mode without null recovery emits exactly one Primary"
2679 );
2680 assert!(matches!(
2681 built.active_penalties[0].info.source,
2682 PenaltySource::Primary
2683 ));
2684 let BasisMetadata::MeasureJet { order_s, .. } = &built.metadata else {
2685 panic!("measure-jet build must return MeasureJet metadata");
2686 };
2687 assert_eq!(*order_s, 1.3, "explicit order must persist verbatim");
2688 }
2689
2690 #[test]
2695 pub(crate) fn single_scale_affine_head_gauge_annihilates_center_cross() {
2696 let n = 90usize;
2697 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2698 let t = i as f64 / (n as f64 - 1.0);
2699 if k == 0 {
2700 3.0 * t
2701 } else {
2702 (2.0 * std::f64::consts::PI * t).sin() + 0.2 * t
2703 }
2704 });
2705 let spec = MeasureJetBasisSpec {
2706 center_strategy: CenterStrategy::FarthestPoint { num_centers: 18 },
2707 double_penalty: false,
2708 multiscale: false,
2709 ..MeasureJetBasisSpec::default()
2710 };
2711 let geom = realize_measure_jet_geometry(data.view(), &spec).expect("realized geometry");
2712 let m = geom.centers.nrows();
2713 let head_rank = geom.head_transform.ncols();
2714 assert!(head_rank > 0, "fixture must realize an affine head");
2715 assert_eq!(
2716 geom.z.ncols(),
2717 m - 1,
2718 "affine gauge replaces duplicated RBF directions without widening the smooth"
2719 );
2720 let rbf_rank = m - (head_rank + 1);
2721 let z_rbf = geom.z.slice(ndarray::s![..m, ..rbf_rank]).to_owned();
2722 let k_cc =
2723 measure_jet_design_matrix(geom.centers.view(), geom.centers.view(), geom.length_scale)
2724 .expect("center kernel");
2725 let affine = measure_jet_affine_value_basis(geom.centers.view(), geom.masses.view());
2726 assert_eq!(affine.ncols(), head_rank + 1);
2727 let mut weighted_affine = affine.clone();
2728 for (i, mut row) in weighted_affine.outer_iter_mut().enumerate() {
2729 row.mapv_inplace(|v| v * geom.masses[i]);
2730 }
2731 let constraint_cross = k_cc.t().dot(&weighted_affine);
2732 let residual = constraint_cross.t().dot(&z_rbf);
2733 let scale = constraint_cross
2734 .iter()
2735 .fold(1.0_f64, |acc, value| acc.max(value.abs()));
2736 assert!(
2737 residual.iter().all(|value| value.abs() <= 1e-10 * scale),
2738 "A^T W Kcc Z_rbf must vanish; max residual {:.3e}",
2739 residual
2740 .iter()
2741 .fold(0.0_f64, |acc, value| acc.max(value.abs()))
2742 );
2743 }
2744
2745 #[test]
2748 pub(crate) fn affine_null_penalty_is_covariant_under_coefficient_reparameterization() {
2749 let centers = array![
2750 [-1.0, 0.2],
2751 [-0.4, -0.3],
2752 [0.1, 0.5],
2753 [0.7, -0.2],
2754 [1.2, 0.4],
2755 [1.8, -0.1],
2756 ];
2757 let masses = array![0.08, 0.12, 0.18, 0.22, 0.17, 0.23];
2758 let evaluation = Array2::<f64>::from_shape_fn((centers.nrows(), 3), |(i, j)| {
2759 ((i + 2 * j + 1) as f64).sin() + 0.15 * (i * (j + 1)) as f64
2760 });
2761 let reparameterization = array![[1.7, 0.2, -0.1], [0.0, 0.6, 0.3], [0.0, 0.0, 1.3]];
2762 let base = affine_function_nullspace_penalty(&evaluation, centers.view(), masses.view())
2763 .expect("base function-space penalty");
2764 let transformed_evaluation = evaluation.dot(&reparameterization);
2765 let transformed = affine_function_nullspace_penalty(
2766 &transformed_evaluation,
2767 centers.view(),
2768 masses.view(),
2769 )
2770 .expect("reparameterized function-space penalty");
2771 let expected = reparameterization.t().dot(&base).dot(&reparameterization);
2772 let scale = expected
2773 .iter()
2774 .fold(1.0_f64, |acc, value| acc.max(value.abs()));
2775 assert!(
2776 transformed
2777 .iter()
2778 .zip(expected.iter())
2779 .all(|(actual, want)| (actual - want).abs() <= 1e-11 * scale),
2780 "S(E R) must equal R^T S(E) R"
2781 );
2782 }
2783
2784 #[test]
2787 pub(crate) fn double_penalty_leaves_primary_matrix_unchanged() {
2788 let n = 64usize;
2789 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2790 let t = i as f64 / (n as f64 - 1.0);
2791 if k == 0 { 2.5 * t } else { (4.0 * t).cos() }
2792 });
2793 let base = MeasureJetBasisSpec {
2794 center_strategy: CenterStrategy::FarthestPoint { num_centers: 16 },
2795 order_s: 1.25,
2796 double_penalty: false,
2797 ..MeasureJetBasisSpec::default()
2798 };
2799 let without = build_measure_jet_basis(data.view(), &base).expect("primary-only build");
2800 let with = build_measure_jet_basis(
2801 data.view(),
2802 &MeasureJetBasisSpec {
2803 double_penalty: true,
2804 ..base.clone()
2805 },
2806 )
2807 .expect("double-penalty build");
2808 assert_eq!(without.active_penalties.len(), 1);
2809 assert_eq!(with.active_penalties.len(), 2);
2810 assert!(matches!(
2811 without.active_penalties[0].info.source,
2812 PenaltySource::Primary
2813 ));
2814 assert!(matches!(
2815 with.active_penalties[0].info.source,
2816 PenaltySource::Primary
2817 ));
2818 assert!(matches!(
2819 with.active_penalties[1].info.source,
2820 PenaltySource::DoublePenaltyNullspace
2821 ));
2822 assert!(
2823 without.active_penalties[0]
2824 .matrix
2825 .iter()
2826 .zip(with.active_penalties[0].matrix.iter())
2827 .all(|(a, b)| (a - b).abs() <= 1e-13),
2828 "turning on null recovery must not modify Primary"
2829 );
2830 }
2831
2832 #[test]
2834 pub(crate) fn householder_sum_to_zero_basis_is_orthonormal() {
2835 let m = 9usize;
2836 let u = householder_sum_to_zero_u(m);
2837 let z = householder_sum_to_zero_z(&u);
2838 for j in 0..(m - 1) {
2839 let col_j = z.column(j);
2840 assert!(col_j.sum().abs() <= 1e-12, "column {j} must sum to zero");
2841 for j2 in j..(m - 1) {
2842 let dot = col_j.dot(&z.column(j2));
2843 let want = if j == j2 { 1.0 } else { 0.0 };
2844 assert!(
2845 (dot - want).abs() <= 1e-12,
2846 "orthonormality failure at ({j}, {j2}): {dot}"
2847 );
2848 }
2849 }
2850 }
2851
2852 pub(crate) fn frozen_spec_fixture(
2857 order_s: f64,
2858 multiscale: bool,
2859 ) -> (Array2<f64>, MeasureJetBasisSpec) {
2860 let n = 140usize;
2865 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2866 let t = i as f64 / (n as f64 - 1.0);
2867 if k == 0 {
2868 t * 3.0
2869 } else {
2870 0.5 * (t * 3.0).cos() + if i % 9 == 0 { 0.8 } else { 0.0 }
2871 }
2872 });
2873 let spec = MeasureJetBasisSpec {
2874 center_strategy: CenterStrategy::FarthestPoint { num_centers: 70 },
2875 order_s,
2876 multiscale,
2877 learn_length_scale: false,
2881 ..MeasureJetBasisSpec::default()
2882 };
2883 let first = build_measure_jet_basis(data.view(), &spec).expect("fixture build");
2884 let BasisMetadata::MeasureJet {
2885 centers,
2886 length_scale,
2887 eps_band,
2888 masses,
2889 support_means,
2890 penalty_normalization_scales,
2891 raw_penalty_normalization_scales,
2892 fused_penalty_normalization_scale,
2893 constraint_transform,
2894 ..
2895 } = &first.metadata
2896 else {
2897 panic!("measure-jet build must return MeasureJet metadata");
2898 };
2899 let frozen = MeasureJetBasisSpec {
2900 center_strategy: CenterStrategy::UserProvided(centers.clone()),
2901 order_s,
2902 alpha: spec.alpha,
2903 tau0: spec.tau0,
2904 num_scales: eps_band.len(),
2905 length_scale: *length_scale,
2906 double_penalty: spec.double_penalty,
2907 learn_length_scale: false,
2908 multiscale,
2909 identifiability: MeasureJetIdentifiability::FrozenTransform {
2910 transform: constraint_transform.clone().expect("fit-time z"),
2911 },
2912 frozen_quadrature: Some(MeasureJetFrozenQuadrature {
2913 masses: masses.clone(),
2914 eps_band: eps_band.clone(),
2915 support_means: support_means.clone(),
2916 penalty_normalization_scales: penalty_normalization_scales.clone(),
2917 raw_penalty_normalization_scales: raw_penalty_normalization_scales.clone(),
2918 fused_penalty_normalization_scale: *fused_penalty_normalization_scale,
2919 sigma_coord: None,
2920 }),
2921 };
2922 (data, frozen)
2923 }
2924
2925 #[test]
2930 pub(crate) fn psi_producer_matches_fd_per_level_mode() {
2931 let (data, frozen) = frozen_spec_fixture(0.0, true);
2932 let derivs =
2933 build_measure_jet_basis_psi_derivatives(data.view(), &frozen).expect("psi derivatives");
2934 let l_count = frozen
2935 .frozen_quadrature
2936 .as_ref()
2937 .expect("frozen quadrature")
2938 .eps_band
2939 .len();
2940 assert_eq!(
2941 derivs.penalties_first.len(),
2942 2,
2943 "per-level coords are (α, lnτ)"
2944 );
2945 assert_eq!(derivs.penalties_first[0].len(), l_count + 1);
2946 assert_eq!(derivs.penalties_cross_pairs, vec![(0, 1)]);
2947 let pen_at = |alpha: f64, tau0: f64| {
2948 let trial = MeasureJetBasisSpec {
2949 alpha,
2950 tau0,
2951 ..frozen.clone()
2952 };
2953 build_measure_jet_basis(data.view(), &trial)
2954 .expect("trial build")
2955 .active_penalties
2956 .into_iter()
2957 .map(|penalty| penalty.matrix)
2958 .collect::<Vec<_>>()
2959 };
2960 let h = 1e-4;
2963 let (a0, t0) = (frozen.alpha, frozen.tau0);
2964 let ap = pen_at(a0 + h, t0);
2965 let am = pen_at(a0 - h, t0);
2966 let tp = pen_at(a0, t0 * h.exp());
2967 let tm = pen_at(a0, t0 * (-h).exp());
2968 assert_eq!(
2969 ap.len(),
2970 l_count + 1,
2971 "fixture must keep every scale active"
2972 );
2973 for level in 0..l_count {
2974 let fd_alpha = (&ap[level] - &am[level]) / (2.0 * h);
2975 let fd_tau = (&tp[level] - &tm[level]) / (2.0 * h);
2976 for (name, analytic, fd) in [
2977 ("alpha", &derivs.penalties_first[0][level], fd_alpha),
2978 ("ln_tau", &derivs.penalties_first[1][level], fd_tau),
2979 ] {
2980 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2981 for (x, y) in analytic.iter().zip(fd.iter()) {
2982 assert!(
2983 (x - y).abs() <= 5e-5 * scale,
2984 "{name} jet of scale-candidate {level}: analytic {x:.6e} vs FD {y:.6e}"
2985 );
2986 }
2987 }
2988 }
2989 for coord in 0..2 {
2991 assert!(
2992 derivs.penalties_first[coord][l_count]
2993 .iter()
2994 .all(|v| *v == 0.0),
2995 "null-component candidate must have zero (α, lnτ) drift"
2996 );
2997 }
2998 let provider = derivs
3000 .penalties_cross_provider
3001 .as_ref()
3002 .expect("cross provider");
3003 let cross = provider.evaluate(0, 1).expect("cross pair (α, lnτ)");
3004 let pp = pen_at(a0 + h, t0 * h.exp());
3005 let pm = pen_at(a0 + h, t0 * (-h).exp());
3006 let mp = pen_at(a0 - h, t0 * h.exp());
3007 let mm = pen_at(a0 - h, t0 * (-h).exp());
3008 for level in 0..l_count {
3009 let fd = (&(&pp[level] - &pm[level]) - &(&mp[level] - &mm[level])) / (4.0 * h * h);
3010 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
3011 for (x, y) in cross[level].iter().zip(fd.iter()) {
3012 assert!(
3013 (x - y).abs() <= 5e-4 * scale,
3014 "cross (α, lnτ) jet of scale-candidate {level}: analytic {x:.6e} vs FD {y:.6e}"
3015 );
3016 }
3017 }
3018 }
3019
3020 #[test]
3026 pub(crate) fn psi_producer_matches_fd_length_scale() {
3027 let (data, mut frozen) = frozen_spec_fixture(0.0, false);
3030 frozen.learn_length_scale = true;
3031 let derivs =
3032 build_measure_jet_basis_psi_derivatives(data.view(), &frozen).expect("psi derivatives");
3033 assert_eq!(
3035 derivs.design_first.len(),
3036 1,
3037 "single-scale + learn_length_scale enrolls exactly the ℓ coordinate"
3038 );
3039 assert_eq!(
3040 derivs.penalties_first[0].len(),
3041 2,
3042 "single-scale double penalty carries Primary + affine/null component"
3043 );
3044 let ell0 = frozen.length_scale;
3048 let build_at = |ell: f64| {
3049 let trial = MeasureJetBasisSpec {
3050 length_scale: ell,
3051 ..frozen.clone()
3052 };
3053 build_measure_jet_basis(data.view(), &trial).expect("trial build")
3054 };
3055 let h: f64 = 1e-4;
3056 let plus = build_at(ell0 * h.exp());
3057 let minus = build_at(ell0 * (-h).exp());
3058 let at = build_at(ell0);
3059 assert_eq!(
3060 plus.active_penalties.len(),
3061 2,
3062 "fixture must keep both candidates active"
3063 );
3064 assert_eq!(
3065 minus.active_penalties.len(),
3066 2,
3067 "fixture must keep both candidates active"
3068 );
3069 assert_eq!(
3070 at.active_penalties.len(),
3071 2,
3072 "fixture must keep both candidates active"
3073 );
3074
3075 let x_plus = plus.design.to_dense();
3076 let x_minus = minus.design.to_dense();
3077 let x_0 = at.design.to_dense();
3078 let fd_first = (&x_plus - &x_minus) / (2.0 * h);
3079 let fd_second = (&x_plus - &(&x_0 * 2.0) + &x_minus) / (h * h);
3080 let scale1 = fd_first.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
3081 for (x, y) in derivs.design_first[0].iter().zip(fd_first.iter()) {
3082 assert!(
3083 (x - y).abs() <= 5e-5 * scale1,
3084 "∂X/∂lnℓ: analytic {x:.6e} vs FD {y:.6e}"
3085 );
3086 }
3087 let scale2 = fd_second.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
3088 for (x, y) in derivs.design_second_diag[0].iter().zip(fd_second.iter()) {
3089 assert!(
3090 (x - y).abs() <= 1e-3 * scale2,
3091 "∂²X/∂lnℓ²: analytic {x:.6e} vs FD {y:.6e}"
3092 );
3093 }
3094
3095 for candidate in 0..2 {
3096 let fd_penalty_first = (&plus.active_penalties[candidate].matrix
3097 - &minus.active_penalties[candidate].matrix)
3098 / (2.0 * h);
3099 let fd_penalty_second = (&plus.active_penalties[candidate].matrix
3100 - &(&at.active_penalties[candidate].matrix * 2.0)
3101 + &minus.active_penalties[candidate].matrix)
3102 / (h * h);
3103 let first_scale = fd_penalty_first
3104 .iter()
3105 .fold(1e-12_f64, |acc, value| acc.max(value.abs()));
3106 let second_scale = fd_penalty_second
3107 .iter()
3108 .fold(1e-10_f64, |acc, value| acc.max(value.abs()));
3109 for (analytic, finite_difference) in derivs.penalties_first[0][candidate]
3110 .iter()
3111 .zip(fd_penalty_first.iter())
3112 {
3113 assert!(
3114 (analytic - finite_difference).abs() <= 1e-4 * first_scale,
3115 "candidate {candidate} ∂S~/∂lnℓ: analytic {analytic:.6e} vs FD {finite_difference:.6e}"
3116 );
3117 }
3118 for (analytic, finite_difference) in derivs.penalties_second_diag[0][candidate]
3119 .iter()
3120 .zip(fd_penalty_second.iter())
3121 {
3122 assert!(
3123 (analytic - finite_difference).abs() <= 5e-3 * second_scale,
3124 "candidate {candidate} ∂²S~/∂lnℓ²: analytic {analytic:.6e} vs FD {finite_difference:.6e}"
3125 );
3126 }
3127 }
3128 }
3129
3130 #[test]
3134 pub(crate) fn quadrature_nodes_are_cell_barycenters() {
3135 let data = array![
3138 [0.0, 0.2],
3139 [0.4, -0.2],
3140 [0.2, 0.0],
3141 [9.8, 10.1],
3142 [10.2, 9.9],
3143 ];
3144 let seeds = array![[0.1, 0.1], [10.0, 10.0], [-50.0, -50.0]];
3145 let (nodes, masses) =
3146 measure_jet_quadrature_nodes(data.view(), seeds.view()).expect("quadrature nodes");
3147 assert!((masses.sum() - 1.0).abs() <= 1e-15, "masses must sum to 1");
3148 assert!((masses[0] - 0.6).abs() <= 1e-15);
3149 assert!((masses[1] - 0.4).abs() <= 1e-15);
3150 assert_eq!(masses[2], 0.0);
3151 assert_eq!(nodes[(0, 0)], 0.2);
3153 assert_eq!(nodes[(0, 1)], 0.0);
3154 assert_eq!(nodes[(1, 0)], 10.0);
3156 assert_eq!(nodes[(1, 1)], 10.0);
3157 assert_eq!(nodes[(2, 0)], -50.0);
3159 assert_eq!(nodes[(2, 1)], -50.0);
3160 }
3161
3162 #[test]
3166 pub(crate) fn build_replay_roundtrip_reproduces_design_and_penalty() {
3167 let n = 140usize;
3170 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
3171 let t = i as f64 / (n as f64 - 1.0);
3172 if k == 0 {
3173 t * 3.0
3174 } else {
3175 0.5 * (t * 3.0).cos() + if i % 9 == 0 { 0.8 } else { 0.0 }
3176 }
3177 });
3178 let spec = MeasureJetBasisSpec {
3179 center_strategy: CenterStrategy::FarthestPoint { num_centers: 70 },
3180 multiscale: true,
3181 ..MeasureJetBasisSpec::default()
3182 };
3183 let first = build_measure_jet_basis(data.view(), &spec).expect("first build");
3184 let BasisMetadata::MeasureJet {
3185 centers,
3186 length_scale,
3187 eps_band,
3188 order_s,
3189 alpha,
3190 tau0,
3191 masses,
3192 support_means,
3193 penalty_normalization_scales,
3194 raw_penalty_normalization_scales,
3195 fused_penalty_normalization_scale,
3196 constraint_transform,
3197 ..
3198 } = &first.metadata
3199 else {
3200 panic!("measure-jet build must return MeasureJet metadata");
3201 };
3202 let replay_spec = MeasureJetBasisSpec {
3203 center_strategy: CenterStrategy::UserProvided(centers.clone()),
3204 order_s: *order_s,
3205 alpha: *alpha,
3206 tau0: *tau0,
3207 num_scales: eps_band.len(),
3208 length_scale: *length_scale,
3209 double_penalty: spec.double_penalty,
3210 learn_length_scale: spec.learn_length_scale,
3211 multiscale: spec.multiscale,
3212 identifiability: MeasureJetIdentifiability::FrozenTransform {
3213 transform: constraint_transform.clone().expect("fit-time z"),
3214 },
3215 frozen_quadrature: Some(MeasureJetFrozenQuadrature {
3216 masses: masses.clone(),
3217 eps_band: eps_band.clone(),
3218 support_means: support_means.clone(),
3219 penalty_normalization_scales: penalty_normalization_scales.clone(),
3220 raw_penalty_normalization_scales: raw_penalty_normalization_scales.clone(),
3221 fused_penalty_normalization_scale: *fused_penalty_normalization_scale,
3222 sigma_coord: None,
3223 }),
3224 };
3225 assert_eq!(
3228 first.active_penalties.len(),
3229 eps_band.len() + 1,
3230 "per-level mode must emit one candidate per scale + null component"
3231 );
3232 let second = build_measure_jet_basis(data.view(), &replay_spec).expect("replay build");
3233 let x1 = first.design.to_dense();
3234 let x2 = second.design.to_dense();
3235 assert_eq!(x1.shape(), x2.shape());
3236 for (a, b) in x1.iter().zip(x2.iter()) {
3237 assert!((a - b).abs() <= 1e-12, "design replay drift: {a} vs {b}");
3238 }
3239 assert_eq!(first.active_penalties.len(), second.active_penalties.len());
3240 for (p1, p2) in first
3241 .active_penalties
3242 .iter()
3243 .zip(second.active_penalties.iter())
3244 {
3245 for (a, b) in p1.matrix.iter().zip(p2.matrix.iter()) {
3246 assert!((a - b).abs() <= 1e-12, "penalty replay drift: {a} vs {b}");
3247 }
3248 }
3249 }
3250}