Skip to main content

fdars_core/
clustering_advanced.rs

1//! Advanced functional clustering algorithms.
2//!
3//! This module provides four advanced functional clustering paradigms beyond the
4//! basic k-means and fuzzy c-means found in [`clustering`](crate::clustering):
5//!
6//! - **DBSCAN** ([`dbscan_fd`]): Density-based clustering over precomputed functional
7//!   L2 distances. Discovers clusters of arbitrary shape and flags noise curves as
8//!   [`None`] in the assignment vector. No need to specify k.
9//!
10//! - **kCFC** ([`kcfc_cluster`]): K-means-style assignment loop where each cluster's
11//!   centroid is replaced by a per-cluster FPCA model. A curve is assigned to the
12//!   cluster whose FPCA basis produces the smallest reconstruction error.
13//!
14//! - **funFEM** and **Align-and-Cluster**: Discriminative-subspace GMM clustering and
15//!   elastic joint clustering — implemented in plan 33-03.
16//!
17//! All algorithms are strictly additive relative to the existing `clustering` module.
18//! No existing public signature is modified.
19
20use crate::distance::l2_distance_matrix;
21use crate::error::FdarError;
22use crate::helpers::simpsons_weights;
23use crate::matrix::FdMatrix;
24use crate::regression::{fdata_to_pc_1d, FpcaResult};
25use rand::prelude::*;
26
27// ────────────────────────────────────────────────────────────────────────────
28// DBSCAN over functional L2 distances
29// ────────────────────────────────────────────────────────────────────────────
30
31/// Configuration for DBSCAN density clustering over functional data.
32///
33/// DBSCAN discovers clusters of arbitrary shape by expanding dense
34/// regions of curves in functional L2 distance space. Curves that
35/// do not belong to any dense region are flagged as noise ([`None`]
36/// in [`DbscanResult::cluster`]).
37///
38/// ## Distance units
39///
40/// `eps` is in the same units as `l2_distance_matrix` — that is,
41/// functional L2 distance with Simpson's-rule integration weights.
42/// For a constant-1 curve on `argvals` spanning \[0, 1\], the L2 norm
43/// is ≈ 1.0. As a practical starting point, set `eps` to a fraction
44/// (e.g. 0.1–0.5) of the dataset's median pairwise L2 distance:
45///
46/// ```
47/// # use fdars_core::distance::l2_distance_matrix;
48/// # use fdars_core::matrix::FdMatrix;
49/// # let data = FdMatrix::zeros(10, 20);
50/// # let argvals: Vec<f64> = (0..20).map(|i| i as f64 / 19.0).collect();
51/// let dist = l2_distance_matrix(&data, &argvals);
52/// let n = data.nrows();
53/// let mut upper: Vec<f64> = Vec::new();
54/// for i in 0..n {
55///     for j in (i + 1)..n {
56///         upper.push(dist[(i, j)]);
57///     }
58/// }
59/// upper.sort_by(|a, b| a.partial_cmp(b).unwrap());
60/// let median_dist = upper[upper.len() / 2];
61/// let eps = 0.3 * median_dist; // start here, tune as needed
62/// let _ = eps;
63/// ```
64#[derive(Debug, Clone, PartialEq)]
65#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
66#[non_exhaustive]
67pub struct DbscanConfig {
68    /// Neighbourhood radius in functional L2 distance units (default: 0.5).
69    ///
70    /// Must be strictly positive (`eps > 0`). A tiny positive value makes
71    /// every curve a noise point; a very large value merges everything into
72    /// one cluster.
73    pub eps: f64,
74    /// Minimum number of curves (including the point itself) required in the
75    /// `eps`-neighbourhood for a point to be considered a core point (default: 3).
76    ///
77    /// Must be ≥ 1.
78    pub min_points: usize,
79}
80
81impl Default for DbscanConfig {
82    fn default() -> Self {
83        Self {
84            eps: 0.5,
85            min_points: 3,
86        }
87    }
88}
89
90/// Result of DBSCAN density clustering over functional data.
91#[derive(Debug, Clone)]
92#[non_exhaustive]
93pub struct DbscanResult {
94    /// Cluster assignment for each curve.
95    ///
96    /// `None` = noise (not part of any dense cluster).
97    /// `Some(c)` = cluster index `c` (0-based, contiguous).
98    pub cluster: Vec<Option<usize>>,
99    /// Number of discovered clusters (excludes noise points).
100    pub n_clusters: usize,
101    /// Number of noise points (curves assigned `None`).
102    pub n_noise: usize,
103    /// Precomputed n × n pairwise L2 distance matrix used internally.
104    pub distances: FdMatrix,
105}
106
107/// Run DBSCAN over functional L2 distances.
108///
109/// Discovers arbitrarily-shaped clusters in functional data by expanding
110/// dense regions. Curves not reachable from any core point are labelled
111/// noise (`None`).
112///
113/// # Arguments
114///
115/// * `data` — Functional data matrix (n × m, column-major).
116/// * `argvals` — Evaluation grid (length m).
117/// * `config` — Algorithm parameters; see [`DbscanConfig`].
118///
119/// # Errors
120///
121/// Returns [`FdarError::InvalidDimension`] if `n == 0`, `m == 0`, or
122/// `argvals.len() != m`.  Returns [`FdarError::InvalidParameter`] if
123/// `config.eps <= 0` or `config.min_points == 0`.
124///
125/// # Examples
126///
127/// ```
128/// use fdars_core::clustering_advanced::{dbscan_fd, DbscanConfig};
129/// use fdars_core::matrix::FdMatrix;
130/// use std::f64::consts::PI;
131///
132/// // Two tight clusters of 5 sin curves each, well separated
133/// let m = 30;
134/// let n = 10;
135/// let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
136/// let mut col_major = vec![0.0_f64; n * m];
137/// for i in 0..5 {
138///     for (j, &tj) in t.iter().enumerate() {
139///         col_major[i + j * n] = (2.0 * PI * tj).sin();
140///     }
141/// }
142/// for i in 5..10 {
143///     for (j, &tj) in t.iter().enumerate() {
144///         col_major[i + j * n] = (2.0 * PI * tj).sin() + 5.0;
145///     }
146/// }
147/// let data = FdMatrix::from_column_major(col_major, n, m).unwrap();
148///
149/// let mut cfg = DbscanConfig::default();
150/// cfg.eps = 1.0;
151/// cfg.min_points = 2;
152/// let result = dbscan_fd(&data, &t, &cfg).unwrap();
153/// assert_eq!(result.n_clusters, 2);
154/// assert_eq!(result.n_noise, 0);
155/// ```
156#[must_use = "expensive computation whose result should not be discarded"]
157pub fn dbscan_fd(
158    data: &FdMatrix,
159    argvals: &[f64],
160    config: &DbscanConfig,
161) -> Result<DbscanResult, FdarError> {
162    let (n, m) = data.shape();
163
164    // Validation
165    if n == 0 || m == 0 {
166        return Err(FdarError::InvalidDimension {
167            parameter: "data",
168            expected: "at least 1 row and 1 column".to_string(),
169            actual: format!("{n} rows, {m} columns"),
170        });
171    }
172    if argvals.len() != m {
173        return Err(FdarError::InvalidDimension {
174            parameter: "argvals",
175            expected: format!("{m}"),
176            actual: format!("{}", argvals.len()),
177        });
178    }
179    if config.eps <= 0.0 {
180        return Err(FdarError::InvalidParameter {
181            parameter: "eps",
182            message: format!("eps must be > 0, got {}", config.eps),
183        });
184    }
185    if config.min_points == 0 {
186        return Err(FdarError::InvalidParameter {
187            parameter: "min_points",
188            message: "min_points must be >= 1".to_string(),
189        });
190    }
191
192    let dist = l2_distance_matrix(data, argvals);
193
194    // Standard DBSCAN
195    // labels: None = unvisited/noise, Some(c) = cluster c
196    let mut labels: Vec<Option<usize>> = vec![None; n];
197    let mut visited: Vec<bool> = vec![false; n];
198    let mut cluster_id: usize = 0;
199
200    for i in 0..n {
201        if visited[i] {
202            continue;
203        }
204        visited[i] = true;
205
206        // Compute eps-neighbourhood of i (excluding i itself)
207        let neighbors: Vec<usize> = (0..n)
208            .filter(|&j| j != i && dist[(i, j)] <= config.eps)
209            .collect();
210
211        // Core-point rule: i + its neighbors must reach min_points
212        if neighbors.len() + 1 < config.min_points {
213            // Leave as noise for now (may be absorbed as border point later)
214            continue;
215        }
216
217        // i is a core point — start a new cluster
218        labels[i] = Some(cluster_id);
219
220        // BFS expansion
221        let mut queue = neighbors.clone();
222        let mut qi = 0;
223        while qi < queue.len() {
224            let j = queue[qi];
225            qi += 1;
226
227            if !visited[j] {
228                visited[j] = true;
229                let j_neighbors: Vec<usize> = (0..n)
230                    .filter(|&k| k != j && dist[(j, k)] <= config.eps)
231                    .collect();
232                if j_neighbors.len() + 1 >= config.min_points {
233                    // j is also a core point — add its unqueued neighbours
234                    for nb in j_neighbors {
235                        if !queue.contains(&nb) {
236                            queue.push(nb);
237                        }
238                    }
239                }
240            }
241
242            // Absorb j into cluster if not yet assigned
243            if labels[j].is_none() {
244                labels[j] = Some(cluster_id);
245            }
246        }
247
248        cluster_id += 1;
249    }
250
251    let n_clusters = cluster_id;
252    let n_noise = labels.iter().filter(|l| l.is_none()).count();
253
254    Ok(DbscanResult {
255        cluster: labels,
256        n_clusters,
257        n_noise,
258        distances: dist,
259    })
260}
261
262// ────────────────────────────────────────────────────────────────────────────
263// kCFC: per-cluster FPCA reassignment loop
264// ────────────────────────────────────────────────────────────────────────────
265
266/// Configuration for kCFC (k-means-like clustering via Functional Components).
267///
268/// Each cluster is represented by a per-cluster FPCA model rather than a
269/// centroid. A curve is assigned to the cluster whose FPCA basis gives the
270/// smallest L2 reconstruction error.
271///
272/// **Reference:** Chiou & Li (2007), "Functional clustering and identifying
273/// substructures of longitudinal data." R baseline: `fdapace::kCFC`.
274#[derive(Debug, Clone, PartialEq)]
275#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
276#[non_exhaustive]
277pub struct KcfcConfig {
278    /// Number of clusters (default: 2). Must be ≥ 1 and ≤ n.
279    pub k: usize,
280    /// Number of per-cluster FPC components to use for reconstruction (default: 3).
281    ///
282    /// Clamped internally to `min(n_k, m)` where `n_k` is the cluster size.
283    pub ncomp: usize,
284    /// Maximum number of reassignment iterations (default: 50).
285    pub max_iter: usize,
286    /// Random seed for k-means++ initialization (default: 42).
287    pub seed: u64,
288}
289
290impl Default for KcfcConfig {
291    fn default() -> Self {
292        Self {
293            k: 2,
294            ncomp: 3,
295            max_iter: 50,
296            seed: 42,
297        }
298    }
299}
300
301/// Result of kCFC per-cluster FPCA clustering.
302#[derive(Debug, Clone)]
303#[non_exhaustive]
304pub struct KcfcResult {
305    /// Cluster assignment for each curve (0-based, contiguous, length n).
306    pub cluster: Vec<usize>,
307    /// Per-cluster FPCA models.
308    ///
309    /// `fpca_models[k]` is `None` if cluster k was empty at convergence.
310    pub fpca_models: Vec<Option<FpcaResult>>,
311    /// Reconstruction error matrix (n × k).
312    ///
313    /// `reconstruction_errors[(i, k)]` is the L2 squared reconstruction error
314    /// of curve i against cluster k's FPCA model.
315    pub reconstruction_errors: FdMatrix,
316    /// Number of iterations performed.
317    pub iterations: usize,
318    /// Whether the algorithm converged (no label changes in the last iteration).
319    pub converged: bool,
320}
321
322/// Cluster functional data using per-cluster FPCA reconstruction errors (kCFC).
323///
324/// Initialises with k-means++ hard labels, then iterates: fit a per-cluster
325/// FPCA model for each cluster, compute the reconstruction error of every
326/// curve against every cluster's model, and reassign each curve to the cluster
327/// with the smallest error. Repeats until convergence or `config.max_iter`.
328///
329/// # Arguments
330///
331/// * `data` — Functional data matrix (n × m, column-major).
332/// * `argvals` — Evaluation grid (length m).
333/// * `config` — Algorithm parameters; see [`KcfcConfig`].
334///
335/// # Errors
336///
337/// Returns [`FdarError::InvalidDimension`] if `n == 0`, `m == 0`, or
338/// `argvals.len() != m`.  Returns [`FdarError::InvalidParameter`] if
339/// `config.k == 0` or `config.k > n`.
340///
341/// # Examples
342///
343/// ```
344/// use fdars_core::clustering_advanced::{kcfc_cluster, KcfcConfig};
345/// use fdars_core::matrix::FdMatrix;
346/// use std::f64::consts::PI;
347///
348/// let m = 30;
349/// let n = 20;  // 10 per cluster
350/// let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
351/// let mut col_major = vec![0.0_f64; n * m];
352/// for i in 0..10 {
353///     for (j, &tj) in t.iter().enumerate() {
354///         col_major[i + j * n] = (2.0 * PI * tj).sin();
355///     }
356/// }
357/// for i in 10..20 {
358///     for (j, &tj) in t.iter().enumerate() {
359///         col_major[i + j * n] = (2.0 * PI * tj).cos() + 5.0;
360///     }
361/// }
362/// let data = FdMatrix::from_column_major(col_major, n, m).unwrap();
363///
364/// let mut cfg = KcfcConfig::default();
365/// cfg.k = 2;
366/// cfg.ncomp = 2;
367/// let result = kcfc_cluster(&data, &t, &cfg).unwrap();
368/// assert_eq!(result.cluster.len(), n);
369/// ```
370#[must_use = "expensive computation whose result should not be discarded"]
371pub fn kcfc_cluster(
372    data: &FdMatrix,
373    argvals: &[f64],
374    config: &KcfcConfig,
375) -> Result<KcfcResult, FdarError> {
376    let (n, m) = data.shape();
377
378    // Validation
379    if n == 0 || m == 0 {
380        return Err(FdarError::InvalidDimension {
381            parameter: "data",
382            expected: "at least 1 row and 1 column".to_string(),
383            actual: format!("{n} rows, {m} columns"),
384        });
385    }
386    if argvals.len() != m {
387        return Err(FdarError::InvalidDimension {
388            parameter: "argvals",
389            expected: format!("{m}"),
390            actual: format!("{}", argvals.len()),
391        });
392    }
393    if config.k == 0 {
394        return Err(FdarError::InvalidParameter {
395            parameter: "k",
396            message: "k must be >= 1".to_string(),
397        });
398    }
399    if config.k > n {
400        return Err(FdarError::InvalidParameter {
401            parameter: "k",
402            message: format!("k={} exceeds number of curves n={}", config.k, n),
403        });
404    }
405    if config.ncomp == 0 {
406        return Err(FdarError::InvalidParameter {
407            parameter: "ncomp",
408            message: "ncomp must be >= 1".to_string(),
409        });
410    }
411
412    let k = config.k;
413    let weights = simpsons_weights(argvals);
414
415    // ── k-means++ initialisation over raw row curves ──────────────────────
416    let row_major = data.to_row_major(); // n * m, row-major buffer
417    let mut rng = StdRng::seed_from_u64(config.seed);
418
419    // Select first center uniformly
420    let mut center_indices: Vec<usize> = Vec::with_capacity(k);
421    center_indices.push(rng.gen_range(0..n));
422
423    // Maintain min-distance-squared to nearest chosen center
424    let mut min_dist_sq: Vec<f64> = (0..n)
425        .map(|i| {
426            let c0 = center_indices[0];
427            let d = l2_dist_rowmajor(&row_major, i, c0, m, &weights);
428            d * d
429        })
430        .collect();
431
432    while center_indices.len() < k {
433        // Sample proportional to D^2
434        let total: f64 = min_dist_sq.iter().sum();
435        let chosen = if total < 1e-15 {
436            rng.gen_range(0..n)
437        } else {
438            let r = rng.gen::<f64>() * total;
439            let mut cumsum = 0.0;
440            let mut sel = n - 1;
441            for (i, &d) in min_dist_sq.iter().enumerate() {
442                cumsum += d;
443                if cumsum >= r {
444                    sel = i;
445                    break;
446                }
447            }
448            sel
449        };
450        center_indices.push(chosen);
451
452        // Update min_dist_sq with distance to new center
453        for i in 0..n {
454            let d = l2_dist_rowmajor(&row_major, i, chosen, m, &weights);
455            let d2 = d * d;
456            if d2 < min_dist_sq[i] {
457                min_dist_sq[i] = d2;
458            }
459        }
460    }
461
462    // Initial assignment: assign each curve to nearest center
463    let mut cluster: Vec<usize> = (0..n)
464        .map(|i| {
465            center_indices
466                .iter()
467                .enumerate()
468                .min_by(|(_, &c1), (_, &c2)| {
469                    let d1 = l2_dist_rowmajor(&row_major, i, c1, m, &weights);
470                    let d2 = l2_dist_rowmajor(&row_major, i, c2, m, &weights);
471                    d1.partial_cmp(&d2).unwrap_or(std::cmp::Ordering::Equal)
472                })
473                .map(|(ki, _)| ki)
474                .unwrap_or(0)
475        })
476        .collect();
477
478    // ── Reassignment loop ─────────────────────────────────────────────────
479    let mut fpca_models: Vec<Option<FpcaResult>> = vec![None; k];
480    let mut reconstruction_errors = FdMatrix::zeros(n, k);
481    let mut converged = false;
482    let mut iterations = 0;
483
484    for _iter in 0..config.max_iter {
485        iterations += 1;
486
487        // ── Fit per-cluster FPCA models ───────────────────────────────────
488        for ki in 0..k {
489            let member_indices: Vec<usize> = (0..n).filter(|&i| cluster[i] == ki).collect();
490
491            if member_indices.is_empty() {
492                // Keep previous model (or None on first iteration)
493                continue;
494            }
495
496            // Gather member rows into a (n_k x m) FdMatrix
497            let n_k = member_indices.len();
498            let mut col_major_k = vec![0.0_f64; n_k * m];
499            for (row_in_k, &orig_i) in member_indices.iter().enumerate() {
500                for j in 0..m {
501                    col_major_k[row_in_k + j * n_k] = data[(orig_i, j)];
502                }
503            }
504            let data_k = FdMatrix::from_column_major(col_major_k, n_k, m)?;
505
506            // Fit FPCA (ncomp clamped internally to min(n_k, m))
507            match fdata_to_pc_1d(&data_k, config.ncomp, argvals) {
508                Ok(fpca) => {
509                    fpca_models[ki] = Some(fpca);
510                }
511                Err(_) => {
512                    // Degenerate cluster; keep prior model
513                }
514            }
515        }
516
517        // ── Compute reconstruction errors for all curves vs all clusters ──
518        for i in 0..n {
519            let curve_row = data.row(i);
520            let curve_mat = FdMatrix::from_slice(&curve_row, 1, m)?;
521
522            for ki in 0..k {
523                let err = match &fpca_models[ki] {
524                    None => f64::INFINITY,
525                    Some(fpca) => {
526                        let ncomp_eff = fpca.rotation.ncols();
527                        match fpca.project(&curve_mat) {
528                            Ok(scores) => {
529                                match fpca.reconstruct(&scores, ncomp_eff) {
530                                    Ok(recon) => {
531                                        // L2^2 reconstruction error with Simpson weights
532                                        let mut err_sq = 0.0;
533                                        for j in 0..m {
534                                            let diff = curve_row[j] - recon[(0, j)];
535                                            err_sq += diff * diff * weights[j];
536                                        }
537                                        err_sq
538                                    }
539                                    Err(_) => f64::INFINITY,
540                                }
541                            }
542                            Err(_) => f64::INFINITY,
543                        }
544                    }
545                };
546                reconstruction_errors[(i, ki)] = err;
547            }
548        }
549
550        // ── Reassign each curve to the cluster with minimum error ─────────
551        let mut changed = false;
552        for i in 0..n {
553            let best_k = (0..k)
554                .min_by(|&a, &b| {
555                    reconstruction_errors[(i, a)]
556                        .partial_cmp(&reconstruction_errors[(i, b)])
557                        .unwrap_or(std::cmp::Ordering::Equal)
558                })
559                .unwrap_or(0);
560            if cluster[i] != best_k {
561                cluster[i] = best_k;
562                changed = true;
563            }
564        }
565
566        if !changed {
567            converged = true;
568            break;
569        }
570    }
571
572    Ok(KcfcResult {
573        cluster,
574        fpca_models,
575        reconstruction_errors,
576        iterations,
577        converged,
578    })
579}
580
581// ────────────────────────────────────────────────────────────────────────────
582// funFEM: Discriminative-subspace (Fisher-EM) functional clustering
583// ────────────────────────────────────────────────────────────────────────────
584
585/// Configuration for funFEM discriminative-subspace clustering.
586///
587/// funFEM applies the Fisher-EM algorithm to functional data: it first extracts
588/// global FPC scores via [`fdata_to_pc_1d`](crate::regression::fdata_to_pc_1d),
589/// then alternates between finding a discriminative subspace (maximising
590/// between-class vs within-class scatter) and running a GMM E/M step in that
591/// subspace.
592///
593/// ## Divergence from R `funFEM`
594///
595/// This is a simplified Fisher-EM implementation. The discriminative subspace is
596/// found by computing W\_soft^{-1} B\_soft via Cholesky inversion followed by
597/// an SVD (instead of a proper generalized-eigenvalue solver), because no
598/// generalized-eigenvalue crate is used. The multi-pass outer loop re-estimates
599/// the subspace at each iteration. This diverges from the iterative schedule in
600/// the original paper (Bouveyron & Brunet, 2014) but gives practical cluster
601/// recovery on well-separated functional data.
602///
603/// ## Example
604///
605/// ```
606/// use fdars_core::clustering_advanced::{funfem_cluster, FunFemConfig};
607/// use fdars_core::matrix::FdMatrix;
608/// use std::f64::consts::PI;
609///
610/// let m = 30;
611/// let n = 12;
612/// let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
613/// let mut col_major = vec![0.0_f64; n * m];
614/// for i in 0..6 {
615///     for (j, &tj) in t.iter().enumerate() {
616///         col_major[i + j * n] = (2.0 * PI * tj).sin();
617///     }
618/// }
619/// for i in 6..12 {
620///     for (j, &tj) in t.iter().enumerate() {
621///         col_major[i + j * n] = (2.0 * PI * tj).sin() + 5.0;
622///     }
623/// }
624/// let data = FdMatrix::from_column_major(col_major, n, m).unwrap();
625///
626/// let mut cfg = FunFemConfig::default();
627/// cfg.k = 2;
628/// cfg.ncomp = 4;
629/// let result = funfem_cluster(&data, &t, &cfg).unwrap();
630/// assert_eq!(result.cluster.len(), n);
631/// ```
632#[derive(Debug, Clone, PartialEq)]
633#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
634#[non_exhaustive]
635pub struct FunFemConfig {
636    /// Number of clusters (default: 2). Must be ≥ 1 and ≤ n.
637    pub k: usize,
638    /// Number of global FPC components for the score space (default: 10).
639    /// Clamped internally to `min(n, m)`.
640    pub ncomp: usize,
641    /// Discriminative subspace dimension (default: 0 = auto = min(k-1, ncomp_eff)).
642    /// Clamped to `ncomp_eff` if larger.
643    pub p_disc: usize,
644    /// Maximum outer Fisher-EM iterations (default: 50).
645    pub max_iter: usize,
646    /// Log-likelihood convergence tolerance (default: 1e-6).
647    pub tol: f64,
648    /// Random seed for k-means++ initialization (default: 42).
649    pub seed: u64,
650}
651
652impl Default for FunFemConfig {
653    fn default() -> Self {
654        Self {
655            k: 2,
656            ncomp: 10,
657            p_disc: 0,
658            max_iter: 50,
659            tol: 1e-6,
660            seed: 42,
661        }
662    }
663}
664
665/// Result of funFEM discriminative-subspace clustering.
666#[derive(Debug, Clone)]
667#[non_exhaustive]
668pub struct FunFemResult {
669    /// Cluster assignment for each curve (0-based, contiguous, length n).
670    pub cluster: Vec<usize>,
671    /// Soft membership matrix (n × k).
672    pub membership: FdMatrix,
673    /// Discriminative directions (ncomp_eff × p_disc_eff).
674    pub disc_subspace: FdMatrix,
675    /// Final log-likelihood.
676    pub log_likelihood: f64,
677    /// Number of outer iterations performed.
678    pub iterations: usize,
679    /// Whether the algorithm converged.
680    pub converged: bool,
681}
682
683/// Cluster functional data using the Fisher-EM discriminative-subspace GMM (funFEM).
684///
685/// Extracts global FPC scores, then iterates between estimating a discriminative
686/// subspace and running a GMM E/M step within it. Returns cluster assignments
687/// and the discriminative directions used.
688///
689/// # Arguments
690///
691/// * `data` — Functional data matrix (n × m, column-major).
692/// * `argvals` — Evaluation grid (length m).
693/// * `config` — Algorithm parameters; see [`FunFemConfig`].
694///
695/// # Errors
696///
697/// Returns [`FdarError::InvalidDimension`] if `n == 0`, `m == 0`, or
698/// `argvals.len() != m`. Returns [`FdarError::InvalidParameter`] if
699/// `config.k == 0`, `config.k > n`, or `config.ncomp == 0`.
700#[must_use = "expensive computation whose result should not be discarded"]
701pub fn funfem_cluster(
702    data: &FdMatrix,
703    argvals: &[f64],
704    config: &FunFemConfig,
705) -> Result<FunFemResult, FdarError> {
706    let (n, m) = data.shape();
707
708    // Validation
709    if n == 0 || m == 0 {
710        return Err(FdarError::InvalidDimension {
711            parameter: "data",
712            expected: "at least 1 row and 1 column".to_string(),
713            actual: format!("{n} rows, {m} columns"),
714        });
715    }
716    if argvals.len() != m {
717        return Err(FdarError::InvalidDimension {
718            parameter: "argvals",
719            expected: format!("{m}"),
720            actual: format!("{}", argvals.len()),
721        });
722    }
723    if config.k == 0 {
724        return Err(FdarError::InvalidParameter {
725            parameter: "k",
726            message: "k must be >= 1".to_string(),
727        });
728    }
729    if config.k > n {
730        return Err(FdarError::InvalidParameter {
731            parameter: "k",
732            message: format!("k={} exceeds number of curves n={}", config.k, n),
733        });
734    }
735    if config.ncomp == 0 {
736        return Err(FdarError::InvalidParameter {
737            parameter: "ncomp",
738            message: "ncomp must be >= 1".to_string(),
739        });
740    }
741
742    let k = config.k;
743
744    // Step 1: Global FPC scores via fdata_to_pc_1d (clamps ncomp to min(n,m))
745    let fpca = fdata_to_pc_1d(data, config.ncomp, argvals)?;
746    // scores: n x ncomp_eff (row-major semantics via FdMatrix indexing)
747    let scores = &fpca.scores; // FdMatrix n x ncomp_eff
748    let ncomp_eff = scores.ncols();
749
750    // Effective discriminative dimension
751    let p_disc_eff = if config.p_disc == 0 {
752        (k - 1).max(1).min(ncomp_eff)
753    } else {
754        config.p_disc.min(ncomp_eff)
755    };
756
757    // Step 2: K-means++ initialization on the score rows
758    let weights_uniform = vec![1.0; ncomp_eff]; // uniform weight for Euclidean distance in score space
759    let row_major_scores = scores.to_row_major(); // n * ncomp_eff
760    let mut rng = StdRng::seed_from_u64(config.seed);
761
762    let mut center_indices: Vec<usize> = Vec::with_capacity(k);
763    center_indices.push(rng.gen_range(0..n));
764    let mut min_dist_sq: Vec<f64> = (0..n)
765        .map(|i| {
766            let c0 = center_indices[0];
767            let d = l2_dist_rowmajor(&row_major_scores, i, c0, ncomp_eff, &weights_uniform);
768            d * d
769        })
770        .collect();
771    while center_indices.len() < k {
772        let total: f64 = min_dist_sq.iter().sum();
773        let chosen = if total < 1e-15 {
774            rng.gen_range(0..n)
775        } else {
776            let r = rng.gen::<f64>() * total;
777            let mut cumsum = 0.0;
778            let mut sel = n - 1;
779            for (i, &d) in min_dist_sq.iter().enumerate() {
780                cumsum += d;
781                if cumsum >= r {
782                    sel = i;
783                    break;
784                }
785            }
786            sel
787        };
788        center_indices.push(chosen);
789        for i in 0..n {
790            let d = l2_dist_rowmajor(&row_major_scores, i, chosen, ncomp_eff, &weights_uniform);
791            let d2 = d * d;
792            if d2 < min_dist_sq[i] {
793                min_dist_sq[i] = d2;
794            }
795        }
796    }
797
798    // Initial hard labels from nearest center
799    let mut cluster: Vec<usize> = (0..n)
800        .map(|i| {
801            center_indices
802                .iter()
803                .enumerate()
804                .min_by(|(_, &c1), (_, &c2)| {
805                    let d1 =
806                        l2_dist_rowmajor(&row_major_scores, i, c1, ncomp_eff, &weights_uniform);
807                    let d2 =
808                        l2_dist_rowmajor(&row_major_scores, i, c2, ncomp_eff, &weights_uniform);
809                    d1.partial_cmp(&d2).unwrap_or(std::cmp::Ordering::Equal)
810                })
811                .map(|(ki, _)| ki)
812                .unwrap_or(0)
813        })
814        .collect();
815
816    // Initial uniform cluster weights and diagonal covariances in score space
817    let mut pi: Vec<f64> = vec![1.0 / k as f64; k];
818    // Per-cluster means in score space (ncomp_eff-dim)
819    let mut mu_k: Vec<Vec<f64>> = vec![vec![0.0; ncomp_eff]; k];
820    // Per-cluster diagonal variance in score space
821    let mut sigma_k: Vec<Vec<f64>> = vec![vec![1.0; ncomp_eff]; k];
822
823    // Initialize mu_k from initial assignment
824    update_gmm_params_from_hard(
825        &row_major_scores,
826        &cluster,
827        k,
828        ncomp_eff,
829        &mut pi,
830        &mut mu_k,
831        &mut sigma_k,
832    );
833
834    // Discriminative subspace directions (ncomp_eff x p_disc_eff) — identity init
835    let mut disc_dirs: Vec<f64> = {
836        let mut v = vec![0.0_f64; ncomp_eff * p_disc_eff];
837        for d in 0..p_disc_eff {
838            if d < ncomp_eff {
839                v[d + d * ncomp_eff] = 1.0; // column d = e_d
840            }
841        }
842        v
843    };
844
845    let mut prev_ll = f64::NEG_INFINITY;
846    // Initialize responsibilities from hard cluster assignment
847    let mut resp = vec![0.0_f64; n * k];
848    for i in 0..n {
849        let ki = cluster[i].min(k - 1);
850        resp[i * k + ki] = 1.0;
851    }
852    let mut converged = false;
853    let mut iterations = 0;
854
855    for _iter in 0..config.max_iter {
856        iterations += 1;
857
858        // ── Project scores onto current discriminative subspace ───────────
859        // proj_scores[i, d] = sum_j scores[(i,j)] * disc_dirs[j + d*ncomp_eff]
860        // disc_dirs is ncomp_eff x p_disc_eff column-major
861        let mut proj_scores = vec![0.0_f64; n * p_disc_eff]; // n x p_disc_eff row-major
862        for i in 0..n {
863            for d in 0..p_disc_eff {
864                let mut val = 0.0;
865                for j in 0..ncomp_eff {
866                    val += scores[(i, j)] * disc_dirs[j + d * ncomp_eff];
867                }
868                proj_scores[i * p_disc_eff + d] = val;
869            }
870        }
871
872        // ── GMM in projected subspace: per-cluster mean and diag variance ─
873        // Compute cluster means in projected space
874        let mut mu_disc: Vec<Vec<f64>> = vec![vec![0.0; p_disc_eff]; k];
875        let mut n_k_soft: Vec<f64> = vec![0.0; k];
876        for i in 0..n {
877            for ki in 0..k {
878                let r = resp[i * k + ki];
879                n_k_soft[ki] += r;
880                for d in 0..p_disc_eff {
881                    mu_disc[ki][d] += r * proj_scores[i * p_disc_eff + d];
882                }
883            }
884        }
885        for ki in 0..k {
886            if n_k_soft[ki] > 1e-10 {
887                for d in 0..p_disc_eff {
888                    mu_disc[ki][d] /= n_k_soft[ki];
889                }
890            }
891        }
892
893        // Per-cluster diagonal variance in projected space
894        let mut var_disc: Vec<Vec<f64>> = vec![vec![1.0; p_disc_eff]; k];
895        for ki in 0..k {
896            if n_k_soft[ki] > 1e-10 {
897                for d in 0..p_disc_eff {
898                    let mut v = 0.0;
899                    for i in 0..n {
900                        let diff = proj_scores[i * p_disc_eff + d] - mu_disc[ki][d];
901                        v += resp[i * k + ki] * diff * diff;
902                    }
903                    var_disc[ki][d] = (v / n_k_soft[ki]).max(1e-8);
904                }
905            }
906        }
907
908        // ── E-step: compute log-responsibilities in projected space ────────
909        let mut log_resp = vec![0.0_f64; n * k];
910        let mut ll = 0.0;
911        for i in 0..n {
912            let mut log_components = vec![0.0_f64; k];
913            for ki in 0..k {
914                let log_pi = if pi[ki] > 1e-300 { pi[ki].ln() } else { -700.0 };
915                let mut log_lik = log_pi;
916                for d in 0..p_disc_eff {
917                    let diff = proj_scores[i * p_disc_eff + d] - mu_disc[ki][d];
918                    let var = var_disc[ki][d];
919                    log_lik -= 0.5 * (var.ln() + diff * diff / var);
920                }
921                log_lik -= 0.5 * (p_disc_eff as f64) * std::f64::consts::TAU.ln();
922                log_components[ki] = log_lik;
923            }
924            // Log-sum-exp normalization
925            let log_sum = log_sum_exp(&log_components);
926            ll += log_sum;
927            for ki in 0..k {
928                log_resp[i * k + ki] = log_components[ki] - log_sum;
929            }
930        }
931
932        // Exponentiate
933        resp.fill(0.0);
934        for i in 0..n {
935            for ki in 0..k {
936                resp[i * k + ki] = log_resp[i * k + ki].exp().max(1e-300);
937            }
938        }
939
940        // ── Update pi from responsibilities ────────────────────────────────
941        let mut n_k_new: Vec<f64> = vec![0.0; k];
942        for i in 0..n {
943            for ki in 0..k {
944                n_k_new[ki] += resp[i * k + ki];
945            }
946        }
947        let n_total: f64 = n_k_new.iter().sum();
948        for ki in 0..k {
949            pi[ki] = (n_k_new[ki] / n_total).max(1e-300);
950        }
951
952        // Hard assignments from responsibilities (for subspace update)
953        cluster = (0..n)
954            .map(|i| {
955                (0..k)
956                    .max_by(|&a, &b| {
957                        resp[i * k + a]
958                            .partial_cmp(&resp[i * k + b])
959                            .unwrap_or(std::cmp::Ordering::Equal)
960                    })
961                    .unwrap_or(0)
962            })
963            .collect();
964
965        // Update score-space params for next subspace computation
966        update_gmm_params_from_soft(
967            &row_major_scores,
968            &resp,
969            k,
970            ncomp_eff,
971            n,
972            &mut pi,
973            &mut mu_k,
974            &mut sigma_k,
975        );
976
977        // ── Fisher-EM discriminative subspace update ───────────────────────
978        // Compute between-scatter B_soft and within-scatter W_soft in score space
979        let global_mean: Vec<f64> = (0..ncomp_eff)
980            .map(|j| {
981                (0..n)
982                    .map(|i| row_major_scores[i * ncomp_eff + j])
983                    .sum::<f64>()
984                    / n as f64
985            })
986            .collect();
987
988        let mut b_soft = vec![0.0_f64; ncomp_eff * ncomp_eff]; // ncomp_eff x ncomp_eff row-major
989        let mut w_soft = vec![0.0_f64; ncomp_eff * ncomp_eff];
990
991        // Between-scatter: sum_k n_k (mu_k - mu_global)(mu_k - mu_global)^T
992        for ki in 0..k {
993            let nk = n_k_new[ki].max(1.0);
994            for j in 0..ncomp_eff {
995                for l in 0..ncomp_eff {
996                    b_soft[j * ncomp_eff + l] +=
997                        nk * (mu_k[ki][j] - global_mean[j]) * (mu_k[ki][l] - global_mean[l]);
998                }
999            }
1000        }
1001
1002        // Within-scatter: sum_i sum_k r_{ik} (x_i - mu_k)(x_i - mu_k)^T
1003        for i in 0..n {
1004            for ki in 0..k {
1005                let r = resp[i * k + ki];
1006                for j in 0..ncomp_eff {
1007                    let dj = row_major_scores[i * ncomp_eff + j] - mu_k[ki][j];
1008                    for l in 0..ncomp_eff {
1009                        let dl = row_major_scores[i * ncomp_eff + l] - mu_k[ki][l];
1010                        w_soft[j * ncomp_eff + l] += r * dj * dl;
1011                    }
1012                }
1013            }
1014        }
1015
1016        // Add data-scaled regularization floor to W_soft diagonal (Pitfall 3)
1017        let trace_w: f64 = (0..ncomp_eff)
1018            .map(|j| w_soft[j * ncomp_eff + j])
1019            .sum::<f64>();
1020        let reg_floor = (trace_w / ncomp_eff as f64 * 1e-4).max(1e-8);
1021        for j in 0..ncomp_eff {
1022            w_soft[j * ncomp_eff + j] += reg_floor;
1023        }
1024
1025        // Compute W^{-1} B via Cholesky + SVD: new discriminative directions
1026        // W_soft Cholesky, then solve each column of B
1027        match crate::linalg::cholesky_factor(&w_soft, ncomp_eff) {
1028            Ok(l_w) => {
1029                // Form W^{-1} B by solving W * X = B column by column
1030                let mut winv_b = vec![0.0_f64; ncomp_eff * ncomp_eff];
1031                for col in 0..ncomp_eff {
1032                    let b_col: Vec<f64> = (0..ncomp_eff)
1033                        .map(|r| b_soft[r * ncomp_eff + col])
1034                        .collect();
1035                    let x = crate::linalg::cholesky_forward_back(&l_w, &b_col, ncomp_eff);
1036                    for r in 0..ncomp_eff {
1037                        winv_b[r * ncomp_eff + col] = x[r];
1038                    }
1039                }
1040
1041                // SVD of W^{-1}B to get top p_disc_eff directions
1042                use nalgebra::{DMatrix, SVD};
1043                let mat = DMatrix::from_row_slice(ncomp_eff, ncomp_eff, &winv_b);
1044                let svd = SVD::new(mat, true, false);
1045                if let Some(u) = svd.u {
1046                    // u columns = left singular vectors (ncomp_eff x ncomp_eff)
1047                    // Store top p_disc_eff columns as new disc_dirs (ncomp_eff x p_disc_eff column-major)
1048                    let mut new_dirs = vec![0.0_f64; ncomp_eff * p_disc_eff];
1049                    for d in 0..p_disc_eff {
1050                        for r in 0..ncomp_eff {
1051                            new_dirs[r + d * ncomp_eff] = u[(r, d)];
1052                        }
1053                    }
1054                    disc_dirs = new_dirs;
1055                }
1056                // If SVD u unavailable, keep previous disc_dirs
1057            }
1058            Err(_) => {
1059                // Cholesky failed — keep previous discriminative directions (identity fallback)
1060            }
1061        }
1062
1063        // Convergence check
1064        let delta = (ll - prev_ll).abs();
1065        prev_ll = ll;
1066        if _iter > 0 && delta < config.tol {
1067            converged = true;
1068            break;
1069        }
1070    }
1071
1072    // Build membership FdMatrix (n x k)
1073    let mut membership_data = vec![0.0_f64; n * k];
1074    for i in 0..n {
1075        for ki in 0..k {
1076            // column-major: membership[(i, ki)] = membership_data[i + ki*n]
1077            membership_data[i + ki * n] = resp[i * k + ki];
1078        }
1079    }
1080    let membership = FdMatrix::from_column_major(membership_data, n, k)?;
1081
1082    // Build disc_subspace FdMatrix (ncomp_eff x p_disc_eff)
1083    // disc_dirs is already ncomp_eff x p_disc_eff column-major
1084    let disc_subspace = FdMatrix::from_column_major(disc_dirs, ncomp_eff, p_disc_eff)?;
1085
1086    Ok(FunFemResult {
1087        cluster,
1088        membership,
1089        disc_subspace,
1090        log_likelihood: prev_ll,
1091        iterations,
1092        converged,
1093    })
1094}
1095
1096/// Update GMM parameters from hard assignments.
1097fn update_gmm_params_from_hard(
1098    scores_rm: &[f64],
1099    cluster: &[usize],
1100    k: usize,
1101    d: usize,
1102    pi: &mut [f64],
1103    mu_k: &mut [Vec<f64>],
1104    sigma_k: &mut [Vec<f64>],
1105) {
1106    let n = cluster.len();
1107    let mut counts = vec![0usize; k];
1108    for &c in cluster {
1109        if c < k {
1110            counts[c] += 1;
1111        }
1112    }
1113    for ki in 0..k {
1114        pi[ki] = (counts[ki] as f64 / n as f64).max(1e-300);
1115        mu_k[ki] = vec![0.0; d];
1116        // Initialize to 0 so that the accumulation below yields Σ(x-μ)²
1117        // (not 1 + Σ(x-μ)², which was the previous incorrect value).
1118        sigma_k[ki] = vec![0.0; d];
1119        for i in 0..n {
1120            if cluster[i] == ki {
1121                for j in 0..d {
1122                    mu_k[ki][j] += scores_rm[i * d + j];
1123                }
1124            }
1125        }
1126        if counts[ki] > 0 {
1127            for j in 0..d {
1128                mu_k[ki][j] /= counts[ki] as f64;
1129            }
1130        }
1131        for i in 0..n {
1132            if cluster[i] == ki {
1133                for j in 0..d {
1134                    let diff = scores_rm[i * d + j] - mu_k[ki][j];
1135                    sigma_k[ki][j] += diff * diff;
1136                }
1137            }
1138        }
1139        if counts[ki] > 1 {
1140            for j in 0..d {
1141                sigma_k[ki][j] = (sigma_k[ki][j] / counts[ki] as f64).max(1e-8);
1142            }
1143        } else {
1144            for j in 0..d {
1145                sigma_k[ki][j] = 1.0;
1146            }
1147        }
1148    }
1149}
1150
1151/// Update GMM parameters from soft responsibilities.
1152fn update_gmm_params_from_soft(
1153    scores_rm: &[f64],
1154    resp: &[f64],
1155    k: usize,
1156    d: usize,
1157    n: usize,
1158    pi: &mut [f64],
1159    mu_k: &mut [Vec<f64>],
1160    sigma_k: &mut [Vec<f64>],
1161) {
1162    let mut n_k = vec![0.0_f64; k];
1163    for i in 0..n {
1164        for ki in 0..k {
1165            n_k[ki] += resp[i * k + ki];
1166        }
1167    }
1168    let total: f64 = n_k.iter().sum();
1169    for ki in 0..k {
1170        pi[ki] = (n_k[ki] / total).max(1e-300);
1171        mu_k[ki] = vec![0.0; d];
1172        sigma_k[ki] = vec![1.0; d];
1173        if n_k[ki] > 1e-10 {
1174            for i in 0..n {
1175                let r = resp[i * k + ki];
1176                for j in 0..d {
1177                    mu_k[ki][j] += r * scores_rm[i * d + j];
1178                }
1179            }
1180            for j in 0..d {
1181                mu_k[ki][j] /= n_k[ki];
1182            }
1183            let mut var_j = vec![0.0_f64; d];
1184            for i in 0..n {
1185                let r = resp[i * k + ki];
1186                for j in 0..d {
1187                    let diff = scores_rm[i * d + j] - mu_k[ki][j];
1188                    var_j[j] += r * diff * diff;
1189                }
1190            }
1191            for j in 0..d {
1192                sigma_k[ki][j] = (var_j[j] / n_k[ki]).max(1e-8);
1193            }
1194        }
1195    }
1196}
1197
1198/// Log-sum-exp of a slice (numerically stable).
1199fn log_sum_exp(v: &[f64]) -> f64 {
1200    let max_v = v.iter().cloned().fold(f64::NEG_INFINITY, f64::max);
1201    if max_v == f64::NEG_INFINITY {
1202        return f64::NEG_INFINITY;
1203    }
1204    max_v + v.iter().map(|&x| (x - max_v).exp()).sum::<f64>().ln()
1205}
1206
1207/// Compute L2 distance between two rows in a flat row-major buffer.
1208///
1209/// `buf[i * m .. (i+1) * m]` is row `i`.
1210fn l2_dist_rowmajor(buf: &[f64], i: usize, j: usize, m: usize, weights: &[f64]) -> f64 {
1211    let mut sq = 0.0;
1212    for t in 0..m {
1213        let d = buf[i * m + t] - buf[j * m + t];
1214        sq += d * d * weights[t];
1215    }
1216    sq.sqrt()
1217}
1218
1219// ────────────────────────────────────────────────────────────────────────────
1220// Align-and-cluster: elastic k-means via Karcher-mean templates
1221// ────────────────────────────────────────────────────────────────────────────
1222
1223/// Configuration for joint align-and-cluster (elastic k-means).
1224///
1225/// Alternates between updating per-cluster template curves (via Karcher mean
1226/// in the elastic metric) and reassigning each curve to the nearest template
1227/// by elastic distance. This is the only clusterer in this module that is
1228/// shape-invariant — it finds clusters based on amplitude shape, not phase.
1229///
1230/// **Reference:** Inspired by Sangalli et al. (2010) joint clustering and alignment
1231/// and the `fdasrvf` elastic k-means variant (Srivastava et al.).
1232///
1233/// ## Example
1234///
1235/// ```
1236/// use fdars_core::clustering_advanced::{align_cluster_fd, AlignClusterConfig};
1237/// use fdars_core::matrix::FdMatrix;
1238/// use std::f64::consts::PI;
1239///
1240/// let m = 25;
1241/// let n = 10;
1242/// let t: Vec<f64> = (0..m).map(|i| i as f64 / (m - 1) as f64).collect();
1243/// let mut col_major = vec![0.0_f64; n * m];
1244/// for i in 0..5 {
1245///     for (j, &tj) in t.iter().enumerate() {
1246///         col_major[i + j * n] = (2.0 * PI * tj).sin();
1247///     }
1248/// }
1249/// for i in 5..10 {
1250///     for (j, &tj) in t.iter().enumerate() {
1251///         col_major[i + j * n] = (2.0 * PI * tj).sin() + 5.0;
1252///     }
1253/// }
1254/// let data = FdMatrix::from_column_major(col_major, n, m).unwrap();
1255///
1256/// let mut cfg = AlignClusterConfig::default();
1257/// cfg.k = 2;
1258/// cfg.max_iter = 10;
1259/// cfg.karcher_max_iter = 8;
1260/// let result = align_cluster_fd(&data, &t, &cfg).unwrap();
1261/// assert_eq!(result.cluster.len(), n);
1262/// assert_eq!(result.templates.len(), 2);
1263/// ```
1264#[derive(Debug, Clone, PartialEq)]
1265#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
1266#[non_exhaustive]
1267pub struct AlignClusterConfig {
1268    /// Number of clusters (default: 2). Must be ≥ 1 and ≤ n.
1269    pub k: usize,
1270    /// Maximum number of outer iterations (default: 20).
1271    pub max_iter: usize,
1272    /// Random seed for initial template selection (default: 42).
1273    pub seed: u64,
1274    /// If `true`, use `amplitude_distance` (shape-invariant); otherwise use
1275    /// `elastic_distance` (full Fisher-Rao distance including phase). Default: `true`.
1276    pub use_amplitude_only: bool,
1277    /// Penalty weight for `elastic_distance` when `use_amplitude_only = false` (default: 0.0).
1278    pub elastic_lambda: f64,
1279    /// Maximum inner iterations for Karcher mean (default: 15).
1280    pub karcher_max_iter: usize,
1281    /// Convergence tolerance for Karcher mean (default: 1e-4).
1282    pub karcher_tol: f64,
1283}
1284
1285impl Default for AlignClusterConfig {
1286    fn default() -> Self {
1287        Self {
1288            k: 2,
1289            max_iter: 20,
1290            seed: 42,
1291            use_amplitude_only: true,
1292            elastic_lambda: 0.0,
1293            karcher_max_iter: 15,
1294            karcher_tol: 1e-4,
1295        }
1296    }
1297}
1298
1299/// Result of joint align-and-cluster (elastic k-means).
1300#[derive(Debug, Clone)]
1301#[non_exhaustive]
1302pub struct AlignClusterResult {
1303    /// Cluster assignment for each curve (0-based, contiguous, length n).
1304    pub cluster: Vec<usize>,
1305    /// Per-cluster template curves (k entries, each of length m).
1306    pub templates: Vec<Vec<f64>>,
1307    /// Distance matrix (n × k) — `distances[(i, ki)]` is the elastic distance
1308    /// from curve i to template ki.
1309    pub distances: FdMatrix,
1310    /// Number of outer iterations performed.
1311    pub iterations: usize,
1312    /// Whether the algorithm converged (no label changes in last iteration).
1313    pub converged: bool,
1314}
1315
1316/// Cluster functional data using elastic k-means with Karcher-mean templates.
1317///
1318/// Initialises k templates by random distinct-curve selection, then alternates
1319/// between reassigning each curve to the nearest template (by elastic distance)
1320/// and updating each template to the Karcher mean of its cluster members.
1321/// Empty clusters are reinitialized to a random non-member curve.
1322///
1323/// # Arguments
1324///
1325/// * `data` — Functional data matrix (n × m, column-major).
1326/// * `argvals` — Evaluation grid (length m).
1327/// * `config` — Algorithm parameters; see [`AlignClusterConfig`].
1328///
1329/// # Errors
1330///
1331/// Returns [`FdarError::InvalidDimension`] if `n == 0`, `m == 0`, or
1332/// `argvals.len() != m`. Returns [`FdarError::InvalidParameter`] if
1333/// `config.k == 0` or `config.k > n`.
1334#[must_use = "expensive computation whose result should not be discarded"]
1335pub fn align_cluster_fd(
1336    data: &FdMatrix,
1337    argvals: &[f64],
1338    config: &AlignClusterConfig,
1339) -> Result<AlignClusterResult, FdarError> {
1340    use crate::alignment::{amplitude_distance, elastic_distance, karcher_mean};
1341
1342    let (n, m) = data.shape();
1343
1344    // Validation
1345    if n == 0 || m == 0 {
1346        return Err(FdarError::InvalidDimension {
1347            parameter: "data",
1348            expected: "at least 1 row and 1 column".to_string(),
1349            actual: format!("{n} rows, {m} columns"),
1350        });
1351    }
1352    if argvals.len() != m {
1353        return Err(FdarError::InvalidDimension {
1354            parameter: "argvals",
1355            expected: format!("{m}"),
1356            actual: format!("{}", argvals.len()),
1357        });
1358    }
1359    if config.k == 0 {
1360        return Err(FdarError::InvalidParameter {
1361            parameter: "k",
1362            message: "k must be >= 1".to_string(),
1363        });
1364    }
1365    if config.k > n {
1366        return Err(FdarError::InvalidParameter {
1367            parameter: "k",
1368            message: format!("k={} exceeds number of curves n={}", config.k, n),
1369        });
1370    }
1371
1372    let k = config.k;
1373
1374    // Initialize k templates by seeded random distinct-curve selection.
1375    // Use a shuffled index list and pick every n/k curves for spread,
1376    // then shuffle within the stratum using the seeded RNG to ensure
1377    // reproducibility while covering the data range.
1378    let mut rng = StdRng::seed_from_u64(config.seed);
1379    let mut shuffled: Vec<usize> = (0..n).collect();
1380    // Fisher-Yates shuffle
1381    for i in (1..n).rev() {
1382        let j = rng.gen_range(0..=i);
1383        shuffled.swap(i, j);
1384    }
1385    // Pick k evenly-strided indices from the shuffled list for spread
1386    let step = n / k;
1387    let template_indices: Vec<usize> = (0..k).map(|ki| shuffled[(ki * step).min(n - 1)]).collect();
1388
1389    // Initialize templates as the selected curve values
1390    let mut templates: Vec<Vec<f64>> = template_indices.iter().map(|&ci| data.row(ci)).collect();
1391
1392    let mut cluster: Vec<usize> = vec![0; n];
1393    let mut distances = FdMatrix::zeros(n, k);
1394    let mut converged = false;
1395    let mut iterations = 0;
1396
1397    for _iter in 0..config.max_iter {
1398        iterations += 1;
1399
1400        // ── Reassignment: compute distances to each template ──────────────
1401        for i in 0..n {
1402            let curve_i = data.row(i);
1403            for ki in 0..k {
1404                let dist = if config.use_amplitude_only {
1405                    amplitude_distance(&curve_i, &templates[ki], argvals, config.elastic_lambda)
1406                } else {
1407                    elastic_distance(&curve_i, &templates[ki], argvals, config.elastic_lambda)
1408                };
1409                distances[(i, ki)] = dist;
1410            }
1411        }
1412
1413        // Assign each curve to nearest template
1414        let mut changed = false;
1415        for i in 0..n {
1416            let best_k = (0..k)
1417                .min_by(|&a, &b| {
1418                    distances[(i, a)]
1419                        .partial_cmp(&distances[(i, b)])
1420                        .unwrap_or(std::cmp::Ordering::Equal)
1421                })
1422                .unwrap_or(0);
1423            if cluster[i] != best_k {
1424                cluster[i] = best_k;
1425                changed = true;
1426            }
1427        }
1428
1429        // ── Template update: Karcher mean of cluster members ──────────────
1430        // Track whether any template was changed (including by empty-cluster
1431        // reinit) so convergence is not declared on the same iteration that a
1432        // reinit occurred — the new template must be tested in the next pass.
1433        let mut template_changed = false;
1434        for ki in 0..k {
1435            let member_indices: Vec<usize> = (0..n).filter(|&i| cluster[i] == ki).collect();
1436
1437            if member_indices.is_empty() {
1438                // Empty-cluster fallback (Pitfall 6): reinit to a random non-member curve
1439                // Find all non-members
1440                let non_members: Vec<usize> = (0..n).filter(|&i| cluster[i] != ki).collect();
1441                if !non_members.is_empty() {
1442                    let rand_idx = rng.gen_range(0..non_members.len());
1443                    templates[ki] = data.row(non_members[rand_idx]);
1444                    template_changed = true;
1445                }
1446                // (If somehow all n are in one cluster, keep the old template)
1447                continue;
1448            }
1449
1450            // Gather member rows into (n_k x m) FdMatrix
1451            let n_k = member_indices.len();
1452            let mut col_major_k = vec![0.0_f64; n_k * m];
1453            for (row_in_k, &orig_i) in member_indices.iter().enumerate() {
1454                for j in 0..m {
1455                    col_major_k[row_in_k + j * n_k] = data[(orig_i, j)];
1456                }
1457            }
1458            let data_k = FdMatrix::from_column_major(col_major_k, n_k, m)?;
1459
1460            // Karcher mean — returns KarcherMeanResult, template is `.mean` field
1461            let km = karcher_mean(
1462                &data_k,
1463                argvals,
1464                config.karcher_max_iter,
1465                config.karcher_tol,
1466                config.elastic_lambda,
1467            );
1468            templates[ki] = km.mean;
1469        }
1470
1471        // Converge only when no curve moved AND no template was reinit-ed due to
1472        // an empty cluster.  Including template_changed prevents the algorithm
1473        // from declaring convergence on the same iteration as an empty-cluster
1474        // reinit, before the new template has been used in a subsequent
1475        // reassignment pass.
1476        if !changed && !template_changed {
1477            converged = true;
1478            break;
1479        }
1480    }
1481
1482    Ok(AlignClusterResult {
1483        cluster,
1484        templates,
1485        distances,
1486        iterations,
1487        converged,
1488    })
1489}
1490
1491// ────────────────────────────────────────────────────────────────────────────
1492// Tests
1493// ────────────────────────────────────────────────────────────────────────────
1494
1495#[cfg(test)]
1496mod tests {
1497    use super::*;
1498    use crate::test_helpers::{adjusted_rand_index, uniform_grid};
1499    use std::f64::consts::PI;
1500
1501    // ── Synthetic data generators ────────────────────────────────────────
1502
1503    /// Two tight clusters, n_per curves each. Cluster 0: sin wave.
1504    /// Cluster 1: sin wave shifted up by 5.
1505    fn two_tight_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
1506        let t = uniform_grid(m);
1507        let n = 2 * n_per;
1508        let mut col_major = vec![0.0_f64; n * m];
1509        for i in 0..n_per {
1510            for (j, &tj) in t.iter().enumerate() {
1511                col_major[i + j * n] = (2.0 * PI * tj).sin();
1512            }
1513        }
1514        for i in 0..n_per {
1515            for (j, &tj) in t.iter().enumerate() {
1516                col_major[(i + n_per) + j * n] = (2.0 * PI * tj).sin() + 5.0;
1517            }
1518        }
1519        let labels: Vec<usize> = (0..n).map(|i| if i < n_per { 0 } else { 1 }).collect();
1520        (
1521            FdMatrix::from_column_major(col_major, n, m).unwrap(),
1522            t,
1523            labels,
1524        )
1525    }
1526
1527    /// 2 tight clusters (n_per each) + 2 far constant-offset outlier curves.
1528    fn clusters_with_noise(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>) {
1529        let t = uniform_grid(m);
1530        let n = 2 * n_per + 2;
1531        let mut col_major = vec![0.0_f64; n * m];
1532        // Cluster 0: sin wave
1533        for i in 0..n_per {
1534            for (j, &tj) in t.iter().enumerate() {
1535                col_major[i + j * n] = (2.0 * PI * tj).sin();
1536            }
1537        }
1538        // Cluster 1: sin wave + 5
1539        for i in 0..n_per {
1540            for (j, &tj) in t.iter().enumerate() {
1541                col_major[(i + n_per) + j * n] = (2.0 * PI * tj).sin() + 5.0;
1542            }
1543        }
1544        // Outlier 0: constant 100
1545        let o0 = 2 * n_per;
1546        for j in 0..m {
1547            col_major[o0 + j * n] = 100.0;
1548        }
1549        // Outlier 1: constant -100
1550        let o1 = 2 * n_per + 1;
1551        for j in 0..m {
1552            col_major[o1 + j * n] = -100.0;
1553        }
1554        (FdMatrix::from_column_major(col_major, n, m).unwrap(), t)
1555    }
1556
1557    /// Two well-separated clusters for kCFC testing.
1558    fn two_separated_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
1559        let t = uniform_grid(m);
1560        let n = 2 * n_per;
1561        let mut col_major = vec![0.0_f64; n * m];
1562        // Cluster 0: sin waves
1563        for i in 0..n_per {
1564            for (j, &tj) in t.iter().enumerate() {
1565                col_major[i + j * n] = (2.0 * PI * tj).sin() + 0.05 * (i as f64 / n_per as f64);
1566            }
1567        }
1568        // Cluster 1: cos waves shifted up by 8 (very different shape)
1569        for i in 0..n_per {
1570            for (j, &tj) in t.iter().enumerate() {
1571                col_major[(i + n_per) + j * n] =
1572                    (2.0 * PI * tj).cos() + 8.0 + 0.05 * (i as f64 / n_per as f64);
1573            }
1574        }
1575        let labels: Vec<usize> = (0..n).map(|i| if i < n_per { 0 } else { 1 }).collect();
1576        (
1577            FdMatrix::from_column_major(col_major, n, m).unwrap(),
1578            t,
1579            labels,
1580        )
1581    }
1582
1583    // ── DBSCAN tests ─────────────────────────────────────────────────────
1584
1585    #[test]
1586    fn test_dbscan_core_points() {
1587        let m = 30;
1588        let n_per = 5;
1589        let (data, t, _labels) = two_tight_clusters(n_per, m);
1590        let result = dbscan_fd(
1591            &data,
1592            &t,
1593            &DbscanConfig {
1594                eps: 1.0,
1595                min_points: 2,
1596                ..Default::default()
1597            },
1598        )
1599        .unwrap();
1600        assert_eq!(result.n_clusters, 2, "expected 2 clusters");
1601        assert_eq!(result.n_noise, 0, "expected 0 noise points");
1602        assert_eq!(result.cluster.len(), 2 * n_per);
1603    }
1604
1605    #[test]
1606    fn test_dbscan_noise_flagging() {
1607        let m = 30;
1608        let n_per = 5;
1609        let (data, t) = clusters_with_noise(n_per, m);
1610        let result = dbscan_fd(
1611            &data,
1612            &t,
1613            &DbscanConfig {
1614                eps: 1.5,
1615                min_points: 2,
1616                ..Default::default()
1617            },
1618        )
1619        .unwrap();
1620        // The 2 outlier curves should be noise
1621        assert_eq!(
1622            result.n_noise, 2,
1623            "expected exactly 2 noise points, got {}",
1624            result.n_noise
1625        );
1626        assert_eq!(result.n_clusters, 2, "expected 2 clusters");
1627        // Verify outlier indices (last 2) are None
1628        let n = data.nrows();
1629        assert!(result.cluster[n - 2].is_none(), "outlier 0 should be noise");
1630        assert!(result.cluster[n - 1].is_none(), "outlier 1 should be noise");
1631    }
1632
1633    #[test]
1634    fn test_dbscan_zero_eps_returns_err() {
1635        let m = 20;
1636        let (data, t, _) = two_tight_clusters(5, m);
1637        assert!(
1638            dbscan_fd(
1639                &data,
1640                &t,
1641                &DbscanConfig {
1642                    eps: 0.0,
1643                    ..Default::default()
1644                }
1645            )
1646            .is_err(),
1647            "eps=0 should return Err"
1648        );
1649    }
1650
1651    #[test]
1652    fn test_dbscan_negative_eps_returns_err() {
1653        let m = 20;
1654        let (data, t, _) = two_tight_clusters(5, m);
1655        assert!(
1656            dbscan_fd(
1657                &data,
1658                &t,
1659                &DbscanConfig {
1660                    eps: -1.0,
1661                    ..Default::default()
1662                }
1663            )
1664            .is_err(),
1665            "eps=-1 should return Err"
1666        );
1667    }
1668
1669    #[test]
1670    fn test_dbscan_invalid_min_points_zero() {
1671        let m = 20;
1672        let (data, t, _) = two_tight_clusters(5, m);
1673        assert!(
1674            dbscan_fd(
1675                &data,
1676                &t,
1677                &DbscanConfig {
1678                    min_points: 0,
1679                    ..Default::default()
1680                }
1681            )
1682            .is_err(),
1683            "min_points=0 should return Err"
1684        );
1685    }
1686
1687    #[test]
1688    fn test_dbscan_empty_data() {
1689        let data = FdMatrix::zeros(0, 0);
1690        let t: Vec<f64> = vec![];
1691        assert!(
1692            dbscan_fd(&data, &t, &DbscanConfig::default()).is_err(),
1693            "empty data should return Err"
1694        );
1695    }
1696
1697    #[test]
1698    fn test_dbscan_mismatched_argvals() {
1699        let m = 20;
1700        let (data, _t, _) = two_tight_clusters(5, m);
1701        let wrong_t = uniform_grid(m + 1);
1702        assert!(
1703            dbscan_fd(&data, &wrong_t, &DbscanConfig::default()).is_err(),
1704            "mismatched argvals should return Err"
1705        );
1706    }
1707
1708    #[test]
1709    fn test_dbscan_distances_shape() {
1710        let m = 20;
1711        let n_per = 4;
1712        let (data, t, _) = two_tight_clusters(n_per, m);
1713        let result = dbscan_fd(
1714            &data,
1715            &t,
1716            &DbscanConfig {
1717                eps: 1.0,
1718                min_points: 2,
1719                ..Default::default()
1720            },
1721        )
1722        .unwrap();
1723        let n = 2 * n_per;
1724        assert_eq!(
1725            result.distances.shape(),
1726            (n, n),
1727            "distance matrix must be n x n"
1728        );
1729    }
1730
1731    // ── kCFC tests ───────────────────────────────────────────────────────
1732
1733    #[test]
1734    fn test_kcfc_recovery() {
1735        let m = 40;
1736        let n_per = 10;
1737        let (data, t, ground_truth) = two_separated_clusters(n_per, m);
1738        let result = kcfc_cluster(
1739            &data,
1740            &t,
1741            &KcfcConfig {
1742                k: 2,
1743                ncomp: 3,
1744                max_iter: 50,
1745                seed: 42,
1746                ..Default::default()
1747            },
1748        )
1749        .unwrap();
1750        let ari = adjusted_rand_index(&result.cluster, &ground_truth);
1751        assert!(
1752            ari >= 0.90,
1753            "kCFC ARI={ari:.3} should be >= 0.90 on well-separated data"
1754        );
1755    }
1756
1757    #[test]
1758    fn test_kcfc_errors_ordering() {
1759        // Curves in cluster 0 should have smaller error against cluster 0's FPCA
1760        // than against cluster 1's FPCA, and vice versa.
1761        let m = 40;
1762        let n_per = 10;
1763        let (data, t, ground_truth) = two_separated_clusters(n_per, m);
1764        let result = kcfc_cluster(
1765            &data,
1766            &t,
1767            &KcfcConfig {
1768                k: 2,
1769                ncomp: 3,
1770                max_iter: 50,
1771                seed: 42,
1772                ..Default::default()
1773            },
1774        )
1775        .unwrap();
1776
1777        // Determine which result cluster corresponds to ground truth 0 and 1.
1778        // Use an explicit lookup instead of arithmetic (1 - gt0_cluster) so the
1779        // test is not k==2-specific and does not risk wrapping on usize.
1780        let gt0_cluster = result.cluster[0]; // first curve is in ground truth 0
1781        let gt1_cluster = result.cluster[n_per]; // first curve in ground truth 1
1782
1783        let mut correct_ordering = 0;
1784        let mut total = 0;
1785        for i in 0..data.nrows() {
1786            let expected_cluster = if ground_truth[i] == 0 {
1787                gt0_cluster
1788            } else {
1789                gt1_cluster
1790            };
1791            // Find the "other" cluster by scanning — avoids k==2-only arithmetic.
1792            let other_cluster = (0..2).find(|&c| c != expected_cluster).unwrap_or(0);
1793            let err_own = result.reconstruction_errors[(i, expected_cluster)];
1794            let err_other = result.reconstruction_errors[(i, other_cluster)];
1795            if err_own.is_finite() && err_other.is_finite() {
1796                if err_own < err_other {
1797                    correct_ordering += 1;
1798                }
1799                total += 1;
1800            }
1801        }
1802        // At least 80% should have correct error ordering
1803        assert!(
1804            correct_ordering * 10 >= total * 8,
1805            "only {correct_ordering}/{total} curves had smaller error for their true cluster"
1806        );
1807    }
1808
1809    #[test]
1810    fn test_kcfc_deterministic() {
1811        let m = 30;
1812        let n_per = 8;
1813        let (data, t, _) = two_separated_clusters(n_per, m);
1814        let cfg = KcfcConfig {
1815            k: 2,
1816            ncomp: 2,
1817            max_iter: 30,
1818            seed: 7,
1819            ..Default::default()
1820        };
1821        let r1 = kcfc_cluster(&data, &t, &cfg).unwrap();
1822        let r2 = kcfc_cluster(&data, &t, &cfg).unwrap();
1823        assert_eq!(
1824            r1.cluster, r2.cluster,
1825            "identical seed must produce identical assignments"
1826        );
1827    }
1828
1829    #[test]
1830    fn test_kcfc_invalid_k_zero() {
1831        let m = 20;
1832        let (data, t, _) = two_tight_clusters(5, m);
1833        assert!(
1834            kcfc_cluster(
1835                &data,
1836                &t,
1837                &KcfcConfig {
1838                    k: 0,
1839                    ..Default::default()
1840                }
1841            )
1842            .is_err(),
1843            "k=0 should return Err"
1844        );
1845    }
1846
1847    #[test]
1848    fn test_kcfc_invalid_k_gt_n() {
1849        let m = 20;
1850        let n = 4;
1851        let (data, t, _) = two_tight_clusters(n / 2, m);
1852        assert!(
1853            kcfc_cluster(
1854                &data,
1855                &t,
1856                &KcfcConfig {
1857                    k: n + 1,
1858                    ..Default::default()
1859                }
1860            )
1861            .is_err(),
1862            "k>n should return Err"
1863        );
1864    }
1865
1866    #[test]
1867    fn test_kcfc_empty_data() {
1868        let data = FdMatrix::zeros(0, 0);
1869        let t: Vec<f64> = vec![];
1870        assert!(
1871            kcfc_cluster(&data, &t, &KcfcConfig::default()).is_err(),
1872            "empty data should return Err"
1873        );
1874    }
1875
1876    #[test]
1877    fn test_kcfc_mismatched_argvals() {
1878        let m = 20;
1879        let (data, _t, _) = two_tight_clusters(5, m);
1880        let wrong_t = uniform_grid(m + 3);
1881        assert!(
1882            kcfc_cluster(&data, &wrong_t, &KcfcConfig::default()).is_err(),
1883            "mismatched argvals should return Err"
1884        );
1885    }
1886
1887    #[test]
1888    fn test_kcfc_result_shapes() {
1889        let m = 20;
1890        let n_per = 5;
1891        let (data, t, _) = two_separated_clusters(n_per, m);
1892        let n = 2 * n_per;
1893        let result = kcfc_cluster(
1894            &data,
1895            &t,
1896            &KcfcConfig {
1897                k: 2,
1898                ncomp: 2,
1899                ..Default::default()
1900            },
1901        )
1902        .unwrap();
1903        assert_eq!(result.cluster.len(), n);
1904        assert_eq!(result.fpca_models.len(), 2);
1905        assert_eq!(result.reconstruction_errors.shape(), (n, 2));
1906    }
1907
1908    #[test]
1909    fn test_kcfc_ncomp_zero_returns_err() {
1910        // ncomp == 0 must be caught at entry rather than silently assigning all
1911        // curves to cluster 0 (which was the pre-fix behavior when fdata_to_pc_1d
1912        // returned Err and was swallowed by the degenerate-cluster arm).
1913        let m = 20;
1914        let (data, t, _) = two_tight_clusters(5, m);
1915        let result = kcfc_cluster(
1916            &data,
1917            &t,
1918            &KcfcConfig {
1919                k: 2,
1920                ncomp: 0,
1921                ..Default::default()
1922            },
1923        );
1924        assert!(result.is_err(), "ncomp=0 must return Err, got Ok");
1925        if let Err(FdarError::InvalidParameter { parameter, .. }) = result {
1926            assert_eq!(parameter, "ncomp");
1927        } else {
1928            panic!("expected InvalidParameter {{ parameter: \"ncomp\" }}");
1929        }
1930    }
1931
1932    // ── funFEM tests ─────────────────────────────────────────────────────────
1933
1934    #[test]
1935    fn test_funfem_recovery() {
1936        let m = 40;
1937        let n_per = 10;
1938        let (data, t, ground_truth) = two_separated_clusters(n_per, m);
1939        let result = funfem_cluster(
1940            &data,
1941            &t,
1942            &FunFemConfig {
1943                k: 2,
1944                ncomp: 5,
1945                p_disc: 1,
1946                max_iter: 30,
1947                tol: 1e-5,
1948                seed: 42,
1949            },
1950        )
1951        .unwrap();
1952        let ari = adjusted_rand_index(&result.cluster, &ground_truth);
1953        assert!(
1954            ari >= 0.90,
1955            "funFEM ARI={ari:.3} should be >= 0.90 on well-separated data"
1956        );
1957    }
1958
1959    #[test]
1960    fn test_funfem_deterministic() {
1961        let m = 30;
1962        let n_per = 8;
1963        let (data, t, _) = two_separated_clusters(n_per, m);
1964        let cfg = FunFemConfig {
1965            k: 2,
1966            ncomp: 4,
1967            p_disc: 1,
1968            max_iter: 20,
1969            tol: 1e-5,
1970            seed: 99,
1971        };
1972        let r1 = funfem_cluster(&data, &t, &cfg).unwrap();
1973        let r2 = funfem_cluster(&data, &t, &cfg).unwrap();
1974        assert_eq!(
1975            r1.cluster, r2.cluster,
1976            "same seed must produce identical assignments"
1977        );
1978    }
1979
1980    #[test]
1981    fn test_funfem_invalid_k_zero() {
1982        let m = 20;
1983        let (data, t, _) = two_tight_clusters(5, m);
1984        assert!(
1985            funfem_cluster(
1986                &data,
1987                &t,
1988                &FunFemConfig {
1989                    k: 0,
1990                    ..FunFemConfig::default()
1991                }
1992            )
1993            .is_err(),
1994            "k=0 must return Err"
1995        );
1996    }
1997
1998    #[test]
1999    fn test_funfem_invalid_k_gt_n() {
2000        let m = 20;
2001        let (data, t, _) = two_tight_clusters(3, m);
2002        assert!(
2003            funfem_cluster(
2004                &data,
2005                &t,
2006                &FunFemConfig {
2007                    k: 10,
2008                    ..FunFemConfig::default()
2009                }
2010            )
2011            .is_err(),
2012            "k>n must return Err"
2013        );
2014    }
2015
2016    #[test]
2017    fn test_funfem_invalid_ncomp_zero() {
2018        let m = 20;
2019        let (data, t, _) = two_tight_clusters(5, m);
2020        assert!(
2021            funfem_cluster(
2022                &data,
2023                &t,
2024                &FunFemConfig {
2025                    ncomp: 0,
2026                    ..FunFemConfig::default()
2027                }
2028            )
2029            .is_err(),
2030            "ncomp=0 must return Err"
2031        );
2032    }
2033
2034    #[test]
2035    fn test_funfem_invalid_empty_data() {
2036        let data = FdMatrix::zeros(0, 0);
2037        let t: Vec<f64> = vec![];
2038        assert!(
2039            funfem_cluster(&data, &t, &FunFemConfig::default()).is_err(),
2040            "empty data must return Err"
2041        );
2042    }
2043
2044    #[test]
2045    fn test_funfem_invalid_argvals_mismatch() {
2046        let m = 20;
2047        let (data, _t, _) = two_tight_clusters(5, m);
2048        let wrong_t = uniform_grid(m + 2);
2049        assert!(
2050            funfem_cluster(&data, &wrong_t, &FunFemConfig::default()).is_err(),
2051            "argvals mismatch must return Err"
2052        );
2053    }
2054
2055    #[test]
2056    fn test_funfem_output_shapes() {
2057        let m = 30;
2058        let n_per = 6;
2059        let (data, t, _) = two_separated_clusters(n_per, m);
2060        let n = 2 * n_per;
2061        let result = funfem_cluster(
2062            &data,
2063            &t,
2064            &FunFemConfig {
2065                k: 2,
2066                ncomp: 4,
2067                p_disc: 1,
2068                max_iter: 10,
2069                tol: 1e-4,
2070                seed: 1,
2071            },
2072        )
2073        .unwrap();
2074        assert_eq!(result.cluster.len(), n);
2075        assert_eq!(result.membership.shape(), (n, 2));
2076    }
2077
2078    // ── Align-and-cluster tests ───────────────────────────────────────────
2079
2080    /// Two clusters differing in fundamental shape:
2081    /// Group 0 = flat curves (constant value near 0)
2082    /// Group 1 = sin curves (oscillating shape, amplitude offset +5)
2083    ///
2084    /// Each group has small within-group phase/warp variability via different
2085    /// time-warps (g(t) = t^alpha with alpha varying slightly across members).
2086    /// The shapes are so distinct that elastic k-means trivially recovers them,
2087    /// demonstrating shape-invariant clustering on data with within-group warps.
2088    fn time_warped_clusters(n_per: usize, m: usize) -> (FdMatrix, Vec<f64>, Vec<usize>) {
2089        let t = uniform_grid(m);
2090        let n = 2 * n_per;
2091        let mut col_major = vec![0.0_f64; n * m];
2092        // Cluster 0: sin(2πt) evaluated at time-warped grid t^alpha (alpha near 1)
2093        // Small phase warp within cluster — but all are sin-shaped
2094        for i in 0..n_per {
2095            let alpha = 1.0 + 0.1 * (i as f64 / n_per as f64); // alpha in [1.0, 1.1]
2096            for (j, &tj) in t.iter().enumerate() {
2097                let warped = tj.powf(alpha);
2098                col_major[i + j * n] = (2.0 * PI * warped).sin();
2099            }
2100        }
2101        // Cluster 1: flat curves (constant value 8) — completely different shape
2102        for i in 0..n_per {
2103            for j in 0..m {
2104                col_major[(i + n_per) + j * n] = 8.0 + 0.05 * i as f64;
2105            }
2106        }
2107        let labels: Vec<usize> = (0..n).map(|i| if i < n_per { 0 } else { 1 }).collect();
2108        (
2109            FdMatrix::from_column_major(col_major, n, m).unwrap(),
2110            t,
2111            labels,
2112        )
2113    }
2114
2115    #[test]
2116    fn test_align_cluster_shape_shift() {
2117        // Use small n_per and m as per the anti-stall note (elastic DP is slow)
2118        let m = 30;
2119        let n_per = 6;
2120        let (data, t, ground_truth) = time_warped_clusters(n_per, m);
2121        let result = align_cluster_fd(
2122            &data,
2123            &t,
2124            &AlignClusterConfig {
2125                k: 2,
2126                max_iter: 15,
2127                seed: 42,
2128                use_amplitude_only: true,
2129                elastic_lambda: 0.0,
2130                karcher_max_iter: 10,
2131                karcher_tol: 1e-3,
2132            },
2133        )
2134        .unwrap();
2135        let ari = adjusted_rand_index(&result.cluster, &ground_truth);
2136        assert!(
2137            ari >= 0.90,
2138            "align_cluster ARI={ari:.3} on shape-distinct data should be >= 0.90"
2139        );
2140    }
2141
2142    #[test]
2143    fn test_align_cluster_recovery() {
2144        let m = 30;
2145        let n_per = 6;
2146        let (data, t, ground_truth) = two_separated_clusters(n_per, m);
2147        let result = align_cluster_fd(
2148            &data,
2149            &t,
2150            &AlignClusterConfig {
2151                k: 2,
2152                max_iter: 15,
2153                seed: 7,
2154                use_amplitude_only: true,
2155                elastic_lambda: 0.0,
2156                karcher_max_iter: 10,
2157                karcher_tol: 1e-3,
2158            },
2159        )
2160        .unwrap();
2161        let ari = adjusted_rand_index(&result.cluster, &ground_truth);
2162        assert!(
2163            ari >= 0.90,
2164            "align_cluster ARI={ari:.3} on amplitude-separated data should be >= 0.90"
2165        );
2166    }
2167
2168    #[test]
2169    fn test_align_cluster_invalid_k_zero() {
2170        let m = 20;
2171        let (data, t, _) = two_tight_clusters(5, m);
2172        assert!(
2173            align_cluster_fd(
2174                &data,
2175                &t,
2176                &AlignClusterConfig {
2177                    k: 0,
2178                    ..AlignClusterConfig::default()
2179                }
2180            )
2181            .is_err(),
2182            "k=0 must return Err"
2183        );
2184    }
2185
2186    #[test]
2187    fn test_align_cluster_invalid_k_gt_n() {
2188        let m = 20;
2189        let (data, t, _) = two_tight_clusters(3, m);
2190        assert!(
2191            align_cluster_fd(
2192                &data,
2193                &t,
2194                &AlignClusterConfig {
2195                    k: 20,
2196                    ..AlignClusterConfig::default()
2197                }
2198            )
2199            .is_err(),
2200            "k>n must return Err"
2201        );
2202    }
2203
2204    #[test]
2205    fn test_align_cluster_invalid_empty_data() {
2206        let data = FdMatrix::zeros(0, 0);
2207        let t: Vec<f64> = vec![];
2208        assert!(
2209            align_cluster_fd(&data, &t, &AlignClusterConfig::default()).is_err(),
2210            "empty data must return Err"
2211        );
2212    }
2213
2214    #[test]
2215    fn test_align_cluster_invalid_argvals_mismatch() {
2216        let m = 20;
2217        let (data, _t, _) = two_tight_clusters(5, m);
2218        let wrong_t = uniform_grid(m + 5);
2219        assert!(
2220            align_cluster_fd(&data, &wrong_t, &AlignClusterConfig::default()).is_err(),
2221            "argvals mismatch must return Err"
2222        );
2223    }
2224
2225    #[test]
2226    fn test_align_cluster_output_shapes() {
2227        let m = 20;
2228        let n_per = 4;
2229        let (data, t, _) = two_separated_clusters(n_per, m);
2230        let n = 2 * n_per;
2231        let result = align_cluster_fd(
2232            &data,
2233            &t,
2234            &AlignClusterConfig {
2235                k: 2,
2236                max_iter: 5,
2237                seed: 1,
2238                karcher_max_iter: 5,
2239                karcher_tol: 1e-2,
2240                ..AlignClusterConfig::default()
2241            },
2242        )
2243        .unwrap();
2244        assert_eq!(result.cluster.len(), n);
2245        assert_eq!(result.templates.len(), 2);
2246        assert!(result.templates.iter().all(|t| t.len() == m));
2247        assert_eq!(result.distances.shape(), (n, 2));
2248    }
2249}