Skip to main content

stats_claw/algorithms/clustering/
hierarchical.rs

1//! Agglomerative (bottom-up) hierarchical clustering.
2//!
3//! Every point starts in its own cluster; the two closest clusters are merged
4//! repeatedly until `k` remain. Inter-cluster distances are maintained by the
5//! Lance–Williams recurrence, so a single update rule covers Ward, single,
6//! complete, and average linkage. The result is a flat `k`-cluster labelling
7//! compared to `scikit-learn`'s `AgglomerativeClustering` by adjusted Rand index.
8
9use crate::algorithms::{count_to_f64, euclidean_sq};
10
11/// Inter-cluster distance rule driving the agglomerative merges.
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum Linkage {
14    /// Minimises the increase in total within-cluster variance (`scikit-learn`'s
15    /// default). Operates on squared Euclidean distances.
16    Ward,
17    /// Distance between the two nearest members of the clusters.
18    Single,
19    /// Distance between the two farthest members of the clusters.
20    Complete,
21    /// Mean pairwise distance between members of the two clusters (UPGMA).
22    Average,
23}
24
25/// Clusters `data` into `k` groups by agglomerative merging under `linkage`.
26///
27/// Deterministic — ties are broken by the lower cluster index, so repeated runs
28/// match. Returns one label per point; an empty input or `k == 0` yields an empty
29/// labelling, and `k` larger than the point count leaves every point singleton.
30///
31/// # Arguments
32///
33/// * `data` — observations; each inner slice is one point of equal dimension.
34/// * `k` — number of flat clusters to cut the dendrogram into.
35/// * `linkage` — the inter-cluster distance rule (see [`Linkage`]).
36///
37/// # Returns
38///
39/// A label vector with contiguous ids `0, 1, …` (no noise label).
40///
41/// # Examples
42///
43/// ```
44/// use stats_claw::algorithms::clustering::{agglomerative, Linkage};
45///
46/// let data = vec![vec![0.0], vec![0.1], vec![9.0], vec![9.1]];
47/// let labels = agglomerative(&data, 2, Linkage::Ward);
48/// assert_eq!(labels.first(), labels.get(1), "near points split");
49/// assert_ne!(labels.first(), labels.get(2), "far points merged");
50/// ```
51#[must_use]
52pub fn agglomerative(data: &[Vec<f64>], k: usize, linkage: Linkage) -> Vec<usize> {
53    let n = data.len();
54    if n == 0 || k == 0 {
55        return Vec::new();
56    }
57    let mut state = State::new(data, linkage);
58    while state.active_count() > k.min(n) {
59        let Some((a, b)) = state.closest_pair() else {
60            break;
61        };
62        state.merge(a, b, linkage);
63    }
64    state.flat_labels()
65}
66
67/// Mutable agglomeration state: cluster membership, sizes, and the working
68/// pairwise-distance matrix updated in place via Lance–Williams.
69struct State {
70    /// Cluster id assigned to each point (`members[i]` = cluster of point `i`).
71    members: Vec<usize>,
72    /// Sizes of every cluster slot (0 once a slot has been merged away).
73    sizes: Vec<usize>,
74    /// Whether each cluster slot is still an active cluster.
75    active: Vec<bool>,
76    /// Row-major distance matrix between cluster slots (`dist[i*n + j]`).
77    dist: Vec<f64>,
78    /// Number of original points (= number of slots).
79    n: usize,
80}
81
82impl State {
83    /// Initialises one singleton cluster per point with the base distance matrix.
84    /// Ward stores squared distances; the other linkages store plain distances.
85    fn new(data: &[Vec<f64>], linkage: Linkage) -> Self {
86        let n = data.len();
87        let mut dist = vec![0.0_f64; n * n];
88        for i in 0..n {
89            for j in (i + 1)..n {
90                let base = base_distance(data, i, j, linkage);
91                set_dist(&mut dist, n, i, j, base);
92            }
93        }
94        Self {
95            members: (0..n).collect(),
96            sizes: vec![1; n],
97            active: vec![true; n],
98            dist,
99            n,
100        }
101    }
102
103    /// Number of clusters still active.
104    fn active_count(&self) -> usize {
105        self.active.iter().filter(|&&a| a).count()
106    }
107
108    /// Finds the closest active cluster pair `(a, b)` with `a < b`.
109    fn closest_pair(&self) -> Option<(usize, usize)> {
110        let mut best: Option<(usize, usize)> = None;
111        let mut best_d = f64::INFINITY;
112        for i in 0..self.n {
113            if !active_at(&self.active, i) {
114                continue;
115            }
116            for j in (i + 1)..self.n {
117                if !active_at(&self.active, j) {
118                    continue;
119                }
120                let d = get_dist(&self.dist, self.n, i, j);
121                if d < best_d {
122                    best_d = d;
123                    best = Some((i, j));
124                }
125            }
126        }
127        best
128    }
129
130    /// Merges cluster `b` into `a`, updating distances to all other clusters by the
131    /// Lance–Williams recurrence for `linkage`.
132    fn merge(&mut self, a: usize, b: usize, linkage: Linkage) {
133        let size_a = self.sizes.get(a).copied().unwrap_or(0);
134        let size_b = self.sizes.get(b).copied().unwrap_or(0);
135        for other in 0..self.n {
136            if other == a || other == b || !active_at(&self.active, other) {
137                continue;
138            }
139            let to_a = get_dist(&self.dist, self.n, a, other);
140            let to_b = get_dist(&self.dist, self.n, b, other);
141            let between = get_dist(&self.dist, self.n, a, b);
142            let size_i = self.sizes.get(other).copied().unwrap_or(0);
143            let updated = lance_williams(linkage, to_a, to_b, between, size_a, size_b, size_i);
144            set_dist(&mut self.dist, self.n, a, other, updated);
145        }
146        if let Some(slot) = self.sizes.get_mut(a) {
147            *slot = size_a + size_b;
148        }
149        if let Some(slot) = self.active.get_mut(b) {
150            *slot = false;
151        }
152        for m in &mut self.members {
153            if *m == b {
154                *m = a;
155            }
156        }
157    }
158
159    /// Renumbers the surviving cluster ids to a contiguous `0..k` labelling.
160    fn flat_labels(&self) -> Vec<usize> {
161        super::relabel_contiguous(&self.members)
162    }
163}
164
165/// Base distance between singleton points `i` and `j` (squared for Ward, plain
166/// Euclidean otherwise).
167fn base_distance(data: &[Vec<f64>], i: usize, j: usize, linkage: Linkage) -> f64 {
168    let (Some(pi), Some(pj)) = (data.get(i), data.get(j)) else {
169        return f64::INFINITY;
170    };
171    let sq = euclidean_sq(pi, pj);
172    match linkage {
173        Linkage::Ward => sq,
174        _ => sq.sqrt(),
175    }
176}
177
178/// Lance–Williams update for the distance between a merged cluster `a∪b` and a
179/// third cluster `i`, given the prior pairwise distances and cluster sizes.
180fn lance_williams(
181    linkage: Linkage,
182    to_a: f64,
183    to_b: f64,
184    between: f64,
185    size_a: usize,
186    size_b: usize,
187    size_i: usize,
188) -> f64 {
189    match linkage {
190        Linkage::Single => to_a.min(to_b),
191        Linkage::Complete => to_a.max(to_b),
192        Linkage::Average => {
193            let (na, nb) = (count_to_f64(size_a), count_to_f64(size_b));
194            na.mul_add(to_a, nb * to_b) / (na + nb)
195        }
196        Linkage::Ward => {
197            let (na, nb, ni) = (
198                count_to_f64(size_a),
199                count_to_f64(size_b),
200                count_to_f64(size_i),
201            );
202            let total = na + nb + ni;
203            let weighted = (na + ni).mul_add(to_a, (nb + ni) * to_b);
204            ni.mul_add(-between, weighted) / total
205        }
206    }
207}
208
209/// Reads the active flag for slot `i`, defaulting to `false` out of range.
210fn active_at(active: &[bool], i: usize) -> bool {
211    active.get(i).copied().unwrap_or(false)
212}
213
214/// Reads the symmetric distance between slots `i` and `j`.
215fn get_dist(dist: &[f64], n: usize, i: usize, j: usize) -> f64 {
216    dist.get(i * n + j).copied().unwrap_or(f64::INFINITY)
217}
218
219/// Writes the distance between slots `i` and `j` symmetrically.
220fn set_dist(dist: &mut [f64], n: usize, i: usize, j: usize, value: f64) {
221    if let Some(slot) = dist.get_mut(i * n + j) {
222        *slot = value;
223    }
224    if let Some(slot) = dist.get_mut(j * n + i) {
225        *slot = value;
226    }
227}
228
229#[cfg(test)]
230mod tests {
231    use super::*;
232
233    #[test]
234    fn ward_separates_two_far_pairs() {
235        let data = vec![vec![0.0], vec![0.1], vec![9.0], vec![9.1]];
236        let labels = agglomerative(&data, 2, Linkage::Ward);
237        assert_eq!(labels.first(), labels.get(1), "near pair split");
238        assert_ne!(labels.first(), labels.get(2), "far pair merged");
239    }
240
241    #[test]
242    fn single_linkage_matches_ward_on_separated_blobs() {
243        let data = vec![vec![0.0], vec![0.2], vec![5.0], vec![5.2]];
244        let single = agglomerative(&data, 2, Linkage::Single);
245        assert_eq!(single.first(), single.get(1), "near pair split");
246        assert_ne!(single.first(), single.get(2), "far pair merged");
247    }
248
249    #[test]
250    fn deterministic_for_fixed_inputs() {
251        let data = vec![vec![0.0], vec![1.0], vec![10.0], vec![11.0]];
252        assert_eq!(
253            agglomerative(&data, 2, Linkage::Average),
254            agglomerative(&data, 2, Linkage::Average)
255        );
256    }
257}