Skip to main content

openquant/
onc.rs

1//! Optimal Number of Clusters (ONC): partition a correlation matrix with k-means, choosing the
2//! number of clusters by silhouette quality.
3//!
4//! Not from AFML. References: López de Prado, *Machine Learning for Asset Managers* (2020),
5//! Chapter 4, §4.4 (Snippet 4.1, base clustering; Snippet 4.2, higher-level clustering);
6//! López de Prado and Lewis (2019), *Detection of false investment strategies using
7//! unsupervised learning methods*; Rousseeuw (1987) for the silhouette.
8//!
9//! The algorithm:
10//! 1. Convert correlations to distances `d_ij = sqrt((1 - rho_ij) / 2)` (inputs clamped to
11//!    `[-1, 1]`) and represent each item by its row of that distance matrix.
12//! 2. Run k-means for every `k` from 2 to `max(N - 1, 2)`, `repeat` times each, and keep the
13//!    partition with the highest t-statistic of the silhouettes, `mean(S) / std(S)`; equal
14//!    t-statistics go to the higher mean silhouette. Each run is one k-means++ initialisation
15//!    (the greedy variant scikit-learn uses) followed by Lloyd's algorithm, as Snippet 4.1 runs
16//!    scikit-learn's `KMeans(n_init=1)` `n_init` times.
17//! 3. Compute that t-statistic per cluster. If more than two clusters score below the average,
18//!    pool their members, re-run the whole procedure on them, and keep the result only if its
19//!    mean cluster t-statistic beats that of the clusters it replaced
20//!    ([`check_improve_clusters`]).
21//!
22//! Conventions:
23//! - The input is an `N x N` correlation matrix, `N >= 2`; rows and columns are items in the
24//!   same order. Symmetry and a unit diagonal are not checked.
25//! - Negative correlation is distance, not similarity: `rho = -1` is maximally far apart. Take
26//!   absolute correlations first if a series and its mirror image should cluster together.
27//! - Member indices in [`OncResult::clusters`] and the order of
28//!   [`OncResult::silhouette_scores`] refer to the original row order.
29//! - The result has at least two clusters unless every row of the matrix is the same (for
30//!   example all ones), which gives one cluster of every item.
31//! - One random stream, seeded with [`DEFAULT_SEED`] (or the seed given to
32//!   [`get_onc_clusters_with_seed`]), drives every initialisation, so results are deterministic.
33//! - ONC is a random search: the partition kept is the best of `repeat` k-means runs per `k`,
34//!   and a partition with a high t-statistic may be a k-means local optimum that few
35//!   initialisations reach. On clean structure any seed finds the same answer. On real data it
36//!   need not: on the 30 breast-cancer features of `tests/fixtures/onc`, `repeat = 50` returns
37//!   one of a handful of partitions depending on the seed, all of them coarsenings of the same
38//!   eight groups. Compare a few seeds, and raise `repeat`, before reading much into a
39//!   particular partition. Cost grows at least as `N^3` (every `k`, `repeat` times, quadratic
40//!   silhouettes), plus the recursion.
41//!
42//! ```
43//! use nalgebra::DMatrix;
44//! use openquant::onc::{get_onc_clusters, OncError};
45//!
46//! # fn main() -> Result<(), OncError> {
47//! // Two blocks of three: 0.8 within a block, 0.1 across.
48//! let block = |i: usize| i / 3;
49//! let corr = DMatrix::from_fn(6, 6, |i, j| {
50//!     if i == j {
51//!         1.0
52//!     } else if block(i) == block(j) {
53//!         0.8
54//!     } else {
55//!         0.1
56//!     }
57//! });
58//!
59//! let result = get_onc_clusters(&corr, 3)?;
60//! let mut found: Vec<Vec<usize>> = result.clusters.values().cloned().collect();
61//! found.sort();
62//! assert_eq!(found, vec![vec![0, 1, 2], vec![3, 4, 5]]);
63//! assert_eq!(result.silhouette_scores.len(), 6);
64//! assert!(result.silhouette_scores.iter().all(|s| *s > 0.5));
65//! # Ok(())
66//! # }
67//! ```
68
69use crate::util::stats;
70use nalgebra::DMatrix;
71use rand::rngs::StdRng;
72use rand::{RngExt, SeedableRng};
73use std::collections::BTreeMap;
74
75/// Seed of the random stream [`get_onc_clusters`] uses.
76pub const DEFAULT_SEED: u64 = 42;
77
78/// Errors returned by [`get_onc_clusters`].
79#[derive(Debug, Clone, PartialEq, thiserror::Error)]
80pub enum OncError {
81    /// The correlation matrix is not square or has fewer than two rows.
82    #[error("the correlation matrix must be square with at least two rows")]
83    InvalidCorrelationMatrix,
84    /// `repeat` is zero.
85    #[error("repeat must be positive")]
86    InvalidRepeat,
87    /// No candidate partition could be selected. Not expected on a finite correlation matrix;
88    /// it can occur when `NaN` entries make every candidate's quality score `NaN`.
89    #[error("clustering failed to produce a partition")]
90    ClusteringFailed,
91}
92
93/// Partition found by [`get_onc_clusters`].
94#[derive(Debug, Clone)]
95pub struct OncResult {
96    /// The input correlation matrix with rows and columns permuted so that the members of each
97    /// cluster are contiguous, clusters in label order.
98    pub ordered_correlation: DMatrix<f64>,
99    /// Cluster label (`0..number of clusters`) to the indices of its members, in the original
100    /// row order of the input.
101    pub clusters: BTreeMap<usize, Vec<usize>>,
102    /// Silhouette score of every item, indexed by the original row order; a singleton
103    /// cluster's member scores 0.
104    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
114/// Keep the re-clustered partition only if its mean cluster t-stat beats the mean t-stat of the
115/// clusters that were re-clustered (MLAM Snippet 4.2); otherwise keep the old partition.
116///
117/// Returns `new_cluster` when `new_tstat_mean > mean_redo_tstat` and `old_cluster` otherwise
118/// (ties and `NaN` keep the old one). Exposed for parity with mlfinlab; [`get_onc_clusters`]
119/// calls it internally.
120///
121/// ```
122/// use openquant::onc::check_improve_clusters;
123///
124/// assert_eq!(check_improve_clusters(2.0, 1.5, "old", "new"), "new");
125/// assert_eq!(check_improve_clusters(1.5, 1.5, "old", "new"), "old");
126/// ```
127pub 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
140/// Partition the items of a correlation matrix with ONC (MLAM §4.4, Snippets 4.1–4.2).
141///
142/// `corr_mat` is an `N x N` correlation matrix (`N >= 2`; entries clamped to `[-1, 1]`,
143/// symmetry and unit diagonal not checked). `repeat` is the number of k-means initialisations
144/// per candidate `k`. The number of clusters is chosen by the silhouette t-statistic; see the
145/// [module documentation](self) for the full procedure. The search starts at `k = 2`, so a
146/// matrix with no structure still comes back partitioned: a low mean silhouette is the sign
147/// that the clusters are not real. The exception is a matrix whose rows are all identical,
148/// such as all ones: there is nothing to separate, every k-means centroid is the same point,
149/// and the result is one cluster of every item, each scoring a silhouette of 0.
150///
151/// # Errors
152///
153/// - [`OncError::InvalidRepeat`] if `repeat == 0`.
154/// - [`OncError::InvalidCorrelationMatrix`] if `corr_mat` is not square or has fewer than two
155///   rows.
156/// - [`OncError::ClusteringFailed`] if no candidate partition can be selected (only with
157///   `NaN` entries).
158///
159/// ```
160/// use nalgebra::DMatrix;
161/// use openquant::onc::{get_onc_clusters, OncError};
162///
163/// // Items 0, 2 and 4 move together, as do 1, 3 and 5.
164/// let corr = DMatrix::from_fn(6, 6, |i, j| {
165///     if i == j {
166///         1.0
167///     } else if i % 2 == j % 2 {
168///         0.9
169///     } else {
170///         0.0
171///     }
172/// });
173/// let result = get_onc_clusters(&corr, 2).unwrap();
174/// let mut found: Vec<Vec<usize>> = result.clusters.values().cloned().collect();
175/// found.sort();
176/// assert_eq!(found, vec![vec![0, 2, 4], vec![1, 3, 5]]);
177/// // The ordered matrix puts each block on the diagonal.
178/// let first = &result.clusters[&0];
179/// assert_eq!(result.ordered_correlation[(0, 1)], corr[(first[0], first[1])]);
180///
181/// assert_eq!(get_onc_clusters(&corr, 0).unwrap_err(), OncError::InvalidRepeat);
182/// assert_eq!(
183///     get_onc_clusters(&DMatrix::from_element(1, 1, 1.0), 1).unwrap_err(),
184///     OncError::InvalidCorrelationMatrix
185/// );
186/// ```
187pub 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
191/// [`get_onc_clusters`] with the seed of its random stream given explicitly.
192///
193/// `get_onc_clusters(corr, repeat)` is `get_onc_clusters_with_seed(corr, repeat, DEFAULT_SEED)`.
194/// The seed drives every k-means++ initialisation, including those of the re-clustering step, so
195/// a given `(corr_mat, repeat, seed)` always gives the same partition. Running a few seeds is
196/// the way to check whether a partition is a property of the data or of one random stream: see
197/// the [module documentation](self) on how the answer depends on `repeat`.
198///
199/// # Errors
200///
201/// As [`get_onc_clusters`].
202///
203/// ```
204/// use nalgebra::DMatrix;
205/// use openquant::onc::{get_onc_clusters, get_onc_clusters_with_seed, DEFAULT_SEED};
206///
207/// let block = |i: usize| i / 4;
208/// let corr = DMatrix::from_fn(12, 12, |i, j| {
209///     if i == j {
210///         1.0
211///     } else if block(i) == block(j) {
212///         0.7
213///     } else {
214///         0.1
215///     }
216/// });
217/// let default = get_onc_clusters(&corr, 5).unwrap();
218/// let seeded = get_onc_clusters_with_seed(&corr, 5, DEFAULT_SEED).unwrap();
219/// assert_eq!(seeded.clusters, default.clusters);
220/// // Clean blocks come back the same under any seed.
221/// for seed in 0..5 {
222///     let result = get_onc_clusters_with_seed(&corr, 5, seed).unwrap();
223///     let mut found: Vec<Vec<usize>> = result.clusters.values().cloned().collect();
224///     found.sort();
225///     assert_eq!(found, vec![vec![0, 1, 2, 3], vec![4, 5, 6, 7], vec![8, 9, 10, 11]]);
226/// }
227/// ```
228pub 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            // A perfect clustering has zero silhouette variance, so its t-stat is +inf and only
379            // another +inf may replace it. Only a NaN incumbent is replaced unconditionally.
380            // Equal t-stats go to the higher mean silhouette: with equal-sized exact blocks,
381            // merging whole blocks also gives every item the same silhouette (+inf), and which
382            // of the two k-means happened to find first must not decide.
383            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    // Population deviation (ddof = 0).
419    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
466/// Items as rows of a row-major buffer, for the k-means inner loops.
467struct 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
489/// k-means++ seeding (Arthur and Vassilvitskii 2007) in its greedy form, as scikit-learn does:
490/// each new centre is the best of `2 + ln k` candidates drawn with probability proportional to
491/// the squared distance to the nearest centre chosen so far. Returns the centres' row indices.
492fn 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                // Every point sits on a centre already: any choice is as good.
515                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
535/// Lloyd's algorithm from the given centres, until no label changes (at most 300 passes).
536fn 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        // An empty cluster takes the point farthest from its centroid, as in scikit-learn.
572        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
593/// One k-means run: k-means++ seeding, then Lloyd's algorithm.
594fn 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
605/// Euclidean distances between the rows.
606fn 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        // A point alone in its cluster has no intra-cluster distance; its silhouette is 0 by
632        // definition (Rousseeuw 1987; scikit-learn does the same). Scoring it (b - 0) / b = 1
633        // makes "every point its own cluster" look like the perfect clustering.
634        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        // Points 0 and 1 sit together, point 2 is alone and far away.
669        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        // Two identical pairs: k = 2 is perfect (every silhouette is 1, zero variance, t-stat
678        // +inf). The search goes on to try k = 3 and must keep k = 2.
679        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}