Skip to main content

kestrel_chartkit/
clustering.rs

1//! Deterministic clustering and robust adaptive-threshold primitives for regime/tradability
2//! engines, so each one stops hand-rolling its own bucketing/outlier-sensitive threshold logic.
3
4use std::collections::VecDeque;
5
6use crate::stats::rolling_median;
7
8/// Result of [`kmeans_1d`]: final centroids and, for every input value (same order/length as the
9/// input slice), which centroid it was assigned to.
10#[derive(Debug, Clone, PartialEq)]
11pub struct KMeansResult {
12    pub centroids: Vec<f64>,
13    pub assignments: Vec<usize>,
14    pub iterations: usize,
15}
16
17/// Deterministic 1-D k-means: `k` initial centroids are the values at evenly spaced quantiles of
18/// the *sorted* input (not a random seed), so the same input always produces the same clustering
19/// — no RNG dependency, fully reproducible. Lloyd's algorithm then runs until assignments stop
20/// changing or `max_iterations` is reached.
21///
22/// Returns `None` for degenerate input: empty `values`, `k == 0`, or `k > values.len()`.
23pub 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/// A robust rolling threshold band: median +/- `k` scaled median absolute deviations (MAD),
97/// rather than mean +/- k*stddev, so a handful of outlier bars in the window do not blow out the
98/// band the way a plain stddev-based threshold would.
99#[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
113/// MAD-to-stddev consistency constant under a normal distribution (`1 / Phi^-1(3/4)`), the
114/// standard scaling so a MAD-based band is comparable in width to a stddev-based one.
115const MAD_CONSISTENCY_CONSTANT: f64 = 1.482_602_218_505_602;
116
117/// Streaming robust threshold over a fixed-size trailing window.
118#[derive(Debug, Clone)]
119pub struct RollingRobustThreshold {
120    window_len: usize,
121    k: f64,
122    buffer: VecDeque<f64>,
123}
124
125impl RollingRobustThreshold {
126    /// `window_len` bars of history, `k` scaled-MAD multiplier for the band width.
127    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    /// Feeds one value. Returns `None` until the window has `window_len` values.
141    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        // Points 0..3 must share one cluster, 3..6 the other.
175        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        // A tight cluster around 100 plus one wild outlier.
204        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        // Median stays near 100 despite the outlier; a mean-based threshold would be dragged
211        // toward the 1000.0 outlier instead.
212        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}