Skip to main content

fdars_core/alignment/
nd.rs

1//! Multidimensional (R^d) SRSF transforms and elastic alignment.
2
3use super::srsf::reparameterize_curve;
4use super::{
5    dp_alignment_core, dp_edge_weight, dp_grid_solve, dp_lambda_penalty, dp_path_to_gamma,
6};
7use crate::error::FdarError;
8use crate::helpers::{cumulative_trapz, l2_distance, simpsons_weights};
9use crate::iter_maybe_parallel;
10use crate::matrix::{FdCurveSet, FdMatrix};
11#[cfg(feature = "parallel")]
12use rayon::iter::ParallelIterator;
13
14/// Result of aligning multidimensional (R^d) curves.
15#[derive(Debug, Clone, PartialEq)]
16#[non_exhaustive]
17pub struct AlignmentResultNd {
18    /// Optimal warping function (length m), same for all dimensions.
19    pub gamma: Vec<f64>,
20    /// Aligned curve: d vectors, each length m.
21    pub f_aligned: Vec<Vec<f64>>,
22    /// Elastic distance after alignment.
23    pub distance: f64,
24}
25
26/// Scale derivative vector at one point by 1/√‖f'‖, writing into result_dims.
27#[inline]
28fn srsf_scale_point(derivs: &[FdMatrix], result_dims: &mut [FdMatrix], i: usize, j: usize) {
29    let d = derivs.len();
30    let norm_sq: f64 = derivs.iter().map(|dd| dd[(i, j)].powi(2)).sum();
31    let norm = norm_sq.sqrt();
32    if norm < 1e-15 {
33        for k in 0..d {
34            result_dims[k][(i, j)] = 0.0;
35        }
36    } else {
37        let scale = 1.0 / norm.sqrt();
38        for k in 0..d {
39            result_dims[k][(i, j)] = derivs[k][(i, j)] * scale;
40        }
41    }
42}
43
44/// Compute the SRSF transform for multidimensional (R^d) curves.
45///
46/// For f: \[0,1\] → R^d, the SRSF is q(t) = f'(t) / √‖f'(t)‖ where ‖·‖ is the
47/// Euclidean norm in R^d. For d=1 this reduces to `sign(f') · √|f'|`.
48///
49/// # Arguments
50/// * `data` — Set of n curves in R^d, each with m evaluation points
51/// * `argvals` — Evaluation points (length m)
52///
53/// # Returns
54/// `FdCurveSet` of SRSF values with the same shape as input.
55pub fn srsf_transform_nd(data: &FdCurveSet, argvals: &[f64]) -> FdCurveSet {
56    let d = data.ndim();
57    let n = data.ncurves();
58    let m = data.npoints();
59
60    if d == 0 || n == 0 || m == 0 || argvals.len() != m {
61        return FdCurveSet {
62            dims: (0..d).map(|_| FdMatrix::zeros(n, m)).collect(),
63        };
64    }
65
66    let derivs: Vec<FdMatrix> = data
67        .dims
68        .iter()
69        .map(|dim_mat| {
70            match crate::fdata::deriv(
71                dim_mat,
72                crate::fdata::DerivDomain::OneD { argvals, nderiv: 1 },
73            ) {
74                crate::fdata::DerivResult::OneD(m) => m,
75                _ => unreachable!("1D domain yields a 1D result"),
76            }
77        })
78        .collect();
79
80    let mut result_dims: Vec<FdMatrix> = (0..d).map(|_| FdMatrix::zeros(n, m)).collect();
81    for i in 0..n {
82        for j in 0..m {
83            srsf_scale_point(&derivs, &mut result_dims, i, j);
84        }
85    }
86
87    FdCurveSet { dims: result_dims }
88}
89
90/// Reconstruct an R^d curve from its SRSF.
91///
92/// Given d-dimensional SRSF vectors and initial point f0, reconstructs:
93/// `f_k(t) = f0_k + ∫₀ᵗ q_k(s) · ‖q(s)‖ ds` for each dimension k.
94///
95/// # Arguments
96/// * `q` — SRSF: d vectors, each length m
97/// * `argvals` — Evaluation points (length m)
98/// * `f0` — Initial values in R^d (length d)
99///
100/// # Returns
101/// Reconstructed curve: d vectors, each length m.
102pub fn srsf_inverse_nd(q: &[Vec<f64>], argvals: &[f64], f0: &[f64]) -> Vec<Vec<f64>> {
103    let d = q.len();
104    if d == 0 {
105        return Vec::new();
106    }
107    let m = q[0].len();
108    if m == 0 {
109        return vec![Vec::new(); d];
110    }
111
112    // Compute ||q(t)|| at each time point
113    let norms: Vec<f64> = (0..m)
114        .map(|j| {
115            let norm_sq: f64 = q.iter().map(|qk| qk[j].powi(2)).sum();
116            norm_sq.sqrt()
117        })
118        .collect();
119
120    // For each dimension, integrand = q_k(t) * ||q(t)||
121    let mut result = Vec::with_capacity(d);
122    for k in 0..d {
123        let integrand: Vec<f64> = (0..m).map(|j| q[k][j] * norms[j]).collect();
124        let integral = cumulative_trapz(&integrand, argvals);
125        let curve: Vec<f64> = integral.iter().map(|&v| f0[k] + v).collect();
126        result.push(curve);
127    }
128
129    result
130}
131
132/// Core DP alignment for R^d SRSFs.
133///
134/// Same DP grid and coprime neighborhood as `dp_alignment_core`, but edge weight
135/// is the sum of `dp_edge_weight` over d dimensions.
136fn dp_alignment_core_nd(
137    q1: &[Vec<f64>],
138    q2: &[Vec<f64>],
139    argvals: &[f64],
140    lambda: f64,
141) -> Vec<f64> {
142    let d = q1.len();
143    let m = argvals.len();
144    if m < 2 || d == 0 {
145        return argvals.to_vec();
146    }
147
148    // For d=1, delegate to existing implementation for exact backward compat
149    if d == 1 {
150        return dp_alignment_core(&q1[0], &q2[0], argvals, lambda);
151    }
152
153    // Normalize each dimension's SRSF to unit L2 norm
154    let q1n: Vec<Vec<f64>> = q1
155        .iter()
156        .map(|qk| {
157            let norm = qk.iter().map(|&v| v * v).sum::<f64>().sqrt().max(1e-10);
158            qk.iter().map(|&v| v / norm).collect()
159        })
160        .collect();
161    let q2n: Vec<Vec<f64>> = q2
162        .iter()
163        .map(|qk| {
164            let norm = qk.iter().map(|&v| v * v).sum::<f64>().sqrt().max(1e-10);
165            qk.iter().map(|&v| v / norm).collect()
166        })
167        .collect();
168
169    let path = dp_grid_solve(m, m, |sr, sc, tr, tc| {
170        let w: f64 = (0..d)
171            .map(|k| dp_edge_weight(&q1n[k], &q2n[k], argvals, sc, tc, sr, tr))
172            .sum();
173        w + dp_lambda_penalty(argvals, sc, tc, sr, tr, lambda)
174    });
175
176    dp_path_to_gamma(&path, argvals)
177}
178
179/// Align an R^d curve f2 to f1 using the elastic framework.
180///
181/// Finds the optimal warping γ (shared across all dimensions) such that
182/// f2∘γ is as close as possible to f1 in the elastic metric.
183///
184/// # Arguments
185/// * `f1` — Target curves (d dimensions)
186/// * `f2` — Curves to align (d dimensions)
187/// * `argvals` — Evaluation points (length m)
188/// * `lambda` — Penalty weight (0.0 = no penalty)
189pub fn elastic_align_pair_nd(
190    f1: &FdCurveSet,
191    f2: &FdCurveSet,
192    argvals: &[f64],
193    lambda: f64,
194) -> AlignmentResultNd {
195    let d = f1.ndim();
196    let m = f1.npoints();
197
198    // Compute SRSFs
199    let q1_set = srsf_transform_nd(f1, argvals);
200    let q2_set = srsf_transform_nd(f2, argvals);
201
202    // Extract first curve from each dimension
203    let q1: Vec<Vec<f64>> = q1_set.dims.iter().map(|dm| dm.row(0)).collect();
204    let q2: Vec<Vec<f64>> = q2_set.dims.iter().map(|dm| dm.row(0)).collect();
205
206    // DP alignment using summed cost over dimensions
207    let gamma = dp_alignment_core_nd(&q1, &q2, argvals, lambda);
208
209    // Apply warping to f2 in each dimension
210    let f_aligned: Vec<Vec<f64>> = f2
211        .dims
212        .iter()
213        .map(|dm| {
214            let row = dm.row(0);
215            reparameterize_curve(&row, argvals, &gamma)
216        })
217        .collect();
218
219    // Compute elastic distance: sum of squared L2 distances between aligned SRSFs
220    let f_aligned_set = {
221        let dims: Vec<FdMatrix> = f_aligned
222            .iter()
223            .map(|fa| {
224                FdMatrix::from_slice(fa, 1, m).expect("dimension invariant: data.len() == n * m")
225            })
226            .collect();
227        FdCurveSet { dims }
228    };
229    let q_aligned = srsf_transform_nd(&f_aligned_set, argvals);
230    let weights = simpsons_weights(argvals);
231
232    let mut dist_sq = 0.0;
233    for k in 0..d {
234        let q1k = q1_set.dims[k].row(0);
235        let qak = q_aligned.dims[k].row(0);
236        let d_k = l2_distance(&q1k, &qak, &weights);
237        dist_sq += d_k * d_k;
238    }
239
240    AlignmentResultNd {
241        gamma,
242        f_aligned,
243        distance: dist_sq.sqrt(),
244    }
245}
246
247/// Elastic distance between two R^d curves.
248///
249/// Aligns f2 to f1 and returns the post-alignment SRSF distance.
250pub fn elastic_distance_nd(f1: &FdCurveSet, f2: &FdCurveSet, argvals: &[f64], lambda: f64) -> f64 {
251    elastic_align_pair_nd(f1, f2, argvals, lambda).distance
252}
253
254// ─── Karcher Mean for N-d Curves ─────────────────────────────────────────
255
256/// Result of the Karcher mean computation for multidimensional (R^d) curves.
257#[derive(Debug, Clone, PartialEq)]
258#[non_exhaustive]
259pub struct KarcherMeanResultNd {
260    /// Karcher mean curve: d vectors of length m.
261    pub mean: Vec<Vec<f64>>,
262    /// SRSF of the Karcher mean: d vectors of length m.
263    pub mean_srsf: Vec<Vec<f64>>,
264    /// Final warping functions (n x m).
265    pub gammas: FdMatrix,
266    /// Curves aligned to the mean: d matrices, each n x m.
267    pub aligned_data: Vec<FdMatrix>,
268    /// Number of iterations used.
269    pub n_iter: usize,
270    /// Whether the algorithm converged.
271    pub converged: bool,
272}
273
274/// Result of PCA on aligned multidimensional (R^d) curves.
275#[derive(Debug, Clone, PartialEq)]
276#[non_exhaustive]
277pub struct PcaNdResult {
278    /// PC scores (n x ncomp).
279    pub scores: FdMatrix,
280    /// Principal components per dimension: d matrices, each ncomp x m.
281    pub components: Vec<FdMatrix>,
282    /// Explained variance for each component.
283    pub explained_variance: Vec<f64>,
284    /// Cumulative proportion of variance explained.
285    pub cumulative_variance: Vec<f64>,
286    /// Covariance eigenvalues (same as explained_variance for convenience).
287    pub covariance_eigenvalues: Vec<f64>,
288}
289
290/// Compute SRSF for a single R^d curve (d vectors of length m).
291fn srsf_single_nd(curve: &[Vec<f64>], argvals: &[f64]) -> Vec<Vec<f64>> {
292    let m = argvals.len();
293    let dims: Vec<FdMatrix> = curve
294        .iter()
295        .map(|c| FdMatrix::from_slice(c, 1, m).expect("dimension invariant: data.len() == n * m"))
296        .collect();
297    let cs = FdCurveSet { dims };
298    let q_set = srsf_transform_nd(&cs, argvals);
299    q_set.dims.iter().map(|dm| dm.row(0)).collect()
300}
301
302/// Compute the relative change between two N-d mean SRSFs.
303fn relative_change_nd(old: &[Vec<f64>], new: &[Vec<f64>]) -> f64 {
304    let mut diff_sq = 0.0;
305    let mut old_sq = 0.0;
306    for (qo, qn) in old.iter().zip(new.iter()) {
307        for (&a, &b) in qo.iter().zip(qn.iter()) {
308            diff_sq += (a - b).powi(2);
309            old_sq += a * a;
310        }
311    }
312    diff_sq.sqrt() / old_sq.sqrt().max(1e-10)
313}
314
315/// Select the curve whose SRSF is closest to the pointwise mean SRSF.
316///
317/// Returns the index of the template curve.
318fn select_template_nd(data: &[FdCurveSet], srsfs: &[Vec<Vec<f64>>]) -> usize {
319    let n = data.len();
320    let d = srsfs[0].len();
321    let m = srsfs[0][0].len();
322
323    // Compute pointwise mean SRSF
324    let mut mean_q: Vec<Vec<f64>> = vec![vec![0.0; m]; d];
325    for q in srsfs {
326        for k in 0..d {
327            for j in 0..m {
328                mean_q[k][j] += q[k][j];
329            }
330        }
331    }
332    for k in 0..d {
333        for j in 0..m {
334            mean_q[k][j] /= n as f64;
335        }
336    }
337
338    // Find curve closest to mean
339    let mut min_dist = f64::INFINITY;
340    let mut min_idx = 0;
341    for (i, q) in srsfs.iter().enumerate() {
342        let mut dist_sq = 0.0;
343        for k in 0..d {
344            for j in 0..m {
345                dist_sq += (q[k][j] - mean_q[k][j]).powi(2);
346            }
347        }
348        if dist_sq < min_dist {
349            min_dist = dist_sq;
350            min_idx = i;
351        }
352    }
353    min_idx
354}
355
356/// Compute the Karcher (Frechet) mean for multidimensional (R^d) curves.
357///
358/// Iteratively aligns all N-d curves to the current mean estimate in SRSF space,
359/// computes the pointwise mean of aligned SRSFs per dimension, and reconstructs
360/// the mean curve.
361///
362/// # Arguments
363/// * `data` — Slice of n `FdCurveSet`s, each with d dimensions and m evaluation points
364/// * `argvals` — Evaluation points (length m)
365/// * `max_iter` — Maximum number of iterations
366/// * `tol` — Convergence tolerance (relative SRSF change)
367/// * `lambda` — Roughness penalty weight (0.0 = no penalty)
368///
369/// # Errors
370/// Returns `FdarError::InvalidDimension` if inputs are inconsistent.
371#[must_use = "expensive computation whose result should not be discarded"]
372pub fn karcher_mean_nd(
373    data: &[FdCurveSet],
374    argvals: &[f64],
375    max_iter: usize,
376    tol: f64,
377    lambda: f64,
378) -> Result<KarcherMeanResultNd, FdarError> {
379    let n = data.len();
380    if n < 2 {
381        return Err(FdarError::InvalidDimension {
382            parameter: "data",
383            expected: "at least 2 curves".to_string(),
384            actual: format!("{n}"),
385        });
386    }
387
388    let d = data[0].ndim();
389    let m = data[0].npoints();
390    if d == 0 || m < 2 || argvals.len() != m {
391        return Err(FdarError::InvalidDimension {
392            parameter: "data/argvals",
393            expected: format!("d > 0, m >= 2, argvals.len() == m (m={m})"),
394            actual: format!("d={d}, m={m}, argvals.len()={}", argvals.len()),
395        });
396    }
397
398    // Verify all curves have the same dimensions
399    for (i, cs) in data.iter().enumerate() {
400        if cs.ndim() != d || cs.npoints() != m {
401            return Err(FdarError::InvalidDimension {
402                parameter: "data",
403                expected: format!("all curves d={d}, m={m}"),
404                actual: format!("curve {i}: d={}, m={}", cs.ndim(), cs.npoints()),
405            });
406        }
407    }
408
409    // Extract curves as Vec<Vec<f64>> (d vectors) per observation
410    let curves: Vec<Vec<Vec<f64>>> = (0..n)
411        .map(|i| data[i].dims.iter().map(|dm| dm.row(0)).collect())
412        .collect();
413
414    // Compute SRSFs for all curves
415    let srsfs: Vec<Vec<Vec<f64>>> = curves.iter().map(|c| srsf_single_nd(c, argvals)).collect();
416
417    // Select template (closest to mean SRSF)
418    let template_idx = select_template_nd(data, &srsfs);
419    let mut mu_q = srsfs[template_idx].clone();
420    let mut mu_f = curves[template_idx].clone();
421
422    // Iterative alignment loop
423    let mut converged = false;
424    let mut n_iter = 0;
425    let mut gammas = FdMatrix::zeros(n, m);
426
427    for iter in 0..max_iter {
428        n_iter = iter + 1;
429
430        // Align all curves to current mean (parallel)
431        let align_results: Vec<(Vec<f64>, Vec<Vec<f64>>)> = iter_maybe_parallel!(0..n)
432            .map(|i| {
433                // Build single-curve FdCurveSet for mean and curve i
434                let mean_cs = {
435                    let dims: Vec<FdMatrix> = mu_f
436                        .iter()
437                        .map(|v| {
438                            FdMatrix::from_slice(v, 1, m)
439                                .expect("dimension invariant: data.len() == n * m")
440                        })
441                        .collect();
442                    FdCurveSet { dims }
443                };
444                let curve_cs = {
445                    let dims: Vec<FdMatrix> = curves[i]
446                        .iter()
447                        .map(|v| {
448                            FdMatrix::from_slice(v, 1, m)
449                                .expect("dimension invariant: data.len() == n * m")
450                        })
451                        .collect();
452                    FdCurveSet { dims }
453                };
454
455                let result = elastic_align_pair_nd(&mean_cs, &curve_cs, argvals, lambda);
456                (result.gamma, result.f_aligned)
457            })
458            .collect();
459
460        // Store gammas and compute aligned SRSFs
461        let mut new_mu_q: Vec<Vec<f64>> = vec![vec![0.0; m]; d];
462        for (i, (gamma, f_aligned)) in align_results.iter().enumerate() {
463            for j in 0..m {
464                gammas[(i, j)] = gamma[j];
465            }
466
467            // Compute SRSF of aligned curve
468            let q_aligned = srsf_single_nd(f_aligned, argvals);
469            for k in 0..d {
470                for j in 0..m {
471                    new_mu_q[k][j] += q_aligned[k][j];
472                }
473            }
474        }
475        for k in 0..d {
476            for j in 0..m {
477                new_mu_q[k][j] /= n as f64;
478            }
479        }
480
481        // Check convergence
482        let rel = relative_change_nd(&mu_q, &new_mu_q);
483        mu_q = new_mu_q;
484
485        // Reconstruct mean curve from mean SRSF
486        let f0: Vec<f64> = mu_f.iter().map(|v| v[0]).collect();
487        mu_f = srsf_inverse_nd(&mu_q, argvals, &f0);
488
489        if rel < tol {
490            converged = true;
491            break;
492        }
493    }
494
495    // Post-centering: center the warps via sqrt_mean_inverse
496    let gam_inv = super::sqrt_mean_inverse(&gammas, argvals);
497    for i in 0..n {
498        let gam_i: Vec<f64> = (0..m).map(|j| gammas[(i, j)]).collect();
499        let gam_centered = reparameterize_curve(&gam_i, argvals, &gam_inv);
500        for j in 0..m {
501            gammas[(i, j)] = gam_centered[j];
502        }
503    }
504
505    // Recompute aligned data using final centered warps
506    let mut aligned_data: Vec<FdMatrix> = (0..d).map(|_| FdMatrix::zeros(n, m)).collect();
507    for i in 0..n {
508        let gamma_i: Vec<f64> = (0..m).map(|j| gammas[(i, j)]).collect();
509        for k in 0..d {
510            let f_aligned = reparameterize_curve(&curves[i][k], argvals, &gamma_i);
511            for j in 0..m {
512                aligned_data[k][(i, j)] = f_aligned[j];
513            }
514        }
515    }
516
517    // Recompute mean from final aligned data
518    let mut mean: Vec<Vec<f64>> = vec![vec![0.0; m]; d];
519    for k in 0..d {
520        for j in 0..m {
521            for i in 0..n {
522                mean[k][j] += aligned_data[k][(i, j)];
523            }
524            mean[k][j] /= n as f64;
525        }
526    }
527
528    // Recompute mean SRSF
529    let mean_srsf = srsf_single_nd(&mean, argvals);
530
531    Ok(KarcherMeanResultNd {
532        mean,
533        mean_srsf,
534        gammas,
535        aligned_data,
536        n_iter,
537        converged,
538    })
539}
540
541/// Compute the cross-dimensional covariance matrix of aligned N-d curves.
542///
543/// Stacks all d dimensions of aligned curves into a single (n x d*m) matrix,
544/// centers columns, and returns X^T X / (n-1) as a (d*m x d*m) covariance matrix.
545///
546/// # Errors
547/// Returns `FdarError::InvalidDimension` if d*m exceeds 10000 (to prevent
548/// excessive memory usage) or if input dimensions are inconsistent.
549#[must_use = "expensive computation whose result should not be discarded"]
550pub fn karcher_covariance_nd(
551    result: &KarcherMeanResultNd,
552    argvals: &[f64],
553) -> Result<FdMatrix, FdarError> {
554    let d = result.aligned_data.len();
555    if d == 0 {
556        return Err(FdarError::InvalidDimension {
557            parameter: "aligned_data",
558            expected: "d > 0".to_string(),
559            actual: "0".to_string(),
560        });
561    }
562    let (n, m) = result.aligned_data[0].shape();
563    if argvals.len() != m {
564        return Err(FdarError::InvalidDimension {
565            parameter: "argvals",
566            expected: format!("{m}"),
567            actual: format!("{}", argvals.len()),
568        });
569    }
570
571    let dm = d * m;
572    if dm > 10_000 {
573        return Err(FdarError::InvalidParameter {
574            parameter: "d*m",
575            message: format!(
576                "d*m = {dm} exceeds limit of 10000; covariance matrix would be too large"
577            ),
578        });
579    }
580
581    if n < 2 {
582        return Err(FdarError::InvalidDimension {
583            parameter: "aligned_data",
584            expected: "n >= 2".to_string(),
585            actual: format!("{n}"),
586        });
587    }
588
589    // Build (n x dm) stacked matrix
590    let mut stacked = FdMatrix::zeros(n, dm);
591    for k in 0..d {
592        for i in 0..n {
593            for j in 0..m {
594                stacked[(i, k * m + j)] = result.aligned_data[k][(i, j)];
595            }
596        }
597    }
598
599    // Center columns
600    let mut col_mean = vec![0.0; dm];
601    for j in 0..dm {
602        for i in 0..n {
603            col_mean[j] += stacked[(i, j)];
604        }
605        col_mean[j] /= n as f64;
606    }
607    for i in 0..n {
608        for j in 0..dm {
609            stacked[(i, j)] -= col_mean[j];
610        }
611    }
612
613    // Compute covariance: X^T X / (n-1)
614    let nf = (n - 1) as f64;
615    let mut cov = FdMatrix::zeros(dm, dm);
616    for p in 0..dm {
617        for q in p..dm {
618            let mut s = 0.0;
619            for i in 0..n {
620                s += stacked[(i, p)] * stacked[(i, q)];
621            }
622            s /= nf;
623            cov[(p, q)] = s;
624            cov[(q, p)] = s;
625        }
626    }
627
628    Ok(cov)
629}
630
631/// Perform PCA on aligned multidimensional (R^d) curves.
632///
633/// Stacks aligned data from all d dimensions into an (n x d*m) matrix, centers,
634/// computes the SVD, and extracts principal components and scores.
635///
636/// # Arguments
637/// * `result` — Pre-computed Karcher mean result for N-d curves
638/// * `argvals` — Evaluation points (length m)
639/// * `ncomp` — Number of principal components to extract
640///
641/// # Errors
642/// Returns `FdarError` if inputs are invalid or SVD fails.
643#[must_use = "expensive computation whose result should not be discarded"]
644pub fn pca_nd(
645    result: &KarcherMeanResultNd,
646    argvals: &[f64],
647    ncomp: usize,
648) -> Result<PcaNdResult, FdarError> {
649    let d = result.aligned_data.len();
650    if d == 0 {
651        return Err(FdarError::InvalidDimension {
652            parameter: "aligned_data",
653            expected: "d > 0".to_string(),
654            actual: "0".to_string(),
655        });
656    }
657    let (n, m) = result.aligned_data[0].shape();
658    if n < 2 || m < 2 || ncomp < 1 || argvals.len() != m {
659        return Err(FdarError::InvalidDimension {
660            parameter: "aligned_data/argvals/ncomp",
661            expected: "n >= 2, m >= 2, ncomp >= 1, argvals.len() == m".to_string(),
662            actual: format!(
663                "n={n}, m={m}, ncomp={ncomp}, argvals.len()={}",
664                argvals.len()
665            ),
666        });
667    }
668    let ncomp = ncomp.min(n - 1);
669    let dm = d * m;
670
671    // Build (n x dm) stacked matrix and center
672    let mut stacked = FdMatrix::zeros(n, dm);
673    for k in 0..d {
674        for i in 0..n {
675            for j in 0..m {
676                stacked[(i, k * m + j)] = result.aligned_data[k][(i, j)];
677            }
678        }
679    }
680
681    // Center columns
682    let mut col_mean = vec![0.0; dm];
683    for j in 0..dm {
684        for i in 0..n {
685            col_mean[j] += stacked[(i, j)];
686        }
687        col_mean[j] /= n as f64;
688    }
689    for i in 0..n {
690        for j in 0..dm {
691            stacked[(i, j)] -= col_mean[j];
692        }
693    }
694
695    // Economy SVD: compute Gram matrix G = X X^T / (n-1), size n x n
696    // (much smaller than dm x dm when dm >> n)
697    let nf = (n - 1) as f64;
698    let mut gram = FdMatrix::zeros(n, n);
699    for i in 0..n {
700        for j in i..n {
701            let mut s = 0.0;
702            for p in 0..dm {
703                s += stacked[(i, p)] * stacked[(j, p)];
704            }
705            s /= nf;
706            gram[(i, j)] = s;
707            gram[(j, i)] = s;
708        }
709    }
710
711    // Eigen-decompose Gram matrix via nalgebra SVD (symmetric, so SVD = eigendecomposition)
712    use nalgebra::SVD;
713    let svd = SVD::new(gram.to_dmatrix(), true, true);
714    let u = svd.u.as_ref().ok_or_else(|| FdarError::ComputationFailed {
715        operation: "SVD",
716        detail: "SVD failed to compute U matrix for Gram matrix".to_string(),
717    })?;
718
719    // Eigenvalues of Gram = singular values of Gram = eigenvalues of X X^T / (n-1)
720    // These are also the eigenvalues of the covariance matrix (for the top n components)
721    let eigenvalues: Vec<f64> = svd.singular_values.iter().take(ncomp).copied().collect();
722
723    // Scores: score_ik = u_ik * sqrt(lambda_k * (n-1))
724    // Since Gram = X X^T / (n-1), and SVD(Gram) = U S U^T,
725    // the scores of the data are: X V = U * sqrt(S * (n-1))
726    let mut scores = FdMatrix::zeros(n, ncomp);
727    for k in 0..ncomp {
728        let scale = (eigenvalues[k] * nf).sqrt();
729        for i in 0..n {
730            scores[(i, k)] = u[(i, k)] * scale;
731        }
732    }
733
734    // Loadings: V_k = X^T U_k / sqrt(lambda_k * (n-1))
735    // Reshape back into d matrices of (ncomp x m)
736    let mut components: Vec<FdMatrix> = (0..d).map(|_| FdMatrix::zeros(ncomp, m)).collect();
737    for k in 0..ncomp {
738        let scale = (eigenvalues[k] * nf).sqrt().max(1e-15);
739        let mut loading = vec![0.0; dm];
740        for p in 0..dm {
741            let mut s = 0.0;
742            for i in 0..n {
743                s += stacked[(i, p)] * u[(i, k)];
744            }
745            loading[p] = s / scale;
746        }
747
748        // Distribute into per-dimension matrices
749        for dim in 0..d {
750            for j in 0..m {
751                components[dim][(k, j)] = loading[dim * m + j];
752            }
753        }
754    }
755
756    // Cumulative variance
757    let total_var: f64 = svd.singular_values.iter().sum();
758    let mut cumulative_variance = Vec::with_capacity(ncomp);
759    let mut running = 0.0;
760    for ev in &eigenvalues {
761        running += ev;
762        cumulative_variance.push(if total_var > 0.0 {
763            running / total_var
764        } else {
765            0.0
766        });
767    }
768
769    // Explained variance = eigenvalues
770    let explained_variance = eigenvalues.clone();
771    let covariance_eigenvalues = eigenvalues;
772
773    Ok(PcaNdResult {
774        scores,
775        components,
776        explained_variance,
777        cumulative_variance,
778        covariance_eigenvalues,
779    })
780}
781
782#[cfg(test)]
783mod tests {
784    use super::*;
785    use std::f64::consts::PI;
786
787    /// Build n identical R^2 curves (circle-like) as FdCurveSets.
788    fn make_identical_curves(n: usize, m: usize) -> (Vec<FdCurveSet>, Vec<f64>) {
789        let t: Vec<f64> = (0..m).map(|j| j as f64 / (m - 1) as f64).collect();
790        let dim0: Vec<f64> = t.iter().map(|&ti| (2.0 * PI * ti).sin()).collect();
791        let dim1: Vec<f64> = t.iter().map(|&ti| (2.0 * PI * ti).cos()).collect();
792
793        let data: Vec<FdCurveSet> = (0..n)
794            .map(|_| {
795                let m0 = FdMatrix::from_slice(&dim0, 1, m)
796                    .expect("dimension invariant: data.len() == n * m");
797                let m1 = FdMatrix::from_slice(&dim1, 1, m)
798                    .expect("dimension invariant: data.len() == n * m");
799                FdCurveSet { dims: vec![m0, m1] }
800            })
801            .collect();
802        (data, t)
803    }
804
805    /// Build n shifted R^2 sine curves.
806    fn make_shifted_curves(n: usize, m: usize) -> (Vec<FdCurveSet>, Vec<f64>) {
807        let t: Vec<f64> = (0..m).map(|j| j as f64 / (m - 1) as f64).collect();
808        let data: Vec<FdCurveSet> = (0..n)
809            .map(|i| {
810                let shift = 0.05 * (i as f64 - n as f64 / 2.0);
811                let dim0: Vec<f64> = t
812                    .iter()
813                    .map(|&ti| (2.0 * PI * (ti + shift)).sin())
814                    .collect();
815                let dim1: Vec<f64> = t
816                    .iter()
817                    .map(|&ti| (2.0 * PI * (ti + shift)).cos())
818                    .collect();
819                let m0 = FdMatrix::from_slice(&dim0, 1, m)
820                    .expect("dimension invariant: data.len() == n * m");
821                let m1 = FdMatrix::from_slice(&dim1, 1, m)
822                    .expect("dimension invariant: data.len() == n * m");
823                FdCurveSet { dims: vec![m0, m1] }
824            })
825            .collect();
826        (data, t)
827    }
828
829    #[test]
830    fn karcher_mean_nd_identical_curves() {
831        let (data, t) = make_identical_curves(5, 31);
832        let result = karcher_mean_nd(&data, &t, 10, 1e-4, 0.0).expect("should succeed");
833
834        let d = 2;
835        let m = 31;
836
837        // Mean should be close to the input curves
838        let input_dim0: Vec<f64> = t.iter().map(|&ti| (2.0 * PI * ti).sin()).collect();
839        let input_dim1: Vec<f64> = t.iter().map(|&ti| (2.0 * PI * ti).cos()).collect();
840
841        let max_diff_0: f64 = result.mean[0]
842            .iter()
843            .zip(input_dim0.iter())
844            .map(|(&a, &b)| (a - b).abs())
845            .fold(0.0_f64, f64::max);
846        let max_diff_1: f64 = result.mean[1]
847            .iter()
848            .zip(input_dim1.iter())
849            .map(|(&a, &b)| (a - b).abs())
850            .fold(0.0_f64, f64::max);
851
852        assert!(
853            max_diff_0 < 0.3,
854            "Mean dim 0 should be close to input, max diff = {max_diff_0}"
855        );
856        assert!(
857            max_diff_1 < 0.3,
858            "Mean dim 1 should be close to input, max diff = {max_diff_1}"
859        );
860
861        // Gammas should be near-identity
862        let n = 5;
863        for i in 0..n {
864            for j in 0..m {
865                let diff = (result.gammas[(i, j)] - t[j]).abs();
866                assert!(
867                    diff < 0.15,
868                    "Warp for identical curves should be near identity: gamma[{i},{j}] diff = {diff}"
869                );
870            }
871        }
872
873        // Correct number of dimensions
874        assert_eq!(result.mean.len(), d);
875        assert_eq!(result.mean_srsf.len(), d);
876        assert_eq!(result.aligned_data.len(), d);
877    }
878
879    #[test]
880    fn karcher_mean_nd_output_dimensions() {
881        let (data, t) = make_shifted_curves(8, 25);
882        let result = karcher_mean_nd(&data, &t, 5, 1e-3, 0.0).expect("should succeed");
883
884        let n = 8;
885        let m = 25;
886        let d = 2;
887
888        assert_eq!(result.mean.len(), d);
889        assert_eq!(result.mean_srsf.len(), d);
890        for k in 0..d {
891            assert_eq!(result.mean[k].len(), m);
892            assert_eq!(result.mean_srsf[k].len(), m);
893        }
894        assert_eq!(result.gammas.shape(), (n, m));
895        assert_eq!(result.aligned_data.len(), d);
896        for k in 0..d {
897            assert_eq!(result.aligned_data[k].shape(), (n, m));
898        }
899        assert!(result.n_iter <= 5);
900    }
901
902    #[test]
903    fn karcher_mean_nd_convergence() {
904        let (data, t) = make_shifted_curves(10, 31);
905        let result = karcher_mean_nd(&data, &t, 20, 1e-3, 0.0).expect("should succeed");
906
907        // With well-behaved shifted sine curves, algorithm should converge
908        assert!(
909            result.converged,
910            "Algorithm should converge for shifted sine curves, n_iter={}",
911            result.n_iter
912        );
913    }
914
915    #[test]
916    fn pca_nd_basic_properties() {
917        let (data, t) = make_shifted_curves(10, 31);
918        let km = karcher_mean_nd(&data, &t, 10, 1e-3, 0.0).expect("karcher_mean should succeed");
919        let pca = pca_nd(&km, &t, 3).expect("pca_nd should succeed");
920
921        let n = 10;
922        let ncomp = 3;
923        let m = 31;
924
925        // Scores shape
926        assert_eq!(pca.scores.shape(), (n, ncomp));
927
928        // Components shape: d=2, each ncomp x m
929        assert_eq!(pca.components.len(), 2);
930        for comp in &pca.components {
931            assert_eq!(comp.shape(), (ncomp, m));
932        }
933
934        // Explained variance: non-negative
935        for ev in &pca.explained_variance {
936            assert!(
937                *ev >= -1e-10,
938                "Explained variance should be non-negative: {ev}"
939            );
940        }
941
942        // Explained variance should be approximately decreasing
943        for i in 1..pca.explained_variance.len() {
944            assert!(
945                pca.explained_variance[i] <= pca.explained_variance[i - 1] + 1e-8,
946                "Explained variance should be decreasing: {} > {}",
947                pca.explained_variance[i],
948                pca.explained_variance[i - 1]
949            );
950        }
951
952        // Cumulative variance should be increasing
953        for i in 1..pca.cumulative_variance.len() {
954            assert!(
955                pca.cumulative_variance[i] >= pca.cumulative_variance[i - 1] - 1e-10,
956                "Cumulative variance should be increasing"
957            );
958        }
959    }
960
961    #[test]
962    fn karcher_covariance_nd_symmetric() {
963        let (data, t) = make_shifted_curves(8, 21);
964        let km = karcher_mean_nd(&data, &t, 5, 1e-3, 0.0).expect("karcher_mean should succeed");
965        let cov = karcher_covariance_nd(&km, &t).expect("covariance should succeed");
966
967        let dm = 2 * 21;
968        assert_eq!(cov.shape(), (dm, dm));
969
970        // Verify symmetry
971        for p in 0..dm {
972            for q in p..dm {
973                let diff = (cov[(p, q)] - cov[(q, p)]).abs();
974                assert!(
975                    diff < 1e-12,
976                    "Covariance should be symmetric at ({p},{q}): diff = {diff}"
977                );
978            }
979        }
980    }
981}