Skip to main content

qualia_core_db/solvers/learning/clustering/
hierarchical.rs

1//! Agglomerative hierarchical clustering (ISL ch 12.4.2) — bottom-up merging of the
2//! two closest clusters under a linkage rule, producing a dendrogram that can be cut
3//! into any number of clusters. Kernel-class `AllPairs` (the cluster distances).
4
5use crate::solvers::learning::LearningError;
6
7/// How the distance between two clusters is defined from member point distances.
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum Linkage {
10    /// Nearest pair (min).
11    Single,
12    /// Farthest pair (max).
13    Complete,
14    /// Mean over all cross pairs.
15    Average,
16}
17
18/// A fitted dendrogram: the sequence of `n−1` merges, each recorded as two
19/// representative original-point indices.
20#[derive(Debug, Clone)]
21pub struct Hierarchical {
22    merges: Vec<(usize, usize)>,
23    n: usize,
24}
25
26fn sq_dist(x: &[f64], p: usize, i: usize, j: usize) -> f64 {
27    let a = &x[i * p..(i + 1) * p];
28    let b = &x[j * p..(j + 1) * p];
29    a.iter().zip(b).map(|(u, v)| (u - v) * (u - v)).sum()
30}
31
32fn cluster_distance(x: &[f64], p: usize, ca: &[usize], cb: &[usize], linkage: Linkage) -> f64 {
33    let mut acc = match linkage {
34        Linkage::Single => f64::INFINITY,
35        Linkage::Complete => f64::NEG_INFINITY,
36        Linkage::Average => 0.0,
37    };
38    for &i in ca {
39        for &j in cb {
40            let d = sq_dist(x, p, i, j);
41            match linkage {
42                Linkage::Single => acc = acc.min(d),
43                Linkage::Complete => acc = acc.max(d),
44                Linkage::Average => acc += d,
45            }
46        }
47    }
48    if linkage == Linkage::Average {
49        acc /= (ca.len() * cb.len()) as f64;
50    }
51    acc
52}
53
54impl Hierarchical {
55    /// Build the full dendrogram for a row-major `n × p` matrix under `linkage`.
56    pub fn fit(x: &[f64], n: usize, p: usize, linkage: Linkage) -> Result<Self, LearningError> {
57        if n == 0 || p == 0 || x.len() != n * p {
58            return Err(LearningError::InvalidDimension);
59        }
60        let mut clusters: Vec<Vec<usize>> = (0..n).map(|i| vec![i]).collect();
61        let mut merges = Vec::with_capacity(n.saturating_sub(1));
62        while clusters.len() > 1 {
63            // Find the closest pair of active clusters.
64            let mut best = f64::INFINITY;
65            let (mut bi, mut bj) = (0, 1);
66            for a in 0..clusters.len() {
67                for b in (a + 1)..clusters.len() {
68                    let d = cluster_distance(x, p, &clusters[a], &clusters[b], linkage);
69                    if d < best {
70                        best = d;
71                        bi = a;
72                        bj = b;
73                    }
74                }
75            }
76            // Record the merge by a representative point from each cluster.
77            merges.push((clusters[bi][0], clusters[bj][0]));
78            // Merge bj into bi, then drop bj.
79            let moved = clusters.remove(bj);
80            clusters[bi].extend(moved);
81        }
82        Ok(Self { merges, n })
83    }
84
85    /// Cut the dendrogram into `k` clusters and return a label per original point
86    /// (labels are `0..k`, assigned in first-appearance order). `k` is clamped to
87    /// `1..=n`.
88    pub fn labels(&self, k: usize) -> Vec<usize> {
89        let k = k.clamp(1, self.n);
90        // Union-find over the first (n − k) merges.
91        let mut parent: Vec<usize> = (0..self.n).collect();
92        fn find(parent: &mut [usize], mut x: usize) -> usize {
93            while parent[x] != x {
94                parent[x] = parent[parent[x]];
95                x = parent[x];
96            }
97            x
98        }
99        let n_unions = self.n.saturating_sub(k);
100        for &(a, b) in self.merges.iter().take(n_unions) {
101            let ra = find(&mut parent, a);
102            let rb = find(&mut parent, b);
103            if ra != rb {
104                parent[ra] = rb;
105            }
106        }
107        // Relabel roots to 0..k in first-appearance order.
108        let mut label_of = std::collections::HashMap::new();
109        let mut next = 0;
110        let mut out = vec![0usize; self.n];
111        for i in 0..self.n {
112            let r = find(&mut parent, i);
113            let l = *label_of.entry(r).or_insert_with(|| {
114                let v = next;
115                next += 1;
116                v
117            });
118            out[i] = l;
119        }
120        out
121    }
122
123    pub fn n_merges(&self) -> usize {
124        self.merges.len()
125    }
126}
127
128#[cfg(test)]
129mod tests {
130    use super::*;
131
132    #[test]
133    fn separates_two_obvious_groups() {
134        // Two tight groups far apart → cutting into 2 clusters splits them.
135        let x = [
136            0.0, 0.0, 0.2, 0.1, -0.1, 0.2, 10.0, 10.0, 10.2, 9.9, 9.8, 10.1,
137        ];
138        let h = Hierarchical::fit(&x, 6, 2, Linkage::Average).unwrap();
139        assert_eq!(h.n_merges(), 5);
140        let labels = h.labels(2);
141        // First 3 share a label; last 3 share the other.
142        assert!(labels[0] == labels[1] && labels[1] == labels[2]);
143        assert!(labels[3] == labels[4] && labels[4] == labels[5]);
144        assert_ne!(labels[0], labels[3]);
145    }
146
147    #[test]
148    fn k_equals_n_is_all_singletons_and_k_one_is_all_together() {
149        let x = [0.0, 1.0, 2.0, 3.0]; // 4 points in 1-D
150        let h = Hierarchical::fit(&x, 4, 1, Linkage::Single).unwrap();
151        let singletons = h.labels(4);
152        let mut sorted = singletons.clone();
153        sorted.sort_unstable();
154        sorted.dedup();
155        assert_eq!(sorted.len(), 4); // every point its own cluster
156        let one = h.labels(1);
157        assert!(one.iter().all(|&l| l == 0)); // all together
158    }
159
160    #[test]
161    fn complete_and_single_linkage_both_run() {
162        let x = [0.0, 0.0, 1.0, 1.0, 5.0, 5.0, 6.0, 6.0];
163        for linkage in [Linkage::Single, Linkage::Complete, Linkage::Average] {
164            let h = Hierarchical::fit(&x, 4, 2, linkage).unwrap();
165            let labels = h.labels(2);
166            assert_eq!(labels[0], labels[1]);
167            assert_eq!(labels[2], labels[3]);
168            assert_ne!(labels[0], labels[2]);
169        }
170    }
171
172    #[test]
173    fn guards() {
174        assert_eq!(
175            Hierarchical::fit(&[1.0, 2.0, 3.0], 2, 2, Linkage::Single).unwrap_err(),
176            LearningError::InvalidDimension
177        );
178    }
179}