1use crate::error::FdarError;
31use crate::matrix::FdMatrix;
32use rand::rngs::StdRng;
33use rand::{Rng, SeedableRng};
34
35#[non_exhaustive]
41#[derive(Debug, Clone, PartialEq)]
42pub struct KMedoidsConfig {
43 pub k: usize,
45 pub max_iter: usize,
47 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
63#[non_exhaustive]
64pub enum Linkage {
65 #[default]
67 Single,
68 Complete,
70 Average,
72}
73
74#[derive(Debug, Clone, PartialEq)]
76#[non_exhaustive]
77pub struct KMedoidsResult {
78 pub labels: Vec<usize>,
80 pub medoid_indices: Vec<usize>,
82 pub within_distances: Vec<f64>,
84 pub total_within_distance: f64,
86 pub n_iter: usize,
88 pub converged: bool,
90}
91
92#[derive(Debug, Clone, PartialEq)]
94#[non_exhaustive]
95pub struct Dendrogram {
96 pub merges: Vec<(usize, usize, f64)>,
99 pub n: usize,
101}
102
103fn 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#[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 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 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 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 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#[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
367pub 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 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#[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}