kestrel_chartkit/
clustering.rs1use std::collections::VecDeque;
5
6use crate::stats::rolling_median;
7
8#[derive(Debug, Clone, PartialEq)]
11pub struct KMeansResult {
12 pub centroids: Vec<f64>,
13 pub assignments: Vec<usize>,
14 pub iterations: usize,
15}
16
17pub fn kmeans_1d(values: &[f64], k: usize, max_iterations: usize) -> Option<KMeansResult> {
24 if values.is_empty() || k == 0 || k > values.len() {
25 return None;
26 }
27
28 let mut sorted = values.to_vec();
29 sorted.sort_by(f64::total_cmp);
30 let mut centroids: Vec<f64> = (0..k)
31 .map(|i| {
32 let idx = if k == 1 {
33 0
34 } else {
35 i * (sorted.len() - 1) / (k - 1)
36 };
37 sorted.get(idx).copied().unwrap_or(0.0)
38 })
39 .collect();
40
41 let mut assignments = vec![0usize; values.len()];
42 let mut iterations = 0;
43
44 for _ in 0..max_iterations.max(1) {
45 iterations += 1;
46 let mut changed = false;
47
48 for (i, &v) in values.iter().enumerate() {
49 let mut best = 0;
50 let mut best_dist = f64::INFINITY;
51 for (c_idx, &c) in centroids.iter().enumerate() {
52 let dist = (v - c).abs();
53 if dist < best_dist {
54 best_dist = dist;
55 best = c_idx;
56 }
57 }
58 if let Some(assign_ref) = assignments.get_mut(i) {
59 if *assign_ref != best {
60 changed = true;
61 *assign_ref = best;
62 }
63 }
64 }
65
66 let mut sums = vec![0.0; k];
67 let mut counts = vec![0usize; k];
68 for (i, &v) in values.iter().enumerate() {
69 let cluster_idx = assignments.get(i).copied().unwrap_or(0);
70 if let (Some(s), Some(cnt)) = (sums.get_mut(cluster_idx), counts.get_mut(cluster_idx)) {
71 *s += v;
72 *cnt += 1;
73 }
74 }
75 for c in 0..k {
76 let cnt = counts.get(c).copied().unwrap_or(0);
77 if cnt > 0 {
78 if let (Some(centroid), Some(&sum)) = (centroids.get_mut(c), sums.get(c)) {
79 *centroid = sum / cnt as f64;
80 }
81 }
82 }
83
84 if !changed {
85 break;
86 }
87 }
88
89 Some(KMeansResult {
90 centroids,
91 assignments,
92 iterations,
93 })
94}
95
96#[derive(Debug, Clone, Copy, PartialEq)]
100pub struct RobustBand {
101 pub median: f64,
102 pub mad: f64,
103 pub upper: f64,
104 pub lower: f64,
105}
106
107impl RobustBand {
108 pub fn contains(&self, value: f64) -> bool {
109 value >= self.lower && value <= self.upper
110 }
111}
112
113const MAD_CONSISTENCY_CONSTANT: f64 = 1.482_602_218_505_602;
116
117#[derive(Debug, Clone)]
119pub struct RollingRobustThreshold {
120 window_len: usize,
121 k: f64,
122 buffer: VecDeque<f64>,
123}
124
125impl RollingRobustThreshold {
126 pub fn new(window_len: usize, k: f64) -> Self {
128 let window_len = window_len.max(1);
129 Self {
130 window_len,
131 k,
132 buffer: VecDeque::with_capacity(window_len),
133 }
134 }
135
136 pub fn reset(&mut self) {
137 self.buffer.clear();
138 }
139
140 pub fn update(&mut self, value: f64) -> Option<RobustBand> {
142 if self.buffer.len() >= self.window_len {
143 self.buffer.pop_front();
144 }
145 self.buffer.push_back(value);
146 if self.buffer.len() < self.window_len {
147 return None;
148 }
149
150 let values: Vec<f64> = self.buffer.iter().copied().collect();
151 let median = rolling_median(&values);
152 let abs_deviations: Vec<f64> = values.iter().map(|v| (v - median).abs()).collect();
153 let mad = rolling_median(&abs_deviations) * MAD_CONSISTENCY_CONSTANT;
154
155 let offset = self.k * mad;
156 Some(RobustBand {
157 median,
158 mad,
159 upper: median + offset,
160 lower: median - offset,
161 })
162 }
163}
164
165#[cfg(test)]
166mod tests {
167 use super::*;
168
169 #[test]
170 fn test_kmeans_1d_separates_clean_clusters() {
171 let values = [1.0, 1.1, 0.9, 10.0, 10.2, 9.8];
172 let result = kmeans_1d(&values, 2, 20).unwrap();
173
174 let low_cluster = result.assignments[0];
176 assert_eq!(result.assignments[1], low_cluster);
177 assert_eq!(result.assignments[2], low_cluster);
178
179 let high_cluster = result.assignments[3];
180 assert_ne!(low_cluster, high_cluster);
181 assert_eq!(result.assignments[4], high_cluster);
182 assert_eq!(result.assignments[5], high_cluster);
183 }
184
185 #[test]
186 fn test_kmeans_1d_is_deterministic_across_runs() {
187 let values = [3.0, 7.0, 1.0, 9.0, 2.0, 8.0, 4.0, 6.0];
188 let a = kmeans_1d(&values, 3, 50).unwrap();
189 let b = kmeans_1d(&values, 3, 50).unwrap();
190 assert_eq!(a, b);
191 }
192
193 #[test]
194 fn test_kmeans_1d_rejects_degenerate_input() {
195 assert!(kmeans_1d(&[], 1, 10).is_none());
196 assert!(kmeans_1d(&[1.0, 2.0], 0, 10).is_none());
197 assert!(kmeans_1d(&[1.0, 2.0], 3, 10).is_none());
198 }
199
200 #[test]
201 fn test_robust_threshold_ignores_a_single_outlier() {
202 let mut robust = RollingRobustThreshold::new(9, 3.0);
203 let values = [100.0, 101.0, 99.0, 100.5, 99.5, 100.0, 100.0, 99.0, 1000.0];
205 let mut last = None;
206 for v in values {
207 last = robust.update(v);
208 }
209 let band = last.unwrap();
210 assert!((band.median - 100.0).abs() < 1.0);
213 assert!(band.contains(100.5));
214 assert!(
215 !band.contains(1000.0),
216 "outlier must fall outside the robust band"
217 );
218 }
219
220 #[test]
221 fn test_robust_threshold_none_until_window_full() {
222 let mut robust = RollingRobustThreshold::new(5, 2.0);
223 for v in [1.0, 2.0, 3.0, 4.0] {
224 assert!(robust.update(v).is_none());
225 }
226 assert!(robust.update(5.0).is_some());
227 }
228}