stats_claw/algorithms/clustering/
hierarchical.rs1use crate::algorithms::{count_to_f64, euclidean_sq};
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum Linkage {
14 Ward,
17 Single,
19 Complete,
21 Average,
23}
24
25#[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
67struct State {
70 members: Vec<usize>,
72 sizes: Vec<usize>,
74 active: Vec<bool>,
76 dist: Vec<f64>,
78 n: usize,
80}
81
82impl State {
83 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 fn active_count(&self) -> usize {
105 self.active.iter().filter(|&&a| a).count()
106 }
107
108 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 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 fn flat_labels(&self) -> Vec<usize> {
161 super::relabel_contiguous(&self.members)
162 }
163}
164
165fn 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
178fn 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
209fn active_at(active: &[bool], i: usize) -> bool {
211 active.get(i).copied().unwrap_or(false)
212}
213
214fn get_dist(dist: &[f64], n: usize, i: usize, j: usize) -> f64 {
216 dist.get(i * n + j).copied().unwrap_or(f64::INFINITY)
217}
218
219fn 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}