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;
121
122use super::{
123 AnisoBasisPsiDerivatives, AnisoPenaltyCrossProvider, BasisBuildResult, BasisError,
124 BasisMetadata, CenterStrategy, PenaltyCandidate, PenaltySource,
125 filter_active_penalty_candidates_with_ops, normalize_penalty,
126 normalize_penalty_cross_psi_derivative, normalize_penaltywith_psi_derivatives,
127 select_centers_by_strategy, trace_of_product,
128};
129
130pub(crate) const MEASURE_JET_PROFILE_CUTOFF: f64 = 3.0;
136
137pub(crate) const MEASURE_JET_PSEUDOINVERSE_RTOL: f64 = 64.0 * f64::EPSILON;
141
142pub(crate) const MEASURE_JET_DEFAULT_ORDER_S: f64 = 1.5;
148
149pub(crate) const MEASURE_JET_MIN_AUTO_SCALES: usize = 3;
153pub(crate) const MEASURE_JET_MAX_AUTO_SCALES: usize = 8;
154
155pub(crate) const MEASURE_JET_AUTO_LENGTH_SCALE_FACTOR: f64 = 1.0;
170
171pub(crate) const MEASURE_JET_FUSED_RIDGE_FRACTION: f64 = 1e-2;
180
181pub(crate) const MEASURE_JET_PARALLEL_FORM_BUDGET_DOUBLES: usize = 1 << 26;
187
188#[derive(Debug, Clone, Serialize, Deserialize, Default)]
195pub enum MeasureJetIdentifiability {
196 #[default]
199 CenterSumToZero,
200 FrozenTransform { transform: Array2<f64> },
203}
204
205#[derive(Debug, Clone, Serialize, Deserialize)]
210pub struct MeasureJetFrozenQuadrature {
211 pub masses: Array1<f64>,
213 pub eps_band: Vec<f64>,
215 pub support_means: Vec<f64>,
218 pub penalty_normalization_scales: Vec<f64>,
221 pub raw_penalty_normalization_scales: Vec<f64>,
224 pub fused_penalty_normalization_scale: 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
403pub(crate) fn affine_preserving_coefficient_ridge(
411 kz: &Array2<f64>,
412 centers: ArrayView2<'_, f64>,
413 masses: ArrayView1<'_, f64>,
414) -> Result<Array2<f64>, BasisError> {
415 let m = centers.nrows();
416 let d = centers.ncols();
417 let p = kz.ncols();
418 if kz.nrows() != m || masses.len() != m {
419 crate::bail_dim_basis!(
420 "measure-jet affine-preserving ridge shape mismatch: kz {:?}, centers {:?}, masses {}",
421 kz.dim(),
422 centers.dim(),
423 masses.len()
424 );
425 }
426 if p == 0 {
427 return Ok(Array2::<f64>::zeros((0, 0)));
428 }
429 let mut weighted_kz = kz.clone();
430 for (i, mut row) in weighted_kz.outer_iter_mut().enumerate() {
431 row.mapv_inplace(|v| v * masses[i]);
432 }
433 let normal = kz.t().dot(&weighted_kz);
434 let normal_pinv = symmetric_pseudoinverse(&normal, "affine ridge normal")?;
435 let mut affine = Array2::<f64>::ones((m, d + 1));
436 for i in 0..m {
437 for k in 0..d {
438 affine[(i, k + 1)] = centers[(i, k)];
439 }
440 }
441 let mut weighted_affine = affine.clone();
442 for (i, mut row) in weighted_affine.outer_iter_mut().enumerate() {
443 row.mapv_inplace(|v| v * masses[i]);
444 }
445 let rhs = kz.t().dot(&weighted_affine);
446 let beta = normal_pinv.dot(&rhs);
447 let beta_gram = beta.t().dot(&beta);
448 let (evals, evecs) = beta_gram.eigh(Side::Lower).map_err(|e| {
449 BasisError::InvalidInput(format!(
450 "measure-jet affine ridge subspace eigendecomposition failed: {e}"
451 ))
452 })?;
453 let lam_max = evals.iter().fold(0.0_f64, |acc, v| acc.max((*v).max(0.0)));
454 let rank_tol = MEASURE_JET_PSEUDOINVERSE_RTOL * ((d + 1).max(1) as f64) * lam_max;
455 let mut ridge = Array2::<f64>::eye(p);
456 for k in 0..(d + 1) {
457 let lam = evals[k].max(0.0);
458 if lam <= rank_tol {
459 continue;
460 }
461 let dir = beta.dot(&evecs.column(k).to_owned()) / lam.sqrt();
462 for r in 0..p {
463 for c in 0..p {
464 ridge[(r, c)] -= dir[r] * dir[c];
465 }
466 }
467 }
468 Ok((&ridge + &ridge.t()) * 0.5)
469}
470
471pub(crate) fn pairwise_sq_dists(a: ArrayView2<'_, f64>, b: ArrayView2<'_, f64>) -> Array2<f64> {
482 let an: Vec<f64> = a.outer_iter().map(|r| r.dot(&r)).collect();
483 let bn: Vec<f64> = b.outer_iter().map(|r| r.dot(&r)).collect();
484 let mut g = a.dot(&b.t());
485 g.axis_iter_mut(Axis(0))
486 .into_par_iter()
487 .enumerate()
488 .for_each(|(i, mut row)| {
489 for (j, v) in row.iter_mut().enumerate() {
490 *v = (an[i] + bn[j] - 2.0 * *v).max(0.0);
491 }
492 });
493 g
494}
495
496pub(crate) const MEASURE_JET_ASSIGN_BLOCK_ROWS: usize = 65_536;
500
501pub(crate) fn validate_finite_points(
502 points: ArrayView2<'_, f64>,
503 what: &str,
504) -> Result<(), BasisError> {
505 for (i, row) in points.outer_iter().enumerate() {
506 if row.iter().any(|v| !v.is_finite()) {
507 crate::bail_invalid_basis!("measure-jet {what} row {i} has a non-finite coordinate");
508 }
509 }
510 Ok(())
511}
512
513pub(crate) fn median_nearest_center_spacing(dist2: &Array2<f64>) -> Result<f64, BasisError> {
516 let m = dist2.nrows();
517 if m < 2 {
518 return Err(BasisError::InsufficientColumnsForConstraint { found: m });
519 }
520 let mut nearest: Vec<f64> = Vec::with_capacity(m);
521 for i in 0..m {
522 let mut best = f64::INFINITY;
523 for j in 0..m {
524 if j != i && dist2[(i, j)] < best {
525 best = dist2[(i, j)];
526 }
527 }
528 nearest.push(best.sqrt());
529 }
530 nearest.sort_by(|a, b| a.partial_cmp(b).expect("finite center spacings"));
531 let median = nearest[nearest.len() / 2];
532 if !(median.is_finite() && median > 0.0) {
533 crate::bail_invalid_basis!(
534 "measure-jet centers are degenerate (median nearest-center spacing = {median}); \
535 duplicate centers cannot carry a scale band"
536 );
537 }
538 Ok(median)
539}
540
541pub fn measure_jet_band(
549 centers: ArrayView2<'_, f64>,
550 num_scales: usize,
551) -> Result<MeasureJetBand, BasisError> {
552 validate_finite_points(centers, "centers")?;
553 let dist2 = pairwise_sq_dists(centers, centers);
554 let eps_min = median_nearest_center_spacing(&dist2)?;
555 let d = centers.ncols();
557 let mut diag2 = 0.0_f64;
558 for k in 0..d {
559 let col = centers.column(k);
560 let mut lo = f64::INFINITY;
561 let mut hi = f64::NEG_INFINITY;
562 for &v in col.iter() {
563 lo = lo.min(v);
564 hi = hi.max(v);
565 }
566 diag2 += (hi - lo) * (hi - lo);
567 }
568 let eps_max = 0.5 * diag2.sqrt();
569 if !(eps_max.is_finite() && eps_max > eps_min) {
570 return Ok(MeasureJetBand {
571 eps: vec![eps_min],
572 log_step: std::f64::consts::LN_2,
573 });
574 }
575 let auto = ((eps_max / eps_min).log2().ceil() as usize + 1)
576 .clamp(MEASURE_JET_MIN_AUTO_SCALES, MEASURE_JET_MAX_AUTO_SCALES);
577 let count = if num_scales == 0 { auto } else { num_scales };
578 if count == 1 {
579 return Ok(MeasureJetBand {
580 eps: vec![eps_min],
581 log_step: std::f64::consts::LN_2,
582 });
583 }
584 let ratio = (eps_max / eps_min).powf(1.0 / (count as f64 - 1.0));
585 let mut eps = Vec::with_capacity(count);
586 let mut e = eps_min;
587 for _ in 0..count {
588 eps.push(e);
589 e *= ratio;
590 }
591 Ok(MeasureJetBand {
592 eps,
593 log_step: ratio.ln(),
594 })
595}
596
597pub fn measure_jet_quadrature_nodes(
604 data: ArrayView2<'_, f64>,
605 centers: ArrayView2<'_, f64>,
606) -> Result<(Array2<f64>, Array1<f64>), BasisError> {
607 if data.ncols() != centers.ncols() {
608 crate::bail_dim_basis!(
609 "measure-jet mass assignment dimension mismatch: data d={} centers d={}",
610 data.ncols(),
611 centers.ncols()
612 );
613 }
614 validate_finite_points(data, "data")?;
615 validate_finite_points(centers, "centers")?;
616 let n = data.nrows();
617 let m = centers.nrows();
618 let d = centers.ncols();
619 if n == 0 || m == 0 {
620 crate::bail_invalid_basis!("measure-jet mass assignment needs nonempty data and centers");
621 }
622 let cn: Vec<f64> = centers.outer_iter().map(|r| r.dot(&r)).collect();
627 let assignments: Vec<usize> = (0..n)
628 .step_by(MEASURE_JET_ASSIGN_BLOCK_ROWS)
629 .flat_map(|start| {
630 let end = (start + MEASURE_JET_ASSIGN_BLOCK_ROWS).min(n);
631 let g = data.slice(ndarray::s![start..end, ..]).dot(¢ers.t());
632 let block: Vec<usize> = g
633 .axis_iter(Axis(0))
634 .into_par_iter()
635 .map(|row| {
636 let mut best_j = 0usize;
637 let mut best = f64::INFINITY;
638 for (j, &gij) in row.iter().enumerate() {
639 let s = cn[j] - 2.0 * gij;
640 if s < best {
641 best = s;
642 best_j = j;
643 }
644 }
645 best_j
646 })
647 .collect();
648 block
649 })
650 .collect();
651 let mut masses = Array1::<f64>::zeros(m);
652 let mut nodes = centers.to_owned();
653 let mut sums = Array2::<f64>::zeros((m, d));
654 let unit = 1.0 / n as f64;
655 for (i, &j) in assignments.iter().enumerate() {
656 masses[j] += unit;
657 for k in 0..d {
658 sums[(j, k)] += data[(i, k)];
659 }
660 }
661 let mut barycenter = sums;
664 for j in 0..m {
665 let count = masses[j] * n as f64;
666 if count > 0.0 {
667 for k in 0..d {
668 barycenter[(j, k)] /= count;
669 nodes[(j, k)] = barycenter[(j, k)];
670 }
671 }
672 }
673 Ok((nodes, masses))
674}
675
676pub fn measure_jet_center_masses(
679 data: ArrayView2<'_, f64>,
680 centers: ArrayView2<'_, f64>,
681) -> Result<Array1<f64>, BasisError> {
682 measure_jet_quadrature_nodes(data, centers).map(|(_, masses)| masses)
683}
684
685pub(crate) fn assemble_weighted_forms<F>(
708 centers: ArrayView2<'_, f64>,
709 masses: ArrayView1<'_, f64>,
710 band: &MeasureJetBand,
711 order_s: f64,
712 alpha: f64,
713 tau0: f64,
714 n_forms: usize,
715 channels: usize,
716 weights: &F,
717) -> Result<Vec<Array2<f64>>, BasisError>
718where
719 F: Fn(usize, f64, f64, f64, &mut [[f64; 3]]) + Sync,
720{
721 let m = centers.nrows();
722 let d = centers.ncols();
723 if n_forms == 0 || !(1..=3).contains(&channels) {
724 crate::bail_invalid_basis!(
725 "measure-jet assembly needs at least one output form and 1..=3 block channels"
726 );
727 }
728 if masses.len() != m {
729 crate::bail_dim_basis!(
730 "measure-jet energy mass/center mismatch: {} masses for {} centers",
731 masses.len(),
732 m
733 );
734 }
735 if band.eps.is_empty() || band.eps.iter().any(|e| !(e.is_finite() && *e > 0.0)) {
736 crate::bail_invalid_basis!("measure-jet energy needs a nonempty positive scale band");
737 }
738 if !(order_s.is_finite() && order_s > 0.0 && order_s < 2.0) {
739 crate::bail_invalid_basis!(
740 "measure-jet order s must lie in (0, 2) for the affine-jet energy; got {order_s}"
741 );
742 }
743 if !(alpha.is_finite() && tau0.is_finite() && tau0 >= 0.0) {
744 crate::bail_invalid_basis!(
745 "measure-jet energy needs finite alpha and finite tau0 >= 0; got alpha={alpha}, tau0={tau0}"
746 );
747 }
748 if masses.iter().any(|v| !(v.is_finite() && *v >= 0.0)) {
749 crate::bail_invalid_basis!("measure-jet energy needs finite nonnegative center masses");
750 }
751 let dist2 = pairwise_sq_dists(centers, centers);
752
753 let assemble_scale = |scale_idx: usize, eps: f64| -> Result<Vec<Array2<f64>>, BasisError> {
758 let mut out: Vec<Array2<f64>> =
759 (0..n_forms).map(|_| Array2::<f64>::zeros((m, m))).collect();
760 let cutoff2 = (MEASURE_JET_PROFILE_CUTOFF * eps) * (MEASURE_JET_PROFILE_CUTOFF * eps);
761 let inv_two_eps2 = 1.0 / (2.0 * eps * eps);
762 let eta = 2.0 * order_s + (d as f64) * (2.0 - 2.0 * alpha);
763 let scale_weight = band.log_step * eps.powf(-eta);
764 let net_radius2 = 0.25 * eps * eps;
768 let mut outer: Vec<usize> = Vec::new();
769 for i in 0..m {
770 if masses[i] <= 0.0 {
771 continue;
772 }
773 let covered = outer.iter().any(|&o| dist2[(i, o)] <= net_radius2);
774 if !covered {
775 outer.push(i);
776 }
777 }
778 let mut net_mass = vec![0.0_f64; m];
779 for i in 0..m {
780 if masses[i] <= 0.0 {
781 continue;
782 }
783 let mut best = f64::INFINITY;
784 let mut best_o = usize::MAX;
785 for &o in &outer {
786 if dist2[(i, o)] < best {
787 best = dist2[(i, o)];
788 best_o = o;
789 }
790 }
791 if best_o != usize::MAX {
792 net_mass[best_o] += masses[i];
793 }
794 }
795 let mut wbuf = vec![[0.0_f64; 3]; n_forms];
796 for &i in &outer {
797 let mut idx: Vec<usize> = Vec::new();
799 for j in 0..m {
800 if dist2[(i, j)] <= cutoff2 {
801 idx.push(j);
802 }
803 }
804 let ml = idx.len();
805 let mut w = Array1::<f64>::zeros(ml);
807 let mut q = 0.0_f64;
808 for (a, &j) in idx.iter().enumerate() {
809 let wj = masses[j] * (-dist2[(i, j)] * inv_two_eps2).exp();
810 w[a] = wj;
811 q += wj;
812 }
813 if !(q > 0.0) {
814 continue;
815 }
816 let mut phi = Array2::<f64>::zeros((ml, d));
818 for (a, &j) in idx.iter().enumerate() {
819 for k in 0..d {
820 phi[(a, k)] = (centers[(j, k)] - centers[(i, k)]) / eps;
821 }
822 }
823 let a_mean = phi.t().dot(&w) / q;
824 let mut wphi = phi.clone();
826 for (a, mut row) in wphi.outer_iter_mut().enumerate() {
827 row.mapv_inplace(|v| v * w[a]);
828 }
829 let mut b = wphi.clone();
830 for (a, mut row) in b.outer_iter_mut().enumerate() {
831 for k in 0..d {
832 row[k] -= w[a] * a_mean[k];
833 }
834 }
835 let mut g = phi.t().dot(&wphi);
836 g.mapv_inplace(|v| v / q);
837 for r in 0..d {
838 for c in 0..d {
839 g[(r, c)] -= a_mean[r] * a_mean[c];
840 }
841 }
842 let g_pinv = symmetric_pseudoinverse(&g, "local affine Gram")?;
843 let bm = b.dot(&g_pinv);
844 let base = scale_weight * net_mass[i] * q.powf(1.0 - 2.0 * alpha);
845 weights(scale_idx, eps, q, base, &mut wbuf);
846 for (a, &ja) in idx.iter().enumerate() {
849 let bma = bm.row(a);
850 for (c, &jc) in idx.iter().enumerate() {
851 let b_c = b.row(c);
852 let mut val_r = -w[a] * w[c] / q - bma.dot(&b_c) / q;
853 if a == c {
854 val_r += w[a];
855 }
856 for (k, out_k) in out.iter_mut().enumerate() {
857 let wk = wbuf[k];
858 out_k[(ja, jc)] += wk[0] * val_r;
859 }
860 }
861 }
862 }
863 Ok(out)
864 };
865
866 let n_scales = band.eps.len();
867 let parallel_ok = m
868 .saturating_mul(m)
869 .saturating_mul(n_scales)
870 .saturating_mul(n_forms)
871 <= MEASURE_JET_PARALLEL_FORM_BUDGET_DOUBLES;
872 let per_scale: Vec<Vec<Array2<f64>>> = if parallel_ok {
873 band.eps
874 .par_iter()
875 .enumerate()
876 .map(|(scale_idx, &eps)| assemble_scale(scale_idx, eps))
877 .collect::<Result<Vec<_>, BasisError>>()?
878 } else {
879 band.eps
880 .iter()
881 .enumerate()
882 .map(|(scale_idx, &eps)| assemble_scale(scale_idx, eps))
883 .collect::<Result<Vec<_>, BasisError>>()?
884 };
885
886 let mut totals: Vec<Array2<f64>> = (0..n_forms).map(|_| Array2::<f64>::zeros((m, m))).collect();
887 for scale_forms in per_scale {
888 for (total, part) in totals.iter_mut().zip(scale_forms) {
889 *total += ∂
890 }
891 }
892 Ok(totals.into_iter().map(|t| (&t + &t.t()) * 0.5).collect())
894}
895
896pub fn measure_jet_energy_form(
911 centers: ArrayView2<'_, f64>,
912 masses: ArrayView1<'_, f64>,
913 band: &MeasureJetBand,
914 order_s: f64,
915 alpha: f64,
916 tau0: f64,
917) -> Result<Array2<f64>, BasisError> {
918 let mut forms = assemble_weighted_forms(
919 centers,
920 masses,
921 band,
922 order_s,
923 alpha,
924 tau0,
925 1,
926 1,
927 &|_, _, _, base, out: &mut [[f64; 3]]| out[0] = [base, 0.0, 0.0],
928 )?;
929 let q = forms.swap_remove(0);
930 project_symmetric_psd(q, "measure-jet energy form")
938}
939
940pub(crate) fn project_symmetric_psd(
946 a: Array2<f64>,
947 label: &str,
948) -> Result<Array2<f64>, BasisError> {
949 let n = a.nrows();
950 if n == 0 {
951 return Ok(a);
952 }
953 let (evals, evecs) = a.eigh(Side::Lower).map_err(|e| {
954 BasisError::InvalidInput(format!(
955 "measure-jet PSD projection `{label}` eigendecomposition failed: {e}"
956 ))
957 })?;
958 if evals.iter().all(|&lam| lam >= 0.0) {
959 return Ok(a);
960 }
961 let mut scaled = evecs.clone();
962 for (k, mut col) in scaled.axis_iter_mut(Axis(1)).enumerate() {
963 let lam = evals[k].max(0.0);
964 col.mapv_inplace(|v| v * lam);
965 }
966 let psd = scaled.dot(&evecs.t());
967 Ok((&psd + &psd.t()) * 0.5)
968}
969
970pub fn measure_jet_energy_form_with_jets(
985 centers: ArrayView2<'_, f64>,
986 masses: ArrayView1<'_, f64>,
987 band: &MeasureJetBand,
988 order_s: f64,
989 alpha: f64,
990 tau0: f64,
991) -> Result<MeasureJetEnergyJets, BasisError> {
992 if !(tau0.is_finite() && tau0 > 0.0) {
993 crate::bail_invalid_basis!(
994 "measure-jet jets need tau0 > 0 because the retained τ coordinate is ln τ; got {tau0}"
995 );
996 }
997 let mut forms = assemble_weighted_forms(
998 centers,
999 masses,
1000 band,
1001 order_s,
1002 alpha,
1003 tau0,
1004 10,
1005 3,
1006 &|_, eps: f64, q: f64, base: f64, out: &mut [[f64; 3]]| {
1007 let gs = -2.0 * eps.ln();
1008 let intrinsic_dim = centers.ncols() as f64;
1009 let ga = 2.0 * intrinsic_dim * eps.ln() - 2.0 * q.max(f64::MIN_POSITIVE).ln();
1010 out[0] = [base, 0.0, 0.0];
1011 out[1] = [gs * base, 0.0, 0.0];
1012 out[2] = [gs * gs * base, 0.0, 0.0];
1013 out[3] = [ga * base, 0.0, 0.0];
1014 out[4] = [ga * ga * base, 0.0, 0.0];
1015 out[5] = [gs * ga * base, 0.0, 0.0];
1016 out[6] = [0.0, 0.0, 0.0];
1017 out[7] = [0.0, 0.0, 0.0];
1018 out[8] = [0.0, 0.0, 0.0];
1019 out[9] = [0.0, 0.0, 0.0];
1020 },
1021 )?;
1022 let d2q_dalpha_dlogtau = forms.pop().expect("ten assembled forms");
1023 let d2q_ds_dlogtau = forms.pop().expect("ten assembled forms");
1024 let d2q_dlogtau2 = forms.pop().expect("ten assembled forms");
1025 let dq_dlogtau = forms.pop().expect("ten assembled forms");
1026 let d2q_ds_dalpha = forms.pop().expect("ten assembled forms");
1027 let d2q_dalpha2 = forms.pop().expect("ten assembled forms");
1028 let dq_dalpha = forms.pop().expect("ten assembled forms");
1029 let d2q_ds2 = forms.pop().expect("ten assembled forms");
1030 let dq_ds = forms.pop().expect("ten assembled forms");
1031 let q = forms.pop().expect("ten assembled forms");
1032 Ok(MeasureJetEnergyJets {
1033 q,
1034 dq_ds,
1035 d2q_ds2,
1036 dq_dalpha,
1037 d2q_dalpha2,
1038 d2q_ds_dalpha,
1039 dq_dlogtau,
1040 d2q_dlogtau2,
1041 d2q_ds_dlogtau,
1042 d2q_dalpha_dlogtau,
1043 })
1044}
1045
1046pub fn measure_jet_scale_spectrum(
1052 centers: ArrayView2<'_, f64>,
1053 masses: ArrayView1<'_, f64>,
1054 band: &MeasureJetBand,
1055 order_s: f64,
1056 alpha: f64,
1057 tau0: f64,
1058 values: ArrayView1<'_, f64>,
1059) -> Result<Vec<f64>, BasisError> {
1060 if values.len() != centers.nrows() {
1061 crate::bail_dim_basis!(
1062 "measure-jet scale spectrum needs one value per center: {} values for {} centers",
1063 values.len(),
1064 centers.nrows()
1065 );
1066 }
1067 let forms = measure_jet_energy_forms_per_scale(centers, masses, band, order_s, alpha, tau0)?;
1068 Ok(forms
1069 .iter()
1070 .map(|q_l| values.dot(&q_l.dot(&values)))
1071 .collect())
1072}
1073
1074pub fn measure_jet_energy_forms_per_scale(
1080 centers: ArrayView2<'_, f64>,
1081 masses: ArrayView1<'_, f64>,
1082 band: &MeasureJetBand,
1083 order_s: f64,
1084 alpha: f64,
1085 tau0: f64,
1086) -> Result<Vec<Array2<f64>>, BasisError> {
1087 let n_scales = band.eps.len();
1088 assemble_weighted_forms(
1089 centers,
1090 masses,
1091 band,
1092 order_s,
1093 alpha,
1094 tau0,
1095 n_scales,
1096 1,
1097 &|scale_idx, _, _, base, out: &mut [[f64; 3]]| {
1098 for (k, slot) in out.iter_mut().enumerate() {
1099 *slot = if k == scale_idx {
1100 [base, 0.0, 0.0]
1101 } else {
1102 [0.0, 0.0, 0.0]
1103 };
1104 }
1105 },
1106 )
1107}
1108
1109pub fn measure_jet_support_curve(
1117 queries: ArrayView2<'_, f64>,
1118 centers: ArrayView2<'_, f64>,
1119 masses: ArrayView1<'_, f64>,
1120 eps_band: &[f64],
1121) -> Result<Array2<f64>, BasisError> {
1122 if queries.ncols() != centers.ncols() {
1123 crate::bail_dim_basis!(
1124 "measure-jet support curve dimension mismatch: queries d={} centers d={}",
1125 queries.ncols(),
1126 centers.ncols()
1127 );
1128 }
1129 if masses.len() != centers.nrows() {
1130 crate::bail_dim_basis!(
1131 "measure-jet support curve mass/center mismatch: {} masses for {} centers",
1132 masses.len(),
1133 centers.nrows()
1134 );
1135 }
1136 if eps_band.is_empty() || eps_band.iter().any(|e| !(e.is_finite() && *e > 0.0)) {
1137 crate::bail_invalid_basis!("measure-jet support curve needs a nonempty positive band");
1138 }
1139 validate_finite_points(queries, "queries")?;
1140 validate_finite_points(centers, "centers")?;
1141 let nq = queries.nrows();
1142 let nl = eps_band.len();
1143 let d2 = pairwise_sq_dists(queries, centers);
1146 let mut out = Array2::<f64>::zeros((nq, nl));
1147 out.axis_iter_mut(Axis(0))
1148 .into_par_iter()
1149 .enumerate()
1150 .for_each(|(qi, mut row)| {
1151 let d2_row = d2.row(qi);
1152 for (li, &eps) in eps_band.iter().enumerate() {
1153 let inv_two_eps2 = 1.0 / (2.0 * eps * eps);
1154 let mut acc = 0.0_f64;
1155 for (j, &dd) in d2_row.iter().enumerate() {
1156 acc += masses[j] * (-dd * inv_two_eps2).exp();
1157 }
1158 row[li] = acc;
1159 }
1160 });
1161 Ok(out)
1162}
1163
1164pub(crate) fn measure_jet_support_means(
1165 centers: ArrayView2<'_, f64>,
1166 masses: ArrayView1<'_, f64>,
1167 eps_band: &[f64],
1168) -> Result<Vec<f64>, BasisError> {
1169 let total_mass = masses.sum();
1170 if !(total_mass.is_finite() && total_mass > 0.0) {
1171 crate::bail_invalid_basis!(
1172 "measure-jet support means need positive finite total mass; got {total_mass}"
1173 );
1174 }
1175 let support = measure_jet_support_curve(centers, centers, masses, eps_band)?;
1176 let mut means = vec![0.0_f64; eps_band.len()];
1177 for (i, row) in support.rows().into_iter().enumerate() {
1178 let mass = masses[i];
1179 for (mean, &q) in means.iter_mut().zip(row.iter()) {
1180 *mean += mass * q;
1181 }
1182 }
1183 for mean in &mut means {
1184 *mean /= total_mass;
1185 if !(*mean).is_finite() || *mean <= 0.0 {
1186 crate::bail_invalid_basis!(
1187 "measure-jet support mean must be positive and finite; got {mean}"
1188 );
1189 }
1190 }
1191 Ok(means)
1192}
1193
1194pub fn measure_jet_design_matrix(
1196 data: ArrayView2<'_, f64>,
1197 centers: ArrayView2<'_, f64>,
1198 length_scale: f64,
1199) -> Result<Array2<f64>, BasisError> {
1200 if data.ncols() != centers.ncols() {
1201 crate::bail_dim_basis!(
1202 "measure-jet design dimension mismatch: data d={} centers d={}",
1203 data.ncols(),
1204 centers.ncols()
1205 );
1206 }
1207 if !(length_scale.is_finite() && length_scale > 0.0) {
1208 crate::bail_invalid_basis!(
1209 "measure-jet design needs a positive finite length_scale; got {length_scale}"
1210 );
1211 }
1212 validate_finite_points(data, "data")?;
1213 validate_finite_points(centers, "centers")?;
1214 let inv_two_l2 = 1.0 / (2.0 * length_scale * length_scale);
1215 let mut out = pairwise_sq_dists(data, centers);
1218 out.axis_iter_mut(Axis(0))
1219 .into_par_iter()
1220 .for_each(|mut row| {
1221 row.mapv_inplace(|d2| (-d2 * inv_two_l2).exp());
1222 });
1223 Ok(out)
1224}
1225
1226fn measure_jet_affine_head_transform(
1255 centers: ArrayView2<'_, f64>,
1256 masses: ArrayView1<'_, f64>,
1257) -> Array2<f64> {
1258 let m = centers.nrows();
1259 let d = centers.ncols();
1260 let total_mass = masses.sum();
1261 let mdot = |u: &Array1<f64>, v: &Array1<f64>| -> f64 {
1263 let mut acc = 0.0;
1264 for i in 0..m {
1265 acc += masses[i] * u[i] * v[i];
1266 }
1267 acc
1268 };
1269 let cols: Vec<Array1<f64>> = (0..d)
1273 .map(|k| {
1274 let col = centers.column(k).to_owned();
1275 let mean = if total_mass > 0.0 {
1276 mdot(&col, &Array1::ones(m)) / total_mass
1277 } else {
1278 0.0
1279 };
1280 col.mapv(|x| x - mean)
1281 })
1282 .collect();
1283 let max_norm = cols.iter().fold(0.0_f64, |acc, c| acc.max(mdot(c, c).sqrt()));
1285 let drop_below =
1286 (MEASURE_JET_PSEUDOINVERSE_RTOL * (d.max(1) as f64) * max_norm).max(f64::MIN_POSITIVE);
1287 let mut q_cols: Vec<Array1<f64>> = Vec::new();
1291 let mut t_cols: Vec<Array1<f64>> = Vec::new();
1292 for k in 0..d {
1293 let mut v = cols[k].clone();
1294 let mut t = Array1::<f64>::zeros(d);
1295 t[k] = 1.0;
1296 for (q, tq) in q_cols.iter().zip(t_cols.iter()) {
1297 let proj = mdot(q, &v);
1298 v.scaled_add(-proj, q);
1299 t.scaled_add(-proj, tq);
1300 }
1301 let norm = mdot(&v, &v).sqrt();
1302 if norm > drop_below {
1303 v.mapv_inplace(|x| x / norm);
1304 t.mapv_inplace(|x| x / norm);
1305 q_cols.push(v);
1306 t_cols.push(t);
1307 }
1308 }
1309 let head_rank = t_cols.len();
1310 let mut t_mat = Array2::<f64>::zeros((d, head_rank));
1311 for (r, t) in t_cols.into_iter().enumerate() {
1312 t_mat.column_mut(r).assign(&t);
1313 }
1314 t_mat
1315}
1316
1317pub fn realized_measure_jet_length_scale(
1322 centers: ArrayView2<'_, f64>,
1323 spec_length_scale: f64,
1324) -> Result<f64, BasisError> {
1325 if spec_length_scale.is_finite() && spec_length_scale > 0.0 {
1326 return Ok(spec_length_scale);
1327 }
1328 if spec_length_scale != 0.0 {
1329 crate::bail_invalid_basis!(
1330 "measure-jet length_scale must be positive (or 0.0 for auto); got {spec_length_scale}"
1331 );
1332 }
1333 let dist2 = pairwise_sq_dists(centers, centers);
1334 let spacing = median_nearest_center_spacing(&dist2)?;
1335 Ok(MEASURE_JET_AUTO_LENGTH_SCALE_FACTOR * spacing)
1336}
1337
1338pub(crate) struct RealizedMeasureJetGeometry {
1343 pub(crate) centers: Array2<f64>,
1344 pub(crate) masses: Array1<f64>,
1345 pub(crate) eps_band: Vec<f64>,
1346 pub(crate) log_step: f64,
1347 pub(crate) length_scale: f64,
1348 pub(crate) order_s_eval: f64,
1352 pub(crate) per_level: bool,
1354 pub(crate) z: Array2<f64>,
1355 pub(crate) coefficient_gauge: gam_problem::Gauge,
1356 pub(crate) kz: Array2<f64>,
1357 pub(crate) head_transform: Array2<f64>,
1363}
1364
1365pub(crate) fn realize_measure_jet_geometry(
1366 data: ArrayView2<'_, f64>,
1367 spec: &MeasureJetBasisSpec,
1368) -> Result<RealizedMeasureJetGeometry, BasisError> {
1369 if data.ncols() == 0 {
1370 crate::bail_invalid_basis!("measure-jet smooth needs at least one feature column");
1371 }
1372 validate_finite_points(data, "data")?;
1373 let seed_centers = select_centers_by_strategy(data, &spec.center_strategy)?;
1374 let m = seed_centers.nrows();
1375 if m < 3 {
1376 return Err(BasisError::InsufficientColumnsForConstraint { found: m });
1377 }
1378 let order_s = if spec.order_s == 0.0 {
1379 MEASURE_JET_DEFAULT_ORDER_S
1380 } else {
1381 spec.order_s
1382 };
1383 let (centers, masses, eps_band, log_step) = match &spec.frozen_quadrature {
1390 Some(frozen) => {
1391 if frozen.masses.len() != m {
1392 crate::bail_dim_basis!(
1393 "frozen measure-jet quadrature mismatch: {} masses for {} centers",
1394 frozen.masses.len(),
1395 m
1396 );
1397 }
1398 if frozen.eps_band.is_empty() {
1399 crate::bail_invalid_basis!("frozen measure-jet quadrature has an empty band");
1400 }
1401 let log_step = if frozen.eps_band.len() >= 2 {
1402 (frozen.eps_band[1] / frozen.eps_band[0]).ln()
1403 } else {
1404 std::f64::consts::LN_2
1405 };
1406 (
1407 seed_centers,
1408 frozen.masses.clone(),
1409 frozen.eps_band.clone(),
1410 log_step,
1411 )
1412 }
1413 None => {
1414 let (nodes, masses) = measure_jet_quadrature_nodes(data, seed_centers.view())?;
1415 let band = measure_jet_band(nodes.view(), spec.num_scales)?;
1416 (nodes, masses, band.eps, band.log_step)
1417 }
1418 };
1419 let length_scale = realized_measure_jet_length_scale(centers.view(), spec.length_scale)?;
1420 let head_transform = if spec.multiscale {
1431 Array2::<f64>::zeros((centers.ncols(), 0))
1432 } else {
1433 measure_jet_affine_head_transform(centers.view(), masses.view())
1434 };
1435 let head_rank = head_transform.ncols();
1436 let m_aug = m + head_rank;
1437 let (z, coefficient_gauge) = match &spec.identifiability {
1444 MeasureJetIdentifiability::FrozenTransform { transform } => {
1445 if transform.nrows() != m_aug {
1446 crate::bail_dim_basis!(
1447 "frozen measure-jet identifiability transform mismatch: {} representers + {} head columns but transform has {} rows",
1448 m,
1449 head_rank,
1450 transform.nrows()
1451 );
1452 }
1453 (
1454 transform.clone(),
1455 gam_problem::Gauge::from_block_transforms(&[transform.clone()]),
1456 )
1457 }
1458 MeasureJetIdentifiability::CenterSumToZero => {
1459 let u = householder_sum_to_zero_u(m);
1463 let z_rbf = householder_sum_to_zero_z(&u);
1464 let mut z_block = Array2::<f64>::zeros((m_aug, (m - 1) + head_rank));
1465 z_block.slice_mut(ndarray::s![..m, ..m - 1]).assign(&z_rbf);
1466 for r in 0..head_rank {
1467 z_block[(m + r, (m - 1) + r)] = 1.0;
1468 }
1469 (
1470 z_block.clone(),
1471 gam_problem::Gauge::from_block_transforms(&[z_block]),
1472 )
1473 }
1474 };
1475 let k_cc = measure_jet_design_matrix(centers.view(), centers.view(), length_scale)?;
1480 let mut k_aug = Array2::<f64>::zeros((m, m_aug));
1481 k_aug.slice_mut(ndarray::s![.., ..m]).assign(&k_cc);
1482 if head_rank > 0 {
1483 let head_cc = centers.dot(&head_transform);
1484 k_aug.slice_mut(ndarray::s![.., m..]).assign(&head_cc);
1485 }
1486 let kz = coefficient_gauge.restrict_design(&k_aug);
1487 Ok(RealizedMeasureJetGeometry {
1488 centers,
1489 masses,
1490 eps_band,
1491 log_step,
1492 length_scale,
1493 order_s_eval: order_s,
1494 per_level: spec.multiscale,
1499 z,
1500 coefficient_gauge,
1501 kz,
1502 head_transform,
1503 })
1504}
1505
1506pub fn measure_jet_multiscale_mode(spec: &MeasureJetBasisSpec) -> bool {
1513 spec.multiscale
1514}
1515
1516pub fn build_measure_jet_basis(
1523 data: ArrayView2<'_, f64>,
1524 spec: &MeasureJetBasisSpec,
1525) -> Result<BasisBuildResult, BasisError> {
1526 let RealizedMeasureJetGeometry {
1527 centers,
1528 masses,
1529 eps_band,
1530 log_step,
1531 length_scale,
1532 order_s_eval: order_s,
1533 per_level,
1534 z,
1535 coefficient_gauge,
1536 kz,
1537 head_transform,
1538 } = realize_measure_jet_geometry(data, spec)?;
1539 let band = MeasureJetBand {
1540 eps: eps_band.clone(),
1541 log_step,
1542 };
1543 let m = centers.nrows();
1544 let head_rank = head_transform.ncols();
1545 let m_aug = m + head_rank;
1546 let kernel_design = measure_jet_design_matrix(data, centers.view(), length_scale)?;
1551 let mut raw_design = Array2::<f64>::zeros((data.nrows(), m_aug));
1552 raw_design
1553 .slice_mut(ndarray::s![.., ..m])
1554 .assign(&kernel_design);
1555 if head_rank > 0 {
1556 let head_design = data.dot(&head_transform);
1557 raw_design
1558 .slice_mut(ndarray::s![.., m..])
1559 .assign(&head_design);
1560 }
1561 let constrained_design = coefficient_gauge.restrict_design(&raw_design);
1562 let design = gam_linalg::matrix::DesignMatrix::Dense(
1563 gam_linalg::matrix::DenseDesignMatrix::from(constrained_design),
1564 );
1565 let support_means = measure_jet_support_means(centers.view(), masses.view(), &eps_band)?;
1566 let mut candidates = Vec::new();
1579 let mut penalty_normalization_scales = Vec::new();
1580 let mut raw_penalty_normalization_scales = Vec::new();
1581 let mut fused_penalty_normalization_scale = None;
1582 if per_level {
1583 let forms = measure_jet_energy_forms_per_scale(
1584 centers.view(),
1585 masses.view(),
1586 &band,
1587 order_s,
1588 spec.alpha,
1589 spec.tau0,
1590 )?;
1591 for (level, q_l) in forms.into_iter().enumerate() {
1592 let s_l = kz.t().dot(&q_l).dot(&kz);
1593 let (s_norm, c_l) = normalize_penalty(&((&s_l + &s_l.t()) * 0.5));
1594 let intrinsic_dim = centers.ncols() as f64;
1595 let eta = 2.0 * order_s + intrinsic_dim * (2.0 - 2.0 * spec.alpha);
1596 let scale_weight = log_step * eps_band[level].powf(-eta);
1597 penalty_normalization_scales.push(c_l);
1598 raw_penalty_normalization_scales.push(c_l / scale_weight);
1599 candidates.push(PenaltyCandidate {
1600 matrix: s_norm,
1601 nullspace_dim_hint: 0,
1602 source: PenaltySource::Other(format!("measure_jet_scale_{level}")),
1603 normalization_scale: c_l,
1604 kronecker_factors: None,
1605 op: None,
1606 });
1607 }
1608 } else {
1609 let q_form = measure_jet_energy_form(
1610 centers.view(),
1611 masses.view(),
1612 &band,
1613 order_s,
1614 spec.alpha,
1615 spec.tau0,
1616 )?;
1617 let mut penalty = kz.t().dot(&q_form).dot(&kz);
1618 penalty = (&penalty + &penalty.t()) * 0.5;
1619 if spec.double_penalty {
1639 let ridge = if head_rank > 0 {
1653 let mut r_raw = Array2::<f64>::zeros((m_aug, m_aug));
1654 for i in 0..m {
1655 r_raw[(i, i)] = 1.0;
1656 }
1657 coefficient_gauge.restrict_penalty(&r_raw)
1658 } else {
1659 affine_preserving_coefficient_ridge(&kz, centers.view(), masses.view())?
1660 };
1661 let primary_fro = trace_of_product(&penalty, &penalty).sqrt();
1662 let ridge_fro = trace_of_product(&ridge, &ridge).sqrt();
1663 if primary_fro.is_finite()
1664 && primary_fro > 0.0
1665 && ridge_fro.is_finite()
1666 && ridge_fro > 0.0
1667 {
1668 let w = MEASURE_JET_FUSED_RIDGE_FRACTION * primary_fro / ridge_fro;
1669 penalty = &penalty + &(&ridge * w);
1670 }
1671 }
1672 let (penalty_norm, c_primary) = normalize_penalty(&penalty);
1673 fused_penalty_normalization_scale = Some(c_primary);
1674 candidates.push(PenaltyCandidate {
1675 matrix: penalty_norm,
1676 nullspace_dim_hint: 0,
1677 source: PenaltySource::Primary,
1678 normalization_scale: c_primary,
1679 kronecker_factors: None,
1680 op: None,
1681 });
1682 }
1683 if spec.double_penalty && per_level {
1688 let ridge = affine_preserving_coefficient_ridge(&kz, centers.view(), masses.view())?;
1689 let (ridge_norm, c_ridge) = normalize_penalty(&ridge);
1690 candidates.push(PenaltyCandidate {
1691 matrix: ridge_norm,
1692 nullspace_dim_hint: 0,
1693 source: PenaltySource::DoublePenaltyNullspace,
1694 normalization_scale: c_ridge,
1695 kronecker_factors: None,
1696 op: None,
1697 });
1698 }
1699 let (penalties, nullspace_dims, penaltyinfo, null_eigenvectors, ops) =
1700 filter_active_penalty_candidates_with_ops(candidates)?;
1701 Ok(BasisBuildResult {
1702 design,
1703 penalties,
1704 nullspace_dims,
1705 penaltyinfo,
1706 metadata: BasisMetadata::MeasureJet {
1707 centers,
1708 input_scales: None,
1709 length_scale,
1710 eps_band,
1711 order_s: spec.order_s,
1716 alpha: spec.alpha,
1717 tau0: spec.tau0,
1718 masses,
1719 support_means,
1720 penalty_normalization_scales,
1721 raw_penalty_normalization_scales,
1722 fused_penalty_normalization_scale,
1723 constraint_transform: Some(z),
1724 },
1725 kronecker_factored: None,
1726 ops,
1727 null_eigenvectors,
1728 joint_null_rotation: None,
1729 })
1730}
1731
1732pub fn build_measure_jet_basis_psi_derivatives(
1757 data: ArrayView2<'_, f64>,
1758 spec: &MeasureJetBasisSpec,
1759) -> Result<AnisoBasisPsiDerivatives, BasisError> {
1760 if !(spec.tau0.is_finite() && spec.tau0 > 0.0) {
1761 crate::bail_invalid_basis!(
1762 "measure-jet ψ derivatives need tau0 > 0 because the retained τ coordinate is ln τ; got {}",
1763 spec.tau0
1764 );
1765 }
1766 let geom = realize_measure_jet_geometry(data, spec)?;
1767 let band = MeasureJetBand {
1768 eps: geom.eps_band.clone(),
1769 log_step: geom.log_step,
1770 };
1771 let n = data.nrows();
1772 let p = geom.kz.ncols(); let kz = &geom.kz;
1774 let sandwich = |j: &Array2<f64>| {
1775 let s = kz.t().dot(j).dot(kz);
1776 (&s + &s.t()) * 0.5
1777 };
1778 let (n_coords, pairs, raw): (
1782 usize,
1783 Vec<(usize, usize)>,
1784 Vec<(
1785 Array2<f64>,
1786 Vec<Array2<f64>>,
1787 Vec<Array2<f64>>,
1788 Vec<Array2<f64>>,
1789 )>,
1790 ) = if geom.per_level {
1791 let l_count = band.eps.len();
1792 let forms = assemble_weighted_forms(
1795 geom.centers.view(),
1796 geom.masses.view(),
1797 &band,
1798 geom.order_s_eval,
1799 spec.alpha,
1800 spec.tau0,
1801 6 * l_count,
1802 3,
1803 &|scale_idx, eps: f64, q: f64, base: f64, out: &mut [[f64; 3]]| {
1804 for slot in out.iter_mut() {
1805 *slot = [0.0, 0.0, 0.0];
1806 }
1807 let intrinsic_dim = geom.centers.ncols() as f64;
1808 let ga = 2.0 * intrinsic_dim * eps.ln() - 2.0 * q.max(f64::MIN_POSITIVE).ln();
1809 let k0 = 6 * scale_idx;
1810 out[k0] = [base, 0.0, 0.0];
1811 out[k0 + 1] = [ga * base, 0.0, 0.0];
1812 out[k0 + 2] = [ga * ga * base, 0.0, 0.0];
1813 out[k0 + 3] = [0.0, 0.0, 0.0];
1814 out[k0 + 4] = [0.0, 0.0, 0.0];
1815 out[k0 + 5] = [0.0, 0.0, 0.0];
1816 },
1817 )?;
1818 let mut raw = Vec::with_capacity(l_count);
1819 for level in 0..l_count {
1820 let chunk = &forms[6 * level..6 * level + 6];
1821 raw.push((
1822 sandwich(&chunk[0]),
1823 vec![sandwich(&chunk[1]), sandwich(&chunk[3])],
1824 vec![sandwich(&chunk[2]), sandwich(&chunk[4])],
1825 vec![sandwich(&chunk[5])],
1826 ));
1827 }
1828 (2usize, vec![(0usize, 1usize)], raw)
1829 } else {
1830 let q_value = sandwich(&measure_jet_energy_form(
1840 geom.centers.view(),
1841 geom.masses.view(),
1842 &band,
1843 geom.order_s_eval,
1844 spec.alpha,
1845 spec.tau0,
1846 )?);
1847 let raw = vec![(q_value, Vec::new(), Vec::new(), Vec::new())];
1848 (0usize, Vec::new(), raw)
1849 };
1850 let length_scale_design = if spec.learn_length_scale {
1859 let ell = geom.length_scale;
1860 let k = measure_jet_design_matrix(data, geom.centers.view(), ell)?;
1861 let r2 = pairwise_sq_dists(data, geom.centers.view());
1862 let inv_l2 = 1.0 / (ell * ell);
1863 let mut dk = k.clone();
1864 let mut d2k = k.clone();
1865 for ((dk_v, d2k_v), &r2_v) in dk.iter_mut().zip(d2k.iter_mut()).zip(r2.iter()) {
1866 let a = r2_v * inv_l2;
1867 let kij = *dk_v;
1868 *dk_v = kij * a;
1869 *d2k_v = kij * (a * a - 2.0 * a);
1870 }
1871 let m = geom.centers.nrows();
1876 let m_aug = m + geom.head_transform.ncols();
1877 let mut dk_aug = Array2::<f64>::zeros((data.nrows(), m_aug));
1878 let mut d2k_aug = Array2::<f64>::zeros((data.nrows(), m_aug));
1879 dk_aug.slice_mut(ndarray::s![.., ..m]).assign(&dk);
1880 d2k_aug.slice_mut(ndarray::s![.., ..m]).assign(&d2k);
1881 let dx_du = geom.coefficient_gauge.restrict_design(&dk_aug);
1882 let d2x_du2 = geom.coefficient_gauge.restrict_design(&d2k_aug);
1883 Some((dx_du, d2x_du2))
1884 } else {
1885 None
1886 };
1887 let n_active = raw.len();
1888 let ridge = spec.double_penalty && geom.per_level;
1894 let n_cands = n_active + usize::from(ridge);
1895 let zero_p = || Array2::<f64>::zeros((p, p));
1896 let mut penalties_first: Vec<Vec<Array2<f64>>> =
1897 (0..n_coords).map(|_| Vec::with_capacity(n_cands)).collect();
1898 let mut penalties_second_diag: Vec<Vec<Array2<f64>>> =
1899 (0..n_coords).map(|_| Vec::with_capacity(n_cands)).collect();
1900 let mut crosses: Vec<Vec<Array2<f64>>> = (0..pairs.len()).map(|_| Vec::new()).collect();
1904 for (s_raw, firsts, seconds, cross_raw) in &raw {
1905 let fro = trace_of_product(s_raw, s_raw).sqrt();
1914 let c = if fro.is_finite() && fro > 1e-12 {
1915 fro
1916 } else {
1917 1.0
1918 };
1919 for coord in 0..n_coords {
1920 let (_, s_first, s_second, _) =
1921 normalize_penaltywith_psi_derivatives(s_raw, &firsts[coord], &seconds[coord]);
1922 penalties_first[coord].push(s_first);
1923 penalties_second_diag[coord].push(s_second);
1924 }
1925 for (pair_idx, &(a, b)) in pairs.iter().enumerate() {
1926 let cross_raw_mat = normalize_penalty_cross_psi_derivative(
1927 s_raw,
1928 &firsts[a],
1929 &firsts[b],
1930 &cross_raw[pair_idx],
1931 c,
1932 );
1933 crosses[pair_idx].push(cross_raw_mat);
1934 }
1935 }
1936 if ridge {
1937 for coord in 0..n_coords {
1938 penalties_first[coord].push(zero_p());
1939 penalties_second_diag[coord].push(zero_p());
1940 }
1941 for pair_crosses in crosses.iter_mut() {
1942 pair_crosses.push(zero_p());
1943 }
1944 }
1945 let coord_offset = usize::from(length_scale_design.is_some());
1951 if coord_offset == 1 {
1952 penalties_first.insert(0, (0..n_cands).map(|_| zero_p()).collect());
1953 penalties_second_diag.insert(0, (0..n_cands).map(|_| zero_p()).collect());
1954 }
1955 let n_coords_total = n_coords + coord_offset;
1956 let mut all_pairs: Vec<(usize, usize)> = pairs
1958 .iter()
1959 .map(|&(a, b)| (a + coord_offset, b + coord_offset))
1960 .collect();
1961 let mut all_crosses: Vec<Vec<Array2<f64>>> = crosses;
1962 if coord_offset == 1 {
1967 for c in 1..n_coords_total {
1968 all_pairs.push((0, c));
1969 all_crosses.push((0..n_cands).map(|_| zero_p()).collect());
1970 }
1971 }
1972 let pair_index: Vec<((usize, usize), Vec<Array2<f64>>)> = all_pairs
1973 .iter()
1974 .copied()
1975 .zip(all_crosses.into_iter())
1976 .collect();
1977 let shifted_pairs = all_pairs;
1978 let provider = AnisoPenaltyCrossProvider::new(move |a, b| {
1979 pair_index
1980 .iter()
1981 .find(|((pa, pb), _)| (*pa, *pb) == (a, b) || (*pa, *pb) == (b, a))
1982 .map(|(_, mats)| mats.clone())
1983 .ok_or_else(|| {
1984 BasisError::InvalidInput(format!(
1985 "measure-jet ψ cross derivative requested for unknown pair ({a}, {b})"
1986 ))
1987 })
1988 });
1989 let mut design_first: Vec<Array2<f64>> = (0..n_coords_total)
1990 .map(|_| Array2::<f64>::zeros((n, p)))
1991 .collect();
1992 let mut design_second_diag: Vec<Array2<f64>> = (0..n_coords_total)
1993 .map(|_| Array2::<f64>::zeros((n, p)))
1994 .collect();
1995 if let Some((dx_du, d2x_du2)) = length_scale_design {
1996 design_first[0] = dx_du;
1997 design_second_diag[0] = d2x_du2;
1998 }
1999 Ok(AnisoBasisPsiDerivatives {
2000 design_first,
2001 design_second_diag,
2002 design_second_cross: Vec::new(),
2003 design_second_cross_pairs: Vec::new(),
2004 penalties_first,
2005 penalties_second_diag,
2006 penalties_cross_pairs: shifted_pairs,
2007 penalties_cross_provider: Some(provider),
2008 implicit_operator: None,
2009 })
2010}
2011
2012#[cfg(test)]
2013mod tests {
2014 use super::*;
2015 pub(crate) fn two_cluster_centers() -> (ndarray::Array2<f64>, ndarray::Array1<f64>) {
2016 let centers = array![
2017 [0.00, 0.00],
2018 [0.31, 0.05],
2019 [0.58, -0.07],
2020 [0.93, 0.11],
2021 [1.22, 0.02],
2022 [1.49, -0.04],
2023 [3.10, 2.00],
2024 [3.42, 2.13],
2025 [3.71, 1.91],
2026 [4.05, 2.07],
2027 [4.33, 1.96],
2028 [4.61, 2.12],
2029 ];
2030 let m = centers.nrows();
2031 let masses = ndarray::Array1::<f64>::from_elem(m, 1.0 / m as f64);
2032 (centers, masses)
2033 }
2034 use ndarray::array;
2035
2036 pub(crate) fn band_for(centers: &Array2<f64>) -> MeasureJetBand {
2037 measure_jet_band(centers.view(), 0).expect("band")
2038 }
2039
2040 #[test]
2043 pub(crate) fn energy_form_annihilates_constants_exactly() {
2044 let (centers, masses) = two_cluster_centers();
2045 let band = band_for(¢ers);
2046 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2047 .expect("energy form");
2048 let m = q.nrows();
2049 let ones = Array1::<f64>::ones(m);
2050 let qv = q.dot(&ones);
2051 let scale = q.iter().fold(0.0_f64, |acc, v| acc.max(v.abs()));
2052 assert!(scale > 0.0, "energy form is identically zero");
2053 for (i, v) in qv.iter().enumerate() {
2054 assert!(
2055 v.abs() <= 1e-12 * scale,
2056 "Q·1 leak at row {i}: {v:.3e} vs scale {scale:.3e}"
2057 );
2058 }
2059 let vqv = ones.dot(&qv);
2060 assert!(
2061 vqv.abs() <= 1e-12 * scale,
2062 "constant carries energy: 1ᵀQ1 = {vqv:.3e}"
2063 );
2064 }
2065
2066 #[test]
2069 pub(crate) fn energy_form_annihilates_affine_at_default_tau() {
2070 let (centers, masses) = two_cluster_centers();
2071 let band = band_for(¢ers);
2072 let m = centers.nrows();
2073 let mut affine = Array1::<f64>::zeros(m);
2075 let mut rough = Array1::<f64>::zeros(m);
2076 for i in 0..m {
2077 affine[i] = 0.7 + 1.3 * centers[(i, 0)] - 0.4 * centers[(i, 1)];
2078 rough[i] = if i % 2 == 0 { 1.0 } else { -1.0 };
2079 }
2080 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2081 .expect("energy form");
2082 let e_affine = affine.dot(&q.dot(&affine));
2083 let e_rough = rough.dot(&q.dot(&rough));
2084 assert!(e_rough > 0.0, "rough vector must pay energy");
2085 assert!(
2086 e_affine.abs() <= 1e-12 * e_rough,
2087 "default affine energy {e_affine:.3e} vs rough {e_rough:.3e}"
2088 );
2089 }
2090
2091 #[test]
2093 pub(crate) fn energy_form_is_psd() {
2094 let (centers, masses) = two_cluster_centers();
2095 let band = band_for(¢ers);
2096 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2097 .expect("energy form");
2098 let m = q.nrows();
2099 for trial in 0..5usize {
2100 let v = Array1::<f64>::from_shape_fn(m, |i| {
2101 ((i * 7 + trial * 13) % 11) as f64 / 11.0 - 0.5
2102 });
2103 let e = v.dot(&q.dot(&v));
2104 assert!(e >= -1e-10, "vᵀQv = {e:.3e} < 0 on trial {trial}");
2105 }
2106 }
2107
2108 #[test]
2111 pub(crate) fn rough_vector_pays_more_than_smooth() {
2112 let m = 24usize;
2113 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2114 let t = i as f64 / (m as f64 - 1.0);
2115 if k == 0 {
2116 t * 4.0
2117 } else {
2118 0.3 * (t * 4.0).sin()
2119 }
2120 });
2121 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2122 let band = band_for(¢ers);
2123 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2124 .expect("energy form");
2125 let slow = Array1::<f64>::from_shape_fn(m, |i| (i as f64 / (m as f64 - 1.0)).powi(2));
2126 let fast = Array1::<f64>::from_shape_fn(m, |i| if i % 2 == 0 { 0.5 } else { -0.5 });
2127 let e_slow = slow.dot(&q.dot(&slow));
2128 let e_fast = fast.dot(&q.dot(&fast));
2129 assert!(
2130 e_fast > 10.0 * e_slow,
2131 "alternating values must pay >> a slow trend: fast {e_fast:.3e} vs slow {e_slow:.3e}"
2132 );
2133 }
2134
2135 #[test]
2140 pub(crate) fn energy_jets_match_finite_differences() {
2141 let (centers, masses) = two_cluster_centers();
2142 let band = band_for(¢ers);
2143 let (s0, a0, tau) = (1.3, 0.8, 1e-3);
2144 let jets =
2145 measure_jet_energy_form_with_jets(centers.view(), masses.view(), &band, s0, a0, tau)
2146 .expect("jets");
2147 let q_at = |s: f64, a: f64| {
2148 measure_jet_energy_form(centers.view(), masses.view(), &band, s, a, tau)
2149 .expect("energy form")
2150 };
2151 let q_plain = q_at(s0, a0);
2153 for (a, b) in jets.q.iter().zip(q_plain.iter()) {
2154 assert!(
2155 (a - b).abs() <= 1e-14 * (1.0 + b.abs()),
2156 "Q drift {a} vs {b}"
2157 );
2158 }
2159 let lt0 = tau.ln();
2160 let q_at_lt = |lt: f64| {
2161 measure_jet_energy_form(centers.view(), masses.view(), &band, s0, a0, lt.exp())
2162 .expect("energy form")
2163 };
2164 let h = 1e-4;
2170 let checks: [(&str, &Array2<f64>, Array2<f64>); 9] = [
2171 ("dq_ds", &jets.dq_ds, {
2172 let (p, m_) = (q_at(s0 + h, a0), q_at(s0 - h, a0));
2173 (&p - &m_) / (2.0 * h)
2174 }),
2175 ("d2q_ds2", &jets.d2q_ds2, {
2176 let (p, c, m_) = (q_at(s0 + h, a0), q_at(s0, a0), q_at(s0 - h, a0));
2177 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2178 }),
2179 ("dq_dalpha", &jets.dq_dalpha, {
2180 let (p, m_) = (q_at(s0, a0 + h), q_at(s0, a0 - h));
2181 (&p - &m_) / (2.0 * h)
2182 }),
2183 ("d2q_dalpha2", &jets.d2q_dalpha2, {
2184 let (p, c, m_) = (q_at(s0, a0 + h), q_at(s0, a0), q_at(s0, a0 - h));
2185 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2186 }),
2187 ("d2q_ds_dalpha", &jets.d2q_ds_dalpha, {
2188 let pp = q_at(s0 + h, a0 + h);
2189 let pm = q_at(s0 + h, a0 - h);
2190 let mp = q_at(s0 - h, a0 + h);
2191 let mm = q_at(s0 - h, a0 - h);
2192 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2193 }),
2194 ("dq_dlogtau", &jets.dq_dlogtau, {
2195 let (p, m_) = (q_at_lt(lt0 + h), q_at_lt(lt0 - h));
2196 (&p - &m_) / (2.0 * h)
2197 }),
2198 ("d2q_dlogtau2", &jets.d2q_dlogtau2, {
2199 let (p, c, m_) = (q_at_lt(lt0 + h), q_at_lt(lt0), q_at_lt(lt0 - h));
2200 (&(&p + &m_) - &(&c * 2.0)) / (h * h)
2201 }),
2202 ("d2q_ds_dlogtau", &jets.d2q_ds_dlogtau, {
2203 let f = |s: f64, lt: f64| {
2204 measure_jet_energy_form(centers.view(), masses.view(), &band, s, a0, lt.exp())
2205 .expect("energy form")
2206 };
2207 let pp = f(s0 + h, lt0 + h);
2208 let pm = f(s0 + h, lt0 - h);
2209 let mp = f(s0 - h, lt0 + h);
2210 let mm = f(s0 - h, lt0 - h);
2211 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2212 }),
2213 ("d2q_dalpha_dlogtau", &jets.d2q_dalpha_dlogtau, {
2214 let f = |a: f64, lt: f64| {
2215 measure_jet_energy_form(centers.view(), masses.view(), &band, s0, a, lt.exp())
2216 .expect("energy form")
2217 };
2218 let pp = f(a0 + h, lt0 + h);
2219 let pm = f(a0 + h, lt0 - h);
2220 let mp = f(a0 - h, lt0 + h);
2221 let mm = f(a0 - h, lt0 - h);
2222 (&(&pp - &pm) - &(&mp - &mm)) / (4.0 * h * h)
2223 }),
2224 ];
2225 for (name, analytic, fd) in checks.iter() {
2226 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2227 for (a, b) in analytic.iter().zip(fd.iter()) {
2228 assert!(
2229 (a - b).abs() <= 5e-5 * scale,
2230 "{name} jet mismatch: analytic {a:.6e} vs FD {b:.6e} (scale {scale:.3e})"
2231 );
2232 }
2233 }
2234 }
2235
2236 #[test]
2240 pub(crate) fn scale_spectrum_sums_to_total_and_localizes_roughness() {
2241 let m = 24usize;
2242 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2243 let t = i as f64 / (m as f64 - 1.0);
2244 if k == 0 { t * 4.0 } else { 0.0 }
2245 });
2246 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2247 let band = band_for(¢ers);
2248 let q = measure_jet_energy_form(centers.view(), masses.view(), &band, 1.5, 1.0, 1e-3)
2249 .expect("energy form");
2250 let fast = Array1::<f64>::from_shape_fn(m, |i| if i % 2 == 0 { 0.5 } else { -0.5 });
2251 let spec = measure_jet_scale_spectrum(
2252 centers.view(),
2253 masses.view(),
2254 &band,
2255 1.5,
2256 1.0,
2257 1e-3,
2258 fast.view(),
2259 )
2260 .expect("spectrum");
2261 assert_eq!(spec.len(), band.eps.len());
2262 let total = fast.dot(&q.dot(&fast));
2263 let sum: f64 = spec.iter().sum();
2264 assert!(
2265 (sum - total).abs() <= 1e-10 * total.abs().max(1e-30),
2266 "spectrum must sum to vᵀQv: {sum:.6e} vs {total:.6e}"
2267 );
2268 let finest = spec[0];
2270 let coarsest = *spec.last().expect("nonempty spectrum");
2271 assert!(
2272 finest > coarsest,
2273 "alternating values must charge fine scales hardest: fine {finest:.3e} vs coarse {coarsest:.3e}"
2274 );
2275 }
2276
2277 #[test]
2280 pub(crate) fn support_curve_separates_on_web_from_off_web() {
2281 let m = 24usize;
2282 let centers = Array2::<f64>::from_shape_fn((m, 2), |(i, k)| {
2283 let t = i as f64 / (m as f64 - 1.0);
2284 if k == 0 { t * 4.0 } else { 0.0 }
2285 });
2286 let masses = Array1::<f64>::from_elem(m, 1.0 / m as f64);
2287 let band = band_for(¢ers);
2288 let queries = array![[2.0, 0.0], [2.0, 1.5]];
2289 let curves =
2290 measure_jet_support_curve(queries.view(), centers.view(), masses.view(), &band.eps)
2291 .expect("support curve");
2292 assert!(
2294 curves[(0, 0)] > 10.0 * curves[(1, 0)],
2295 "fine-scale support must separate web from void: on {:.3e} vs off {:.3e}",
2296 curves[(0, 0)],
2297 curves[(1, 0)]
2298 );
2299 for qi in 0..2 {
2301 for li in 1..band.eps.len() {
2302 assert!(
2303 curves[(qi, li)] >= curves[(qi, li - 1)] - 1e-15,
2304 "support curve must be monotone in scale (query {qi}, level {li})"
2305 );
2306 }
2307 }
2308 }
2309
2310 #[test]
2317 pub(crate) fn default_stays_single_scale_until_multiscale_opt_in() {
2318 let n = 200usize;
2319 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2320 let t = i as f64 / (n as f64 - 1.0);
2321 if k == 0 {
2322 t * 3.0
2323 } else {
2324 0.4 * (t * 3.0).sin()
2325 }
2326 });
2327 let single = MeasureJetBasisSpec {
2333 center_strategy: CenterStrategy::FarthestPoint { num_centers: 80 },
2334 ..MeasureJetBasisSpec::default()
2335 };
2336 assert!(
2337 !measure_jet_multiscale_mode(&single),
2338 "default must resolve to single-scale at any center count"
2339 );
2340 let built_single =
2341 build_measure_jet_basis(data.view(), &single).expect("single-scale build");
2342 assert_eq!(
2343 built_single.penalties.len(),
2344 1,
2345 "single-scale mode emits one fused penalty (ridge folded in, not a 2nd λ)"
2346 );
2347 let multi = MeasureJetBasisSpec {
2351 center_strategy: CenterStrategy::FarthestPoint { num_centers: 80 },
2352 multiscale: true,
2353 ..MeasureJetBasisSpec::default()
2354 };
2355 assert!(
2356 measure_jet_multiscale_mode(&multi),
2357 "multiscale=true must resolve to multiscale mode"
2358 );
2359 let built_multi = build_measure_jet_basis(data.view(), &multi).expect("multiscale build");
2360 assert!(
2361 built_multi.penalties.len() > built_single.penalties.len(),
2362 "multiscale mode emits the per-scale spectral split plus the ridge, got {} (vs single-scale {})",
2363 built_multi.penalties.len(),
2364 built_single.penalties.len()
2365 );
2366 }
2367
2368 #[test]
2373 pub(crate) fn fused_mode_emits_single_primary_candidate() {
2374 let n = 40usize;
2375 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2376 let t = i as f64 / (n as f64 - 1.0);
2377 if k == 0 {
2378 t * 3.0
2379 } else {
2380 0.4 * (t * 3.0).sin()
2381 }
2382 });
2383 let spec = MeasureJetBasisSpec {
2384 center_strategy: CenterStrategy::FarthestPoint { num_centers: 14 },
2385 order_s: 1.3,
2386 ..MeasureJetBasisSpec::default()
2387 };
2388 let built = build_measure_jet_basis(data.view(), &spec).expect("fused build");
2389 assert_eq!(
2390 built.penalties.len(),
2391 1,
2392 "fused single-scale mode emits exactly one Primary candidate (ridge folded in)"
2393 );
2394 let BasisMetadata::MeasureJet { order_s, .. } = &built.metadata else {
2395 panic!("measure-jet build must return MeasureJet metadata");
2396 };
2397 assert_eq!(*order_s, 1.3, "explicit order must persist verbatim");
2398 }
2399
2400 #[test]
2402 pub(crate) fn householder_sum_to_zero_basis_is_orthonormal() {
2403 let m = 9usize;
2404 let u = householder_sum_to_zero_u(m);
2405 let z = householder_sum_to_zero_z(&u);
2406 for j in 0..(m - 1) {
2407 let col_j = z.column(j);
2408 assert!(col_j.sum().abs() <= 1e-12, "column {j} must sum to zero");
2409 for j2 in j..(m - 1) {
2410 let dot = col_j.dot(&z.column(j2));
2411 let want = if j == j2 { 1.0 } else { 0.0 };
2412 assert!(
2413 (dot - want).abs() <= 1e-12,
2414 "orthonormality failure at ({j}, {j2}): {dot}"
2415 );
2416 }
2417 }
2418 }
2419
2420 pub(crate) fn frozen_spec_fixture(
2425 order_s: f64,
2426 multiscale: bool,
2427 ) -> (Array2<f64>, MeasureJetBasisSpec) {
2428 let n = 140usize;
2433 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2434 let t = i as f64 / (n as f64 - 1.0);
2435 if k == 0 {
2436 t * 3.0
2437 } else {
2438 0.5 * (t * 3.0).cos() + if i % 9 == 0 { 0.8 } else { 0.0 }
2439 }
2440 });
2441 let spec = MeasureJetBasisSpec {
2442 center_strategy: CenterStrategy::FarthestPoint { num_centers: 70 },
2443 order_s,
2444 multiscale,
2445 learn_length_scale: false,
2449 ..MeasureJetBasisSpec::default()
2450 };
2451 let first = build_measure_jet_basis(data.view(), &spec).expect("fixture build");
2452 let BasisMetadata::MeasureJet {
2453 centers,
2454 length_scale,
2455 eps_band,
2456 masses,
2457 support_means,
2458 penalty_normalization_scales,
2459 raw_penalty_normalization_scales,
2460 fused_penalty_normalization_scale,
2461 constraint_transform,
2462 ..
2463 } = &first.metadata
2464 else {
2465 panic!("measure-jet build must return MeasureJet metadata");
2466 };
2467 let frozen = MeasureJetBasisSpec {
2468 center_strategy: CenterStrategy::UserProvided(centers.clone()),
2469 order_s,
2470 alpha: spec.alpha,
2471 tau0: spec.tau0,
2472 num_scales: eps_band.len(),
2473 length_scale: *length_scale,
2474 double_penalty: spec.double_penalty,
2475 learn_length_scale: false,
2476 multiscale,
2477 identifiability: MeasureJetIdentifiability::FrozenTransform {
2478 transform: constraint_transform.clone().expect("fit-time z"),
2479 },
2480 frozen_quadrature: Some(MeasureJetFrozenQuadrature {
2481 masses: masses.clone(),
2482 eps_band: eps_band.clone(),
2483 support_means: support_means.clone(),
2484 penalty_normalization_scales: penalty_normalization_scales.clone(),
2485 raw_penalty_normalization_scales: raw_penalty_normalization_scales.clone(),
2486 fused_penalty_normalization_scale: *fused_penalty_normalization_scale,
2487 }),
2488 };
2489 (data, frozen)
2490 }
2491
2492 #[test]
2497 pub(crate) fn psi_producer_matches_fd_per_level_mode() {
2498 let (data, frozen) = frozen_spec_fixture(0.0, true);
2499 let derivs =
2500 build_measure_jet_basis_psi_derivatives(data.view(), &frozen).expect("psi derivatives");
2501 let l_count = frozen
2502 .frozen_quadrature
2503 .as_ref()
2504 .expect("frozen quadrature")
2505 .eps_band
2506 .len();
2507 assert_eq!(
2508 derivs.penalties_first.len(),
2509 2,
2510 "per-level coords are (α, lnτ)"
2511 );
2512 assert_eq!(derivs.penalties_first[0].len(), l_count + 1);
2513 assert_eq!(derivs.penalties_cross_pairs, vec![(0, 1)]);
2514 let pen_at = |alpha: f64, tau0: f64| {
2515 let trial = MeasureJetBasisSpec {
2516 alpha,
2517 tau0,
2518 ..frozen.clone()
2519 };
2520 build_measure_jet_basis(data.view(), &trial)
2521 .expect("trial build")
2522 .penalties
2523 };
2524 let h = 1e-4;
2527 let (a0, t0) = (frozen.alpha, frozen.tau0);
2528 let ap = pen_at(a0 + h, t0);
2529 let am = pen_at(a0 - h, t0);
2530 let tp = pen_at(a0, t0 * h.exp());
2531 let tm = pen_at(a0, t0 * (-h).exp());
2532 assert_eq!(
2533 ap.len(),
2534 l_count + 1,
2535 "fixture must keep every scale active"
2536 );
2537 for level in 0..l_count {
2538 let fd_alpha = (&ap[level] - &am[level]) / (2.0 * h);
2539 let fd_tau = (&tp[level] - &tm[level]) / (2.0 * h);
2540 for (name, analytic, fd) in [
2541 ("alpha", &derivs.penalties_first[0][level], fd_alpha),
2542 ("ln_tau", &derivs.penalties_first[1][level], fd_tau),
2543 ] {
2544 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2545 for (x, y) in analytic.iter().zip(fd.iter()) {
2546 assert!(
2547 (x - y).abs() <= 5e-5 * scale,
2548 "{name} jet of scale-candidate {level}: analytic {x:.6e} vs FD {y:.6e}"
2549 );
2550 }
2551 }
2552 }
2553 for coord in 0..2 {
2555 assert!(
2556 derivs.penalties_first[coord][l_count]
2557 .iter()
2558 .all(|v| *v == 0.0),
2559 "ridge candidate must have zero ψ drift"
2560 );
2561 }
2562 let provider = derivs
2564 .penalties_cross_provider
2565 .as_ref()
2566 .expect("cross provider");
2567 let cross = provider.evaluate(0, 1).expect("cross pair (α, lnτ)");
2568 let pp = pen_at(a0 + h, t0 * h.exp());
2569 let pm = pen_at(a0 + h, t0 * (-h).exp());
2570 let mp = pen_at(a0 - h, t0 * h.exp());
2571 let mm = pen_at(a0 - h, t0 * (-h).exp());
2572 for level in 0..l_count {
2573 let fd = (&(&pp[level] - &pm[level]) - &(&mp[level] - &mm[level])) / (4.0 * h * h);
2574 let scale = fd.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2575 for (x, y) in cross[level].iter().zip(fd.iter()) {
2576 assert!(
2577 (x - y).abs() <= 5e-4 * scale,
2578 "cross (α, lnτ) jet of scale-candidate {level}: analytic {x:.6e} vs FD {y:.6e}"
2579 );
2580 }
2581 }
2582 }
2583
2584 #[test]
2590 pub(crate) fn psi_producer_matches_fd_length_scale() {
2591 let (data, mut frozen) = frozen_spec_fixture(0.0, false);
2594 frozen.learn_length_scale = true;
2595 let derivs =
2596 build_measure_jet_basis_psi_derivatives(data.view(), &frozen).expect("psi derivatives");
2597 assert_eq!(
2599 derivs.design_first.len(),
2600 1,
2601 "single-scale + learn_length_scale enrolls exactly the ℓ coordinate"
2602 );
2603 assert_eq!(
2606 derivs.penalties_first[0].len(),
2607 1,
2608 "one fitted penalty candidate"
2609 );
2610 assert!(
2611 derivs.penalties_first[0][0].iter().all(|v| *v == 0.0)
2612 && derivs.penalties_second_diag[0][0].iter().all(|v| *v == 0.0),
2613 "the jet-energy penalty must not move with ℓ"
2614 );
2615 let ell0 = frozen.length_scale;
2618 let design_at = |ell: f64| {
2619 let trial = MeasureJetBasisSpec {
2620 length_scale: ell,
2621 ..frozen.clone()
2622 };
2623 build_measure_jet_basis(data.view(), &trial)
2624 .expect("trial build")
2625 .design
2626 .to_dense()
2627 };
2628 let h: f64 = 1e-4;
2629 let x_plus = design_at(ell0 * h.exp());
2630 let x_minus = design_at(ell0 * (-h).exp());
2631 let x_0 = design_at(ell0);
2632 let fd_first = (&x_plus - &x_minus) / (2.0 * h);
2633 let fd_second = (&x_plus - &(&x_0 * 2.0) + &x_minus) / (h * h);
2634 let scale1 = fd_first.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2635 for (x, y) in derivs.design_first[0].iter().zip(fd_first.iter()) {
2636 assert!(
2637 (x - y).abs() <= 5e-5 * scale1,
2638 "∂X/∂lnℓ: analytic {x:.6e} vs FD {y:.6e}"
2639 );
2640 }
2641 let scale2 = fd_second.iter().fold(1e-30_f64, |acc, v| acc.max(v.abs()));
2642 for (x, y) in derivs.design_second_diag[0].iter().zip(fd_second.iter()) {
2643 assert!(
2644 (x - y).abs() <= 1e-3 * scale2,
2645 "∂²X/∂lnℓ²: analytic {x:.6e} vs FD {y:.6e}"
2646 );
2647 }
2648 }
2649
2650 #[test]
2654 pub(crate) fn quadrature_nodes_are_cell_barycenters() {
2655 let data = array![
2658 [0.0, 0.2],
2659 [0.4, -0.2],
2660 [0.2, 0.0],
2661 [9.8, 10.1],
2662 [10.2, 9.9],
2663 ];
2664 let seeds = array![[0.1, 0.1], [10.0, 10.0], [-50.0, -50.0]];
2665 let (nodes, masses) =
2666 measure_jet_quadrature_nodes(data.view(), seeds.view()).expect("quadrature nodes");
2667 assert!((masses.sum() - 1.0).abs() <= 1e-15, "masses must sum to 1");
2668 assert!((masses[0] - 0.6).abs() <= 1e-15);
2669 assert!((masses[1] - 0.4).abs() <= 1e-15);
2670 assert_eq!(masses[2], 0.0);
2671 assert_eq!(nodes[(0, 0)], 0.2);
2673 assert_eq!(nodes[(0, 1)], 0.0);
2674 assert_eq!(nodes[(1, 0)], 10.0);
2676 assert_eq!(nodes[(1, 1)], 10.0);
2677 assert_eq!(nodes[(2, 0)], -50.0);
2679 assert_eq!(nodes[(2, 1)], -50.0);
2680 }
2681
2682 #[test]
2686 pub(crate) fn build_replay_roundtrip_reproduces_design_and_penalty() {
2687 let n = 140usize;
2690 let data = Array2::<f64>::from_shape_fn((n, 2), |(i, k)| {
2691 let t = i as f64 / (n as f64 - 1.0);
2692 if k == 0 {
2693 t * 3.0
2694 } else {
2695 0.5 * (t * 3.0).cos() + if i % 9 == 0 { 0.8 } else { 0.0 }
2696 }
2697 });
2698 let spec = MeasureJetBasisSpec {
2699 center_strategy: CenterStrategy::FarthestPoint { num_centers: 70 },
2700 multiscale: true,
2701 ..MeasureJetBasisSpec::default()
2702 };
2703 let first = build_measure_jet_basis(data.view(), &spec).expect("first build");
2704 let BasisMetadata::MeasureJet {
2705 centers,
2706 length_scale,
2707 eps_band,
2708 order_s,
2709 alpha,
2710 tau0,
2711 masses,
2712 support_means,
2713 penalty_normalization_scales,
2714 raw_penalty_normalization_scales,
2715 fused_penalty_normalization_scale,
2716 constraint_transform,
2717 ..
2718 } = &first.metadata
2719 else {
2720 panic!("measure-jet build must return MeasureJet metadata");
2721 };
2722 let replay_spec = MeasureJetBasisSpec {
2723 center_strategy: CenterStrategy::UserProvided(centers.clone()),
2724 order_s: *order_s,
2725 alpha: *alpha,
2726 tau0: *tau0,
2727 num_scales: eps_band.len(),
2728 length_scale: *length_scale,
2729 double_penalty: spec.double_penalty,
2730 learn_length_scale: spec.learn_length_scale,
2731 multiscale: spec.multiscale,
2732 identifiability: MeasureJetIdentifiability::FrozenTransform {
2733 transform: constraint_transform.clone().expect("fit-time z"),
2734 },
2735 frozen_quadrature: Some(MeasureJetFrozenQuadrature {
2736 masses: masses.clone(),
2737 eps_band: eps_band.clone(),
2738 support_means: support_means.clone(),
2739 penalty_normalization_scales: penalty_normalization_scales.clone(),
2740 raw_penalty_normalization_scales: raw_penalty_normalization_scales.clone(),
2741 fused_penalty_normalization_scale: *fused_penalty_normalization_scale,
2742 }),
2743 };
2744 assert_eq!(
2747 first.penalties.len(),
2748 eps_band.len() + 1,
2749 "per-level mode must emit one candidate per scale + ridge"
2750 );
2751 let second = build_measure_jet_basis(data.view(), &replay_spec).expect("replay build");
2752 let x1 = first.design.to_dense();
2753 let x2 = second.design.to_dense();
2754 assert_eq!(x1.shape(), x2.shape());
2755 for (a, b) in x1.iter().zip(x2.iter()) {
2756 assert!((a - b).abs() <= 1e-12, "design replay drift: {a} vs {b}");
2757 }
2758 assert_eq!(first.penalties.len(), second.penalties.len());
2759 for (p1, p2) in first.penalties.iter().zip(second.penalties.iter()) {
2760 for (a, b) in p1.iter().zip(p2.iter()) {
2761 assert!((a - b).abs() <= 1e-12, "penalty replay drift: {a} vs {b}");
2762 }
2763 }
2764 }
2765}