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
496fn affine_function_nullspace_center_quadratic(
500 centers: ArrayView2<'_, f64>,
501 masses: ArrayView1<'_, f64>,
502) -> Result<ConstructiveQuadratic, BasisError> {
503 ConstructiveQuadratic::try_from_dense_psd(
504 affine_function_nullspace_form(centers, masses)?,
505 "measure-jet affine center-value form",
506 )
507}
508
509fn affine_function_nullspace_quadratic(
510 evaluation: &Array2<f64>,
511 centers: ArrayView2<'_, f64>,
512 masses: ArrayView1<'_, f64>,
513) -> Result<ConstructiveQuadratic, BasisError> {
514 if evaluation.nrows() != centers.nrows() {
515 crate::bail_dim_basis!(
516 "measure-jet affine function-space penalty shape mismatch: evaluation {:?}, centers {:?}",
517 evaluation.dim(),
518 centers.dim()
519 );
520 }
521 let center_quadratic = affine_function_nullspace_center_quadratic(centers, masses)?;
522 ConstructiveQuadratic::from_energy_factor(
523 center_quadratic.factor().dot(evaluation),
524 "measure-jet affine/null coefficient penalty",
525 )
526}
527
528pub(crate) fn pairwise_sq_dists(a: ArrayView2<'_, f64>, b: ArrayView2<'_, f64>) -> Array2<f64> {
539 let an: Vec<f64> = a.outer_iter().map(|r| r.dot(&r)).collect();
540 let bn: Vec<f64> = b.outer_iter().map(|r| r.dot(&r)).collect();
541 let mut g = a.dot(&b.t());
542 g.axis_iter_mut(Axis(0))
543 .into_par_iter()
544 .enumerate()
545 .for_each(|(i, mut row)| {
546 for (j, v) in row.iter_mut().enumerate() {
547 *v = (an[i] + bn[j] - 2.0 * *v).max(0.0);
548 }
549 });
550 g
551}
552
553pub(crate) const MEASURE_JET_ASSIGN_BLOCK_ROWS: usize = 65_536;
557
558pub(crate) fn validate_finite_points(
559 points: ArrayView2<'_, f64>,
560 what: &str,
561) -> Result<(), BasisError> {
562 for (i, row) in points.outer_iter().enumerate() {
563 if row.iter().any(|v| !v.is_finite()) {
564 crate::bail_invalid_basis!("measure-jet {what} row {i} has a non-finite coordinate");
565 }
566 }
567 Ok(())
568}
569
570pub(crate) fn median_nearest_center_spacing(dist2: &Array2<f64>) -> Result<f64, BasisError> {
573 let m = dist2.nrows();
574 if m < 2 {
575 return Err(BasisError::InsufficientColumnsForConstraint { found: m });
576 }
577 let mut nearest: Vec<f64> = Vec::with_capacity(m);
578 for i in 0..m {
579 let mut best = f64::INFINITY;
580 for j in 0..m {
581 if j != i && dist2[(i, j)] < best {
582 best = dist2[(i, j)];
583 }
584 }
585 nearest.push(best.sqrt());
586 }
587 nearest.sort_by(|a, b| a.partial_cmp(b).expect("finite center spacings"));
588 let median = nearest[nearest.len() / 2];
589 if !(median.is_finite() && median > 0.0) {
590 crate::bail_invalid_basis!(
591 "measure-jet centers are degenerate (median nearest-center spacing = {median}); \
592 duplicate centers cannot carry a scale band"
593 );
594 }
595 Ok(median)
596}
597
598pub fn measure_jet_band(
606 centers: ArrayView2<'_, f64>,
607 num_scales: usize,
608) -> Result<MeasureJetBand, BasisError> {
609 validate_finite_points(centers, "centers")?;
610 let dist2 = pairwise_sq_dists(centers, centers);
611 let eps_min = median_nearest_center_spacing(&dist2)?;
612 let d = centers.ncols();
614 let mut diag2 = 0.0_f64;
615 for k in 0..d {
616 let col = centers.column(k);
617 let mut lo = f64::INFINITY;
618 let mut hi = f64::NEG_INFINITY;
619 for &v in col.iter() {
620 lo = lo.min(v);
621 hi = hi.max(v);
622 }
623 diag2 += (hi - lo) * (hi - lo);
624 }
625 let eps_max = 0.5 * diag2.sqrt();
626 if !(eps_max.is_finite() && eps_max > eps_min) {
627 return Ok(MeasureJetBand {
628 eps: vec![eps_min],
629 log_step: std::f64::consts::LN_2,
630 });
631 }
632 let auto = ((eps_max / eps_min).log2().ceil() as usize + 1)
633 .clamp(MEASURE_JET_MIN_AUTO_SCALES, MEASURE_JET_MAX_AUTO_SCALES);
634 let count = if num_scales == 0 { auto } else { num_scales };
635 if count == 1 {
636 return Ok(MeasureJetBand {
637 eps: vec![eps_min],
638 log_step: std::f64::consts::LN_2,
639 });
640 }
641 let ratio = (eps_max / eps_min).powf(1.0 / (count as f64 - 1.0));
642 let mut eps = Vec::with_capacity(count);
643 let mut e = eps_min;
644 for _ in 0..count {
645 eps.push(e);
646 e *= ratio;
647 }
648 Ok(MeasureJetBand {
649 eps,
650 log_step: ratio.ln(),
651 })
652}
653
654pub fn measure_jet_quadrature_nodes(
661 data: ArrayView2<'_, f64>,
662 centers: ArrayView2<'_, f64>,
663) -> Result<(Array2<f64>, Array1<f64>), BasisError> {
664 if data.ncols() != centers.ncols() {
665 crate::bail_dim_basis!(
666 "measure-jet mass assignment dimension mismatch: data d={} centers d={}",
667 data.ncols(),
668 centers.ncols()
669 );
670 }
671 validate_finite_points(data, "data")?;
672 validate_finite_points(centers, "centers")?;
673 let n = data.nrows();
674 let m = centers.nrows();
675 let d = centers.ncols();
676 if n == 0 || m == 0 {
677 crate::bail_invalid_basis!("measure-jet mass assignment needs nonempty data and centers");
678 }
679 let cn: Vec<f64> = centers.outer_iter().map(|r| r.dot(&r)).collect();
684 let assignments: Vec<usize> = (0..n)
685 .step_by(MEASURE_JET_ASSIGN_BLOCK_ROWS)
686 .flat_map(|start| {
687 let end = (start + MEASURE_JET_ASSIGN_BLOCK_ROWS).min(n);
688 let g = data.slice(ndarray::s![start..end, ..]).dot(¢ers.t());
689 let block: Vec<usize> = g
690 .axis_iter(Axis(0))
691 .into_par_iter()
692 .map(|row| {
693 let mut best_j = 0usize;
694 let mut best = f64::INFINITY;
695 for (j, &gij) in row.iter().enumerate() {
696 let s = cn[j] - 2.0 * gij;
697 if s < best {
698 best = s;
699 best_j = j;
700 }
701 }
702 best_j
703 })
704 .collect();
705 block
706 })
707 .collect();
708 let mut masses = Array1::<f64>::zeros(m);
709 let mut nodes = centers.to_owned();
710 let mut sums = Array2::<f64>::zeros((m, d));
711 let unit = 1.0 / n as f64;
712 for (i, &j) in assignments.iter().enumerate() {
713 masses[j] += unit;
714 for k in 0..d {
715 sums[(j, k)] += data[(i, k)];
716 }
717 }
718 let mut barycenter = sums;
721 for j in 0..m {
722 let count = masses[j] * n as f64;
723 if count > 0.0 {
724 for k in 0..d {
725 barycenter[(j, k)] /= count;
726 nodes[(j, k)] = barycenter[(j, k)];
727 }
728 }
729 }
730 Ok((nodes, masses))
731}
732
733pub fn measure_jet_center_masses(
736 data: ArrayView2<'_, f64>,
737 centers: ArrayView2<'_, f64>,
738) -> Result<Array1<f64>, BasisError> {
739 measure_jet_quadrature_nodes(data, centers).map(|(_, masses)| masses)
740}
741
742pub(crate) fn assemble_weighted_forms<F>(
765 centers: ArrayView2<'_, f64>,
766 masses: ArrayView1<'_, f64>,
767 band: &MeasureJetBand,
768 order_s: f64,
769 alpha: f64,
770 tau0: f64,
771 n_forms: usize,
772 channels: usize,
773 weights: &F,
774) -> Result<Vec<Array2<f64>>, BasisError>
775where
776 F: Fn(usize, f64, f64, f64, &mut [[f64; 3]]) + Sync,
777{
778 let m = centers.nrows();
779 let d = centers.ncols();
780 if n_forms == 0 || !(1..=3).contains(&channels) {
781 crate::bail_invalid_basis!(
782 "measure-jet assembly needs at least one output form and 1..=3 block channels"
783 );
784 }
785 if masses.len() != m {
786 crate::bail_dim_basis!(
787 "measure-jet energy mass/center mismatch: {} masses for {} centers",
788 masses.len(),
789 m
790 );
791 }
792 if band.eps.is_empty() || band.eps.iter().any(|e| !(e.is_finite() && *e > 0.0)) {
793 crate::bail_invalid_basis!("measure-jet energy needs a nonempty positive scale band");
794 }
795 if !(order_s.is_finite() && order_s > 0.0 && order_s < 2.0) {
796 crate::bail_invalid_basis!(
797 "measure-jet order s must lie in (0, 2) for the affine-jet energy; got {order_s}"
798 );
799 }
800 if !(alpha.is_finite() && tau0.is_finite() && tau0 >= 0.0) {
801 crate::bail_invalid_basis!(
802 "measure-jet energy needs finite alpha and finite tau0 >= 0; got alpha={alpha}, tau0={tau0}"
803 );
804 }
805 if masses.iter().any(|v| !(v.is_finite() && *v >= 0.0)) {
806 crate::bail_invalid_basis!("measure-jet energy needs finite nonnegative center masses");
807 }
808 let dist2 = pairwise_sq_dists(centers, centers);
809
810 let assemble_scale = |scale_idx: usize, eps: f64| -> Result<Vec<Array2<f64>>, BasisError> {
815 let mut out: Vec<Array2<f64>> =
816 (0..n_forms).map(|_| Array2::<f64>::zeros((m, m))).collect();
817 let cutoff2 = (MEASURE_JET_PROFILE_CUTOFF * eps) * (MEASURE_JET_PROFILE_CUTOFF * eps);
818 let inv_two_eps2 = 1.0 / (2.0 * eps * eps);
819 let eta = 2.0 * order_s + (d as f64) * (2.0 - 2.0 * alpha);
820 let scale_weight = band.log_step * eps.powf(-eta);
821 let net_radius2 = 0.25 * eps * eps;
825 let mut outer: Vec<usize> = Vec::new();
826 for i in 0..m {
827 if masses[i] <= 0.0 {
828 continue;
829 }
830 let covered = outer.iter().any(|&o| dist2[(i, o)] <= net_radius2);
831 if !covered {
832 outer.push(i);
833 }
834 }
835 let mut net_mass = vec![0.0_f64; m];
836 for i in 0..m {
837 if masses[i] <= 0.0 {
838 continue;
839 }
840 let mut best = f64::INFINITY;
841 let mut best_o = usize::MAX;
842 for &o in &outer {
843 if dist2[(i, o)] < best {
844 best = dist2[(i, o)];
845 best_o = o;
846 }
847 }
848 if best_o != usize::MAX {
849 net_mass[best_o] += masses[i];
850 }
851 }
852 let mut wbuf = vec![[0.0_f64; 3]; n_forms];
853 for &i in &outer {
854 let mut idx: Vec<usize> = Vec::new();
856 for j in 0..m {
857 if dist2[(i, j)] <= cutoff2 {
858 idx.push(j);
859 }
860 }
861 let ml = idx.len();
862 let mut w = Array1::<f64>::zeros(ml);
864 let mut q = 0.0_f64;
865 for (a, &j) in idx.iter().enumerate() {
866 let wj = masses[j] * (-dist2[(i, j)] * inv_two_eps2).exp();
867 w[a] = wj;
868 q += wj;
869 }
870 if !(q > 0.0) {
871 continue;
872 }
873 let mut phi = Array2::<f64>::zeros((ml, d));
875 for (a, &j) in idx.iter().enumerate() {
876 for k in 0..d {
877 phi[(a, k)] = (centers[(j, k)] - centers[(i, k)]) / eps;
878 }
879 }
880 let a_mean = phi.t().dot(&w) / q;
881 let mut wphi = phi.clone();
883 for (a, mut row) in wphi.outer_iter_mut().enumerate() {
884 row.mapv_inplace(|v| v * w[a]);
885 }
886 let mut b = wphi.clone();
887 for (a, mut row) in b.outer_iter_mut().enumerate() {
888 for k in 0..d {
889 row[k] -= w[a] * a_mean[k];
890 }
891 }
892 let mut g = phi.t().dot(&wphi);
893 g.mapv_inplace(|v| v / q);
894 for r in 0..d {
895 for c in 0..d {
896 g[(r, c)] -= a_mean[r] * a_mean[c];
897 }
898 }
899 let g_pinv = symmetric_pseudoinverse(&g, "local affine Gram")?;
900 let bm = b.dot(&g_pinv);
901 let base = scale_weight * net_mass[i] * q.powf(1.0 - 2.0 * alpha);
902 weights(scale_idx, eps, q, base, &mut wbuf);
903 for (a, &ja) in idx.iter().enumerate() {
906 let bma = bm.row(a);
907 for (c, &jc) in idx.iter().enumerate() {
908 let b_c = b.row(c);
909 let mut val_r = -w[a] * w[c] / q - bma.dot(&b_c) / q;
910 if a == c {
911 val_r += w[a];
912 }
913 for (k, out_k) in out.iter_mut().enumerate() {
914 let wk = wbuf[k];
915 out_k[(ja, jc)] += wk[0] * val_r;
916 }
917 }
918 }
919 }
920 Ok(out)
921 };
922
923 let n_scales = band.eps.len();
924 let parallel_ok = m
925 .saturating_mul(m)
926 .saturating_mul(n_scales)
927 .saturating_mul(n_forms)
928 <= MEASURE_JET_PARALLEL_FORM_BUDGET_DOUBLES;
929 let per_scale: Vec<Vec<Array2<f64>>> = if parallel_ok {
930 band.eps
931 .par_iter()
932 .enumerate()
933 .map(|(scale_idx, &eps)| assemble_scale(scale_idx, eps))
934 .collect::<Result<Vec<_>, BasisError>>()?
935 } else {
936 band.eps
937 .iter()
938 .enumerate()
939 .map(|(scale_idx, &eps)| assemble_scale(scale_idx, eps))
940 .collect::<Result<Vec<_>, BasisError>>()?
941 };
942
943 let mut totals: Vec<Array2<f64>> = (0..n_forms).map(|_| Array2::<f64>::zeros((m, m))).collect();
944 for scale_forms in per_scale {
945 for (total, part) in totals.iter_mut().zip(scale_forms) {
946 *total += ∂
947 }
948 }
949 Ok(totals.into_iter().map(|t| (&t + &t.t()) * 0.5).collect())
951}
952
953pub fn measure_jet_energy_form(
968 centers: ArrayView2<'_, f64>,
969 masses: ArrayView1<'_, f64>,
970 band: &MeasureJetBand,
971 order_s: f64,
972 alpha: f64,
973 tau0: f64,
974) -> Result<Array2<f64>, BasisError> {
975 let mut forms = assemble_weighted_forms(
976 centers,
977 masses,
978 band,
979 order_s,
980 alpha,
981 tau0,
982 1,
983 1,
984 &|_, _, _, base, out: &mut [[f64; 3]]| out[0] = [base, 0.0, 0.0],
985 )?;
986 let q = forms.swap_remove(0);
987 project_symmetric_psd(q, "measure-jet energy form")
995}
996
997pub(crate) fn project_symmetric_psd(
1003 a: Array2<f64>,
1004 label: &str,
1005) -> Result<Array2<f64>, BasisError> {
1006 let n = a.nrows();
1007 if n == 0 {
1008 return Ok(a);
1009 }
1010 let (evals, evecs) = a.eigh(Side::Lower).map_err(|e| {
1011 BasisError::InvalidInput(format!(
1012 "measure-jet PSD projection `{label}` eigendecomposition failed: {e}"
1013 ))
1014 })?;
1015 if evals.iter().all(|&lam| lam >= 0.0) {
1016 return Ok(a);
1017 }
1018 let mut scaled = evecs.clone();
1019 for (k, mut col) in scaled.axis_iter_mut(Axis(1)).enumerate() {
1020 let lam = evals[k].max(0.0);
1021 col.mapv_inplace(|v| v * lam);
1022 }
1023 let psd = scaled.dot(&evecs.t());
1024 Ok((&psd + &psd.t()) * 0.5)
1025}
1026
1027pub fn measure_jet_energy_form_with_jets(
1042 centers: ArrayView2<'_, f64>,
1043 masses: ArrayView1<'_, f64>,
1044 band: &MeasureJetBand,
1045 order_s: f64,
1046 alpha: f64,
1047 tau0: f64,
1048) -> Result<MeasureJetEnergyJets, BasisError> {
1049 if !(tau0.is_finite() && tau0 > 0.0) {
1050 crate::bail_invalid_basis!(
1051 "measure-jet jets need tau0 > 0 because the retained τ coordinate is ln τ; got {tau0}"
1052 );
1053 }
1054 let mut forms = assemble_weighted_forms(
1055 centers,
1056 masses,
1057 band,
1058 order_s,
1059 alpha,
1060 tau0,
1061 10,
1062 3,
1063 &|_, eps: f64, q: f64, base: f64, out: &mut [[f64; 3]]| {
1064 let gs = -2.0 * eps.ln();
1065 let intrinsic_dim = centers.ncols() as f64;
1066 let ga = 2.0 * intrinsic_dim * eps.ln() - 2.0 * q.max(f64::MIN_POSITIVE).ln();
1067 out[0] = [base, 0.0, 0.0];
1068 out[1] = [gs * base, 0.0, 0.0];
1069 out[2] = [gs * gs * base, 0.0, 0.0];
1070 out[3] = [ga * base, 0.0, 0.0];
1071 out[4] = [ga * ga * base, 0.0, 0.0];
1072 out[5] = [gs * ga * base, 0.0, 0.0];
1073 out[6] = [0.0, 0.0, 0.0];
1074 out[7] = [0.0, 0.0, 0.0];
1075 out[8] = [0.0, 0.0, 0.0];
1076 out[9] = [0.0, 0.0, 0.0];
1077 },
1078 )?;
1079 let d2q_dalpha_dlogtau = forms.pop().expect("ten assembled forms");
1080 let d2q_ds_dlogtau = forms.pop().expect("ten assembled forms");
1081 let d2q_dlogtau2 = forms.pop().expect("ten assembled forms");
1082 let dq_dlogtau = forms.pop().expect("ten assembled forms");
1083 let d2q_ds_dalpha = forms.pop().expect("ten assembled forms");
1084 let d2q_dalpha2 = forms.pop().expect("ten assembled forms");
1085 let dq_dalpha = forms.pop().expect("ten assembled forms");
1086 let d2q_ds2 = forms.pop().expect("ten assembled forms");
1087 let dq_ds = forms.pop().expect("ten assembled forms");
1088 let q = forms.pop().expect("ten assembled forms");
1089 Ok(MeasureJetEnergyJets {
1090 q,
1091 dq_ds,
1092 d2q_ds2,
1093 dq_dalpha,
1094 d2q_dalpha2,
1095 d2q_ds_dalpha,
1096 dq_dlogtau,
1097 d2q_dlogtau2,
1098 d2q_ds_dlogtau,
1099 d2q_dalpha_dlogtau,
1100 })
1101}
1102
1103pub fn measure_jet_scale_spectrum(
1109 centers: ArrayView2<'_, f64>,
1110 masses: ArrayView1<'_, f64>,
1111 band: &MeasureJetBand,
1112 order_s: f64,
1113 alpha: f64,
1114 tau0: f64,
1115 values: ArrayView1<'_, f64>,
1116) -> Result<Vec<f64>, BasisError> {
1117 if values.len() != centers.nrows() {
1118 crate::bail_dim_basis!(
1119 "measure-jet scale spectrum needs one value per center: {} values for {} centers",
1120 values.len(),
1121 centers.nrows()
1122 );
1123 }
1124 let forms = measure_jet_energy_forms_per_scale(centers, masses, band, order_s, alpha, tau0)?;
1125 Ok(forms
1126 .iter()
1127 .map(|q_l| values.dot(&q_l.dot(&values)))
1128 .collect())
1129}
1130
1131pub fn measure_jet_energy_forms_per_scale(
1137 centers: ArrayView2<'_, f64>,
1138 masses: ArrayView1<'_, f64>,
1139 band: &MeasureJetBand,
1140 order_s: f64,
1141 alpha: f64,
1142 tau0: f64,
1143) -> Result<Vec<Array2<f64>>, BasisError> {
1144 let n_scales = band.eps.len();
1145 assemble_weighted_forms(
1146 centers,
1147 masses,
1148 band,
1149 order_s,
1150 alpha,
1151 tau0,
1152 n_scales,
1153 1,
1154 &|scale_idx, _, _, base, out: &mut [[f64; 3]]| {
1155 for (k, slot) in out.iter_mut().enumerate() {
1156 *slot = if k == scale_idx {
1157 [base, 0.0, 0.0]
1158 } else {
1159 [0.0, 0.0, 0.0]
1160 };
1161 }
1162 },
1163 )
1164}
1165
1166pub fn measure_jet_support_curve(
1174 queries: ArrayView2<'_, f64>,
1175 centers: ArrayView2<'_, f64>,
1176 masses: ArrayView1<'_, f64>,
1177 eps_band: &[f64],
1178) -> Result<Array2<f64>, BasisError> {
1179 if queries.ncols() != centers.ncols() {
1180 crate::bail_dim_basis!(
1181 "measure-jet support curve dimension mismatch: queries d={} centers d={}",
1182 queries.ncols(),
1183 centers.ncols()
1184 );
1185 }
1186 if masses.len() != centers.nrows() {
1187 crate::bail_dim_basis!(
1188 "measure-jet support curve mass/center mismatch: {} masses for {} centers",
1189 masses.len(),
1190 centers.nrows()
1191 );
1192 }
1193 if eps_band.is_empty() || eps_band.iter().any(|e| !(e.is_finite() && *e > 0.0)) {
1194 crate::bail_invalid_basis!("measure-jet support curve needs a nonempty positive band");
1195 }
1196 validate_finite_points(queries, "queries")?;
1197 validate_finite_points(centers, "centers")?;
1198 let nq = queries.nrows();
1199 let nl = eps_band.len();
1200 let d2 = pairwise_sq_dists(queries, centers);
1203 let mut out = Array2::<f64>::zeros((nq, nl));
1204 out.axis_iter_mut(Axis(0))
1205 .into_par_iter()
1206 .enumerate()
1207 .for_each(|(qi, mut row)| {
1208 let d2_row = d2.row(qi);
1209 for (li, &eps) in eps_band.iter().enumerate() {
1210 let inv_two_eps2 = 1.0 / (2.0 * eps * eps);
1211 let mut acc = 0.0_f64;
1212 for (j, &dd) in d2_row.iter().enumerate() {
1213 acc += masses[j] * (-dd * inv_two_eps2).exp();
1214 }
1215 row[li] = acc;
1216 }
1217 });
1218 Ok(out)
1219}
1220
1221pub(crate) fn measure_jet_support_means(
1222 centers: ArrayView2<'_, f64>,
1223 masses: ArrayView1<'_, f64>,
1224 eps_band: &[f64],
1225) -> Result<Vec<f64>, BasisError> {
1226 let total_mass = masses.sum();
1227 if !(total_mass.is_finite() && total_mass > 0.0) {
1228 crate::bail_invalid_basis!(
1229 "measure-jet support means need positive finite total mass; got {total_mass}"
1230 );
1231 }
1232 let support = measure_jet_support_curve(centers, centers, masses, eps_band)?;
1233 let mut means = vec![0.0_f64; eps_band.len()];
1234 for (i, row) in support.rows().into_iter().enumerate() {
1235 let mass = masses[i];
1236 for (mean, &q) in means.iter_mut().zip(row.iter()) {
1237 *mean += mass * q;
1238 }
1239 }
1240 for mean in &mut means {
1241 *mean /= total_mass;
1242 if !(*mean).is_finite() || *mean <= 0.0 {
1243 crate::bail_invalid_basis!(
1244 "measure-jet support mean must be positive and finite; got {mean}"
1245 );
1246 }
1247 }
1248 Ok(means)
1249}
1250
1251pub fn measure_jet_design_matrix(
1253 data: ArrayView2<'_, f64>,
1254 centers: ArrayView2<'_, f64>,
1255 length_scale: f64,
1256) -> Result<Array2<f64>, BasisError> {
1257 if data.ncols() != centers.ncols() {
1258 crate::bail_dim_basis!(
1259 "measure-jet design dimension mismatch: data d={} centers d={}",
1260 data.ncols(),
1261 centers.ncols()
1262 );
1263 }
1264 if !(length_scale.is_finite() && length_scale > 0.0) {
1265 crate::bail_invalid_basis!(
1266 "measure-jet design needs a positive finite length_scale; got {length_scale}"
1267 );
1268 }
1269 validate_finite_points(data, "data")?;
1270 validate_finite_points(centers, "centers")?;
1271 let inv_two_l2 = 1.0 / (2.0 * length_scale * length_scale);
1272 let mut out = pairwise_sq_dists(data, centers);
1275 out.axis_iter_mut(Axis(0))
1276 .into_par_iter()
1277 .for_each(|mut row| {
1278 row.mapv_inplace(|d2| (-d2 * inv_two_l2).exp());
1279 });
1280 Ok(out)
1281}
1282
1283fn measure_jet_design_log_length_jets(
1286 data: ArrayView2<'_, f64>,
1287 centers: ArrayView2<'_, f64>,
1288 length_scale: f64,
1289) -> Result<(Array2<f64>, Array2<f64>), BasisError> {
1290 let kernel = measure_jet_design_matrix(data, centers, length_scale)?;
1291 let squared_distances = pairwise_sq_dists(data, centers);
1292 let inv_l2 = 1.0 / (length_scale * length_scale);
1293 let mut first = kernel.clone();
1294 let mut second = kernel;
1295 for ((first_value, second_value), &distance_squared) in first
1296 .iter_mut()
1297 .zip(second.iter_mut())
1298 .zip(squared_distances.iter())
1299 {
1300 let a = distance_squared * inv_l2;
1301 let kernel_value = *first_value;
1302 *first_value = kernel_value * a;
1303 *second_value = kernel_value * (a * a - 2.0 * a);
1304 }
1305 Ok((first, second))
1306}
1307
1308pub fn measure_jet_affine_head_transform(
1339 centers: ArrayView2<'_, f64>,
1340 masses: ArrayView1<'_, f64>,
1341) -> Array2<f64> {
1342 let m = centers.nrows();
1343 let d = centers.ncols();
1344 let total_mass = masses.sum();
1345 let mdot = |u: &Array1<f64>, v: &Array1<f64>| -> f64 {
1347 let mut acc = 0.0;
1348 for i in 0..m {
1349 acc += masses[i] * u[i] * v[i];
1350 }
1351 acc
1352 };
1353 let cols: Vec<Array1<f64>> = (0..d)
1357 .map(|k| {
1358 let col = centers.column(k).to_owned();
1359 let mean = if total_mass > 0.0 {
1360 mdot(&col, &Array1::ones(m)) / total_mass
1361 } else {
1362 0.0
1363 };
1364 col.mapv(|x| x - mean)
1365 })
1366 .collect();
1367 let max_norm = cols
1369 .iter()
1370 .fold(0.0_f64, |acc, c| acc.max(mdot(c, c).sqrt()));
1371 let drop_below =
1372 (MEASURE_JET_PSEUDOINVERSE_RTOL * (d.max(1) as f64) * max_norm).max(f64::MIN_POSITIVE);
1373 let mut q_cols: Vec<Array1<f64>> = Vec::new();
1377 let mut t_cols: Vec<Array1<f64>> = Vec::new();
1378 for k in 0..d {
1379 let mut v = cols[k].clone();
1380 let mut t = Array1::<f64>::zeros(d);
1381 t[k] = 1.0;
1382 for (q, tq) in q_cols.iter().zip(t_cols.iter()) {
1383 let proj = mdot(q, &v);
1384 v.scaled_add(-proj, q);
1385 t.scaled_add(-proj, tq);
1386 }
1387 let norm = mdot(&v, &v).sqrt();
1388 if norm > drop_below {
1389 v.mapv_inplace(|x| x / norm);
1390 t.mapv_inplace(|x| x / norm);
1391 q_cols.push(v);
1392 t_cols.push(t);
1393 }
1394 }
1395 let head_rank = t_cols.len();
1396 let mut t_mat = Array2::<f64>::zeros((d, head_rank));
1397 for (r, t) in t_cols.into_iter().enumerate() {
1398 t_mat.column_mut(r).assign(&t);
1399 }
1400 t_mat
1401}
1402
1403pub fn realized_measure_jet_length_scale(
1408 centers: ArrayView2<'_, f64>,
1409 spec_length_scale: f64,
1410) -> Result<f64, BasisError> {
1411 if spec_length_scale.is_finite() && spec_length_scale > 0.0 {
1412 return Ok(spec_length_scale);
1413 }
1414 if spec_length_scale != 0.0 {
1415 crate::bail_invalid_basis!(
1416 "measure-jet length_scale must be positive (or 0.0 for auto); got {spec_length_scale}"
1417 );
1418 }
1419 let dist2 = pairwise_sq_dists(centers, centers);
1420 let spacing = median_nearest_center_spacing(&dist2)?;
1421 Ok(MEASURE_JET_AUTO_LENGTH_SCALE_FACTOR * spacing)
1422}
1423
1424pub(crate) struct RealizedMeasureJetGeometry {
1429 pub(crate) centers: Array2<f64>,
1430 pub(crate) masses: Array1<f64>,
1431 pub(crate) eps_band: Vec<f64>,
1432 pub(crate) log_step: f64,
1433 pub(crate) length_scale: f64,
1434 pub(crate) order_s_eval: f64,
1438 pub(crate) per_level: bool,
1440 pub(crate) z: Array2<f64>,
1441 pub(crate) coefficient_gauge: gam_problem::Gauge,
1442 pub(crate) kz: Array2<f64>,
1443 pub(crate) head_transform: Array2<f64>,
1449}
1450
1451pub(crate) fn realize_measure_jet_geometry(
1452 data: ArrayView2<'_, f64>,
1453 spec: &MeasureJetBasisSpec,
1454) -> Result<RealizedMeasureJetGeometry, BasisError> {
1455 if data.ncols() == 0 {
1456 crate::bail_invalid_basis!("measure-jet smooth needs at least one feature column");
1457 }
1458 validate_finite_points(data, "data")?;
1459 let seed_centers = select_centers_by_strategy(data, &spec.center_strategy)?;
1460 let m = seed_centers.nrows();
1461 if m < 3 {
1462 return Err(BasisError::InsufficientColumnsForConstraint { found: m });
1463 }
1464 let order_s = if spec.order_s == 0.0 {
1465 MEASURE_JET_DEFAULT_ORDER_S
1466 } else {
1467 spec.order_s
1468 };
1469 let (centers, masses, eps_band, log_step) = match &spec.frozen_quadrature {
1476 Some(frozen) => {
1477 if frozen.masses.len() != m {
1478 crate::bail_dim_basis!(
1479 "frozen measure-jet quadrature mismatch: {} masses for {} centers",
1480 frozen.masses.len(),
1481 m
1482 );
1483 }
1484 if frozen.eps_band.is_empty() {
1485 crate::bail_invalid_basis!("frozen measure-jet quadrature has an empty band");
1486 }
1487 let log_step = if frozen.eps_band.len() >= 2 {
1488 (frozen.eps_band[1] / frozen.eps_band[0]).ln()
1489 } else {
1490 std::f64::consts::LN_2
1491 };
1492 (
1493 seed_centers,
1494 frozen.masses.clone(),
1495 frozen.eps_band.clone(),
1496 log_step,
1497 )
1498 }
1499 None => {
1500 let (nodes, masses) = measure_jet_quadrature_nodes(data, seed_centers.view())?;
1501 let band = measure_jet_band(nodes.view(), spec.num_scales)?;
1502 (nodes, masses, band.eps, band.log_step)
1503 }
1504 };
1505 let length_scale = realized_measure_jet_length_scale(centers.view(), spec.length_scale)?;
1506 let head_transform = if spec.multiscale {
1517 Array2::<f64>::zeros((centers.ncols(), 0))
1518 } else {
1519 measure_jet_affine_head_transform(centers.view(), masses.view())
1520 };
1521 let head_rank = head_transform.ncols();
1522 let m_aug = m + head_rank;
1523 let k_cc = measure_jet_design_matrix(centers.view(), centers.view(), length_scale)?;
1524 let head_cc = centers.dot(&head_transform);
1525 let (z, coefficient_gauge) = match &spec.identifiability {
1540 MeasureJetIdentifiability::FrozenTransform { transform } => {
1541 if transform.nrows() != m_aug {
1542 crate::bail_dim_basis!(
1543 "frozen measure-jet identifiability transform mismatch: {} representers + {} head columns but transform has {} rows",
1544 m,
1545 head_rank,
1546 transform.nrows()
1547 );
1548 }
1549 (
1550 transform.clone(),
1551 gam_problem::Gauge::from_block_transforms(&[transform.clone()]),
1552 )
1553 }
1554 MeasureJetIdentifiability::CenterSumToZero => {
1555 let z_rbf = if head_rank > 0 {
1556 let affine = measure_jet_affine_value_basis(centers.view(), masses.view());
1557 let mut weighted_affine = affine.clone();
1558 for (i, mut row) in weighted_affine.outer_iter_mut().enumerate() {
1559 row.mapv_inplace(|v| v * masses[i]);
1560 }
1561 let constraint_cross = k_cc.t().dot(&weighted_affine);
1565 rrqr_nullspace_basis(&constraint_cross, default_rrqr_rank_alpha())
1566 .map_err(BasisError::LinalgError)?
1567 .0
1568 } else {
1569 let u = householder_sum_to_zero_u(m);
1570 householder_sum_to_zero_z(&u)
1571 };
1572 let rbf_rank = z_rbf.ncols();
1573 let mut z_block = Array2::<f64>::zeros((m_aug, rbf_rank + head_rank));
1574 z_block
1575 .slice_mut(ndarray::s![..m, ..rbf_rank])
1576 .assign(&z_rbf);
1577 for r in 0..head_rank {
1578 z_block[(m + r, rbf_rank + r)] = 1.0;
1579 }
1580 (
1581 z_block.clone(),
1582 gam_problem::Gauge::from_block_transforms(&[z_block]),
1583 )
1584 }
1585 };
1586 let mut k_aug = Array2::<f64>::zeros((m, m_aug));
1591 k_aug.slice_mut(ndarray::s![.., ..m]).assign(&k_cc);
1592 if head_rank > 0 {
1593 k_aug.slice_mut(ndarray::s![.., m..]).assign(&head_cc);
1594 }
1595 let kz = coefficient_gauge.restrict_design(&k_aug);
1596 Ok(RealizedMeasureJetGeometry {
1597 centers,
1598 masses,
1599 eps_band,
1600 log_step,
1601 length_scale,
1602 order_s_eval: order_s,
1603 per_level: spec.multiscale,
1608 z,
1609 coefficient_gauge,
1610 kz,
1611 head_transform,
1612 })
1613}
1614
1615pub fn measure_jet_input_noise_scale(
1638 data: ArrayView2<'_, f64>,
1639 centers: ArrayView2<'_, f64>,
1640) -> Result<Option<f64>, BasisError> {
1641 let d = data.ncols();
1642 let m = centers.nrows();
1643 if d == 0 || m == 0 || data.nrows() == 0 {
1644 return Ok(None);
1645 }
1646 if centers.ncols() != d {
1647 crate::bail_dim_basis!(
1648 "measure-jet input-noise estimate: data d={d} disagrees with centers d={}",
1649 centers.ncols()
1650 );
1651 }
1652 validate_finite_points(data, "data")?;
1653 validate_finite_points(centers, "centers")?;
1654 let sq = pairwise_sq_dists(data, centers);
1657 let mut members: Vec<Vec<usize>> = vec![Vec::new(); m];
1658 for (j, row) in sq.axis_iter(Axis(0)).enumerate() {
1659 let mut best = 0usize;
1660 let mut best_d = f64::INFINITY;
1661 for (i, &dij) in row.iter().enumerate() {
1662 if dij < best_d {
1663 best_d = dij;
1664 best = i;
1665 }
1666 }
1667 members[best].push(j);
1668 }
1669 let mut weighted_sum = 0.0_f64;
1670 let mut weight = 0.0_f64;
1671 for cell in &members {
1672 let n_i = cell.len();
1673 if n_i < d + 1 {
1676 continue;
1677 }
1678 let mut mean = Array1::<f64>::zeros(d);
1680 for &j in cell {
1681 mean += &data.row(j);
1682 }
1683 mean /= n_i as f64;
1684 let mut cov = Array2::<f64>::zeros((d, d));
1685 for &j in cell {
1686 let mut centered = data.row(j).to_owned();
1687 centered -= &mean;
1688 for a in 0..d {
1689 for b in 0..d {
1690 cov[(a, b)] += centered[a] * centered[b];
1691 }
1692 }
1693 }
1694 cov /= n_i as f64;
1695 let cov_sym = (&cov + &cov.t()) * 0.5;
1698 let (evals, _) = cov_sym.eigh(Side::Lower).map_err(|e| {
1699 BasisError::InvalidInput(format!(
1700 "measure-jet input-noise estimate: local covariance eigendecomposition failed: {e}"
1701 ))
1702 })?;
1703 let smallest = evals
1704 .iter()
1705 .copied()
1706 .fold(f64::INFINITY, |acc, v| acc.min(v))
1707 .max(0.0);
1708 if smallest.is_finite() {
1709 weighted_sum += n_i as f64 * smallest;
1710 weight += n_i as f64;
1711 }
1712 }
1713 if weight <= 0.0 {
1714 return Ok(None);
1715 }
1716 let sigma2 = weighted_sum / weight;
1717 if !(sigma2.is_finite() && sigma2 > 0.0) {
1718 return Ok(None);
1719 }
1720 Ok(Some(sigma2.sqrt()))
1721}
1722
1723pub fn measure_jet_multiscale_mode(spec: &MeasureJetBasisSpec) -> bool {
1731 spec.multiscale
1732}
1733
1734pub fn build_measure_jet_basis(
1742 data: ArrayView2<'_, f64>,
1743 spec: &MeasureJetBasisSpec,
1744) -> Result<BasisBuildResult, BasisError> {
1745 let RealizedMeasureJetGeometry {
1746 centers,
1747 masses,
1748 eps_band,
1749 log_step,
1750 length_scale,
1751 order_s_eval: order_s,
1752 per_level,
1753 z,
1754 coefficient_gauge,
1755 kz,
1756 head_transform,
1757 } = realize_measure_jet_geometry(data, spec)?;
1758 let band = MeasureJetBand {
1759 eps: eps_band.clone(),
1760 log_step,
1761 };
1762 let m = centers.nrows();
1763 let head_rank = head_transform.ncols();
1764 let m_aug = m + head_rank;
1765 let kernel_design = measure_jet_design_matrix(data, centers.view(), length_scale)?;
1770 let mut raw_design = Array2::<f64>::zeros((data.nrows(), m_aug));
1771 raw_design
1772 .slice_mut(ndarray::s![.., ..m])
1773 .assign(&kernel_design);
1774 if head_rank > 0 {
1775 let head_design = data.dot(&head_transform);
1776 raw_design
1777 .slice_mut(ndarray::s![.., m..])
1778 .assign(&head_design);
1779 }
1780 let constrained_design = coefficient_gauge.restrict_design(&raw_design);
1781 let design = gam_linalg::matrix::DesignMatrix::Dense(
1782 gam_linalg::matrix::DenseDesignMatrix::from(constrained_design),
1783 );
1784 let support_means = measure_jet_support_means(centers.view(), masses.view(), &eps_band)?;
1785 let mut candidates = Vec::new();
1798 let mut penalty_normalization_scales = Vec::new();
1799 let mut raw_penalty_normalization_scales = Vec::new();
1800 let mut fused_penalty_normalization_scale = None;
1801 if per_level {
1802 let forms = measure_jet_energy_forms_per_scale(
1803 centers.view(),
1804 masses.view(),
1805 &band,
1806 order_s,
1807 spec.alpha,
1808 spec.tau0,
1809 )?;
1810 for (level, q_l) in forms.into_iter().enumerate() {
1811 let s_l = kz.t().dot(&q_l).dot(&kz);
1812 let (s_norm, c_l) = normalize_penalty(&((&s_l + &s_l.t()) * 0.5));
1813 let intrinsic_dim = centers.ncols() as f64;
1814 let eta = 2.0 * order_s + intrinsic_dim * (2.0 - 2.0 * spec.alpha);
1815 let scale_weight = log_step * eps_band[level].powf(-eta);
1816 penalty_normalization_scales.push(c_l);
1817 raw_penalty_normalization_scales.push(c_l / scale_weight);
1818 candidates.push(PenaltyCandidate {
1819 matrix: ConstructiveQuadratic::try_from_dense_psd(
1820 s_norm,
1821 "measure-jet scale penalty",
1822 )?,
1823 source: PenaltySource::Other(format!("measure_jet_scale_{level}")),
1824 normalization_scale: c_l,
1825 kronecker_factors: None,
1826 op: None,
1827 });
1828 }
1829 } else {
1830 let q_form = measure_jet_energy_form(
1831 centers.view(),
1832 masses.view(),
1833 &band,
1834 order_s,
1835 spec.alpha,
1836 spec.tau0,
1837 )?;
1838 let penalty = pullback_center_form(&kz, &q_form);
1843 let (penalty_norm, c_primary) = normalize_penalty(&penalty);
1844 fused_penalty_normalization_scale = Some(c_primary);
1845 candidates.push(PenaltyCandidate {
1846 matrix: ConstructiveQuadratic::try_from_dense_psd(
1847 penalty_norm,
1848 "measure-jet primary penalty",
1849 )?,
1850 source: PenaltySource::Primary,
1851 normalization_scale: c_primary,
1852 kronecker_factors: None,
1853 op: None,
1854 });
1855 }
1856 if spec.double_penalty {
1862 let null_penalty = affine_function_nullspace_quadratic(&kz, centers.view(), masses.view())?;
1863 let (_, c_null) = normalize_penalty(null_penalty.dense());
1864 candidates.push(PenaltyCandidate {
1865 matrix: null_penalty
1866 .scaled(1.0 / c_null, "normalized measure-jet null-function penalty")?,
1867 source: PenaltySource::DoublePenaltyNullspace,
1868 normalization_scale: c_null,
1869 kronecker_factors: None,
1870 op: None,
1871 });
1872 }
1873 let filtered = filter_penalty_candidates(candidates)?;
1874 let sigma_coord = measure_jet_input_noise_scale(data, centers.view())?;
1877 Ok(BasisBuildResult {
1878 design,
1879 affine_offset: None,
1880 active_penalties: filtered.active,
1881 dropped_penalties: filtered.dropped,
1882 metadata: BasisMetadata::MeasureJet {
1883 centers,
1884 input_scale: crate::IsotropicScale::ONE,
1885 length_scale,
1886 eps_band,
1887 order_s: spec.order_s,
1892 alpha: spec.alpha,
1893 tau0: spec.tau0,
1894 masses,
1895 support_means,
1896 penalty_normalization_scales,
1897 raw_penalty_normalization_scales,
1898 fused_penalty_normalization_scale,
1899 constraint_transform: Some(z),
1900 sigma_coord,
1903 },
1904 kronecker_factored: None,
1905 joint_null_rotation: None,
1906 })
1907}
1908
1909pub fn build_measure_jet_basis_psi_derivatives(
1932 data: ArrayView2<'_, f64>,
1933 spec: &MeasureJetBasisSpec,
1934) -> Result<AnisoBasisPsiDerivatives, BasisError> {
1935 if !(spec.tau0.is_finite() && spec.tau0 > 0.0) {
1936 crate::bail_invalid_basis!(
1937 "measure-jet ψ derivatives need tau0 > 0 because the retained τ coordinate is ln τ; got {}",
1938 spec.tau0
1939 );
1940 }
1941 let geom = realize_measure_jet_geometry(data, spec)?;
1942 let band = MeasureJetBand {
1943 eps: geom.eps_band.clone(),
1944 log_step: geom.log_step,
1945 };
1946 let n = data.nrows();
1947 let p = geom.kz.ncols();
1948 let m = geom.centers.nrows();
1949 let m_aug = m + geom.head_transform.ncols();
1950
1951 struct LengthScaleJets {
1952 evaluation_first: Array2<f64>,
1953 evaluation_second: Array2<f64>,
1954 design_first: Array2<f64>,
1955 design_second: Array2<f64>,
1956 }
1957
1958 let length_scale_jets = if spec.learn_length_scale {
1965 let (dk_data, d2k_data) =
1966 measure_jet_design_log_length_jets(data, geom.centers.view(), geom.length_scale)?;
1967 let mut dk_data_aug = Array2::<f64>::zeros((n, m_aug));
1968 let mut d2k_data_aug = Array2::<f64>::zeros((n, m_aug));
1969 dk_data_aug.slice_mut(ndarray::s![.., ..m]).assign(&dk_data);
1970 d2k_data_aug
1971 .slice_mut(ndarray::s![.., ..m])
1972 .assign(&d2k_data);
1973
1974 let (dk_centers, d2k_centers) = measure_jet_design_log_length_jets(
1975 geom.centers.view(),
1976 geom.centers.view(),
1977 geom.length_scale,
1978 )?;
1979 let mut dk_centers_aug = Array2::<f64>::zeros((m, m_aug));
1980 let mut d2k_centers_aug = Array2::<f64>::zeros((m, m_aug));
1981 dk_centers_aug
1982 .slice_mut(ndarray::s![.., ..m])
1983 .assign(&dk_centers);
1984 d2k_centers_aug
1985 .slice_mut(ndarray::s![.., ..m])
1986 .assign(&d2k_centers);
1987
1988 Some(LengthScaleJets {
1989 evaluation_first: geom.coefficient_gauge.restrict_design(&dk_centers_aug),
1990 evaluation_second: geom.coefficient_gauge.restrict_design(&d2k_centers_aug),
1991 design_first: geom.coefficient_gauge.restrict_design(&dk_data_aug),
1992 design_second: geom.coefficient_gauge.restrict_design(&d2k_data_aug),
1993 })
1994 } else {
1995 None
1996 };
1997
1998 let coord_offset = usize::from(length_scale_jets.is_some());
1999 let n_coords = coord_offset + if geom.per_level { 2 } else { 0 };
2000 let pairs: Vec<(usize, usize)> = (0..n_coords)
2001 .flat_map(|a| ((a + 1)..n_coords).map(move |b| (a, b)))
2002 .collect();
2003 let zero_p = || Array2::<f64>::zeros((p, p));
2004
2005 struct RawPenaltyJets {
2006 value: Array2<f64>,
2007 first: Vec<Array2<f64>>,
2008 second_diag: Vec<Array2<f64>>,
2009 cross: Vec<Array2<f64>>,
2010 }
2011
2012 let sandwich = |form: &Array2<f64>| pullback_center_form(&geom.kz, form);
2013 let length_diag = |form: &Array2<f64>| {
2014 let jets = length_scale_jets
2015 .as_ref()
2016 .expect("length-scale form jets require an enrolled length coordinate");
2017 pullback_center_form_log_length_jets(
2018 &geom.kz,
2019 &jets.evaluation_first,
2020 &jets.evaluation_second,
2021 form,
2022 )
2023 };
2024 let length_cross = |form_first: &Array2<f64>| {
2025 let jets = length_scale_jets
2026 .as_ref()
2027 .expect("length-scale cross jets require an enrolled length coordinate");
2028 pullback_center_form_log_length_cross(&geom.kz, &jets.evaluation_first, form_first)
2029 };
2030
2031 let mut raw: Vec<RawPenaltyJets> = if geom.per_level {
2037 let l_count = band.eps.len();
2038 let forms = assemble_weighted_forms(
2041 geom.centers.view(),
2042 geom.masses.view(),
2043 &band,
2044 geom.order_s_eval,
2045 spec.alpha,
2046 spec.tau0,
2047 6 * l_count,
2048 3,
2049 &|scale_idx, eps: f64, q: f64, base: f64, out: &mut [[f64; 3]]| {
2050 for slot in out.iter_mut() {
2051 *slot = [0.0, 0.0, 0.0];
2052 }
2053 let intrinsic_dim = geom.centers.ncols() as f64;
2054 let ga = 2.0 * intrinsic_dim * eps.ln() - 2.0 * q.max(f64::MIN_POSITIVE).ln();
2055 let k0 = 6 * scale_idx;
2056 out[k0] = [base, 0.0, 0.0];
2057 out[k0 + 1] = [ga * base, 0.0, 0.0];
2058 out[k0 + 2] = [ga * ga * base, 0.0, 0.0];
2059 out[k0 + 3] = [0.0, 0.0, 0.0];
2060 out[k0 + 4] = [0.0, 0.0, 0.0];
2061 out[k0 + 5] = [0.0, 0.0, 0.0];
2062 },
2063 )?;
2064 let alpha_coord = coord_offset;
2065 let tau_coord = coord_offset + 1;
2066 let mut raw = Vec::with_capacity(l_count + usize::from(spec.double_penalty));
2067 for level in 0..l_count {
2068 let chunk = &forms[6 * level..6 * level + 6];
2069 let mut first: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2070 let mut second_diag: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2071 first[alpha_coord] = sandwich(&chunk[1]);
2072 first[tau_coord] = sandwich(&chunk[3]);
2073 second_diag[alpha_coord] = sandwich(&chunk[2]);
2074 second_diag[tau_coord] = sandwich(&chunk[4]);
2075 if coord_offset == 1 {
2076 let (ell_first, ell_second) = length_diag(&chunk[0]);
2077 first[0] = ell_first;
2078 second_diag[0] = ell_second;
2079 }
2080 let mut cross: Vec<Array2<f64>> = (0..pairs.len()).map(|_| zero_p()).collect();
2081 for (pair_idx, &(a, b)) in pairs.iter().enumerate() {
2082 cross[pair_idx] = if coord_offset == 1 && a == 0 && b == alpha_coord {
2083 length_cross(&chunk[1])
2084 } else if coord_offset == 1 && a == 0 && b == tau_coord {
2085 length_cross(&chunk[3])
2086 } else if a == alpha_coord && b == tau_coord {
2087 sandwich(&chunk[5])
2088 } else {
2089 zero_p()
2090 };
2091 }
2092 raw.push(RawPenaltyJets {
2093 value: sandwich(&chunk[0]),
2094 first,
2095 second_diag,
2096 cross,
2097 });
2098 }
2099 raw
2100 } else {
2101 let q_form = measure_jet_energy_form(
2105 geom.centers.view(),
2106 geom.masses.view(),
2107 &band,
2108 geom.order_s_eval,
2109 spec.alpha,
2110 spec.tau0,
2111 )?;
2112 let mut first: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2113 let mut second_diag: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2114 if coord_offset == 1 {
2115 let (ell_first, ell_second) = length_diag(&q_form);
2116 first[0] = ell_first;
2117 second_diag[0] = ell_second;
2118 }
2119 vec![RawPenaltyJets {
2120 value: sandwich(&q_form),
2121 first,
2122 second_diag,
2123 cross: Vec::new(),
2124 }]
2125 };
2126
2127 if spec.double_penalty {
2128 let null_form =
2129 affine_function_nullspace_center_quadratic(geom.centers.view(), geom.masses.view())?;
2130 let null_form = null_form.dense();
2131 let mut first: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2132 let mut second_diag: Vec<Array2<f64>> = (0..n_coords).map(|_| zero_p()).collect();
2133 if coord_offset == 1 {
2134 let (ell_first, ell_second) = length_diag(null_form);
2135 first[0] = ell_first;
2136 second_diag[0] = ell_second;
2137 }
2138 raw.push(RawPenaltyJets {
2139 value: sandwich(null_form),
2140 first,
2141 second_diag,
2142 cross: (0..pairs.len()).map(|_| zero_p()).collect(),
2145 });
2146 }
2147
2148 let n_cands = raw.len();
2149 let mut penalties_first: Vec<Vec<Array2<f64>>> =
2150 (0..n_coords).map(|_| Vec::with_capacity(n_cands)).collect();
2151 let mut penalties_second_diag: Vec<Vec<Array2<f64>>> =
2152 (0..n_coords).map(|_| Vec::with_capacity(n_cands)).collect();
2153 let mut crosses: Vec<Vec<Array2<f64>>> = (0..pairs.len()).map(|_| Vec::new()).collect();
2157 for candidate in &raw {
2158 let s_raw = &candidate.value;
2159 let fro = trace_of_product(s_raw, s_raw).sqrt();
2168 let c = if fro.is_finite() && fro > 1e-12 {
2169 fro
2170 } else {
2171 1.0
2172 };
2173 for coord in 0..n_coords {
2174 let (_, s_first, s_second, _) = normalize_penaltywith_psi_derivatives(
2175 s_raw,
2176 &candidate.first[coord],
2177 &candidate.second_diag[coord],
2178 );
2179 penalties_first[coord].push(s_first);
2180 penalties_second_diag[coord].push(s_second);
2181 }
2182 for (pair_idx, &(a, b)) in pairs.iter().enumerate() {
2183 let cross_raw_mat = normalize_penalty_cross_psi_derivative(
2184 s_raw,
2185 &candidate.first[a],
2186 &candidate.first[b],
2187 &candidate.cross[pair_idx],
2188 c,
2189 );
2190 crosses[pair_idx].push(cross_raw_mat);
2191 }
2192 }
2193
2194 let pair_index: Vec<((usize, usize), Vec<Array2<f64>>)> =
2195 pairs.iter().copied().zip(crosses.into_iter()).collect();
2196 let provider = AnisoPenaltyCrossProvider::new(move |a, b| {
2197 pair_index
2198 .iter()
2199 .find(|((pa, pb), _)| (*pa, *pb) == (a, b) || (*pa, *pb) == (b, a))
2200 .map(|(_, mats)| mats.clone())
2201 .ok_or_else(|| {
2202 BasisError::InvalidInput(format!(
2203 "measure-jet ψ cross derivative requested for unknown pair ({a}, {b})"
2204 ))
2205 })
2206 });
2207 let mut design_first: Vec<Array2<f64>> = (0..n_coords)
2208 .map(|_| Array2::<f64>::zeros((n, p)))
2209 .collect();
2210 let mut design_second_diag: Vec<Array2<f64>> = (0..n_coords)
2211 .map(|_| Array2::<f64>::zeros((n, p)))
2212 .collect();
2213 if let Some(jets) = &length_scale_jets {
2214 design_first[0] = jets.design_first.clone();
2215 design_second_diag[0] = jets.design_second.clone();
2216 }
2217 Ok(AnisoBasisPsiDerivatives {
2218 design_first,
2219 design_second_diag,
2220 design_second_cross: Vec::new(),
2221 design_second_cross_pairs: Vec::new(),
2222 penalties_first,
2223 penalties_second_diag,
2224 penalties_cross_pairs: pairs,
2225 penalties_cross_provider: Some(provider),
2226 implicit_operator: None,
2227 })
2228}
2229
2230#[cfg(test)]
2231mod tests {
2232 use super::*;
2233
2234 fn lcg_normal(state: &mut u64) -> f64 {
2237 let mut next = || {
2238 *state = state
2239 .wrapping_mul(6364136223846793005)
2240 .wrapping_add(1442695040888963407);
2241 (((*state >> 11) as f64) + 0.5) / (1u64 << 53) as f64
2243 };
2244 let u1 = next();
2245 let u2 = next();
2246 (-2.0 * u1.ln()).sqrt() * (std::f64::consts::TAU * u2).cos()
2247 }
2248
2249 #[test]
2255 pub(crate) fn input_noise_scale_recovers_known_perpendicular_sigma() {
2256 let tang = [1.0 / 5f64.sqrt(), 2.0 / 5f64.sqrt()];
2258 let perp = [2.0 / 5f64.sqrt(), -1.0 / 5f64.sqrt()];
2259 let sigma = 0.05_f64;
2260 let n = 600usize;
2261 let mut state = 0x1234_5678_9abc_def0u64;
2262 let mut data = Array2::<f64>::zeros((n, 2));
2263 for j in 0..n {
2264 let t = 3.0 * (j as f64) / (n as f64 - 1.0);
2266 let noise = sigma * lcg_normal(&mut state);
2267 for a in 0..2 {
2268 data[(j, a)] = t * tang[a] + noise * perp[a];
2269 }
2270 }
2271 let n_centers = 8usize;
2274 let mut centers = Array2::<f64>::zeros((n_centers, 2));
2275 for i in 0..n_centers {
2276 let t = 3.0 * (i as f64 + 0.5) / (n_centers as f64);
2277 for a in 0..2 {
2278 centers[(i, a)] = t * tang[a];
2279 }
2280 }
2281 let est = measure_jet_input_noise_scale(data.view(), centers.view())
2282 .expect("estimate ok")
2283 .expect("noise scale present");
2284 assert!(
2287 (est - sigma).abs() <= 0.4 * sigma,
2288 "estimated σ_coord {est} far from true {sigma}"
2289 );
2290 }
2291
2292 #[test]
2295 pub(crate) fn input_noise_scale_none_when_cells_too_small() {
2296 let data = array![[0.0, 0.0], [1.0, 2.0], [2.0, 4.0]];
2297 let centers = array![[0.0, 0.0], [1.0, 2.0], [2.0, 4.0]];
2298 assert!(
2300 measure_jet_input_noise_scale(data.view(), centers.view())
2301 .expect("estimate ok")
2302 .is_none()
2303 );
2304 }
2305
2306 pub(crate) fn two_cluster_centers() -> (ndarray::Array2<f64>, ndarray::Array1<f64>) {
2307 let centers = array![
2308 [0.00, 0.00],
2309 [0.31, 0.05],
2310 [0.58, -0.07],
2311 [0.93, 0.11],
2312 [1.22, 0.02],
2313 [1.49, -0.04],
2314 [3.10, 2.00],
2315 [3.42, 2.13],
2316 [3.71, 1.91],
2317 [4.05, 2.07],
2318 [4.33, 1.96],
2319 [4.61, 2.12],
2320 ];
2321 let m = centers.nrows();
2322 let masses = ndarray::Array1::<f64>::from_elem(m, 1.0 / m as f64);
2323 (centers, masses)
2324 }
2325 use ndarray::array;
2326
2327 pub(crate) fn band_for(centers: &Array2<f64>) -> MeasureJetBand {
2328 measure_jet_band(centers.view(), 0).expect("band")
2329 }
2330
2331 #[test]
2334 pub(crate) fn energy_form_annihilates_constants_exactly() {
2335 let (centers, masses) = two_cluster_centers();
2336 let band = band_for(¢ers);
2337 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2338 .expect("energy form");
2339 let m = q.nrows();
2340 let ones = Array1::<f64>::ones(m);
2341 let qv = q.dot(&ones);
2342 let scale = q.iter().fold(0.0_f64, |acc, v| acc.max(v.abs()));
2343 assert!(scale > 0.0, "energy form is identically zero");
2344 for (i, v) in qv.iter().enumerate() {
2345 assert!(
2346 v.abs() <= 1e-12 * scale,
2347 "Q·1 leak at row {i}: {v:.3e} vs scale {scale:.3e}"
2348 );
2349 }
2350 let vqv = ones.dot(&qv);
2351 assert!(
2352 vqv.abs() <= 1e-12 * scale,
2353 "constant carries energy: 1ᵀQ1 = {vqv:.3e}"
2354 );
2355 }
2356
2357 #[test]
2360 pub(crate) fn energy_form_annihilates_affine_at_default_tau() {
2361 let (centers, masses) = two_cluster_centers();
2362 let band = band_for(¢ers);
2363 let m = centers.nrows();
2364 let mut affine = Array1::<f64>::zeros(m);
2366 let mut rough = Array1::<f64>::zeros(m);
2367 for i in 0..m {
2368 affine[i] = 0.7 + 1.3 * centers[(i, 0)] - 0.4 * centers[(i, 1)];
2369 rough[i] = if i % 2 == 0 { 1.0 } else { -1.0 };
2370 }
2371 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2372 .expect("energy form");
2373 let e_affine = affine.dot(&q.dot(&affine));
2374 let e_rough = rough.dot(&q.dot(&rough));
2375 assert!(e_rough > 0.0, "rough vector must pay energy");
2376 assert!(
2377 e_affine.abs() <= 1e-12 * e_rough,
2378 "default affine energy {e_affine:.3e} vs rough {e_rough:.3e}"
2379 );
2380 }
2381
2382 #[test]
2384 pub(crate) fn energy_form_is_psd() {
2385 let (centers, masses) = two_cluster_centers();
2386 let band = band_for(¢ers);
2387 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2388 .expect("energy form");
2389 let m = q.nrows();
2390 for trial in 0..5usize {
2391 let v = Array1::<f64>::from_shape_fn(m, |i| {
2392 ((i * 7 + trial * 13) % 11) as f64 / 11.0 - 0.5
2393 });
2394 let e = v.dot(&q.dot(&v));
2395 assert!(e >= -1e-10, "vᵀQv = {e:.3e} < 0 on trial {trial}");
2396 }
2397 }
2398
2399 #[test]
2402 pub(crate) fn rough_vector_pays_more_than_smooth() {
2403 let m = 24usize;
2404 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2405 let t = i as f64 / (m as f64 - 1.0);
2406 if k == 0 {
2407 t * 4.0
2408 } else {
2409 0.3 * (t * 4.0).sin()
2410 }
2411 });
2412 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2413 let band = band_for(¢ers);
2414 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2415 .expect("energy form");
2416 let slow = Array1::<f64>::from_shape_fn(m, |i| (i as f64 / (m as f64 - 1.0)).powi(2));
2417 let fast = Array1::<f64>::from_shape_fn(m, |i| if i % 2 == 0 { 0.5 } else { -0.5 });
2418 let e_slow = slow.dot(&q.dot(&slow));
2419 let e_fast = fast.dot(&q.dot(&fast));
2420 assert!(
2421 e_fast > 10.0 * e_slow,
2422 "alternating values must pay >> a slow trend: fast {e_fast:.3e} vs slow {e_slow:.3e}"
2423 );
2424 }
2425
2426 #[test]
2431 pub(crate) fn energy_jets_match_finite_differences() {
2432 let (centers, masses) = two_cluster_centers();
2433 let band = band_for(¢ers);
2434 let (s0, a0, tau) = (1.3, 0.8, 1e-3);
2435 let jets =
2436 measure_jet_energy_form_with_jets(centers.view(), masses.view(), &band, s0, a0, tau)
2437 .expect("jets");
2438 let q_at = |s: f64, a: f64| {
2439 measure_jet_energy_form(centers.view(), masses.view(), &band, s, a, tau)
2440 .expect("energy form")
2441 };
2442 let q_plain = q_at(s0, a0);
2444 for (a, b) in jets.q.iter().zip(q_plain.iter()) {
2445 assert!(
2446 (a - b).abs() <= 1e-14 * (1.0 + b.abs()),
2447 "Q drift {a} vs {b}"
2448 );
2449 }
2450 let lt0 = tau.ln();
2451 let q_at_lt = |lt: f64| {
2452 measure_jet_energy_form(centers.view(), masses.view(), &band, s0, a0, lt.exp())
2453 .expect("energy form")
2454 };
2455 let h = 1e-4;
2461 let checks: [(&str, &Array2<f64>, Array2<f64>); 9] = [
2462 ("dq_ds", &jets.dq_ds, {
2463 let (p, m_) = (q_at(s0 + h, a0), q_at(s0 - h, a0));
2464 (&p - &m_) / (2.0 * h)
2465 }),
2466 ("d2q_ds2", &jets.d2q_ds2, {
2467 let (p, c, m_) = (q_at(s0 + h, a0), q_at(s0, a0), q_at(s0 - h, a0));
2468 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2469 }),
2470 ("dq_dalpha", &jets.dq_dalpha, {
2471 let (p, m_) = (q_at(s0, a0 + h), q_at(s0, a0 - h));
2472 (&p - &m_) / (2.0 * h)
2473 }),
2474 ("d2q_dalpha2", &jets.d2q_dalpha2, {
2475 let (p, c, m_) = (q_at(s0, a0 + h), q_at(s0, a0), q_at(s0, a0 - h));
2476 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2477 }),
2478 ("d2q_ds_dalpha", &jets.d2q_ds_dalpha, {
2479 let pp = q_at(s0 + h, a0 + h);
2480 let pm = q_at(s0 + h, a0 - h);
2481 let mp = q_at(s0 - h, a0 + h);
2482 let mm = q_at(s0 - h, a0 - h);
2483 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2484 }),
2485 ("dq_dlogtau", &jets.dq_dlogtau, {
2486 let (p, m_) = (q_at_lt(lt0 + h), q_at_lt(lt0 - h));
2487 (&p - &m_) / (2.0 * h)
2488 }),
2489 ("d2q_dlogtau2", &jets.d2q_dlogtau2, {
2490 let (p, c, m_) = (q_at_lt(lt0 + h), q_at_lt(lt0), q_at_lt(lt0 - h));
2491 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2492 }),
2493 ("d2q_ds_dlogtau", &jets.d2q_ds_dlogtau, {
2494 let f = |s: f64, lt: f64| {
2495 measure_jet_energy_form(centers.view(), masses.view(), &band, s, a0, lt.exp())
2496 .expect("energy form")
2497 };
2498 let pp = f(s0 + h, lt0 + h);
2499 let pm = f(s0 + h, lt0 - h);
2500 let mp = f(s0 - h, lt0 + h);
2501 let mm = f(s0 - h, lt0 - h);
2502 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2503 }),
2504 ("d2q_dalpha_dlogtau", &jets.d2q_dalpha_dlogtau, {
2505 let f = |a: f64, lt: f64| {
2506 measure_jet_energy_form(centers.view(), masses.view(), &band, s0, a, lt.exp())
2507 .expect("energy form")
2508 };
2509 let pp = f(a0 + h, lt0 + h);
2510 let pm = f(a0 + h, lt0 - h);
2511 let mp = f(a0 - h, lt0 + h);
2512 let mm = f(a0 - h, lt0 - h);
2513 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2514 }),
2515 ];
2516 for (name, analytic, fd) in checks.iter() {
2517 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2518 for (a, b) in analytic.iter().zip(fd.iter()) {
2519 assert!(
2520 (a - b).abs() <= 5e-5 * scale,
2521 "{name} jet mismatch: analytic {a:.6e} vs FD {b:.6e} (scale {scale:.3e})"
2522 );
2523 }
2524 }
2525 }
2526
2527 #[test]
2531 pub(crate) fn scale_spectrum_sums_to_total_and_localizes_roughness() {
2532 let m = 24usize;
2533 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2534 let t = i as f64 / (m as f64 - 1.0);
2535 if k == 0 { t * 4.0 } else { 0.0 }
2536 });
2537 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2538 let band = band_for(¢ers);
2539 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2540 .expect("energy form");
2541 let fast = Array1::<f64>::from_shape_fn(m, |i| if i % 2 == 0 { 0.5 } else { -0.5 });
2542 let spec = measure_jet_scale_spectrum(
2543 centers.view(),
2544 masses.view(),
2545 &band,
2546 1.5,
2547 1.0,
2548 1e-3,
2549 fast.view(),
2550 )
2551 .expect("spectrum");
2552 assert_eq!(spec.len(), band.eps.len());
2553 let total = fast.dot(&q.dot(&fast));
2554 let sum: f64 = spec.iter().sum();
2555 assert!(
2556 (sum - total).abs() <= 1e-10 * total.abs().max(1e-30),
2557 "spectrum must sum to vᵀQv: {sum:.6e} vs {total:.6e}"
2558 );
2559 let finest = spec[0];
2561 let coarsest = *spec.last().expect("nonempty spectrum");
2562 assert!(
2563 finest > coarsest,
2564 "alternating values must charge fine scales hardest: fine {finest:.3e} vs coarse {coarsest:.3e}"
2565 );
2566 }
2567
2568 #[test]
2571 pub(crate) fn support_curve_separates_on_web_from_off_web() {
2572 let m = 24usize;
2573 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2574 let t = i as f64 / (m as f64 - 1.0);
2575 if k == 0 { t * 4.0 } else { 0.0 }
2576 });
2577 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2578 let band = band_for(¢ers);
2579 let queries = array![[2.0, 0.0], [2.0, 1.5]];
2580 let curves =
2581 measure_jet_support_curve(queries.view(), centers.view(), masses.view(), &band.eps)
2582 .expect("support curve");
2583 assert!(
2585 curves[(0, 0)] > 10.0 * curves[(1, 0)],
2586 "fine-scale support must separate web from void: on {:.3e} vs off {:.3e}",
2587 curves[(0, 0)],
2588 curves[(1, 0)]
2589 );
2590 for qi in 0..2 {
2592 for li in 1..band.eps.len() {
2593 assert!(
2594 curves[(qi, li)] >= curves[(qi, li - 1)] - 1e-15,
2595 "support curve must be monotone in scale (query {qi}, level {li})"
2596 );
2597 }
2598 }
2599 }
2600
2601 #[test]
2609 pub(crate) fn default_stays_single_scale_until_multiscale_opt_in() {
2610 let n = 200usize;
2611 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2612 let t = i as f64 / (n as f64 - 1.0);
2613 if k == 0 {
2614 t * 3.0
2615 } else {
2616 0.4 * (t * 3.0).sin()
2617 }
2618 });
2619 let single = MeasureJetBasisSpec {
2623 center_strategy: CenterStrategy::FarthestPoint { num_centers: 80 },
2624 ..MeasureJetBasisSpec::default()
2625 };
2626 assert!(
2627 !measure_jet_multiscale_mode(&single),
2628 "default must resolve to single-scale at any center count"
2629 );
2630 let built_single =
2631 build_measure_jet_basis(data.view(), &single).expect("single-scale build");
2632 assert_eq!(
2633 built_single.active_penalties.len(),
2634 2,
2635 "single-scale double-penalty mode emits Primary + affine/null component"
2636 );
2637 assert!(matches!(
2638 built_single.active_penalties[0].info.source,
2639 PenaltySource::Primary
2640 ));
2641 assert!(matches!(
2642 built_single.active_penalties[1].info.source,
2643 PenaltySource::DoublePenaltyNullspace
2644 ));
2645 let multi = MeasureJetBasisSpec {
2649 center_strategy: CenterStrategy::FarthestPoint { num_centers: 80 },
2650 multiscale: true,
2651 ..MeasureJetBasisSpec::default()
2652 };
2653 assert!(
2654 measure_jet_multiscale_mode(&multi),
2655 "multiscale=true must resolve to multiscale mode"
2656 );
2657 let built_multi = build_measure_jet_basis(data.view(), &multi).expect("multiscale build");
2658 assert!(
2659 built_multi.active_penalties.len() > built_single.active_penalties.len(),
2660 "multiscale mode emits the per-scale spectral split plus null selection, got {} (vs single-scale {})",
2661 built_multi.active_penalties.len(),
2662 built_single.active_penalties.len()
2663 );
2664 }
2665
2666 #[test]
2670 pub(crate) fn fused_mode_without_double_penalty_emits_single_primary_candidate() {
2671 let n = 40usize;
2672 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2673 let t = i as f64 / (n as f64 - 1.0);
2674 if k == 0 {
2675 t * 3.0
2676 } else {
2677 0.4 * (t * 3.0).sin()
2678 }
2679 });
2680 let spec = MeasureJetBasisSpec {
2681 center_strategy: CenterStrategy::FarthestPoint { num_centers: 14 },
2682 order_s: 1.3,
2683 double_penalty: false,
2684 ..MeasureJetBasisSpec::default()
2685 };
2686 let built = build_measure_jet_basis(data.view(), &spec).expect("fused build");
2687 assert_eq!(
2688 built.active_penalties.len(),
2689 1,
2690 "single-scale mode without null recovery emits exactly one Primary"
2691 );
2692 assert!(matches!(
2693 built.active_penalties[0].info.source,
2694 PenaltySource::Primary
2695 ));
2696 let BasisMetadata::MeasureJet { order_s, .. } = &built.metadata else {
2697 panic!("measure-jet build must return MeasureJet metadata");
2698 };
2699 assert_eq!(*order_s, 1.3, "explicit order must persist verbatim");
2700 }
2701
2702 #[test]
2707 pub(crate) fn single_scale_affine_head_gauge_annihilates_center_cross() {
2708 let n = 90usize;
2709 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2710 let t = i as f64 / (n as f64 - 1.0);
2711 if k == 0 {
2712 3.0 * t
2713 } else {
2714 (2.0 * std::f64::consts::PI * t).sin() + 0.2 * t
2715 }
2716 });
2717 let spec = MeasureJetBasisSpec {
2718 center_strategy: CenterStrategy::FarthestPoint { num_centers: 18 },
2719 double_penalty: false,
2720 multiscale: false,
2721 ..MeasureJetBasisSpec::default()
2722 };
2723 let geom = realize_measure_jet_geometry(data.view(), &spec).expect("realized geometry");
2724 let m = geom.centers.nrows();
2725 let head_rank = geom.head_transform.ncols();
2726 assert!(head_rank > 0, "fixture must realize an affine head");
2727 assert_eq!(
2728 geom.z.ncols(),
2729 m - 1,
2730 "affine gauge replaces duplicated RBF directions without widening the smooth"
2731 );
2732 let rbf_rank = m - (head_rank + 1);
2733 let z_rbf = geom.z.slice(ndarray::s![..m, ..rbf_rank]).to_owned();
2734 let k_cc =
2735 measure_jet_design_matrix(geom.centers.view(), geom.centers.view(), geom.length_scale)
2736 .expect("center kernel");
2737 let affine = measure_jet_affine_value_basis(geom.centers.view(), geom.masses.view());
2738 assert_eq!(affine.ncols(), head_rank + 1);
2739 let mut weighted_affine = affine.clone();
2740 for (i, mut row) in weighted_affine.outer_iter_mut().enumerate() {
2741 row.mapv_inplace(|v| v * geom.masses[i]);
2742 }
2743 let constraint_cross = k_cc.t().dot(&weighted_affine);
2744 let residual = constraint_cross.t().dot(&z_rbf);
2745 let scale = constraint_cross
2746 .iter()
2747 .fold(1.0_f64, |acc, value| acc.max(value.abs()));
2748 assert!(
2749 residual.iter().all(|value| value.abs() <= 1e-10 * scale),
2750 "A^T W Kcc Z_rbf must vanish; max residual {:.3e}",
2751 residual
2752 .iter()
2753 .fold(0.0_f64, |acc, value| acc.max(value.abs()))
2754 );
2755 }
2756
2757 #[test]
2760 pub(crate) fn affine_null_penalty_is_covariant_under_coefficient_reparameterization() {
2761 let centers = array![
2762 [-1.0, 0.2],
2763 [-0.4, -0.3],
2764 [0.1, 0.5],
2765 [0.7, -0.2],
2766 [1.2, 0.4],
2767 [1.8, -0.1],
2768 ];
2769 let masses = array![0.08, 0.12, 0.18, 0.22, 0.17, 0.23];
2770 let evaluation = Array2::<f64>::from_shape_fn((centers.nrows(), 3), |(i, j)| {
2771 ((i + 2 * j + 1) as f64).sin() + 0.15 * (i * (j + 1)) as f64
2772 });
2773 let reparameterization = array![[1.7, 0.2, -0.1], [0.0, 0.6, 0.3], [0.0, 0.0, 1.3]];
2774 let base = affine_function_nullspace_quadratic(&evaluation, centers.view(), masses.view())
2775 .expect("base function-space penalty")
2776 .into_dense();
2777 let transformed_evaluation = evaluation.dot(&reparameterization);
2778 let transformed = affine_function_nullspace_quadratic(
2779 &transformed_evaluation,
2780 centers.view(),
2781 masses.view(),
2782 )
2783 .expect("reparameterized function-space penalty")
2784 .into_dense();
2785 let expected = reparameterization.t().dot(&base).dot(&reparameterization);
2786 let scale = expected
2787 .iter()
2788 .fold(1.0_f64, |acc, value| acc.max(value.abs()));
2789 assert!(
2790 transformed
2791 .iter()
2792 .zip(expected.iter())
2793 .all(|(actual, want)| (actual - want).abs() <= 1e-11 * scale),
2794 "S(E R) must equal R^T S(E) R"
2795 );
2796 }
2797
2798 #[test]
2801 pub(crate) fn double_penalty_leaves_primary_matrix_unchanged() {
2802 let n = 64usize;
2803 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2804 let t = i as f64 / (n as f64 - 1.0);
2805 if k == 0 { 2.5 * t } else { (4.0 * t).cos() }
2806 });
2807 let base = MeasureJetBasisSpec {
2808 center_strategy: CenterStrategy::FarthestPoint { num_centers: 16 },
2809 order_s: 1.25,
2810 double_penalty: false,
2811 ..MeasureJetBasisSpec::default()
2812 };
2813 let without = build_measure_jet_basis(data.view(), &base).expect("primary-only build");
2814 let with = build_measure_jet_basis(
2815 data.view(),
2816 &MeasureJetBasisSpec {
2817 double_penalty: true,
2818 ..base.clone()
2819 },
2820 )
2821 .expect("double-penalty build");
2822 assert_eq!(without.active_penalties.len(), 1);
2823 assert_eq!(with.active_penalties.len(), 2);
2824 assert!(matches!(
2825 without.active_penalties[0].info.source,
2826 PenaltySource::Primary
2827 ));
2828 assert!(matches!(
2829 with.active_penalties[0].info.source,
2830 PenaltySource::Primary
2831 ));
2832 assert!(matches!(
2833 with.active_penalties[1].info.source,
2834 PenaltySource::DoublePenaltyNullspace
2835 ));
2836 assert!(
2837 without.active_penalties[0]
2838 .matrix
2839 .iter()
2840 .zip(with.active_penalties[0].matrix.iter())
2841 .all(|(a, b)| (a - b).abs() <= 1e-13),
2842 "turning on null recovery must not modify Primary"
2843 );
2844 }
2845
2846 #[test]
2848 pub(crate) fn householder_sum_to_zero_basis_is_orthonormal() {
2849 let m = 9usize;
2850 let u = householder_sum_to_zero_u(m);
2851 let z = householder_sum_to_zero_z(&u);
2852 for j in 0..(m - 1) {
2853 let col_j = z.column(j);
2854 assert!(col_j.sum().abs() <= 1e-12, "column {j} must sum to zero");
2855 for j2 in j..(m - 1) {
2856 let dot = col_j.dot(&z.column(j2));
2857 let want = if j == j2 { 1.0 } else { 0.0 };
2858 assert!(
2859 (dot - want).abs() <= 1e-12,
2860 "orthonormality failure at ({j}, {j2}): {dot}"
2861 );
2862 }
2863 }
2864 }
2865
2866 pub(crate) fn frozen_spec_fixture(
2871 order_s: f64,
2872 multiscale: bool,
2873 ) -> (Array2<f64>, MeasureJetBasisSpec) {
2874 let n = 140usize;
2879 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2880 let t = i as f64 / (n as f64 - 1.0);
2881 if k == 0 {
2882 t * 3.0
2883 } else {
2884 0.5 * (t * 3.0).cos() + if i % 9 == 0 { 0.8 } else { 0.0 }
2885 }
2886 });
2887 let spec = MeasureJetBasisSpec {
2888 center_strategy: CenterStrategy::FarthestPoint { num_centers: 70 },
2889 order_s,
2890 multiscale,
2891 learn_length_scale: false,
2895 ..MeasureJetBasisSpec::default()
2896 };
2897 let first = build_measure_jet_basis(data.view(), &spec).expect("fixture build");
2898 let BasisMetadata::MeasureJet {
2899 centers,
2900 length_scale,
2901 eps_band,
2902 masses,
2903 support_means,
2904 penalty_normalization_scales,
2905 raw_penalty_normalization_scales,
2906 fused_penalty_normalization_scale,
2907 constraint_transform,
2908 ..
2909 } = &first.metadata
2910 else {
2911 panic!("measure-jet build must return MeasureJet metadata");
2912 };
2913 let frozen = MeasureJetBasisSpec {
2914 center_strategy: CenterStrategy::UserProvided(centers.clone()),
2915 order_s,
2916 alpha: spec.alpha,
2917 tau0: spec.tau0,
2918 num_scales: eps_band.len(),
2919 length_scale: *length_scale,
2920 double_penalty: spec.double_penalty,
2921 learn_length_scale: false,
2922 multiscale,
2923 identifiability: MeasureJetIdentifiability::FrozenTransform {
2924 transform: constraint_transform.clone().expect("fit-time z"),
2925 },
2926 frozen_quadrature: Some(MeasureJetFrozenQuadrature {
2927 masses: masses.clone(),
2928 eps_band: eps_band.clone(),
2929 support_means: support_means.clone(),
2930 penalty_normalization_scales: penalty_normalization_scales.clone(),
2931 raw_penalty_normalization_scales: raw_penalty_normalization_scales.clone(),
2932 fused_penalty_normalization_scale: *fused_penalty_normalization_scale,
2933 sigma_coord: None,
2934 }),
2935 };
2936 (data, frozen)
2937 }
2938
2939 #[test]
2944 pub(crate) fn psi_producer_matches_fd_per_level_mode() {
2945 let (data, frozen) = frozen_spec_fixture(0.0, true);
2946 let derivs =
2947 build_measure_jet_basis_psi_derivatives(data.view(), &frozen).expect("psi derivatives");
2948 let l_count = frozen
2949 .frozen_quadrature
2950 .as_ref()
2951 .expect("frozen quadrature")
2952 .eps_band
2953 .len();
2954 assert_eq!(
2955 derivs.penalties_first.len(),
2956 2,
2957 "per-level coords are (α, lnτ)"
2958 );
2959 assert_eq!(derivs.penalties_first[0].len(), l_count + 1);
2960 assert_eq!(derivs.penalties_cross_pairs, vec![(0, 1)]);
2961 let pen_at = |alpha: f64, tau0: f64| {
2962 let trial = MeasureJetBasisSpec {
2963 alpha,
2964 tau0,
2965 ..frozen.clone()
2966 };
2967 build_measure_jet_basis(data.view(), &trial)
2968 .expect("trial build")
2969 .active_penalties
2970 .into_iter()
2971 .map(|penalty| penalty.matrix)
2972 .collect::<Vec<_>>()
2973 };
2974 let h = 1e-4;
2977 let (a0, t0) = (frozen.alpha, frozen.tau0);
2978 let ap = pen_at(a0 + h, t0);
2979 let am = pen_at(a0 - h, t0);
2980 let tp = pen_at(a0, t0 * h.exp());
2981 let tm = pen_at(a0, t0 * (-h).exp());
2982 assert_eq!(
2983 ap.len(),
2984 l_count + 1,
2985 "fixture must keep every scale active"
2986 );
2987 for level in 0..l_count {
2988 let fd_alpha = (&ap[level] - &am[level]) / (2.0 * h);
2989 let fd_tau = (&tp[level] - &tm[level]) / (2.0 * h);
2990 for (name, analytic, fd) in [
2991 ("alpha", &derivs.penalties_first[0][level], fd_alpha),
2992 ("ln_tau", &derivs.penalties_first[1][level], fd_tau),
2993 ] {
2994 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2995 for (x, y) in analytic.iter().zip(fd.iter()) {
2996 assert!(
2997 (x - y).abs() <= 5e-5 * scale,
2998 "{name} jet of scale-candidate {level}: analytic {x:.6e} vs FD {y:.6e}"
2999 );
3000 }
3001 }
3002 }
3003 for coord in 0..2 {
3005 assert!(
3006 derivs.penalties_first[coord][l_count]
3007 .iter()
3008 .all(|v| *v == 0.0),
3009 "null-component candidate must have zero (α, lnτ) drift"
3010 );
3011 }
3012 let provider = derivs
3014 .penalties_cross_provider
3015 .as_ref()
3016 .expect("cross provider");
3017 let cross = provider.evaluate(0, 1).expect("cross pair (α, lnτ)");
3018 let pp = pen_at(a0 + h, t0 * h.exp());
3019 let pm = pen_at(a0 + h, t0 * (-h).exp());
3020 let mp = pen_at(a0 - h, t0 * h.exp());
3021 let mm = pen_at(a0 - h, t0 * (-h).exp());
3022 for level in 0..l_count {
3023 let fd = (&(&pp[level] - &pm[level]) - &(&mp[level] - &mm[level])) / (4.0 * h * h);
3024 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
3025 for (x, y) in cross[level].iter().zip(fd.iter()) {
3026 assert!(
3027 (x - y).abs() <= 5e-4 * scale,
3028 "cross (α, lnτ) jet of scale-candidate {level}: analytic {x:.6e} vs FD {y:.6e}"
3029 );
3030 }
3031 }
3032 }
3033
3034 #[test]
3040 pub(crate) fn psi_producer_matches_fd_length_scale() {
3041 let (data, mut frozen) = frozen_spec_fixture(0.0, false);
3044 frozen.learn_length_scale = true;
3045 let derivs =
3046 build_measure_jet_basis_psi_derivatives(data.view(), &frozen).expect("psi derivatives");
3047 assert_eq!(
3049 derivs.design_first.len(),
3050 1,
3051 "single-scale + learn_length_scale enrolls exactly the ℓ coordinate"
3052 );
3053 assert_eq!(
3054 derivs.penalties_first[0].len(),
3055 2,
3056 "single-scale double penalty carries Primary + affine/null component"
3057 );
3058 let ell0 = frozen.length_scale;
3062 let build_at = |ell: f64| {
3063 let trial = MeasureJetBasisSpec {
3064 length_scale: ell,
3065 ..frozen.clone()
3066 };
3067 build_measure_jet_basis(data.view(), &trial).expect("trial build")
3068 };
3069 let h: f64 = 1e-4;
3070 let plus = build_at(ell0 * h.exp());
3071 let minus = build_at(ell0 * (-h).exp());
3072 let at = build_at(ell0);
3073 assert_eq!(
3074 plus.active_penalties.len(),
3075 2,
3076 "fixture must keep both candidates active"
3077 );
3078 assert_eq!(
3079 minus.active_penalties.len(),
3080 2,
3081 "fixture must keep both candidates active"
3082 );
3083 assert_eq!(
3084 at.active_penalties.len(),
3085 2,
3086 "fixture must keep both candidates active"
3087 );
3088
3089 let x_plus = plus.design.to_dense();
3090 let x_minus = minus.design.to_dense();
3091 let x_0 = at.design.to_dense();
3092 let fd_first = (&x_plus - &x_minus) / (2.0 * h);
3093 let fd_second = (&x_plus - &(&x_0 * 2.0) + &x_minus) / (h * h);
3094 let scale1 = fd_first.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
3095 for (x, y) in derivs.design_first[0].iter().zip(fd_first.iter()) {
3096 assert!(
3097 (x - y).abs() <= 5e-5 * scale1,
3098 "∂X/∂lnℓ: analytic {x:.6e} vs FD {y:.6e}"
3099 );
3100 }
3101 let scale2 = fd_second.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
3102 for (x, y) in derivs.design_second_diag[0].iter().zip(fd_second.iter()) {
3103 assert!(
3104 (x - y).abs() <= 1e-3 * scale2,
3105 "∂²X/∂lnℓ²: analytic {x:.6e} vs FD {y:.6e}"
3106 );
3107 }
3108
3109 for candidate in 0..2 {
3110 let fd_penalty_first = (&plus.active_penalties[candidate].matrix
3111 - &minus.active_penalties[candidate].matrix)
3112 / (2.0 * h);
3113 let fd_penalty_second = (&plus.active_penalties[candidate].matrix
3114 - &(&at.active_penalties[candidate].matrix * 2.0)
3115 + &minus.active_penalties[candidate].matrix)
3116 / (h * h);
3117 let first_scale = fd_penalty_first
3118 .iter()
3119 .fold(1e-12_f64, |acc, value| acc.max(value.abs()));
3120 let second_scale = fd_penalty_second
3121 .iter()
3122 .fold(1e-10_f64, |acc, value| acc.max(value.abs()));
3123 for (analytic, finite_difference) in derivs.penalties_first[0][candidate]
3124 .iter()
3125 .zip(fd_penalty_first.iter())
3126 {
3127 assert!(
3128 (analytic - finite_difference).abs() <= 1e-4 * first_scale,
3129 "candidate {candidate} ∂S~/∂lnℓ: analytic {analytic:.6e} vs FD {finite_difference:.6e}"
3130 );
3131 }
3132 for (analytic, finite_difference) in derivs.penalties_second_diag[0][candidate]
3133 .iter()
3134 .zip(fd_penalty_second.iter())
3135 {
3136 assert!(
3137 (analytic - finite_difference).abs() <= 5e-3 * second_scale,
3138 "candidate {candidate} ∂²S~/∂lnℓ²: analytic {analytic:.6e} vs FD {finite_difference:.6e}"
3139 );
3140 }
3141 }
3142 }
3143
3144 #[test]
3148 pub(crate) fn quadrature_nodes_are_cell_barycenters() {
3149 let data = array![
3152 [0.0, 0.2],
3153 [0.4, -0.2],
3154 [0.2, 0.0],
3155 [9.8, 10.1],
3156 [10.2, 9.9],
3157 ];
3158 let seeds = array![[0.1, 0.1], [10.0, 10.0], [-50.0, -50.0]];
3159 let (nodes, masses) =
3160 measure_jet_quadrature_nodes(data.view(), seeds.view()).expect("quadrature nodes");
3161 assert!((masses.sum() - 1.0).abs() <= 1e-15, "masses must sum to 1");
3162 assert!((masses[0] - 0.6).abs() <= 1e-15);
3163 assert!((masses[1] - 0.4).abs() <= 1e-15);
3164 assert_eq!(masses[2], 0.0);
3165 assert_eq!(nodes[(0, 0)], 0.2);
3167 assert_eq!(nodes[(0, 1)], 0.0);
3168 assert_eq!(nodes[(1, 0)], 10.0);
3170 assert_eq!(nodes[(1, 1)], 10.0);
3171 assert_eq!(nodes[(2, 0)], -50.0);
3173 assert_eq!(nodes[(2, 1)], -50.0);
3174 }
3175
3176 #[test]
3180 pub(crate) fn build_replay_roundtrip_reproduces_design_and_penalty() {
3181 let n = 140usize;
3184 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
3185 let t = i as f64 / (n as f64 - 1.0);
3186 if k == 0 {
3187 t * 3.0
3188 } else {
3189 0.5 * (t * 3.0).cos() + if i % 9 == 0 { 0.8 } else { 0.0 }
3190 }
3191 });
3192 let spec = MeasureJetBasisSpec {
3193 center_strategy: CenterStrategy::FarthestPoint { num_centers: 70 },
3194 multiscale: true,
3195 ..MeasureJetBasisSpec::default()
3196 };
3197 let first = build_measure_jet_basis(data.view(), &spec).expect("first build");
3198 let BasisMetadata::MeasureJet {
3199 centers,
3200 length_scale,
3201 eps_band,
3202 order_s,
3203 alpha,
3204 tau0,
3205 masses,
3206 support_means,
3207 penalty_normalization_scales,
3208 raw_penalty_normalization_scales,
3209 fused_penalty_normalization_scale,
3210 constraint_transform,
3211 ..
3212 } = &first.metadata
3213 else {
3214 panic!("measure-jet build must return MeasureJet metadata");
3215 };
3216 let replay_spec = MeasureJetBasisSpec {
3217 center_strategy: CenterStrategy::UserProvided(centers.clone()),
3218 order_s: *order_s,
3219 alpha: *alpha,
3220 tau0: *tau0,
3221 num_scales: eps_band.len(),
3222 length_scale: *length_scale,
3223 double_penalty: spec.double_penalty,
3224 learn_length_scale: spec.learn_length_scale,
3225 multiscale: spec.multiscale,
3226 identifiability: MeasureJetIdentifiability::FrozenTransform {
3227 transform: constraint_transform.clone().expect("fit-time z"),
3228 },
3229 frozen_quadrature: Some(MeasureJetFrozenQuadrature {
3230 masses: masses.clone(),
3231 eps_band: eps_band.clone(),
3232 support_means: support_means.clone(),
3233 penalty_normalization_scales: penalty_normalization_scales.clone(),
3234 raw_penalty_normalization_scales: raw_penalty_normalization_scales.clone(),
3235 fused_penalty_normalization_scale: *fused_penalty_normalization_scale,
3236 sigma_coord: None,
3237 }),
3238 };
3239 assert_eq!(
3242 first.active_penalties.len(),
3243 eps_band.len() + 1,
3244 "per-level mode must emit one candidate per scale + null component"
3245 );
3246 let second = build_measure_jet_basis(data.view(), &replay_spec).expect("replay build");
3247 let x1 = first.design.to_dense();
3248 let x2 = second.design.to_dense();
3249 assert_eq!(x1.shape(), x2.shape());
3250 for (a, b) in x1.iter().zip(x2.iter()) {
3251 assert!((a - b).abs() <= 1e-12, "design replay drift: {a} vs {b}");
3252 }
3253 assert_eq!(first.active_penalties.len(), second.active_penalties.len());
3254 for (p1, p2) in first
3255 .active_penalties
3256 .iter()
3257 .zip(second.active_penalties.iter())
3258 {
3259 for (a, b) in p1.matrix.iter().zip(p2.matrix.iter()) {
3260 assert!((a - b).abs() <= 1e-12, "penalty replay drift: {a} vs {b}");
3261 }
3262 }
3263 }
3264}