1use crate::util::stats;
70use nalgebra::DMatrix;
71use rand::rngs::StdRng;
72use rand::{RngExt, SeedableRng};
73use std::collections::BTreeMap;
74
75pub const DEFAULT_SEED: u64 = 42;
77
78#[derive(Debug, Clone, PartialEq, thiserror::Error)]
80pub enum OncError {
81 #[error("the correlation matrix must be square with at least two rows")]
83 InvalidCorrelationMatrix,
84 #[error("repeat must be positive")]
86 InvalidRepeat,
87 #[error("clustering failed to produce a partition")]
90 ClusteringFailed,
91}
92
93#[derive(Debug, Clone)]
95pub struct OncResult {
96 pub ordered_correlation: DMatrix<f64>,
99 pub clusters: BTreeMap<usize, Vec<usize>>,
102 pub silhouette_scores: Vec<f64>,
105}
106
107#[derive(Clone)]
108struct ClusterState {
109 ordered_correlation: DMatrix<f64>,
110 clusters: BTreeMap<usize, Vec<usize>>,
111 silhouette_scores: Vec<f64>,
112}
113
114pub fn check_improve_clusters<T: Clone>(
128 new_tstat_mean: f64,
129 mean_redo_tstat: f64,
130 old_cluster: T,
131 new_cluster: T,
132) -> T {
133 if new_tstat_mean > mean_redo_tstat {
134 new_cluster
135 } else {
136 old_cluster
137 }
138}
139
140pub fn get_onc_clusters(corr_mat: &DMatrix<f64>, repeat: usize) -> Result<OncResult, OncError> {
188 get_onc_clusters_with_seed(corr_mat, repeat, DEFAULT_SEED)
189}
190
191pub fn get_onc_clusters_with_seed(
229 corr_mat: &DMatrix<f64>,
230 repeat: usize,
231 seed: u64,
232) -> Result<OncResult, OncError> {
233 if repeat == 0 {
234 return Err(OncError::InvalidRepeat);
235 }
236 if corr_mat.nrows() != corr_mat.ncols() || corr_mat.nrows() < 2 {
237 return Err(OncError::InvalidCorrelationMatrix);
238 }
239
240 let mut rng = StdRng::seed_from_u64(seed);
241 let state = cluster_kmeans_top(corr_mat, repeat, &mut rng)?;
242 Ok(OncResult {
243 ordered_correlation: state.ordered_correlation,
244 clusters: state.clusters,
245 silhouette_scores: state.silhouette_scores,
246 })
247}
248
249fn cluster_kmeans_top(
250 corr_mat: &DMatrix<f64>,
251 repeat: usize,
252 rng: &mut StdRng,
253) -> Result<ClusterState, OncError> {
254 let max_num_clusters = corr_mat.ncols().saturating_sub(1).max(2);
255 let base = cluster_kmeans_base(corr_mat, max_num_clusters, repeat, rng)?;
256
257 let mut cluster_quality: BTreeMap<usize, f64> = BTreeMap::new();
258 for (k, members) in &base.clusters {
259 let scores: Vec<f64> = members.iter().map(|&idx| base.silhouette_scores[idx]).collect();
260 cluster_quality.insert(*k, tstat(&scores));
261 }
262
263 let avg_quality = {
264 let vals: Vec<f64> = cluster_quality.values().copied().collect();
265 if vals.is_empty() {
266 0.0
267 } else {
268 vals.iter().sum::<f64>() / vals.len() as f64
269 }
270 };
271
272 let redo_clusters: Vec<usize> = cluster_quality
273 .iter()
274 .filter_map(|(k, q)| if *q < avg_quality { Some(*k) } else { None })
275 .collect();
276
277 if redo_clusters.len() <= 2 {
278 return Ok(base);
279 }
280
281 let mut keys_redo = Vec::new();
282 for key in &redo_clusters {
283 if let Some(v) = base.clusters.get(key) {
284 keys_redo.extend(v.iter().copied());
285 }
286 }
287
288 if keys_redo.len() < 2 {
289 return Ok(base);
290 }
291
292 let corr_tmp = submatrix(corr_mat, &keys_redo);
293 let mean_redo_tstat = {
294 let vals: Vec<f64> =
295 redo_clusters.iter().filter_map(|k| cluster_quality.get(k).copied()).collect();
296 vals.iter().sum::<f64>() / vals.len() as f64
297 };
298
299 let top_state = cluster_kmeans_top(&corr_tmp, repeat, rng)?;
300 let mut top_clusters_global = BTreeMap::new();
301 for (k, v) in top_state.clusters {
302 let mapped: Vec<usize> = v.into_iter().map(|local_idx| keys_redo[local_idx]).collect();
303 top_clusters_global.insert(k, mapped);
304 }
305
306 let mut kept_clusters = BTreeMap::new();
307 for (k, v) in &base.clusters {
308 if !redo_clusters.contains(k) {
309 kept_clusters.insert(*k, v.clone());
310 }
311 }
312
313 let improved = improve_clusters(corr_mat, &kept_clusters, &top_clusters_global)?;
314
315 let new_tstat_mean = {
316 let mut vals = Vec::new();
317 for members in improved.clusters.values() {
318 let scores: Vec<f64> =
319 members.iter().map(|&idx| improved.silhouette_scores[idx]).collect();
320 vals.push(tstat(&scores));
321 }
322 vals.iter().sum::<f64>() / vals.len() as f64
323 };
324
325 Ok(check_improve_clusters(new_tstat_mean, mean_redo_tstat, base, improved))
326}
327
328fn improve_clusters(
329 corr_mat: &DMatrix<f64>,
330 kept_clusters: &BTreeMap<usize, Vec<usize>>,
331 top_clusters: &BTreeMap<usize, Vec<usize>>,
332) -> Result<ClusterState, OncError> {
333 let mut clusters_new: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
334 for members in kept_clusters.values() {
335 clusters_new.insert(clusters_new.len(), members.clone());
336 }
337 for members in top_clusters.values() {
338 clusters_new.insert(clusters_new.len(), members.clone());
339 }
340
341 let mut new_idx = Vec::new();
342 for members in clusters_new.values() {
343 new_idx.extend(members.iter().copied());
344 }
345
346 let corr_new = submatrix(corr_mat, &new_idx);
347 let labels = labels_from_clusters(corr_mat.nrows(), &clusters_new);
348 let dist = corr_to_distance(corr_mat);
349 let silh_scores_new = silhouette_samples(&dist, &labels);
350
351 Ok(ClusterState {
352 ordered_correlation: corr_new,
353 clusters: clusters_new,
354 silhouette_scores: silh_scores_new,
355 })
356}
357
358fn cluster_kmeans_base(
359 corr_mat: &DMatrix<f64>,
360 max_num_clusters: usize,
361 repeat: usize,
362 rng: &mut StdRng,
363) -> Result<ClusterState, OncError> {
364 let distance = corr_to_distance(corr_mat);
365 let points = Points::new(&distance);
366 let pairwise = pairwise_distances(&points);
367
368 let mut best_labels: Option<Vec<usize>> = None;
369 let mut best_silh: Option<Vec<f64>> = None;
370
371 for _ in 0..repeat {
372 for num_clusters in 2..=max_num_clusters {
373 let labels = kmeans_labels(&points, num_clusters, rng)?;
374 let silh = silhouette_from_pairwise(&pairwise, &labels);
375
376 let stat = tstat(&silh);
377 let best_stat = best_silh.as_ref().map_or(f64::NEG_INFINITY, |s| tstat(s));
378 let tie_better = stat == best_stat
384 && best_silh.as_ref().is_some_and(|b| mean_of(&silh) > mean_of(b));
385 if best_stat.is_nan() || stat > best_stat || tie_better {
386 best_labels = Some(labels);
387 best_silh = Some(silh);
388 }
389 }
390 }
391
392 let labels = best_labels.ok_or(OncError::ClusteringFailed)?;
393 let silh = best_silh.ok_or(OncError::ClusteringFailed)?;
394
395 let mut new_idx: Vec<usize> = (0..labels.len()).collect();
396 new_idx.sort_by_key(|&i| labels[i]);
397
398 let corr1 = submatrix(corr_mat, &new_idx);
399
400 let mut raw_clusters: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
401 for (idx, lbl) in labels.iter().copied().enumerate() {
402 raw_clusters.entry(lbl).or_default().push(idx);
403 }
404
405 let mut clusters = BTreeMap::new();
406 for members in raw_clusters.values() {
407 clusters.insert(clusters.len(), members.clone());
408 }
409
410 Ok(ClusterState { ordered_correlation: corr1, clusters, silhouette_scores: silh })
411}
412
413fn mean_of(values: &[f64]) -> f64 {
414 values.iter().sum::<f64>() / values.len() as f64
415}
416
417fn tstat(values: &[f64]) -> f64 {
418 let (Some(mean), Some(std)) = (stats::mean(values), stats::std_dev(values, 0)) else {
420 return 0.0;
421 };
422 if std <= 1e-12 {
423 if mean > 0.0 {
424 f64::INFINITY
425 } else {
426 0.0
427 }
428 } else {
429 mean / std
430 }
431}
432
433fn labels_from_clusters(n: usize, clusters: &BTreeMap<usize, Vec<usize>>) -> Vec<usize> {
434 let mut labels = vec![0usize; n];
435 for (label, members) in clusters {
436 for &idx in members {
437 labels[idx] = *label;
438 }
439 }
440 labels
441}
442
443fn submatrix(m: &DMatrix<f64>, idx: &[usize]) -> DMatrix<f64> {
444 let n = idx.len();
445 let mut out = DMatrix::zeros(n, n);
446 for (i, &ri) in idx.iter().enumerate() {
447 for (j, &cj) in idx.iter().enumerate() {
448 out[(i, j)] = m[(ri, cj)];
449 }
450 }
451 out
452}
453
454fn corr_to_distance(corr: &DMatrix<f64>) -> DMatrix<f64> {
455 let n = corr.nrows();
456 let mut distance = DMatrix::zeros(n, n);
457 for i in 0..n {
458 for j in 0..n {
459 let c = corr[(i, j)].clamp(-1.0, 1.0);
460 distance[(i, j)] = ((1.0 - c) / 2.0).sqrt();
461 }
462 }
463 distance
464}
465
466struct Points {
468 data: Vec<f64>,
469 n: usize,
470 d: usize,
471}
472
473impl Points {
474 fn new(m: &DMatrix<f64>) -> Self {
475 let (n, d) = m.shape();
476 let data = (0..n).flat_map(|i| (0..d).map(move |j| m[(i, j)])).collect();
477 Self { data, n, d }
478 }
479
480 fn row(&self, i: usize) -> &[f64] {
481 &self.data[i * self.d..(i + 1) * self.d]
482 }
483}
484
485fn squared_distance(a: &[f64], b: &[f64]) -> f64 {
486 a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
487}
488
489fn kmeans_plus_plus(points: &Points, k: usize, rng: &mut StdRng) -> Vec<usize> {
493 let n = points.n;
494 let mut centres = Vec::with_capacity(k);
495 centres.push(rng.random_range(0..n));
496 let mut closest: Vec<f64> =
497 (0..n).map(|i| squared_distance(points.row(i), points.row(centres[0]))).collect();
498 let n_trials = 2 + (k as f64).ln() as usize;
499 for _ in 1..k {
500 let total: f64 = closest.iter().sum();
501 let mut best: Option<(f64, usize, Vec<f64>)> = None;
502 for _ in 0..n_trials {
503 let pick = if total > 0.0 {
504 let target = rng.random::<f64>() * total;
505 let mut acc = 0.0;
506 closest
507 .iter()
508 .position(|w| {
509 acc += w;
510 acc > target
511 })
512 .unwrap_or(n - 1)
513 } else {
514 rng.random_range(0..n)
516 };
517 let candidate = points.row(pick);
518 let updated: Vec<f64> = closest
519 .iter()
520 .enumerate()
521 .map(|(i, c)| c.min(squared_distance(points.row(i), candidate)))
522 .collect();
523 let potential: f64 = updated.iter().sum();
524 if best.as_ref().is_none_or(|(p, _, _)| potential < *p) {
525 best = Some((potential, pick, updated));
526 }
527 }
528 let (_, pick, updated) = best.expect("at least two trials");
529 centres.push(pick);
530 closest = updated;
531 }
532 centres
533}
534
535fn lloyd(points: &Points, centres: &[usize]) -> Vec<usize> {
537 let (n, d, k) = (points.n, points.d, centres.len());
538 let mut centroids: Vec<f64> = centres.iter().flat_map(|&c| points.row(c).to_vec()).collect();
539 let mut labels = vec![usize::MAX; n];
540 let mut dist = vec![0.0; n];
541 for _ in 0..300 {
542 let mut changed = false;
543 for i in 0..n {
544 let row = points.row(i);
545 let mut best_c = 0usize;
546 let mut best_dist = f64::INFINITY;
547 for (c, centroid) in centroids.chunks_exact(d).enumerate() {
548 let s = squared_distance(row, centroid);
549 if s < best_dist {
550 best_dist = s;
551 best_c = c;
552 }
553 }
554 dist[i] = best_dist;
555 if labels[i] != best_c {
556 labels[i] = best_c;
557 changed = true;
558 }
559 }
560 if !changed {
561 break;
562 }
563 let mut sums = vec![0.0; k * d];
564 let mut counts = vec![0usize; k];
565 for (i, &label) in labels.iter().enumerate() {
566 counts[label] += 1;
567 for (s, x) in sums[label * d..(label + 1) * d].iter_mut().zip(points.row(i)) {
568 *s += x;
569 }
570 }
571 let mut taken = vec![false; n];
573 for c in 0..k {
574 let centroid = &mut centroids[c * d..(c + 1) * d];
575 if counts[c] == 0 {
576 let far = (0..n)
577 .filter(|&i| !taken[i])
578 .max_by(|&a, &b| dist[a].total_cmp(&dist[b]).then(b.cmp(&a)))
579 .unwrap_or(0);
580 taken[far] = true;
581 centroid.copy_from_slice(points.row(far));
582 } else {
583 let inv = 1.0 / counts[c] as f64;
584 for (x, s) in centroid.iter_mut().zip(&sums[c * d..(c + 1) * d]) {
585 *x = s * inv;
586 }
587 }
588 }
589 }
590 labels
591}
592
593fn kmeans_labels(points: &Points, k: usize, rng: &mut StdRng) -> Result<Vec<usize>, OncError> {
595 if k < 2 || k > points.n {
596 return Err(OncError::ClusteringFailed);
597 }
598 Ok(lloyd(points, &kmeans_plus_plus(points, k, rng)))
599}
600
601fn silhouette_samples(data: &DMatrix<f64>, labels: &[usize]) -> Vec<f64> {
602 silhouette_from_pairwise(&pairwise_distances(&Points::new(data)), labels)
603}
604
605fn pairwise_distances(points: &Points) -> DMatrix<f64> {
607 let n = points.n;
608 let mut pairwise = DMatrix::zeros(n, n);
609 for i in 0..n {
610 for j in i..n {
611 let d = squared_distance(points.row(i), points.row(j)).sqrt();
612 pairwise[(i, j)] = d;
613 pairwise[(j, i)] = d;
614 }
615 }
616 pairwise
617}
618
619fn silhouette_from_pairwise(pairwise: &DMatrix<f64>, labels: &[usize]) -> Vec<f64> {
620 let n = pairwise.nrows();
621 let mut by_cluster: BTreeMap<usize, Vec<usize>> = BTreeMap::new();
622 for (i, lbl) in labels.iter().copied().enumerate() {
623 by_cluster.entry(lbl).or_default().push(i);
624 }
625
626 let mut scores = vec![0.0; n];
627 for i in 0..n {
628 let own = labels[i];
629 let own_members = &by_cluster[&own];
630
631 if own_members.len() <= 1 {
635 continue;
636 }
637
638 let a = own_members.iter().filter(|&&j| j != i).map(|&j| pairwise[(i, j)]).sum::<f64>()
639 / (own_members.len() - 1) as f64;
640
641 let mut b = f64::INFINITY;
642 for (cluster, members) in &by_cluster {
643 if *cluster == own || members.is_empty() {
644 continue;
645 }
646 let mut s = 0.0;
647 for &j in members {
648 s += pairwise[(i, j)];
649 }
650 let mean = s / members.len() as f64;
651 if mean < b {
652 b = mean;
653 }
654 }
655
656 scores[i] = if !b.is_finite() || (a == 0.0 && b == 0.0) { 0.0 } else { (b - a) / a.max(b) };
657 }
658
659 scores
660}
661
662#[cfg(test)]
663mod tests {
664 use super::*;
665
666 #[test]
667 fn silhouette_of_a_singleton_cluster_is_zero() {
668 let data = DMatrix::from_row_slice(3, 1, &[0.0, 0.1, 10.0]);
670 let scores = silhouette_samples(&data, &[0, 0, 1]);
671 assert!(scores[0] > 0.9 && scores[1] > 0.9);
672 assert_eq!(scores[2], 0.0);
673 }
674
675 #[test]
676 fn a_perfect_clustering_is_not_replaced_by_a_later_candidate() {
677 let corr = DMatrix::from_row_slice(
680 4,
681 4,
682 &[1.0, 0.9, 0.0, 0.0, 0.9, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.9, 0.0, 0.0, 0.9, 1.0],
683 );
684 let state =
685 cluster_kmeans_base(&corr, 3, 3, &mut StdRng::seed_from_u64(DEFAULT_SEED)).unwrap();
686 assert_eq!(state.clusters.len(), 2);
687 }
688}