qualia_core_db/solvers/learning/clustering/
hierarchical.rs1use crate::solvers::learning::LearningError;
6
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum Linkage {
10 Single,
12 Complete,
14 Average,
16}
17
18#[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 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 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 merges.push((clusters[bi][0], clusters[bj][0]));
78 let moved = clusters.remove(bj);
80 clusters[bi].extend(moved);
81 }
82 Ok(Self { merges, n })
83 }
84
85 pub fn labels(&self, k: usize) -> Vec<usize> {
89 let k = k.clamp(1, self.n);
90 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 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 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 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]; 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); let one = h.labels(1);
157 assert!(one.iter().all(|&l| l == 0)); }
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}