Skip to main content

max / audiofiles

13.4 KB · 386 lines History Blame Raw
1 //! Unsupervised clustering bootstrap (cold-start for the tag pipeline).
2 //!
3 //! Groups a library's persisted feature vectors so the user can name clusters and seed
4 //! the first labels before any rules or manual tagging exist. k-means runs over
5 //! z-standardized features (MFCC and spectral scales differ by orders of magnitude) with
6 //! deterministic farthest-first seeding, no RNG dependency, reproducible results.
7 //!
8 //! Clustering itself writes nothing; `apply_cluster_tag` is what a named cluster does,
9 //! recording `source = 'cluster'` provenance (sticky, the rules engine never reconciles
10 //! these away). See [`crate::rules`] for the provenance model.
11
12 use crate::db::Database;
13 use crate::error::Result;
14 use tracing::instrument;
15
16 use super::classify::{FEATURE_VERSION, NUM_FEATURES};
17
18 /// One cluster: a representative (medoid) sample plus its members.
19 #[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
20 pub struct Cluster {
21 pub id: usize,
22 /// Member nearest the centroid, a playable representative for naming.
23 pub medoid_hash: String,
24 pub member_hashes: Vec<String>,
25 }
26
27 /// Result of clustering a library.
28 #[derive(Debug, Clone)]
29 pub struct ClusterResult {
30 pub clusters: Vec<Cluster>,
31 /// Feature-extractor version the vectors were produced with.
32 pub feature_version: u32,
33 }
34
35 /// Suggested cluster count for `n` samples: ~sqrt(n/2), clamped to a sane UI range.
36 pub fn suggest_k(n: usize) -> usize {
37 if n < 2 {
38 return n;
39 }
40 let k = ((n as f64 / 2.0).sqrt()).round() as usize;
41 k.clamp(2, 24).min(n)
42 }
43
44 /// Load all current-version feature vectors as `(hash, vector)` pairs.
45 fn load_vectors(db: &Database) -> Result<Vec<(String, Vec<f64>)>> {
46 let mut stmt = db
47 .conn()
48 .prepare("SELECT hash, vector FROM sample_features WHERE feat_version = ?1")?;
49 let rows = stmt.query_map([FEATURE_VERSION], |row| {
50 Ok((row.get::<_, String>(0)?, row.get::<_, String>(1)?))
51 })?;
52
53 let mut out = Vec::new();
54 for r in rows {
55 let (hash, json) = r?;
56 // A vector that carried a non-finite component was serialized to JSON
57 // `null` (serde_json's non-finite behavior), which won't parse back into
58 // `Vec<f64>`. Skip such a row rather than `?`-aborting the WHOLE clustering
59 // run, one crafted sample must not break clustering for the entire
60 // library. Also drop wrong-dimension / any-non-finite vectors, mirroring
61 // the sibling guards in exemplar.rs / similarity.rs.
62 let Ok(v) = serde_json::from_str::<Vec<f64>>(&json) else {
63 continue;
64 };
65 if v.len() == NUM_FEATURES && v.iter().all(|x| x.is_finite()) {
66 out.push((hash, v));
67 }
68 }
69 Ok(out)
70 }
71
72 /// z-standardize each dimension across the dataset (in place). Zero-variance dims
73 /// collapse to 0.0. Returns nothing, operates on the supplied matrix.
74 fn standardize(vectors: &mut [Vec<f64>]) {
75 if vectors.is_empty() {
76 return;
77 }
78 let dims = vectors[0].len();
79 let n = vectors.len() as f64;
80 for d in 0..dims {
81 let mean = vectors.iter().map(|v| v[d]).sum::<f64>() / n;
82 let var = vectors.iter().map(|v| (v[d] - mean).powi(2)).sum::<f64>() / n;
83 let std = var.sqrt();
84 if std > f64::EPSILON {
85 for v in vectors.iter_mut() {
86 v[d] = (v[d] - mean) / std;
87 }
88 } else {
89 for v in vectors.iter_mut() {
90 v[d] = 0.0;
91 }
92 }
93 }
94 }
95
96 fn sq_dist(a: &[f64], b: &[f64]) -> f64 {
97 a.iter().zip(b).map(|(x, y)| (x - y).powi(2)).sum()
98 }
99
100 /// Deterministic farthest-first seeding (Gonzalez): start from the point nearest the
101 /// dataset mean, then repeatedly take the point maximizing its distance to the chosen
102 /// set. Ties break by lowest index.
103 fn init_centroids(points: &[Vec<f64>], k: usize) -> Vec<Vec<f64>> {
104 let dims = points[0].len();
105 let n = points.len() as f64;
106 let mean: Vec<f64> = (0..dims)
107 .map(|d| points.iter().map(|p| p[d]).sum::<f64>() / n)
108 .collect();
109
110 let first = (0..points.len())
111 .min_by(|&a, &b| sq_dist(&points[a], &mean).total_cmp(&sq_dist(&points[b], &mean)))
112 .expect("points is non-empty (guaranteed by caller)");
113
114 let mut centroids = vec![points[first].clone()];
115 let mut min_d: Vec<f64> = points.iter().map(|p| sq_dist(p, &points[first])).collect();
116
117 while centroids.len() < k {
118 let next = (0..points.len())
119 .max_by(|&a, &b| min_d[a].total_cmp(&min_d[b]))
120 .expect("points is non-empty (guaranteed by caller)");
121 centroids.push(points[next].clone());
122 for (i, p) in points.iter().enumerate() {
123 min_d[i] = min_d[i].min(sq_dist(p, &points[next]));
124 }
125 }
126 centroids
127 }
128
129 /// Cluster the library's feature vectors into `k` groups.
130 ///
131 /// Returns clusters with member hashes and a medoid representative. `k` is clamped to
132 /// the number of vectors. Returns an empty result if nothing is analyzed yet.
133 #[instrument(skip_all)]
134 pub fn cluster_library(db: &Database, k: usize, max_iters: usize) -> Result<ClusterResult> {
135 let loaded = load_vectors(db)?;
136 if loaded.is_empty() || k == 0 {
137 return Ok(ClusterResult {
138 clusters: Vec::new(),
139 feature_version: FEATURE_VERSION,
140 });
141 }
142 let hashes: Vec<String> = loaded.iter().map(|(h, _)| h.clone()).collect();
143 let mut points: Vec<Vec<f64>> = loaded.into_iter().map(|(_, v)| v).collect();
144 standardize(&mut points);
145
146 let k = k.min(points.len());
147 let mut centroids = init_centroids(&points, k);
148 let mut assign = vec![0usize; points.len()];
149
150 for _ in 0..max_iters.max(1) {
151 // Assignment step.
152 let mut changed = false;
153 for (i, p) in points.iter().enumerate() {
154 let best = (0..k)
155 .min_by(|&a, &b| sq_dist(p, &centroids[a]).total_cmp(&sq_dist(p, &centroids[b])))
156 .expect("k >= 1 (empty/zero-k inputs return early)");
157 if best != assign[i] {
158 assign[i] = best;
159 changed = true;
160 }
161 }
162 if !changed {
163 break;
164 }
165 // Update step. Empty clusters keep their previous centroid (no NaN drift).
166 let dims = points[0].len();
167 let mut sums = vec![vec![0.0f64; dims]; k];
168 let mut counts = vec![0usize; k];
169 for (i, p) in points.iter().enumerate() {
170 let c = assign[i];
171 counts[c] += 1;
172 for d in 0..dims {
173 sums[c][d] += p[d];
174 }
175 }
176 for c in 0..k {
177 if counts[c] > 0 {
178 for d in 0..dims {
179 centroids[c][d] = sums[c][d] / counts[c] as f64;
180 }
181 }
182 }
183 }
184
185 // Build clusters + medoids.
186 let mut members: Vec<Vec<usize>> = vec![Vec::new(); k];
187 for (i, &c) in assign.iter().enumerate() {
188 members[c].push(i);
189 }
190 let mut clusters = Vec::new();
191 for (c, idxs) in members.into_iter().enumerate() {
192 if idxs.is_empty() {
193 continue;
194 }
195 let medoid = *idxs
196 .iter()
197 .min_by(|&&a, &&b| {
198 sq_dist(&points[a], &centroids[c]).total_cmp(&sq_dist(&points[b], &centroids[c]))
199 })
200 .expect("idxs is non-empty (empty clusters are skipped above)");
201 clusters.push(Cluster {
202 id: c,
203 medoid_hash: hashes[medoid].clone(),
204 member_hashes: idxs.into_iter().map(|i| hashes[i].clone()).collect(),
205 });
206 }
207
208 Ok(ClusterResult {
209 clusters,
210 feature_version: FEATURE_VERSION,
211 })
212 }
213
214 /// Name a cluster: apply `tag` to every member with `source = 'cluster'` provenance.
215 /// Returns the number of samples that gained the tag.
216 #[instrument(skip_all)]
217 pub fn apply_cluster_tag(db: &Database, member_hashes: &[String], tag: &str) -> Result<usize> {
218 // One transaction for the whole cluster, not an autocommit per member.
219 db.transaction_core(|tx| {
220 let mut added = 0;
221 for hash in member_hashes {
222 if crate::rules::apply_tag_sourced(db, tx, hash, tag, "cluster")? {
223 added += 1;
224 }
225 }
226 Ok(added)
227 })
228 }
229
230 #[cfg(test)]
231 mod tests {
232 use super::*;
233
234 fn db_with_features(rows: &[(&str, [f64; NUM_FEATURES])]) -> Database {
235 let db = Database::open_in_memory().unwrap();
236 for (hash, vec) in rows {
237 db.conn()
238 .execute(
239 "INSERT INTO samples (hash, original_name, file_extension, file_size, import_date, last_modified) \
240 VALUES (?1, ?1, 'wav', 1, 0, 0)",
241 [hash],
242 )
243 .unwrap();
244 let json = serde_json::to_string(&vec.to_vec()).unwrap();
245 db.conn()
246 .execute(
247 "INSERT INTO sample_features (hash, feat_version, vector, computed_at) VALUES (?1, ?2, ?3, 0)",
248 rusqlite::params![hash, FEATURE_VERSION, json],
249 )
250 .unwrap();
251 }
252 db
253 }
254
255 fn vec_at(base: f64) -> [f64; NUM_FEATURES] {
256 [base; NUM_FEATURES]
257 }
258
259 #[test]
260 fn suggest_k_is_sane() {
261 assert_eq!(suggest_k(0), 0);
262 assert_eq!(suggest_k(1), 1);
263 assert!((2..=24).contains(&suggest_k(100)));
264 assert_eq!(suggest_k(3), 2); // sqrt(1.5) rounds to 1, clamped up to 2
265 }
266
267 #[test]
268 fn empty_library_yields_no_clusters() {
269 let db = Database::open_in_memory().unwrap();
270 let res = cluster_library(&db, 4, 20).unwrap();
271 assert!(res.clusters.is_empty());
272 }
273
274 #[test]
275 fn separates_two_obvious_groups() {
276 // Two tight groups far apart in feature space.
277 let rows = [
278 ("a1", vec_at(0.0)),
279 ("a2", vec_at(0.1)),
280 ("a3", vec_at(-0.1)),
281 ("b1", vec_at(100.0)),
282 ("b2", vec_at(100.1)),
283 ("b3", vec_at(99.9)),
284 ];
285 let db = db_with_features(&rows);
286 let res = cluster_library(&db, 2, 50).unwrap();
287 assert_eq!(res.clusters.len(), 2);
288
289 // Each cluster must be all-a or all-b.
290 for cluster in &res.clusters {
291 let all_a = cluster.member_hashes.iter().all(|h| h.starts_with('a'));
292 let all_b = cluster.member_hashes.iter().all(|h| h.starts_with('b'));
293 assert!(
294 all_a || all_b,
295 "cluster mixed groups: {:?}",
296 cluster.member_hashes
297 );
298 assert_eq!(cluster.member_hashes.len(), 3);
299 }
300 }
301
302 #[test]
303 fn deterministic_across_runs() {
304 let rows = [
305 ("a1", vec_at(0.0)),
306 ("b1", vec_at(50.0)),
307 ("c1", vec_at(100.0)),
308 ("a2", vec_at(1.0)),
309 ("b2", vec_at(51.0)),
310 ];
311 let db = db_with_features(&rows);
312 let r1 = cluster_library(&db, 3, 50).unwrap();
313 let r2 = cluster_library(&db, 3, 50).unwrap();
314 let sig = |r: &ClusterResult| {
315 let mut v: Vec<Vec<String>> = r
316 .clusters
317 .iter()
318 .map(|c| {
319 let mut m = c.member_hashes.clone();
320 m.sort();
321 m
322 })
323 .collect();
324 v.sort();
325 v
326 };
327 assert_eq!(sig(&r1), sig(&r2));
328 }
329
330 #[test]
331 fn apply_cluster_tag_writes_sticky_provenance() {
332 let rows = [("a1", vec_at(0.0)), ("a2", vec_at(0.1))];
333 let db = db_with_features(&rows);
334 let members = vec!["a1".to_string(), "a2".to_string()];
335 assert_eq!(apply_cluster_tag(&db, &members, "kit.my-kicks").unwrap(), 2);
336
337 assert!(
338 crate::tags::get_sample_tags(&db, "a1")
339 .unwrap()
340 .contains(&"kit.my-kicks".to_string())
341 );
342 let prov = crate::rules::sample_tag_provenance(&db, "a1").unwrap();
343 assert_eq!(prov[0].1, "cluster");
344
345 // Removable as a batch (undo the pass).
346 let removed = crate::rules::remove_tags_by_source(&db, "cluster").unwrap();
347 assert_eq!(removed, 2);
348 assert!(crate::tags::get_sample_tags(&db, "a1").unwrap().is_empty());
349 }
350
351 #[test]
352 fn load_vectors_drops_non_finite_vectors() {
353 let mut nan_vec = vec_at(0.5);
354 nan_vec[0] = f64::NAN;
355 let mut inf_vec = vec_at(0.5);
356 inf_vec[3] = f64::INFINITY;
357 let db = db_with_features(&[("good", vec_at(0.2)), ("nan", nan_vec), ("inf", inf_vec)]);
358
359 let loaded = load_vectors(&db).unwrap();
360 // Only the finite vector survives; NaN/Inf rows are dropped like the
361 // sibling `exemplar::is_usable_vector` guard does.
362 let hashes: Vec<&str> = loaded.iter().map(|(h, _)| h.as_str()).collect();
363 assert_eq!(hashes, vec!["good"]);
364 }
365
366 #[test]
367 fn clustering_ignores_non_finite_and_does_not_nan() {
368 // A poisoned vector must not garbage the medoids: the run completes and
369 // only clusters the finite members.
370 let mut bad = vec_at(5.0);
371 bad[1] = f64::INFINITY;
372 let db = db_with_features(&[
373 ("a1", vec_at(0.0)),
374 ("a2", vec_at(0.1)),
375 ("b1", vec_at(9.0)),
376 ("bad", bad),
377 ]);
378 let res = cluster_library(&db, 2, 20).unwrap();
379 let clustered: usize = res.clusters.iter().map(|c| c.member_hashes.len()).sum();
380 assert_eq!(clustered, 3, "the Inf-poisoned vector should be excluded");
381 for c in &res.clusters {
382 assert!(!c.member_hashes.iter().any(|m| m == "bad"));
383 }
384 }
385 }
386