1use crate::chart_canonicalization::CanonicalChartTopology;
62use faer::Side;
63use gam_linalg::faer_ndarray::FaerEigh;
64use gam_terms::basis::{
65 BasisOptions, Dense, KnotSource, PeriodicBSplineBasisSpec, build_periodic_bspline_basis_1d,
66 create_basis, create_cyclic_difference_penalty_matrix, create_difference_penalty_matrix,
67 periodic_bspline_first_derivative_nd,
68};
69use gam_terms::inference::smooth_test::{SmoothTestInput, SmoothTestScale, wood_smooth_test};
70use ndarray::{Array1, Array2, ArrayView1, Axis};
71use statrs::distribution::{ContinuousCDF, Normal};
72use std::f64::consts::{PI, TAU};
73
74const TRANSPORT_SPLINE_DEGREE: usize = 3;
76const TRANSPORT_PENALTY_ORDER: usize = 2;
80const MIN_TRANSPORT_OBS: usize = 16;
82const OBS_PER_BASIS: usize = 8;
84const MIN_PERIODIC_BASIS: usize = 8;
86const MAX_PERIODIC_BASIS: usize = 20;
87const MIN_OPEN_INTERNAL_KNOTS: usize = 4;
89const MAX_OPEN_INTERNAL_KNOTS: usize = 12;
90const DEGREE_CANDIDATES: [i32; 5] = [-2, -1, 0, 1, 2];
92const FOLD_CHECK_GRID: usize = 512;
94pub const DEFAULT_COMPOSITION_GRID: usize = 256;
96const COMPOSITION_DEFECT_REL_VAR_FLOOR: f64 = 1e-4;
117const REML_LAMBDA_GRID_POINTS: usize = 41;
119const REML_GOLDEN_ITERATIONS: usize = 40;
120const REML_LAMBDA_SPAN_DECADES: f64 = 8.0;
121
122#[derive(Debug, Clone, Copy, PartialEq)]
124pub enum ChartTopology {
125 Circle,
127 Interval { lo: f64, hi: f64 },
129}
130
131impl ChartTopology {
132 pub fn name(&self) -> &'static str {
134 match self {
135 ChartTopology::Circle => "circle",
136 ChartTopology::Interval { .. } => "interval",
137 }
138 }
139
140 fn validate(&self) -> Result<(), String> {
141 match *self {
142 ChartTopology::Circle => Ok(()),
143 ChartTopology::Interval { lo, hi } => {
144 if !(lo.is_finite() && hi.is_finite()) || hi <= lo {
145 Err(format!(
146 "interval chart bounds must be finite and ordered; got [{lo}, {hi}]"
147 ))
148 } else {
149 Ok(())
150 }
151 }
152 }
153 }
154}
155
156impl From<&CanonicalChartTopology> for ChartTopology {
167 fn from(src: &CanonicalChartTopology) -> Self {
168 match src {
169 CanonicalChartTopology::Circle { .. } => ChartTopology::Circle,
170 CanonicalChartTopology::Interval => ChartTopology::Interval { lo: 0.0, hi: 1.0 },
171 }
172 }
173}
174
175impl From<CanonicalChartTopology> for ChartTopology {
176 fn from(src: CanonicalChartTopology) -> Self {
177 ChartTopology::from(&src)
178 }
179}
180
181fn wrap_tau(x: f64) -> f64 {
183 x.rem_euclid(TAU)
184}
185
186fn wrap_pi(x: f64) -> f64 {
188 let w = (x + PI).rem_euclid(TAU) - PI;
189 if w <= -PI { w + TAU } else { w }
190}
191
192fn circular_mean(angles: &[f64]) -> f64 {
194 let mut s = 0.0_f64;
195 let mut c = 0.0_f64;
196 for &a in angles {
197 s += a.sin();
198 c += a.cos();
199 }
200 if s.hypot(c) <= f64::EPSILON * angles.len().max(1) as f64 {
201 0.0
202 } else {
203 s.atan2(c)
204 }
205}
206
207fn resultant_length(angles: &[f64]) -> f64 {
209 if angles.is_empty() {
210 return 0.0;
211 }
212 let mut s = 0.0_f64;
213 let mut c = 0.0_f64;
214 for &a in angles {
215 s += a.sin();
216 c += a.cos();
217 }
218 s.hypot(c) / angles.len() as f64
219}
220
221#[derive(Debug, Clone)]
225enum DomainBasis {
226 Periodic(PeriodicBSplineBasisSpec),
227 Open { knots: Array1<f64>, degree: usize },
228}
229
230impl DomainBasis {
231 fn build(topology: ChartTopology, coords: ArrayView1<'_, f64>) -> Result<Self, String> {
232 let n = coords.len();
233 match topology {
234 ChartTopology::Circle => {
235 let num_basis = (n / OBS_PER_BASIS).clamp(MIN_PERIODIC_BASIS, MAX_PERIODIC_BASIS);
236 Ok(DomainBasis::Periodic(PeriodicBSplineBasisSpec {
237 degree: TRANSPORT_SPLINE_DEGREE,
238 num_basis,
239 period: TAU,
240 origin: 0.0,
241 penalty_order: TRANSPORT_PENALTY_ORDER,
242 }))
243 }
244 ChartTopology::Interval { lo, hi } => {
245 let num_internal =
246 (n / OBS_PER_BASIS).clamp(MIN_OPEN_INTERNAL_KNOTS, MAX_OPEN_INTERNAL_KNOTS);
247 let (seed, knots) = create_basis::<Dense>(
248 coords.mapv(|v| v.clamp(lo, hi)).view(),
249 KnotSource::Generate {
250 data_range: (lo, hi),
251 num_internal_knots: num_internal,
252 },
253 TRANSPORT_SPLINE_DEGREE,
254 BasisOptions::value(),
255 )
256 .map_err(|e| format!("layer transport open basis construction failed: {e}"))?;
257 if seed.nrows() != n {
258 return Err(format!(
259 "layer transport open basis returned {} rows for {n} inputs",
260 seed.nrows()
261 ));
262 }
263 Ok(DomainBasis::Open {
264 knots,
265 degree: TRANSPORT_SPLINE_DEGREE,
266 })
267 }
268 }
269 }
270
271 fn num_basis(&self) -> usize {
272 match self {
273 DomainBasis::Periodic(spec) => spec.num_basis,
274 DomainBasis::Open { knots, degree } => knots.len() - degree - 1,
275 }
276 }
277
278 fn penalty_rank(&self) -> usize {
282 match self {
283 DomainBasis::Periodic(spec) => spec.num_basis - 1,
284 DomainBasis::Open { .. } => self.num_basis() - TRANSPORT_PENALTY_ORDER,
285 }
286 }
287
288 fn penalty(&self) -> Result<Array2<f64>, String> {
289 match self {
290 DomainBasis::Periodic(spec) => {
291 create_cyclic_difference_penalty_matrix(spec.num_basis, TRANSPORT_PENALTY_ORDER)
292 .map_err(|e| format!("cyclic transport penalty failed: {e}"))
293 }
294 DomainBasis::Open { .. } => {
295 create_difference_penalty_matrix(self.num_basis(), TRANSPORT_PENALTY_ORDER, None)
296 .map_err(|e| format!("open transport penalty failed: {e}"))
297 }
298 }
299 }
300
301 fn project(&self, t: f64) -> f64 {
303 match self {
304 DomainBasis::Periodic(_) => wrap_tau(t),
305 DomainBasis::Open { knots, degree } => {
306 let lo = knots[*degree];
307 let hi = knots[knots.len() - 1 - degree];
308 t.clamp(lo, hi)
309 }
310 }
311 }
312
313 fn value_rows(&self, t: ArrayView1<'_, f64>) -> Result<Array2<f64>, String> {
314 let projected = t.mapv(|v| self.project(v));
315 match self {
316 DomainBasis::Periodic(spec) => build_periodic_bspline_basis_1d(projected.view(), spec)
317 .map_err(|e| format!("periodic transport basis evaluation failed: {e}")),
318 DomainBasis::Open { knots, degree } => {
319 let (rows, used_knots) = create_basis::<Dense>(
320 projected.view(),
321 KnotSource::Provided(knots.view()),
322 *degree,
323 BasisOptions::value(),
324 )
325 .map_err(|e| format!("open transport basis evaluation failed: {e}"))?;
326 if used_knots.len() != knots.len() {
327 return Err("open transport basis knot vector drifted".to_string());
328 }
329 Ok(rows.as_ref().to_owned())
330 }
331 }
332 }
333
334 fn derivative_poly_degree(&self) -> usize {
337 let degree = match self {
338 DomainBasis::Periodic(spec) => spec.degree,
339 DomainBasis::Open { degree, .. } => *degree,
340 };
341 degree.saturating_sub(1)
342 }
343
344 fn derivative_breakpoints(&self) -> Vec<f64> {
352 match self {
353 DomainBasis::Periodic(spec) => {
354 let n_seg = spec.num_basis.max(1);
358 (0..=n_seg)
359 .map(|k| spec.origin + spec.period * k as f64 / n_seg as f64)
360 .collect()
361 }
362 DomainBasis::Open { knots, degree } => {
363 let lo = knots[*degree];
364 let hi = knots[knots.len() - 1 - degree];
365 let mut breaks: Vec<f64> = Vec::with_capacity(knots.len());
366 for &k in knots.iter() {
367 if k > lo + 0.0 && k < hi {
368 breaks.push(k);
369 }
370 }
371 breaks.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
372 breaks.dedup_by(|a, b| (*a - *b).abs() <= f64::EPSILON * hi.abs().max(1.0));
373 let mut out = Vec::with_capacity(breaks.len() + 2);
374 out.push(lo);
375 out.extend(breaks.into_iter().filter(|&k| k > lo && k < hi));
376 out.push(hi);
377 out
378 }
379 }
380 }
381
382 fn derivative_rows(&self, t: ArrayView1<'_, f64>) -> Result<Array2<f64>, String> {
383 let projected = t.mapv(|v| self.project(v));
384 match self {
385 DomainBasis::Periodic(spec) => {
386 let n = projected.len();
387 let mut col = Array2::<f64>::zeros((n, 1));
388 for (i, &v) in projected.iter().enumerate() {
389 col[[i, 0]] = v;
390 }
391 let jet = periodic_bspline_first_derivative_nd(
392 col.view(),
393 (0.0, TAU),
394 spec.degree,
395 spec.num_basis,
396 )
397 .map_err(|e| format!("periodic transport derivative failed: {e}"))?;
398 Ok(jet.index_axis(Axis(2), 0).to_owned())
399 }
400 DomainBasis::Open { knots, degree } => {
401 let (rows, used_knots) = create_basis::<Dense>(
402 projected.view(),
403 KnotSource::Provided(knots.view()),
404 *degree,
405 BasisOptions::first_derivative(),
406 )
407 .map_err(|e| format!("open transport derivative failed: {e}"))?;
408 if used_knots.len() != knots.len() {
409 return Err("open transport derivative knot vector drifted".to_string());
410 }
411 Ok(rows.as_ref().to_owned())
412 }
413 }
414 }
415}
416
417struct Penalized1dFit {
422 beta: Array1<f64>,
423 covariance: Array2<f64>,
426 influence: Array2<f64>,
429 lambda: f64,
430 edf: f64,
431 sigma2: f64,
432 residual_rms: f64,
433}
434
435fn fit_penalized_1d(
445 design: &Array2<f64>,
446 penalty: &Array2<f64>,
447 response: ArrayView1<'_, f64>,
448 weights: Option<ArrayView1<'_, f64>>,
449 penalty_rank: usize,
450 known_scale: bool,
451) -> Result<Penalized1dFit, String> {
452 let n = design.nrows();
453 let m = design.ncols();
454 if response.len() != n || penalty.nrows() != m || penalty.ncols() != m {
455 return Err(format!(
456 "penalized 1-D fit shape mismatch: X is {n}×{m}, y has {}, S is {}×{}",
457 response.len(),
458 penalty.nrows(),
459 penalty.ncols()
460 ));
461 }
462 if let Some(w) = weights {
463 if w.len() != n {
464 return Err(format!(
465 "penalized 1-D fit weight length {} does not match n = {n}",
466 w.len()
467 ));
468 }
469 if w.iter().any(|&v| !v.is_finite() || v <= 0.0) {
470 return Err("penalized 1-D fit weights must be finite and positive".to_string());
471 }
472 }
473
474 let mut xtwx = Array2::<f64>::zeros((m, m));
475 let mut xtwy = Array1::<f64>::zeros(m);
476 let mut ytwy = 0.0_f64;
477 let mut sum_w = 0.0_f64;
478 for r in 0..n {
479 let w = weights.map_or(1.0, |wv| wv[r]);
480 let y = response[r];
481 ytwy += w * y * y;
482 sum_w += w;
483 for j in 0..m {
484 let xj = design[[r, j]];
485 if xj == 0.0 {
486 continue;
487 }
488 xtwy[j] += w * xj * y;
489 for k in j..m {
490 xtwx[[j, k]] += w * xj * design[[r, k]];
491 }
492 }
493 }
494 for j in 0..m {
495 for k in 0..j {
496 xtwx[[j, k]] = xtwx[[k, j]];
497 }
498 }
499
500 let trace_scale = (0..m).map(|i| xtwx[[i, i]]).sum::<f64>() / m as f64;
501 let anchor = trace_scale.max(f64::MIN_POSITIVE);
502 let nullspace_dim = m.saturating_sub(penalty_rank);
503 let dof = ((n as f64) - nullspace_dim as f64).max(1.0);
504 let rank_f = penalty_rank as f64;
505
506 let solve_at = |lambda: f64| -> Result<(Array1<f64>, Array1<f64>, Array2<f64>), String> {
507 let mut a = xtwx.clone();
508 for j in 0..m {
509 for k in 0..m {
510 a[[j, k]] += lambda * penalty[[j, k]];
511 }
512 }
513 let diag_scale = (0..m).map(|i| a[[i, i]].abs()).fold(1.0_f64, f64::max);
515 for i in 0..m {
516 a[[i, i]] += 1e-12 * diag_scale;
517 }
518 let (evals, evecs) = a
519 .eigh(Side::Lower)
520 .map_err(|e| format!("penalized 1-D fit eigendecomposition failed: {e:?}"))?;
521 Ok((evals, evecs.t().dot(&xtwy), evecs))
522 };
523
524 let criterion = |lambda: f64| -> f64 {
525 let Ok(parts) = solve_at(lambda) else {
526 return f64::INFINITY;
527 };
528 let (evals, rotated) = (&parts.0, &parts.1);
529 let floor = evals.iter().copied().fold(0.0_f64, f64::max) * 1e-14;
530 let mut prss = ytwy;
531 let mut logdet = 0.0_f64;
532 for i in 0..m {
533 let d = evals[i].max(floor).max(f64::MIN_POSITIVE);
534 prss -= rotated[i] * rotated[i] / d;
535 logdet += d.ln();
536 }
537 let prss = prss.max(f64::MIN_POSITIVE);
538 let fit_term = if known_scale { prss } else { dof * prss.ln() };
539 fit_term + logdet - rank_f * lambda.ln()
540 };
541
542 let lo = anchor * 10f64.powf(-REML_LAMBDA_SPAN_DECADES);
543 let hi = anchor * 10f64.powf(REML_LAMBDA_SPAN_DECADES);
544 let grid: Vec<f64> = (0..REML_LAMBDA_GRID_POINTS)
545 .map(|i| {
546 let t = i as f64 / (REML_LAMBDA_GRID_POINTS - 1) as f64;
547 lo * (hi / lo).powf(t)
548 })
549 .collect();
550 let mut best_idx = 0usize;
551 let mut best_val = f64::INFINITY;
552 for (i, &lam) in grid.iter().enumerate() {
553 let v = criterion(lam);
554 if v < best_val {
555 best_val = v;
556 best_idx = i;
557 }
558 }
559 let mut a_log = grid[best_idx.saturating_sub(1)].ln();
560 let mut c_log = grid[(best_idx + 1).min(REML_LAMBDA_GRID_POINTS - 1)].ln();
561 let golden = (5.0_f64.sqrt() - 1.0) / 2.0;
562 let mut x1 = c_log - golden * (c_log - a_log);
563 let mut x2 = a_log + golden * (c_log - a_log);
564 let mut f1 = criterion(x1.exp());
565 let mut f2 = criterion(x2.exp());
566 for _ in 0..REML_GOLDEN_ITERATIONS {
567 if f1 <= f2 {
568 c_log = x2;
569 x2 = x1;
570 f2 = f1;
571 x1 = c_log - golden * (c_log - a_log);
572 f1 = criterion(x1.exp());
573 } else {
574 a_log = x1;
575 x1 = x2;
576 f1 = f2;
577 x2 = a_log + golden * (c_log - a_log);
578 f2 = criterion(x2.exp());
579 }
580 }
581 let lambda = (0.5 * (a_log + c_log)).exp();
582
583 let (evals, rotated, evecs) = solve_at(lambda)?;
584 let floor = evals.iter().copied().fold(0.0_f64, f64::max) * 1e-14;
585 let mut a_inv = Array2::<f64>::zeros((m, m));
586 let mut beta = Array1::<f64>::zeros(m);
587 for i in 0..m {
588 let d = evals[i].max(floor).max(f64::MIN_POSITIVE);
589 let coeff = rotated[i] / d;
590 for j in 0..m {
591 beta[j] += evecs[[j, i]] * coeff;
592 for k in 0..m {
593 a_inv[[j, k]] += evecs[[j, i]] * evecs[[k, i]] / d;
594 }
595 }
596 }
597 let influence = a_inv.dot(&xtwx);
598 let edf = (0..m).map(|i| influence[[i, i]]).sum::<f64>();
599
600 let fitted = design.dot(&beta);
601 let mut rss = 0.0_f64;
602 for r in 0..n {
603 let w = weights.map_or(1.0, |wv| wv[r]);
604 let e = response[r] - fitted[r];
605 rss += w * e * e;
606 }
607 let sigma2 = if known_scale {
608 1.0
609 } else {
610 (rss / ((n as f64) - edf).max(1.0)).max(f64::MIN_POSITIVE)
611 };
612 let covariance = a_inv.mapv(|v| v * sigma2);
613 let residual_rms = (rss / sum_w.max(f64::MIN_POSITIVE)).sqrt();
614
615 if beta.iter().any(|v| !v.is_finite()) {
616 return Err("penalized 1-D fit produced non-finite coefficients".to_string());
617 }
618 Ok(Penalized1dFit {
619 beta,
620 covariance,
621 influence,
622 lambda,
623 edf,
624 sigma2,
625 residual_rms,
626 })
627}
628
629#[derive(Debug, Clone)]
639pub struct FittedTransport {
640 pub topology_from: ChartTopology,
641 pub topology_to: ChartTopology,
642 pub degree: Option<i32>,
644 pub degree_concentration: Option<f64>,
647 pub rotation_offset: f64,
652 pub beta: Array1<f64>,
654 pub covariance: Array2<f64>,
656 pub smoothing_lambda: f64,
657 pub edf: f64,
659 pub noise_variance: f64,
661 pub n_obs: usize,
662 pub isometry_defect: f64,
664 pub isometry_defect_se: f64,
666 pub topology_preserved: bool,
669 pub min_directional_derivative: f64,
671 pub residual_rms: f64,
673 basis: DomainBasis,
674}
675
676impl FittedTransport {
677 fn linear_slope(&self) -> f64 {
678 self.degree.map_or(0.0, f64::from)
679 }
680
681 pub fn eval(&self, t: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
683 let rows = self.basis.value_rows(t)?;
684 let smooth = rows.dot(&self.beta);
685 let slope = self.linear_slope();
686 let mut out = Array1::<f64>::zeros(t.len());
687 for i in 0..t.len() {
688 let raw = slope * t[i] + self.rotation_offset + smooth[i];
689 out[i] = match self.topology_to {
690 ChartTopology::Circle => wrap_tau(raw),
691 ChartTopology::Interval { .. } => raw,
692 };
693 }
694 Ok(out)
695 }
696
697 pub fn eval_with_variance(
699 &self,
700 t: ArrayView1<'_, f64>,
701 ) -> Result<(Array1<f64>, Array1<f64>), String> {
702 let rows = self.basis.value_rows(t)?;
703 let values = self.eval(t)?;
704 let mut variances = Array1::<f64>::zeros(t.len());
705 for i in 0..t.len() {
706 let row = rows.row(i);
707 variances[i] = row.dot(&self.covariance.dot(&row)).max(0.0);
708 }
709 Ok((values, variances))
710 }
711
712 pub fn derivative(&self, t: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
714 let rows = self.basis.derivative_rows(t)?;
715 let slope = self.linear_slope();
716 Ok(rows.dot(&self.beta).mapv(|v| v + slope))
717 }
718
719 fn raw_at(&self, t: f64) -> Result<f64, String> {
723 let arr = Array1::from_elem(1, t);
724 let smooth = self.basis.value_rows(arr.view())?.dot(&self.beta)[0];
725 Ok(self.linear_slope() * t + self.rotation_offset + smooth)
726 }
727
728 fn oriented_derivative_at(&self, t: &[f64], orientation: f64) -> Result<Vec<f64>, String> {
730 let arr = Array1::from_vec(t.to_vec());
731 let rows = self.basis.derivative_rows(arr.view())?;
732 let slope = self.linear_slope();
733 Ok((0..t.len())
734 .map(|i| orientation * (rows.row(i).dot(&self.beta) + slope))
735 .collect())
736 }
737
738 fn certify_strict_monotonicity(&self) -> Result<f64, String> {
756 let (lo, hi) = match self.topology_from {
757 ChartTopology::Circle => (0.0, TAU),
758 ChartTopology::Interval { lo, hi } => (lo, hi),
759 };
760 let raw_lo = self.raw_at(lo)?;
763 let raw_hi = self.raw_at(hi)?;
764 let orientation = if raw_hi >= raw_lo { 1.0 } else { -1.0 };
765
766 let deg = self.basis.derivative_poly_degree().max(1);
767 let breaks = self.basis.derivative_breakpoints();
768 for window in breaks.windows(2) {
771 let (a, b) = (window[0], window[1]);
772 if !(b > a) {
773 continue;
774 }
775 let span = b - a;
776 let pad = span * 1.0e-9;
780 let n_nodes = deg + 1;
781 let nodes: Vec<f64> = (0..n_nodes)
782 .map(|i| {
783 let s = if n_nodes == 1 {
784 0.5
785 } else {
786 i as f64 / (n_nodes - 1) as f64
787 };
788 (a + pad) + (span - 2.0 * pad) * s
789 })
790 .collect();
791 let values = self.oriented_derivative_at(&nodes, orientation)?;
792
793 let step = if n_nodes > 1 {
798 nodes[1] - nodes[0]
799 } else {
800 span
801 };
802 let coeffs = monomial_from_equispaced(&values);
803
804 let probe_t = a + 0.37 * span;
810 let probe_u = (probe_t - nodes[0]) / step;
811 let probe_recon = eval_monomial(&coeffs, probe_u);
812 let probe_actual = self.oriented_derivative_at(&[probe_t], orientation)?[0];
813 let scale = probe_actual.abs().max(1.0);
814 if (probe_recon - probe_actual).abs() > 1.0e-6 * scale {
815 return Err(format!(
816 "transport monotonicity certificate could not reconstruct h′ on the \
817 span [{a}, {b}] (reconstruction {probe_recon} vs actual {probe_actual}); \
818 refusing to certify"
819 ));
820 }
821
822 for &edge in &[a, b] {
824 let u = (edge - nodes[0]) / step;
825 let v = eval_monomial(&coeffs, u);
826 if !(v > 0.0) {
827 return Err(format!(
828 "transport map is not strictly monotone: orientation·h′ = {v} ≤ 0 at \
829 t = {edge}"
830 ));
831 }
832 }
833 for u_crit in monomial_critical_points(&coeffs) {
836 let t_crit = nodes[0] + u_crit * step;
837 if t_crit > a && t_crit < b {
838 let v = eval_monomial(&coeffs, u_crit);
839 if !(v > 0.0) {
840 return Err(format!(
841 "transport map folds: orientation·h′ = {v} ≤ 0 at interior \
842 extremum t = {t_crit}"
843 ));
844 }
845 }
846 }
847 }
848 Ok(orientation)
849 }
850
851 pub fn invert(&self, y: ArrayView1<'_, f64>) -> Result<Array1<f64>, String> {
871 if y.iter().any(|v| !v.is_finite()) {
872 return Err("transport inverse targets must be finite".to_string());
873 }
874 self.certify_strict_monotonicity()?;
877 let (lo, hi) = match self.topology_from {
878 ChartTopology::Circle => (0.0, TAU),
879 ChartTopology::Interval { lo, hi } => (lo, hi),
880 };
881 let raw_lo = self.raw_at(lo)?;
884 let raw_hi = self.raw_at(hi)?;
885 let increasing = raw_hi > raw_lo;
886 let (raw_min, raw_max) = if increasing {
887 (raw_lo, raw_hi)
888 } else {
889 (raw_hi, raw_lo)
890 };
891 let scale = raw_min.abs().max(raw_max.abs()).max(1.0);
894 let tol = 32.0 * f64::EPSILON * scale;
895
896 let mut probe = Array1::<f64>::zeros(1);
899 let mut raw_at_into = |t: f64| -> Result<f64, String> {
900 probe[0] = t;
901 let smooth = self.basis.value_rows(probe.view())?.dot(&self.beta)[0];
902 Ok(self.linear_slope() * t + self.rotation_offset + smooth)
903 };
904
905 let mut out = Array1::<f64>::zeros(y.len());
906 for (idx, &yi) in y.iter().enumerate() {
907 let target = match self.topology_to {
909 ChartTopology::Interval { .. } => {
910 if yi < raw_min - tol || yi > raw_max + tol {
911 return Err(format!(
912 "transport inverse target {yi} is outside the fitted image \
913 [{raw_min}, {raw_max}]"
914 ));
915 }
916 yi.clamp(raw_min, raw_max)
917 }
918 ChartTopology::Circle => {
919 let ywrapped = wrap_tau(yi);
922 let m = ((raw_min - ywrapped) / TAU).ceil();
923 ywrapped + TAU * m
924 }
925 };
926 let (mut a, mut b) = (lo, hi);
930 let width_floor = f64::EPSILON * hi.abs().max(lo.abs()).max(1.0);
931 for _ in 0..100 {
932 if (b - a) <= width_floor {
933 break;
934 }
935 let mid = 0.5 * (a + b);
936 let rm = raw_at_into(mid)?;
937 let go_right = if increasing { rm < target } else { rm > target };
938 if go_right {
939 a = mid;
940 } else {
941 b = mid;
942 }
943 }
944 out[idx] = 0.5 * (a + b);
945 }
946 Ok(out)
947 }
948
949 pub fn report(&self, layer_from: usize, layer_to: usize) -> LayerTransportReport {
952 LayerTransportReport {
953 layer_from,
954 layer_to,
955 topology_from: self.topology_from,
956 topology_to: self.topology_to,
957 topology_preserved: self.topology_preserved,
958 degree: self.degree,
959 degree_concentration: self.degree_concentration,
960 rotation_offset: self.rotation_offset,
961 isometry_defect: self.isometry_defect,
962 isometry_defect_se: self.isometry_defect_se,
963 min_directional_derivative: self.min_directional_derivative,
964 transport_edf: self.edf,
965 smoothing_lambda: self.smoothing_lambda,
966 noise_variance: self.noise_variance,
967 residual_rms: self.residual_rms,
968 n_obs: self.n_obs,
969 composition_defect: None,
970 composition_max_studentized: None,
971 composition_p_value: None,
972 composition_gauge_reflected: None,
973 }
974 }
975}
976
977#[derive(Debug, Clone)]
979pub struct LayerTransportReport {
980 pub layer_from: usize,
981 pub layer_to: usize,
982 pub topology_from: ChartTopology,
983 pub topology_to: ChartTopology,
984 pub topology_preserved: bool,
986 pub degree: Option<i32>,
988 pub degree_concentration: Option<f64>,
990 pub rotation_offset: f64,
992 pub isometry_defect: f64,
994 pub isometry_defect_se: f64,
996 pub min_directional_derivative: f64,
998 pub transport_edf: f64,
1000 pub smoothing_lambda: f64,
1001 pub noise_variance: f64,
1002 pub residual_rms: f64,
1003 pub n_obs: usize,
1004 pub composition_defect: Option<f64>,
1007 pub composition_max_studentized: Option<f64>,
1009 pub composition_p_value: Option<f64>,
1012 pub composition_gauge_reflected: Option<bool>,
1014}
1015
1016impl LayerTransportReport {
1017 pub fn with_composition(mut self, composition: &CompositionDefectReport) -> Self {
1019 self.composition_defect = Some(composition.rms_defect);
1020 self.composition_max_studentized = Some(composition.max_studentized_defect);
1021 self.composition_p_value = Some(composition.p_value);
1022 self.composition_gauge_reflected = Some(composition.gauge_reflected);
1023 self
1024 }
1025}
1026
1027pub fn fit_transport_map(
1034 coords_from: ArrayView1<'_, f64>,
1035 coords_to: ArrayView1<'_, f64>,
1036 topology_from: ChartTopology,
1037 topology_to: ChartTopology,
1038) -> Result<FittedTransport, String> {
1039 let n = coords_from.len();
1040 if coords_to.len() != n {
1041 return Err(format!(
1042 "layer transport coordinate lengths disagree: {} vs {}",
1043 n,
1044 coords_to.len()
1045 ));
1046 }
1047 if n < MIN_TRANSPORT_OBS {
1048 return Err(format!(
1049 "layer transport needs at least {MIN_TRANSPORT_OBS} paired observations, got {n}"
1050 ));
1051 }
1052 if coords_from
1053 .iter()
1054 .chain(coords_to.iter())
1055 .any(|v| !v.is_finite())
1056 {
1057 return Err("layer transport coordinates must all be finite".to_string());
1058 }
1059 topology_from.validate()?;
1060 topology_to.validate()?;
1061
1062 let (degree, degree_concentration, rotation_offset, response): (
1064 Option<i32>,
1065 Option<f64>,
1066 f64,
1067 Array1<f64>,
1068 ) = match (topology_from, topology_to) {
1069 (ChartTopology::Circle, ChartTopology::Circle) => {
1070 let mut best_degree = DEGREE_CANDIDATES[0];
1077 let mut best_r = f64::NEG_INFINITY;
1078 for &d in DEGREE_CANDIDATES.iter() {
1079 let residual: Vec<f64> = (0..n)
1080 .map(|i| coords_to[i] - f64::from(d) * coords_from[i])
1081 .collect();
1082 let r = resultant_length(&residual);
1083 if r > best_r {
1084 best_r = r;
1085 best_degree = d;
1086 }
1087 }
1088 let residual: Vec<f64> = (0..n)
1089 .map(|i| coords_to[i] - f64::from(best_degree) * coords_from[i])
1090 .collect();
1091 let mu = circular_mean(&residual);
1092 let response = Array1::from_iter(residual.iter().map(|&r| wrap_pi(r - mu)));
1093 (Some(best_degree), Some(best_r), mu, response)
1094 }
1095 (_, ChartTopology::Circle) => {
1096 let angles: Vec<f64> = coords_to.iter().copied().collect();
1100 let mu = circular_mean(&angles);
1101 let response = Array1::from_iter(angles.iter().map(|&a| wrap_pi(a - mu)));
1102 (None, None, mu, response)
1103 }
1104 (_, ChartTopology::Interval { .. }) => (None, None, 0.0, coords_to.to_owned()),
1105 };
1106
1107 let basis = DomainBasis::build(topology_from, coords_from)?;
1109 let design = basis.value_rows(coords_from)?;
1110 let penalty = basis.penalty()?;
1111 let fit = fit_penalized_1d(
1112 &design,
1113 &penalty,
1114 response.view(),
1115 None,
1116 basis.penalty_rank(),
1117 false,
1118 )?;
1119
1120 let slope = degree.map_or(0.0, f64::from);
1122 let deriv_rows = basis.derivative_rows(coords_from)?;
1123 let deriv = deriv_rows.dot(&fit.beta).mapv(|v| v + slope);
1124 let m = basis.num_basis();
1125 let mut defect = 0.0_f64;
1126 let mut grad = Array1::<f64>::zeros(m);
1127 for i in 0..n {
1128 let speed = deriv[i].abs();
1129 let gap = speed - 1.0;
1130 defect += gap * gap;
1131 let sgn = if deriv[i] >= 0.0 { 1.0 } else { -1.0 };
1132 for j in 0..m {
1133 grad[j] += 2.0 * gap * sgn * deriv_rows[[i, j]];
1134 }
1135 }
1136 defect /= n as f64;
1137 grad.mapv_inplace(|v| v / n as f64);
1138 let isometry_defect_se = grad.dot(&fit.covariance.dot(&grad)).max(0.0).sqrt();
1139
1140 let grid = domain_grid(topology_from, FOLD_CHECK_GRID);
1142 let grid_deriv = basis
1143 .derivative_rows(grid.view())?
1144 .dot(&fit.beta)
1145 .mapv(|v| v + slope);
1146 let orientation = if slope != 0.0 {
1147 slope.signum()
1148 } else {
1149 let mean = grid_deriv.iter().sum::<f64>() / grid_deriv.len() as f64;
1150 if mean < 0.0 { -1.0 } else { 1.0 }
1151 };
1152 let min_directional_derivative = grid_deriv
1153 .iter()
1154 .map(|&v| orientation * v)
1155 .fold(f64::INFINITY, f64::min);
1156 let topology_preserved = match (topology_from, topology_to) {
1157 (ChartTopology::Circle, ChartTopology::Circle) => {
1158 matches!(degree, Some(1) | Some(-1)) && min_directional_derivative > 0.0
1159 }
1160 (ChartTopology::Interval { .. }, ChartTopology::Interval { .. }) => {
1161 min_directional_derivative > 0.0
1162 }
1163 _ => false,
1164 };
1165
1166 Ok(FittedTransport {
1167 topology_from,
1168 topology_to,
1169 degree,
1170 degree_concentration,
1171 rotation_offset,
1172 beta: fit.beta,
1173 covariance: fit.covariance,
1174 smoothing_lambda: fit.lambda,
1175 edf: fit.edf,
1176 noise_variance: fit.sigma2,
1177 n_obs: n,
1178 isometry_defect: defect,
1179 isometry_defect_se,
1180 topology_preserved,
1181 min_directional_derivative,
1182 residual_rms: fit.residual_rms,
1183 basis,
1184 })
1185}
1186
1187pub fn fit_layer_transport(
1189 layer_from: usize,
1190 layer_to: usize,
1191 coords_from: ArrayView1<'_, f64>,
1192 coords_to: ArrayView1<'_, f64>,
1193 topology_from: ChartTopology,
1194 topology_to: ChartTopology,
1195) -> Result<LayerTransportReport, String> {
1196 Ok(
1197 fit_transport_map(coords_from, coords_to, topology_from, topology_to)?
1198 .report(layer_from, layer_to),
1199 )
1200}
1201
1202#[derive(Debug, Clone)]
1204pub struct CompositionDefectReport {
1205 pub n_grid: usize,
1206 pub gauge_rotation: f64,
1208 pub gauge_reflected: bool,
1210 pub mean_abs_defect: f64,
1211 pub rms_defect: f64,
1212 pub max_abs_defect: f64,
1213 pub max_studentized_defect: f64,
1215 pub max_studentized_p_value: f64,
1218 pub defect_edf: f64,
1220 pub statistic: f64,
1222 pub ref_df: f64,
1223 pub p_value: f64,
1225}
1226
1227fn monomial_from_equispaced(values: &[f64]) -> Vec<f64> {
1233 let n = values.len();
1234 if n == 0 {
1235 return Vec::new();
1236 }
1237 let mut diffs: Vec<f64> = values.to_vec();
1239 let mut fwd = vec![0.0_f64; n];
1240 fwd[0] = diffs[0];
1241 for k in 1..n {
1242 for i in 0..(n - k) {
1243 diffs[i] = diffs[i + 1] - diffs[i];
1244 }
1245 fwd[k] = diffs[0];
1246 }
1247 let mut coeffs = vec![0.0_f64; n];
1250 let mut poly = vec![0.0_f64; n];
1253 poly[0] = 1.0;
1254 let mut poly_len = 1usize;
1255 let mut factorial = 1.0_f64;
1256 for k in 0..n {
1257 if k > 0 {
1258 factorial *= k as f64;
1259 }
1260 let scale = fwd[k] / factorial;
1261 for (i, &p) in poly.iter().take(poly_len).enumerate() {
1262 coeffs[i] += scale * p;
1263 }
1264 if k + 1 < n {
1266 let mut next = vec![0.0_f64; poly_len + 1];
1267 for i in 0..poly_len {
1268 next[i + 1] += poly[i]; next[i] -= (k as f64) * poly[i]; }
1271 for i in 0..=poly_len {
1272 poly[i] = next[i];
1273 }
1274 poly_len += 1;
1275 }
1276 }
1277 coeffs
1278}
1279
1280fn eval_monomial(coeffs: &[f64], u: f64) -> f64 {
1282 coeffs.iter().rev().fold(0.0_f64, |acc, &c| acc * u + c)
1283}
1284
1285fn monomial_critical_points(coeffs: &[f64]) -> Vec<f64> {
1294 let n = coeffs.len();
1296 if n <= 1 {
1297 return Vec::new();
1298 }
1299 let deriv: Vec<f64> = (1..n).map(|k| k as f64 * coeffs[k]).collect();
1300 match deriv.len() {
1302 0 => Vec::new(),
1303 1 => Vec::new(), 2 => {
1305 let (b, a) = (deriv[0], deriv[1]);
1307 if a.abs() <= f64::MIN_POSITIVE {
1308 Vec::new()
1309 } else {
1310 vec![-b / a]
1311 }
1312 }
1313 3 => {
1314 let (c, b, a) = (deriv[0], deriv[1], deriv[2]);
1316 if a.abs() <= f64::MIN_POSITIVE {
1317 if b.abs() <= f64::MIN_POSITIVE {
1318 Vec::new()
1319 } else {
1320 vec![-c / b]
1321 }
1322 } else {
1323 let disc = b * b - 4.0 * a * c;
1324 if disc < 0.0 {
1325 Vec::new()
1326 } else {
1327 let s = disc.sqrt();
1328 vec![(-b + s) / (2.0 * a), (-b - s) / (2.0 * a)]
1329 }
1330 }
1331 }
1332 _ => {
1333 let lo = 0.0;
1336 let hi = (coeffs.len() - 1) as f64;
1337 let steps = 256;
1338 let mut roots = Vec::new();
1339 let f = |u: f64| eval_monomial(&deriv, u);
1340 let mut prev_u = lo;
1341 let mut prev_v = f(lo);
1342 for i in 1..=steps {
1343 let u = lo + (hi - lo) * i as f64 / steps as f64;
1344 let v = f(u);
1345 if prev_v == 0.0 {
1346 roots.push(prev_u);
1347 } else if prev_v * v < 0.0 {
1348 let (mut a, mut b) = (prev_u, u);
1349 for _ in 0..60 {
1350 let m = 0.5 * (a + b);
1351 if f(a) * f(m) <= 0.0 {
1352 b = m;
1353 } else {
1354 a = m;
1355 }
1356 }
1357 roots.push(0.5 * (a + b));
1358 }
1359 prev_u = u;
1360 prev_v = v;
1361 }
1362 roots
1363 }
1364 }
1365}
1366
1367fn domain_grid(topology: ChartTopology, n: usize) -> Array1<f64> {
1369 match topology {
1370 ChartTopology::Circle => Array1::from_iter((0..n).map(|i| TAU * i as f64 / n as f64)),
1371 ChartTopology::Interval { lo, hi } => {
1372 Array1::from_iter((0..n).map(|i| lo + (hi - lo) * i as f64 / (n - 1).max(1) as f64))
1373 }
1374 }
1375}
1376
1377pub fn composition_defect(
1392 h_ab: &FittedTransport,
1393 h_bc: &FittedTransport,
1394 h_ac: &FittedTransport,
1395 n_grid: usize,
1396) -> Result<CompositionDefectReport, String> {
1397 if h_ab.topology_from != h_ac.topology_from
1398 || h_ab.topology_to != h_bc.topology_from
1399 || h_bc.topology_to != h_ac.topology_to
1400 {
1401 return Err("composition defect requires chart-compatible transports: \
1402 h_ab: A→B, h_bc: B→C, h_ac: A→C"
1403 .to_string());
1404 }
1405 if n_grid < MIN_TRANSPORT_OBS {
1406 return Err(format!(
1407 "composition defect grid must have at least {MIN_TRANSPORT_OBS} points, got {n_grid}"
1408 ));
1409 }
1410
1411 let grid = domain_grid(h_ab.topology_from, n_grid);
1412 let (direct, var_direct) = h_ac.eval_with_variance(grid.view())?;
1413 let (mid, var_mid) = h_ab.eval_with_variance(grid.view())?;
1414 let (composed, var_bc) = h_bc.eval_with_variance(mid.view())?;
1415 let mid_slope = h_bc.derivative(mid.view())?;
1416 let mut variance = Array1::<f64>::zeros(n_grid);
1417 for i in 0..n_grid {
1418 variance[i] = var_direct[i] + var_bc[i] + mid_slope[i] * mid_slope[i] * var_mid[i];
1419 }
1420
1421 let circle_target = matches!(h_ac.topology_to, ChartTopology::Circle);
1423 let mut gauge_reflected = false;
1424 let mut gauge_rotation = 0.0_f64;
1425 let mut defect = Array1::<f64>::zeros(n_grid);
1426 let mut best_sse = f64::INFINITY;
1427 for reflected in [false, true] {
1428 let composed_oriented: Array1<f64> = match (h_ac.topology_to, reflected) {
1429 (_, false) => composed.clone(),
1430 (ChartTopology::Circle, true) => composed.mapv(|v| wrap_tau(-v)),
1431 (ChartTopology::Interval { lo, hi }, true) => composed.mapv(|v| lo + hi - v),
1432 };
1433 let (rotation, candidate): (f64, Array1<f64>) = if circle_target {
1434 let raw: Vec<f64> = (0..n_grid)
1435 .map(|i| wrap_pi(direct[i] - composed_oriented[i]))
1436 .collect();
1437 let rot = circular_mean(&raw);
1438 (
1439 rot,
1440 Array1::from_iter(raw.iter().map(|&d| wrap_pi(d - rot))),
1441 )
1442 } else {
1443 (
1444 0.0,
1445 Array1::from_iter((0..n_grid).map(|i| direct[i] - composed_oriented[i])),
1446 )
1447 };
1448 let sse = candidate.iter().map(|&d| d * d).sum::<f64>();
1449 if sse < best_sse {
1450 best_sse = sse;
1451 gauge_reflected = reflected;
1452 gauge_rotation = rotation;
1453 defect = candidate;
1454 }
1455 }
1456
1457 let coord_scale = match h_ac.topology_to {
1469 ChartTopology::Circle => TAU,
1470 ChartTopology::Interval { lo, hi } => (hi - lo).abs(),
1471 };
1472 let repr_var_floor = (coord_scale * COMPOSITION_DEFECT_REL_VAR_FLOOR).powi(2);
1473 let max_var = variance.iter().copied().fold(0.0_f64, f64::max);
1474 let var_floor = (max_var * 1e-10).max(repr_var_floor).max(f64::MIN_POSITIVE);
1475 let mut max_abs = 0.0_f64;
1476 let mut sum_abs = 0.0_f64;
1477 let mut sum_sq = 0.0_f64;
1478 let mut max_z = 0.0_f64;
1479 for i in 0..n_grid {
1480 let d = defect[i];
1481 let a = d.abs();
1482 max_abs = max_abs.max(a);
1483 sum_abs += a;
1484 sum_sq += d * d;
1485 let z = a / variance[i].max(var_floor).sqrt();
1486 max_z = max_z.max(z);
1487 }
1488 let mean_abs_defect = sum_abs / n_grid as f64;
1489 let rms_defect = (sum_sq / n_grid as f64).sqrt();
1490
1491 let basis = DomainBasis::build(h_ab.topology_from, grid.view())?;
1493 let design = basis.value_rows(grid.view())?;
1494 let penalty = basis.penalty()?;
1495 let weights = variance.mapv(|v| 1.0 / v.max(var_floor));
1496 let fit = fit_penalized_1d(
1497 &design,
1498 &penalty,
1499 defect.view(),
1500 Some(weights.view()),
1501 basis.penalty_rank(),
1502 true,
1503 )?;
1504 let m = basis.num_basis();
1505 let test = wood_smooth_test(SmoothTestInput {
1506 beta: fit.beta.view(),
1507 covariance: &fit.covariance,
1508 influence_matrix: Some(&fit.influence),
1509 whitening_gram: None,
1510 coeff_range: 0..m,
1511 edf: fit.edf,
1512 nullspace_dim: 0,
1513 residual_df: (n_grid as f64 - fit.edf).max(1.0),
1514 scale: SmoothTestScale::Known,
1515 })
1516 .ok_or_else(|| "composition defect smooth test degenerated".to_string())?;
1517
1518 let normal =
1521 Normal::new(0.0, 1.0).map_err(|e| format!("standard normal construction failed: {e}"))?;
1522 let pointwise: f64 = (2.0 * (1.0 - normal.cdf(max_z))).clamp(0.0, 1.0);
1523 let max_studentized_p_value = (n_grid as f64 * pointwise).min(1.0);
1524
1525 Ok(CompositionDefectReport {
1526 n_grid,
1527 gauge_rotation,
1528 gauge_reflected,
1529 mean_abs_defect,
1530 rms_defect,
1531 max_abs_defect: max_abs,
1532 max_studentized_defect: max_z,
1533 max_studentized_p_value,
1534 defect_edf: fit.edf,
1535 statistic: test.statistic,
1536 ref_df: test.ref_df,
1537 p_value: test.p_value,
1538 })
1539}
1540
1541#[derive(Debug, Clone)]
1544pub struct TransportLadderReport {
1545 pub adjacent: Vec<LayerTransportReport>,
1547 pub two_hop: Vec<LayerTransportReport>,
1550 pub circle_transports: Vec<crate::inference::transport_class::CircleTransportReport>,
1555}
1556
1557pub fn transport_ladder(
1563 layers: &[usize],
1564 coords: &[Array1<f64>],
1565 topologies: &[ChartTopology],
1566) -> Result<TransportLadderReport, String> {
1567 let depth = layers.len();
1568 if coords.len() != depth || topologies.len() != depth {
1569 return Err(format!(
1570 "transport ladder inputs disagree: {depth} layers, {} coordinate vectors, {} topologies",
1571 coords.len(),
1572 topologies.len()
1573 ));
1574 }
1575 if depth < 2 {
1576 return Err("transport ladder needs at least two layers".to_string());
1577 }
1578
1579 let mut adjacent_fits: Vec<FittedTransport> = Vec::with_capacity(depth - 1);
1580 let mut adjacent: Vec<LayerTransportReport> = Vec::with_capacity(depth - 1);
1581 for k in 0..depth - 1 {
1582 let fit = fit_transport_map(
1583 coords[k].view(),
1584 coords[k + 1].view(),
1585 topologies[k],
1586 topologies[k + 1],
1587 )
1588 .map_err(|e| {
1589 format!(
1590 "adjacent transport {}→{} failed: {e}",
1591 layers[k],
1592 layers[k + 1]
1593 )
1594 })?;
1595 adjacent.push(fit.report(layers[k], layers[k + 1]));
1596 adjacent_fits.push(fit);
1597 }
1598
1599 let mut two_hop: Vec<LayerTransportReport> = Vec::with_capacity(depth.saturating_sub(2));
1600 for k in 0..depth.saturating_sub(2) {
1601 let direct = fit_transport_map(
1602 coords[k].view(),
1603 coords[k + 2].view(),
1604 topologies[k],
1605 topologies[k + 2],
1606 )
1607 .map_err(|e| {
1608 format!(
1609 "two-hop transport {}→{} failed: {e}",
1610 layers[k],
1611 layers[k + 2]
1612 )
1613 })?;
1614 let composition = composition_defect(
1615 &adjacent_fits[k],
1616 &adjacent_fits[k + 1],
1617 &direct,
1618 DEFAULT_COMPOSITION_GRID,
1619 )
1620 .map_err(|e| {
1621 format!(
1622 "composition test {}→{}→{} failed: {e}",
1623 layers[k],
1624 layers[k + 1],
1625 layers[k + 2]
1626 )
1627 })?;
1628 two_hop.push(
1629 direct
1630 .report(layers[k], layers[k + 2])
1631 .with_composition(&composition),
1632 );
1633 }
1634
1635 let mut circle_transports = Vec::new();
1639 for k in 0..depth - 1 {
1640 if let Some(report) = crate::inference::transport_class::classify_circle_transport_fit(
1641 &adjacent_fits[k],
1642 topologies[k],
1643 topologies[k + 1],
1644 layers[k],
1645 layers[k + 1],
1646 DEFAULT_COMPOSITION_GRID,
1647 ) {
1648 circle_transports.push(report);
1649 }
1650 }
1651
1652 Ok(TransportLadderReport {
1653 adjacent,
1654 two_hop,
1655 circle_transports,
1656 })
1657}
1658
1659#[cfg(test)]
1660mod invert_tests {
1661 use super::*;
1662 use ndarray::Array1;
1663
1664 fn interval(lo: f64, hi: f64) -> ChartTopology {
1665 ChartTopology::Interval { lo, hi }
1666 }
1667
1668 #[test]
1669 fn invert_round_trips_interval_transport() {
1670 let n = 64;
1674 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| i as f64 / (n as f64 - 1.0)));
1675 let to: Array1<f64> = from.mapv(|t| t + 0.25 * (TAU * t).sin() / TAU);
1676 let ft = fit_transport_map(
1677 from.view(),
1678 to.view(),
1679 interval(0.0, 1.0),
1680 interval(0.0, 1.0),
1681 )
1682 .expect("fit");
1683 assert!(
1684 ft.topology_preserved,
1685 "monotone warp should preserve topology"
1686 );
1687
1688 let probe = Array1::from_iter((1..10).map(|i| i as f64 / 10.0));
1689 let fwd = ft.eval(probe.view()).expect("eval");
1691 let back = ft.invert(fwd.view()).expect("invert");
1692 for i in 0..probe.len() {
1693 assert!(
1694 (back[i] - probe[i]).abs() < 1e-6,
1695 "round-trip failed: t={} back={}",
1696 probe[i],
1697 back[i]
1698 );
1699 }
1700 let re_eval = ft.eval(back.view()).expect("eval");
1701 for i in 0..fwd.len() {
1702 assert!((re_eval[i] - fwd[i]).abs() < 1e-9);
1703 }
1704 }
1705
1706 #[test]
1707 fn invert_round_trips_decreasing_interval_transport() {
1708 let n = 64;
1711 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| i as f64 / (n as f64 - 1.0)));
1712 let to: Array1<f64> = from.mapv(|t| 1.0 - 0.5 * t - 0.5 * t * t);
1713 let ft = fit_transport_map(
1714 from.view(),
1715 to.view(),
1716 interval(0.0, 1.0),
1717 interval(0.0, 1.0),
1718 )
1719 .expect("fit");
1720 assert!(ft.topology_preserved);
1721 let probe = Array1::from_iter((1..10).map(|i| i as f64 / 10.0));
1722 let fwd = ft.eval(probe.view()).expect("eval");
1723 let back = ft.invert(fwd.view()).expect("invert");
1724 for i in 0..probe.len() {
1725 assert!(
1726 (back[i] - probe[i]).abs() < 1e-6,
1727 "t={} back={}",
1728 probe[i],
1729 back[i]
1730 );
1731 }
1732 }
1733
1734 #[test]
1735 fn invert_round_trips_circle_transport() {
1736 let n = 128;
1738 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| TAU * i as f64 / n as f64));
1739 let to: Array1<f64> = from.mapv(|t| wrap_tau(t + 0.3 + 0.2 * t.sin()));
1740 let ft = fit_transport_map(
1741 from.view(),
1742 to.view(),
1743 ChartTopology::Circle,
1744 ChartTopology::Circle,
1745 )
1746 .expect("fit");
1747 assert!(ft.topology_preserved, "degree {:?}", ft.degree);
1748
1749 let probe = Array1::from_iter((0..7).map(|i| TAU * (i as f64 + 0.5) / 7.0));
1750 let fwd = ft.eval(probe.view()).expect("eval");
1751 let back = ft.invert(fwd.view()).expect("invert");
1752 for i in 0..probe.len() {
1753 let d = wrap_pi(back[i] - probe[i]).abs();
1755 assert!(d < 1e-5, "probe={} back={} d={}", probe[i], back[i], d);
1756 }
1757 }
1758
1759 #[test]
1760 fn invert_rejects_target_outside_interval_image() {
1761 let n = 32;
1763 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| i as f64 / (n as f64 - 1.0)));
1764 let to: Array1<f64> = from.mapv(|t| 0.5 * t);
1765 let ft = fit_transport_map(
1766 from.view(),
1767 to.view(),
1768 interval(0.0, 1.0),
1769 interval(0.0, 1.0),
1770 )
1771 .expect("fit");
1772 assert!(ft.invert(Array1::from_elem(1, 0.9).view()).is_err());
1773 }
1774
1775 fn fitted_from_target(
1781 from: ArrayView1<'_, f64>,
1782 target: ArrayView1<'_, f64>,
1783 lo: f64,
1784 hi: f64,
1785 ) -> FittedTransport {
1786 let basis = DomainBasis::build(interval(lo, hi), from).expect("basis");
1787 let design = basis.value_rows(from).expect("design");
1788 let m = design.ncols();
1789 let mut xtx = design.t().dot(&design);
1791 let xty = design.t().dot(&target);
1792 let diag = (0..m).map(|i| xtx[[i, i]].abs()).fold(1.0_f64, f64::max);
1793 for i in 0..m {
1794 xtx[[i, i]] += 1e-10 * diag;
1795 }
1796 let (evals, evecs) = xtx.eigh(Side::Lower).expect("eigh");
1797 let rotated = evecs.t().dot(&xty);
1798 let mut beta = Array1::<f64>::zeros(m);
1799 for i in 0..m {
1800 let d = evals[i].max(f64::MIN_POSITIVE);
1801 let c = rotated[i] / d;
1802 for j in 0..m {
1803 beta[j] += evecs[[j, i]] * c;
1804 }
1805 }
1806 FittedTransport {
1807 topology_from: interval(lo, hi),
1808 topology_to: interval(lo, hi),
1809 degree: None,
1810 degree_concentration: None,
1811 rotation_offset: 0.0,
1812 beta,
1813 covariance: Array2::<f64>::zeros((m, m)),
1814 smoothing_lambda: 0.0,
1815 edf: 0.0,
1816 noise_variance: 1.0,
1817 n_obs: from.len(),
1818 isometry_defect: 0.0,
1819 isometry_defect_se: 0.0,
1820 topology_preserved: true,
1821 min_directional_derivative: 1.0,
1822 residual_rms: 0.0,
1823 basis,
1824 }
1825 }
1826
1827 #[test]
1833 fn invert_rejects_between_grid_fold() {
1834 let n = 256;
1835 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| i as f64 / (n as f64 - 1.0)));
1836 let eps = 0.4 / 511.0;
1837 let target: Array1<f64> = from.mapv(|t| (t - 0.5).powi(3) / 3.0 - eps * eps * t);
1838 let mut ft = fitted_from_target(from.view(), target.view(), 0.0, 1.0);
1839
1840 let grid = domain_grid(interval(0.0, 1.0), FOLD_CHECK_GRID);
1843 let grid_d = ft.derivative(grid.view()).expect("grid deriv");
1844 let mean = grid_d.iter().sum::<f64>() / grid_d.len() as f64;
1845 let orientation = if mean < 0.0 { -1.0 } else { 1.0 };
1846 let min_grid = grid_d
1847 .iter()
1848 .map(|&v| orientation * v)
1849 .fold(f64::INFINITY, f64::min);
1850 let dense = Array1::from_iter((0..5120).map(|i| i as f64 / 5119.0));
1852 let dense_d = ft.derivative(dense.view()).expect("dense deriv");
1853 let min_dense = dense_d
1854 .iter()
1855 .map(|&v| orientation * v)
1856 .fold(f64::INFINITY, f64::min);
1857 ft.topology_preserved = min_grid > 0.0;
1858 ft.min_directional_derivative = min_grid;
1859 assert!(
1860 min_grid > 0.0 && min_dense < 0.0,
1861 "fixture must hide a between-grid fold: min on 512-grid={min_grid}, \
1862 min on dense grid={min_dense}"
1863 );
1864
1865 let res = ft.invert(Array1::from_elem(1, 0.0).view());
1868 assert!(
1869 res.is_err(),
1870 "between-grid fold must be rejected by the span-exact certificate \
1871 (topology_preserved={}, min_grid={min_grid}, min_dense={min_dense})",
1872 ft.topology_preserved
1873 );
1874 }
1875
1876 #[test]
1877 fn invert_rejects_non_finite_targets() {
1878 let n = 64;
1879 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| i as f64 / (n as f64 - 1.0)));
1880 let to: Array1<f64> = from.mapv(|t| 0.5 * t);
1881 let ft = fit_transport_map(
1882 from.view(),
1883 to.view(),
1884 interval(0.0, 1.0),
1885 interval(0.0, 1.0),
1886 )
1887 .expect("fit");
1888 for bad in [f64::NAN, f64::INFINITY, f64::NEG_INFINITY] {
1889 assert!(
1890 ft.invert(Array1::from_elem(1, bad).view()).is_err(),
1891 "non-finite target {bad} must be rejected"
1892 );
1893 }
1894 }
1895
1896 #[test]
1897 fn invert_image_tolerance_is_scale_aware() {
1898 let n = 64;
1902 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| i as f64 / (n as f64 - 1.0)));
1903 let scale = 1.0e-8;
1904 let to: Array1<f64> = from.mapv(|t| scale * t);
1905 let ft = fit_transport_map(
1906 from.view(),
1907 to.view(),
1908 interval(0.0, 1.0),
1909 interval(0.0, 1.0),
1910 )
1911 .expect("fit");
1912 let outside = 1.05e-8;
1913 assert!(
1914 ft.invert(Array1::from_elem(1, outside).view()).is_err(),
1915 "target {outside} is 5% outside the [0, {scale}] image and must be rejected"
1916 );
1917 let inside = 0.5e-8;
1919 let t = ft
1920 .invert(Array1::from_elem(1, inside).view())
1921 .expect("invert inside");
1922 let re = ft.eval(t.view()).expect("eval");
1923 assert!((re[0] - inside).abs() < 1e-3 * scale);
1924 }
1925
1926 #[test]
1927 fn invert_round_trips_degree_minus_one_circle() {
1928 let n = 128;
1931 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| TAU * i as f64 / n as f64));
1932 let to: Array1<f64> = from.mapv(|t| wrap_tau(-t + 0.4 + 0.15 * t.sin()));
1933 let ft = fit_transport_map(
1934 from.view(),
1935 to.view(),
1936 ChartTopology::Circle,
1937 ChartTopology::Circle,
1938 )
1939 .expect("fit");
1940 assert_eq!(ft.degree, Some(-1), "expected a degree −1 cover");
1941 assert!(ft.topology_preserved, "degree {:?}", ft.degree);
1942 let probe = Array1::from_iter((0..7).map(|i| TAU * (i as f64 + 0.5) / 7.0));
1943 let fwd = ft.eval(probe.view()).expect("eval");
1944 let back = ft.invert(fwd.view()).expect("invert");
1945 for i in 0..probe.len() {
1946 let d = wrap_pi(back[i] - probe[i]).abs();
1947 assert!(d < 1e-5, "probe={} back={} d={}", probe[i], back[i], d);
1948 }
1949 }
1950
1951 #[test]
1952 fn invert_round_trips_circle_seam_and_interval_endpoints() {
1953 let n = 128;
1955 let from: Array1<f64> = Array1::from_iter((0..n).map(|i| TAU * i as f64 / n as f64));
1956 let to: Array1<f64> = from.mapv(|t| wrap_tau(t + 0.3 + 0.2 * t.sin()));
1957 let ft = fit_transport_map(
1958 from.view(),
1959 to.view(),
1960 ChartTopology::Circle,
1961 ChartTopology::Circle,
1962 )
1963 .expect("fit");
1964 assert!(ft.topology_preserved);
1965 for seam in [1e-9, TAU - 1e-9, 0.0] {
1966 let t = ft
1967 .invert(Array1::from_elem(1, seam).view())
1968 .expect("invert seam");
1969 let re = ft.eval(t.view()).expect("eval");
1970 let d = wrap_pi(re[0] - wrap_tau(seam)).abs();
1971 assert!(d < 1e-6, "seam={seam} re={} d={d}", re[0]);
1972 }
1973
1974 let m = 64;
1976 let ifrom: Array1<f64> = Array1::from_iter((0..m).map(|i| i as f64 / (m as f64 - 1.0)));
1977 let ito: Array1<f64> = ifrom.mapv(|t| t + 0.25 * (TAU * t).sin() / TAU);
1978 let ift = fit_transport_map(
1979 ifrom.view(),
1980 ito.view(),
1981 interval(0.0, 1.0),
1982 interval(0.0, 1.0),
1983 )
1984 .expect("fit");
1985 let raw_lo = ift.raw_at(0.0).expect("raw lo");
1986 let raw_hi = ift.raw_at(1.0).expect("raw hi");
1987 for &edge in &[raw_lo, raw_hi] {
1988 let t = ift
1989 .invert(Array1::from_elem(1, edge).view())
1990 .expect("invert endpoint");
1991 assert!(t[0] >= -1e-9 && t[0] <= 1.0 + 1e-9, "endpoint t={}", t[0]);
1992 let re = ift.eval(t.view()).expect("eval");
1993 assert!((re[0] - edge).abs() < 1e-6, "edge={edge} re={}", re[0]);
1994 }
1995 }
1996
1997 #[test]
1998 fn monomial_reconstruction_is_exact_for_quadratic() {
1999 let coeffs_true = [0.7_f64, -1.3, 2.1]; let values: Vec<f64> = (0..3)
2003 .map(|i| eval_monomial(&coeffs_true, i as f64))
2004 .collect();
2005 let recon = monomial_from_equispaced(&values);
2006 for (a, b) in recon.iter().zip(coeffs_true.iter()) {
2007 assert!((a - b).abs() < 1e-12, "recon {a} vs {b}");
2008 }
2009 let crit = monomial_critical_points(&recon);
2011 assert_eq!(crit.len(), 1);
2012 assert!((crit[0] - 1.3 / 4.2).abs() < 1e-12);
2013 }
2014
2015 #[test]
2023 fn composition_defect_var_floor_is_calibrated() {
2024 for coord_scale in [std::f64::consts::TAU, 1.0_f64, (5.0_f64 - (-3.0_f64)).abs()] {
2025 let floor_std = coord_scale * COMPOSITION_DEFECT_REL_VAR_FLOOR;
2026 let repr_defect = coord_scale * 1e-5; let violation_defect = coord_scale * 1e-2; assert!(
2029 floor_std > 3.0 * repr_defect,
2030 "floor std {floor_std:.3e} does not clear the representation \
2031 defect {repr_defect:.3e} (span {coord_scale})"
2032 );
2033 assert!(
2034 floor_std < 0.1 * violation_defect,
2035 "floor std {floor_std:.3e} is too close to a real violation \
2036 {violation_defect:.3e} (span {coord_scale}) — would cost power"
2037 );
2038 }
2039 }
2040}