stats_claw/algorithms/clustering/
dbscan.rs1use crate::algorithms::clustering::NOISE;
10use crate::algorithms::euclidean_sq;
11
12#[derive(Debug, Clone)]
14pub struct DbscanResult {
15 pub labels: Vec<usize>,
17 pub core_samples: Vec<usize>,
19}
20
21#[must_use]
53pub fn dbscan(data: &[Vec<f64>], eps: f64, min_samples: usize) -> DbscanResult {
54 let n = data.len();
55 let eps_sq = eps * eps;
56 let neighbours: Vec<Vec<usize>> = (0..n).map(|i| region_query(data, i, eps_sq)).collect();
57 let is_core: Vec<bool> = neighbours
58 .iter()
59 .map(|nb| nb.len() >= min_samples)
60 .collect();
61
62 let mut labels = vec![UNVISITED; n];
63 let mut next_cluster = 0_usize;
64 for seed in 0..n {
65 if labels.get(seed).copied() != Some(UNVISITED) || !core_at(&is_core, seed) {
66 continue;
67 }
68 expand_cluster(seed, next_cluster, &neighbours, &is_core, &mut labels);
69 next_cluster += 1;
70 }
71
72 let labels: Vec<usize> = labels
73 .into_iter()
74 .map(|l| if l == UNVISITED { NOISE } else { l })
75 .collect();
76 let core_samples = (0..n).filter(|&i| core_at(&is_core, i)).collect();
77 DbscanResult {
78 labels,
79 core_samples,
80 }
81}
82
83const UNVISITED: usize = usize::MAX - 1;
85
86fn region_query(data: &[Vec<f64>], centre: usize, eps_sq: f64) -> Vec<usize> {
89 let Some(origin) = data.get(centre) else {
90 return Vec::new();
91 };
92 data.iter()
93 .enumerate()
94 .filter(|(_, p)| euclidean_sq(origin, p) <= eps_sq)
95 .map(|(i, _)| i)
96 .collect()
97}
98
99fn core_at(is_core: &[bool], i: usize) -> bool {
101 is_core.get(i).copied().unwrap_or(false)
102}
103
104fn expand_cluster(
108 seed: usize,
109 cluster: usize,
110 neighbours: &[Vec<usize>],
111 is_core: &[bool],
112 labels: &mut [usize],
113) {
114 let mut queue = std::collections::VecDeque::new();
115 queue.push_back(seed);
116 if let Some(slot) = labels.get_mut(seed) {
117 *slot = cluster;
118 }
119 while let Some(current) = queue.pop_front() {
120 if !core_at(is_core, current) {
121 continue;
122 }
123 let Some(reachable) = neighbours.get(current) else {
124 continue;
125 };
126 for &point in reachable {
127 if labels.get(point).copied() == Some(UNVISITED) {
128 if let Some(slot) = labels.get_mut(point) {
129 *slot = cluster;
130 }
131 queue.push_back(point);
132 }
133 }
134 }
135}
136
137#[cfg(test)]
138mod tests {
139 use super::*;
140
141 #[test]
142 fn flags_isolated_point_as_noise() {
143 let data = vec![
144 vec![0.0],
145 vec![0.1],
146 vec![0.2],
147 vec![5.0],
148 vec![5.1],
149 vec![5.2],
150 vec![100.0],
151 ];
152 let r = dbscan(&data, 0.5, 3);
153 assert_eq!(r.labels.last(), Some(&NOISE), "outlier not noise");
154 assert_eq!(r.labels.first(), r.labels.get(2), "dense triple split");
155 }
156
157 #[test]
158 fn core_samples_exclude_noise() {
159 let data = vec![vec![0.0], vec![0.1], vec![0.2], vec![100.0]];
160 let r = dbscan(&data, 0.5, 3);
161 assert!(!r.core_samples.contains(&3), "outlier marked core");
162 assert_eq!(r.core_samples.len(), 3, "core count");
163 }
164
165 #[test]
166 fn deterministic_for_fixed_inputs() {
167 let data = vec![vec![0.0], vec![0.1], vec![5.0], vec![5.1]];
168 assert_eq!(dbscan(&data, 0.5, 2).labels, dbscan(&data, 0.5, 2).labels);
169 }
170}