Skip to main content

fdars_core/alignment/
clustering.rs

1//! Distance-based clustering: k-means (k-medoids) and hierarchical.
2//!
3//! These algorithms work with **any** precomputed distance matrix — elastic
4//! (Fisher-Rao), DTW, Lp, amplitude-only, phase-only, or user-defined.
5//!
6//! # Examples
7//!
8//! ```
9//! use fdars_core::alignment::{
10//!     elastic_self_distance_matrix, hierarchical_from_distances,
11//!     kmedoids_from_distances, cut_dendrogram, Linkage, KMedoidsConfig,
12//! };
13//! use fdars_core::matrix::FdMatrix;
14//!
15//! // Compute any distance matrix
16//! let t: Vec<f64> = (0..20).map(|i| i as f64 / 19.0).collect();
17//! let data = FdMatrix::zeros(5, 20);
18//! let dist = elastic_self_distance_matrix(&data, &t, 0.0);
19//!
20//! // Hierarchical clustering — works with any distance matrix
21//! let dendro = hierarchical_from_distances(&dist, Linkage::Complete).unwrap();
22//! let labels = cut_dendrogram(&dendro, 2).unwrap();
23//!
24//! // K-medoids — works with any distance matrix
25//! let mut config = KMedoidsConfig::default();
26//! config.k = 2;
27//! let result = kmedoids_from_distances(&dist, &config).unwrap();
28//! ```
29
30use crate::error::FdarError;
31use crate::matrix::FdMatrix;
32use rand::rngs::StdRng;
33use rand::{Rng, SeedableRng};
34
35// ─── Types ──────────────────────────────────────────────────────────────────
36
37/// Configuration for k-medoids clustering.
38///
39/// Construct via `KMedoidsConfig::default()`, then assign the fields you need (e.g. `let mut c = KMedoidsConfig::default(); c.field = …;`). This struct is `#[non_exhaustive]`, so external crates cannot build it with a struct literal — not even functional-update `..Default::default()` form.
40#[non_exhaustive]
41#[derive(Debug, Clone, PartialEq)]
42pub struct KMedoidsConfig {
43    /// Number of clusters.
44    pub k: usize,
45    /// Maximum number of iterations.
46    pub max_iter: usize,
47    /// Random seed for k-means++ initialization.
48    pub seed: u64,
49}
50
51impl Default for KMedoidsConfig {
52    fn default() -> Self {
53        Self {
54            k: 2,
55            max_iter: 100,
56            seed: 42,
57        }
58    }
59}
60
61/// Linkage method for hierarchical clustering.
62#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
63#[non_exhaustive]
64pub enum Linkage {
65    /// Minimum distance between clusters.
66    #[default]
67    Single,
68    /// Maximum distance between clusters.
69    Complete,
70    /// Weighted average distance (UPGMA).
71    Average,
72}
73
74/// Result of k-medoids clustering.
75#[derive(Debug, Clone, PartialEq)]
76#[non_exhaustive]
77pub struct KMedoidsResult {
78    /// Cluster label for each observation (0-indexed, length n).
79    pub labels: Vec<usize>,
80    /// Medoid index for each cluster (length k).
81    pub medoid_indices: Vec<usize>,
82    /// Within-cluster sum of distances for each cluster.
83    pub within_distances: Vec<f64>,
84    /// Total within-cluster distance.
85    pub total_within_distance: f64,
86    /// Number of iterations performed.
87    pub n_iter: usize,
88    /// Whether the algorithm converged (labels stabilized).
89    pub converged: bool,
90}
91
92/// Result of hierarchical clustering (dendrogram).
93#[derive(Debug, Clone, PartialEq)]
94#[non_exhaustive]
95pub struct Dendrogram {
96    /// Merge history: each entry `(i, j, distance)` records merging cluster
97    /// indices i and j at the given distance.
98    pub merges: Vec<(usize, usize, f64)>,
99    /// Number of observations.
100    pub n: usize,
101}
102
103// ─── K-Means++ Initialization ───────────────────────────────────────────────
104
105/// Select k initial center indices using k-means++ on a precomputed distance matrix.
106fn kmeans_pp_init(dist_mat: &FdMatrix, k: usize, rng: &mut StdRng) -> Vec<usize> {
107    let n = dist_mat.nrows();
108    let mut centers = Vec::with_capacity(k);
109
110    centers.push(rng.gen_range(0..n));
111
112    let mut min_dist_sq: Vec<f64> = (0..n)
113        .map(|i| {
114            let d = dist_mat[(i, centers[0])];
115            d * d
116        })
117        .collect();
118
119    for _ in 1..k {
120        let total: f64 = min_dist_sq.iter().sum();
121        if total <= 0.0 {
122            for i in 0..n {
123                if !centers.contains(&i) {
124                    centers.push(i);
125                    break;
126                }
127            }
128        } else {
129            let threshold = rng.gen::<f64>() * total;
130            let mut cum = 0.0;
131            let mut chosen = n - 1;
132            for i in 0..n {
133                cum += min_dist_sq[i];
134                if cum >= threshold {
135                    chosen = i;
136                    break;
137                }
138            }
139            centers.push(chosen);
140        }
141
142        let new_center = *centers.last().unwrap();
143        for i in 0..n {
144            let d = dist_mat[(i, new_center)];
145            let d2 = d * d;
146            if d2 < min_dist_sq[i] {
147                min_dist_sq[i] = d2;
148            }
149        }
150    }
151
152    centers
153}
154
155// ─── K-Medoids ─────────────────────────────────────────────────────────────
156
157/// K-medoids (PAM-style) clustering from a precomputed distance matrix.
158///
159/// Uses k-means++ initialization, then alternates between assigning each
160/// observation to its nearest medoid and selecting the medoid that minimizes
161/// within-cluster distances.
162///
163/// Works with **any** distance matrix — elastic, DTW, Lp, or user-defined.
164///
165/// # Arguments
166/// * `dist_mat` — Symmetric n x n distance matrix.
167/// * `config`   — Clustering configuration.
168///
169/// # Errors
170/// Returns [`FdarError::InvalidParameter`] if `k < 1` or `k > n`.
171/// Returns [`FdarError::InvalidDimension`] if `dist_mat` is not square.
172#[must_use = "expensive computation whose result should not be discarded"]
173pub fn kmedoids_from_distances(
174    dist_mat: &FdMatrix,
175    config: &KMedoidsConfig,
176) -> Result<KMedoidsResult, FdarError> {
177    let n = dist_mat.nrows();
178    if dist_mat.ncols() != n {
179        return Err(FdarError::InvalidDimension {
180            parameter: "dist_mat",
181            expected: format!("{n} x {n} (square)"),
182            actual: format!("{} x {}", n, dist_mat.ncols()),
183        });
184    }
185    if config.k < 1 {
186        return Err(FdarError::InvalidParameter {
187            parameter: "k",
188            message: "k must be >= 1".to_string(),
189        });
190    }
191    if config.k > n {
192        return Err(FdarError::InvalidParameter {
193            parameter: "k",
194            message: format!("k ({}) must be <= n ({})", config.k, n),
195        });
196    }
197
198    let k = config.k;
199    let mut rng = StdRng::seed_from_u64(config.seed);
200    let mut medoids = kmeans_pp_init(dist_mat, k, &mut rng);
201
202    // Assign each point to nearest medoid.
203    let mut labels = assign_to_medoids(dist_mat, &medoids, n);
204
205    let mut converged = false;
206    let mut n_iter = 0;
207
208    for iter in 0..config.max_iter {
209        n_iter = iter + 1;
210
211        // Update medoids: for each cluster, pick the member minimizing total distance.
212        for c in 0..k {
213            let members: Vec<usize> = (0..n).filter(|&i| labels[i] == c).collect();
214            if members.is_empty() {
215                continue;
216            }
217            let mut best_cost = f64::INFINITY;
218            let mut best_m = medoids[c];
219            for &candidate in &members {
220                let cost: f64 = members.iter().map(|&j| dist_mat[(candidate, j)]).sum();
221                if cost < best_cost {
222                    best_cost = cost;
223                    best_m = candidate;
224                }
225            }
226            medoids[c] = best_m;
227        }
228
229        // Reassign.
230        let new_labels = assign_to_medoids(dist_mat, &medoids, n);
231        if new_labels == labels {
232            converged = true;
233            labels = new_labels;
234            break;
235        }
236        labels = new_labels;
237    }
238
239    // Compute within-cluster distances.
240    let mut within_distances = vec![0.0; k];
241    for i in 0..n {
242        within_distances[labels[i]] += dist_mat[(i, medoids[labels[i]])];
243    }
244    let total_within_distance: f64 = within_distances.iter().sum();
245
246    Ok(KMedoidsResult {
247        labels,
248        medoid_indices: medoids,
249        within_distances,
250        total_within_distance,
251        n_iter,
252        converged,
253    })
254}
255
256fn assign_to_medoids(dist_mat: &FdMatrix, medoids: &[usize], n: usize) -> Vec<usize> {
257    (0..n)
258        .map(|i| {
259            let mut best_d = f64::INFINITY;
260            let mut best_c = 0;
261            for (c, &med) in medoids.iter().enumerate() {
262                let d = dist_mat[(i, med)];
263                if d < best_d {
264                    best_d = d;
265                    best_c = c;
266                }
267            }
268            best_c
269        })
270        .collect()
271}
272
273// ─── Hierarchical Clustering ───────────────────────────────────────────────
274
275/// Hierarchical agglomerative clustering from a precomputed distance matrix.
276///
277/// Builds a [`Dendrogram`] by iteratively merging the closest pair of clusters.
278/// Works with **any** distance matrix — elastic, DTW, Lp, or user-defined.
279///
280/// # Arguments
281/// * `dist_mat` — Symmetric n x n distance matrix.
282/// * `linkage`  — Linkage criterion.
283///
284/// # Errors
285/// Returns [`FdarError::InvalidDimension`] if `dist_mat` is not square or `n < 2`.
286#[must_use = "expensive computation whose result should not be discarded"]
287pub fn hierarchical_from_distances(
288    dist_mat: &FdMatrix,
289    linkage: Linkage,
290) -> Result<Dendrogram, FdarError> {
291    let n = dist_mat.nrows();
292    if dist_mat.ncols() != n {
293        return Err(FdarError::InvalidDimension {
294            parameter: "dist_mat",
295            expected: format!("{n} x {n} (square)"),
296            actual: format!("{} x {}", n, dist_mat.ncols()),
297        });
298    }
299    if n < 2 {
300        return Err(FdarError::InvalidDimension {
301            parameter: "dist_mat",
302            expected: "at least 2 rows".to_string(),
303            actual: format!("{n} rows"),
304        });
305    }
306
307    let mut active = vec![true; n];
308    let mut cluster_sizes = vec![1usize; n];
309    let mut cluster_dist = FdMatrix::zeros(n, n);
310    for i in 0..n {
311        for j in 0..n {
312            cluster_dist[(i, j)] = dist_mat[(i, j)];
313        }
314    }
315
316    let mut merges: Vec<(usize, usize, f64)> = Vec::with_capacity(n - 1);
317
318    for _ in 0..(n - 1) {
319        let mut min_d = f64::INFINITY;
320        let mut min_i = 0;
321        let mut min_j = 1;
322        for i in 0..n {
323            if !active[i] {
324                continue;
325            }
326            for j in (i + 1)..n {
327                if !active[j] {
328                    continue;
329                }
330                if cluster_dist[(i, j)] < min_d {
331                    min_d = cluster_dist[(i, j)];
332                    min_i = i;
333                    min_j = j;
334                }
335            }
336        }
337
338        merges.push((min_i, min_j, min_d));
339
340        let size_i = cluster_sizes[min_i];
341        let size_j = cluster_sizes[min_j];
342        for k in 0..n {
343            if !active[k] || k == min_i || k == min_j {
344                continue;
345            }
346            let d_ik = cluster_dist[(min_i.min(k), min_i.max(k))];
347            let d_jk = cluster_dist[(min_j.min(k), min_j.max(k))];
348            let new_d = match linkage {
349                Linkage::Single => d_ik.min(d_jk),
350                Linkage::Complete => d_ik.max(d_jk),
351                Linkage::Average => {
352                    (d_ik * size_i as f64 + d_jk * size_j as f64) / (size_i + size_j) as f64
353                }
354            };
355            let (lo, hi) = (min_i.min(k), min_i.max(k));
356            cluster_dist[(lo, hi)] = new_d;
357            cluster_dist[(hi, lo)] = new_d;
358        }
359
360        cluster_sizes[min_i] = size_i + size_j;
361        active[min_j] = false;
362    }
363
364    Ok(Dendrogram { merges, n })
365}
366
367// ─── Cut Dendrogram ─────────────────────────────────────────────────────────
368
369/// Cut a dendrogram to produce k clusters.
370///
371/// Replays the merge history, stopping after `n - k` merges, and returns
372/// cluster labels for each original observation.
373///
374/// # Arguments
375/// * `dendrogram` — Result of [`hierarchical_from_distances`].
376/// * `k`          — Number of clusters desired.
377///
378/// # Errors
379/// Returns [`FdarError::InvalidParameter`] if `k < 1` or `k > n`.
380pub fn cut_dendrogram(dendrogram: &Dendrogram, k: usize) -> Result<Vec<usize>, FdarError> {
381    let n = dendrogram.n;
382
383    if k < 1 {
384        return Err(FdarError::InvalidParameter {
385            parameter: "k",
386            message: "k must be >= 1".to_string(),
387        });
388    }
389    if k > n {
390        return Err(FdarError::InvalidParameter {
391            parameter: "k",
392            message: format!("k ({k}) must be <= n ({n})"),
393        });
394    }
395
396    let mut cluster_of: Vec<usize> = (0..n).collect();
397    let merges_to_apply = n - k;
398
399    for &(ci, cj, _) in dendrogram.merges.iter().take(merges_to_apply) {
400        let target = cluster_of[ci];
401        let source = cluster_of[cj];
402        for label in cluster_of.iter_mut() {
403            if *label == source {
404                *label = target;
405            }
406        }
407    }
408
409    // Compress labels to 0..k-1.
410    let mut unique: Vec<usize> = cluster_of.clone();
411    unique.sort_unstable();
412    unique.dedup();
413    let labels = cluster_of
414        .iter()
415        .map(|&l| unique.iter().position(|&u| u == l).unwrap())
416        .collect();
417
418    Ok(labels)
419}
420
421// ─── Tests ──────────────────────────────────────────────────────────────────
422
423#[cfg(test)]
424mod tests {
425    use super::*;
426    use crate::alignment::elastic_self_distance_matrix;
427    use crate::simulation::{sim_fundata, EFunType, EValType};
428    use crate::test_helpers::uniform_grid;
429
430    fn make_dist_mat(n: usize, m: usize) -> FdMatrix {
431        let t = uniform_grid(m);
432        let data = sim_fundata(n, &t, 3, EFunType::Fourier, EValType::Exponential, Some(42));
433        elastic_self_distance_matrix(&data, &t, 0.0)
434    }
435
436    #[test]
437    fn kmedoids_smoke() {
438        let dist = make_dist_mat(8, 20);
439        let config = KMedoidsConfig {
440            k: 2,
441            max_iter: 10,
442            ..Default::default()
443        };
444        let result = kmedoids_from_distances(&dist, &config).unwrap();
445        assert_eq!(result.labels.len(), 8);
446        assert_eq!(result.medoid_indices.len(), 2);
447        assert_eq!(result.within_distances.len(), 2);
448        assert!(result.total_within_distance >= 0.0);
449        assert!(result.n_iter >= 1);
450    }
451
452    #[test]
453    fn kmedoids_single_cluster() {
454        let dist = make_dist_mat(5, 20);
455        let config = KMedoidsConfig {
456            k: 1,
457            max_iter: 10,
458            ..Default::default()
459        };
460        let result = kmedoids_from_distances(&dist, &config).unwrap();
461        assert!(result.labels.iter().all(|&l| l == 0));
462        assert_eq!(result.medoid_indices.len(), 1);
463    }
464
465    #[test]
466    fn kmedoids_k_too_large() {
467        let dist = make_dist_mat(3, 20);
468        let config = KMedoidsConfig {
469            k: 5,
470            ..Default::default()
471        };
472        assert!(kmedoids_from_distances(&dist, &config).is_err());
473    }
474
475    #[test]
476    fn kmedoids_k_zero() {
477        let dist = make_dist_mat(5, 20);
478        let config = KMedoidsConfig {
479            k: 0,
480            ..Default::default()
481        };
482        assert!(kmedoids_from_distances(&dist, &config).is_err());
483    }
484
485    #[test]
486    fn hierarchical_single_smoke() {
487        let dist = make_dist_mat(5, 20);
488        let dendro = hierarchical_from_distances(&dist, Linkage::Single).unwrap();
489        assert_eq!(dendro.merges.len(), 4);
490        for w in dendro.merges.windows(2) {
491            assert!(
492                w[1].2 >= w[0].2 - 1e-10,
493                "single linkage should be non-decreasing"
494            );
495        }
496    }
497
498    #[test]
499    fn hierarchical_complete_smoke() {
500        let dist = make_dist_mat(5, 20);
501        let dendro = hierarchical_from_distances(&dist, Linkage::Complete).unwrap();
502        assert_eq!(dendro.merges.len(), 4);
503    }
504
505    #[test]
506    fn hierarchical_average_smoke() {
507        let dist = make_dist_mat(5, 20);
508        let dendro = hierarchical_from_distances(&dist, Linkage::Average).unwrap();
509        assert_eq!(dendro.merges.len(), 4);
510    }
511
512    #[test]
513    fn hierarchical_too_few() {
514        let dist = FdMatrix::zeros(1, 1);
515        assert!(hierarchical_from_distances(&dist, Linkage::Single).is_err());
516    }
517
518    #[test]
519    fn cut_dendrogram_all_singletons() {
520        let dist = make_dist_mat(5, 20);
521        let dendro = hierarchical_from_distances(&dist, Linkage::Single).unwrap();
522        let labels = cut_dendrogram(&dendro, 5).unwrap();
523        let mut sorted = labels.clone();
524        sorted.sort_unstable();
525        assert_eq!(sorted, vec![0, 1, 2, 3, 4]);
526    }
527
528    #[test]
529    fn cut_dendrogram_one_cluster() {
530        let dist = make_dist_mat(5, 20);
531        let dendro = hierarchical_from_distances(&dist, Linkage::Single).unwrap();
532        let labels = cut_dendrogram(&dendro, 1).unwrap();
533        assert!(labels.iter().all(|&l| l == 0));
534    }
535
536    #[test]
537    fn cut_dendrogram_k_too_large() {
538        let dist = make_dist_mat(5, 20);
539        let dendro = hierarchical_from_distances(&dist, Linkage::Single).unwrap();
540        assert!(cut_dendrogram(&dendro, 10).is_err());
541    }
542
543    #[test]
544    fn cut_dendrogram_two_clusters() {
545        let dist = make_dist_mat(6, 20);
546        let dendro = hierarchical_from_distances(&dist, Linkage::Single).unwrap();
547        let labels = cut_dendrogram(&dendro, 2).unwrap();
548        assert_eq!(labels.len(), 6);
549        let unique: std::collections::HashSet<usize> = labels.iter().copied().collect();
550        assert_eq!(unique.len(), 2);
551    }
552
553    #[test]
554    fn default_config_values() {
555        let cfg = KMedoidsConfig::default();
556        assert_eq!(cfg.k, 2);
557        assert_eq!(cfg.max_iter, 100);
558        assert_eq!(cfg.seed, 42);
559    }
560
561    #[test]
562    fn default_linkage() {
563        assert_eq!(Linkage::default(), Linkage::Single);
564    }
565
566    #[test]
567    fn non_square_dist_mat_error() {
568        let dist = FdMatrix::zeros(3, 4);
569        assert!(hierarchical_from_distances(&dist, Linkage::Single).is_err());
570        let config = KMedoidsConfig::default();
571        assert!(kmedoids_from_distances(&dist, &config).is_err());
572    }
573}