use std::collections::VecDeque;
use crate::stats::rolling_median;
#[derive(Debug, Clone, PartialEq)]
pub struct KMeansResult {
pub centroids: Vec<f64>,
pub assignments: Vec<usize>,
pub iterations: usize,
}
pub fn kmeans_1d(values: &[f64], k: usize, max_iterations: usize) -> Option<KMeansResult> {
if values.is_empty() || k == 0 || k > values.len() {
return None;
}
let mut sorted = values.to_vec();
sorted.sort_by(|a, b| a.partial_cmp(b).unwrap());
let mut centroids: Vec<f64> = (0..k)
.map(|i| {
let idx = if k == 1 {
0
} else {
i * (sorted.len() - 1) / (k - 1)
};
sorted[idx]
})
.collect();
let mut assignments = vec![0usize; values.len()];
let mut iterations = 0;
for _ in 0..max_iterations.max(1) {
iterations += 1;
let mut changed = false;
for (i, &v) in values.iter().enumerate() {
let mut best = 0;
let mut best_dist = f64::INFINITY;
for (c_idx, &c) in centroids.iter().enumerate() {
let dist = (v - c).abs();
if dist < best_dist {
best_dist = dist;
best = c_idx;
}
}
if assignments[i] != best {
changed = true;
}
assignments[i] = best;
}
let mut sums = vec![0.0; k];
let mut counts = vec![0usize; k];
for (i, &v) in values.iter().enumerate() {
sums[assignments[i]] += v;
counts[assignments[i]] += 1;
}
for c in 0..k {
if counts[c] > 0 {
centroids[c] = sums[c] / counts[c] as f64;
}
}
if !changed {
break;
}
}
Some(KMeansResult {
centroids,
assignments,
iterations,
})
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct RobustBand {
pub median: f64,
pub mad: f64,
pub upper: f64,
pub lower: f64,
}
impl RobustBand {
pub fn contains(&self, value: f64) -> bool {
value >= self.lower && value <= self.upper
}
}
const MAD_CONSISTENCY_CONSTANT: f64 = 1.482_602_218_505_602;
#[derive(Debug, Clone)]
pub struct RollingRobustThreshold {
window_len: usize,
k: f64,
buffer: VecDeque<f64>,
}
impl RollingRobustThreshold {
pub fn new(window_len: usize, k: f64) -> Self {
let window_len = window_len.max(1);
Self {
window_len,
k,
buffer: VecDeque::with_capacity(window_len),
}
}
pub fn reset(&mut self) {
self.buffer.clear();
}
pub fn update(&mut self, value: f64) -> Option<RobustBand> {
if self.buffer.len() >= self.window_len {
self.buffer.pop_front();
}
self.buffer.push_back(value);
if self.buffer.len() < self.window_len {
return None;
}
let values: Vec<f64> = self.buffer.iter().copied().collect();
let median = rolling_median(&values);
let abs_deviations: Vec<f64> = values.iter().map(|v| (v - median).abs()).collect();
let mad = rolling_median(&abs_deviations) * MAD_CONSISTENCY_CONSTANT;
let offset = self.k * mad;
Some(RobustBand {
median,
mad,
upper: median + offset,
lower: median - offset,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kmeans_1d_separates_clean_clusters() {
let values = [1.0, 1.1, 0.9, 10.0, 10.2, 9.8];
let result = kmeans_1d(&values, 2, 20).unwrap();
let low_cluster = result.assignments[0];
assert_eq!(result.assignments[1], low_cluster);
assert_eq!(result.assignments[2], low_cluster);
let high_cluster = result.assignments[3];
assert_ne!(low_cluster, high_cluster);
assert_eq!(result.assignments[4], high_cluster);
assert_eq!(result.assignments[5], high_cluster);
}
#[test]
fn test_kmeans_1d_is_deterministic_across_runs() {
let values = [3.0, 7.0, 1.0, 9.0, 2.0, 8.0, 4.0, 6.0];
let a = kmeans_1d(&values, 3, 50).unwrap();
let b = kmeans_1d(&values, 3, 50).unwrap();
assert_eq!(a, b);
}
#[test]
fn test_kmeans_1d_rejects_degenerate_input() {
assert!(kmeans_1d(&[], 1, 10).is_none());
assert!(kmeans_1d(&[1.0, 2.0], 0, 10).is_none());
assert!(kmeans_1d(&[1.0, 2.0], 3, 10).is_none());
}
#[test]
fn test_robust_threshold_ignores_a_single_outlier() {
let mut robust = RollingRobustThreshold::new(9, 3.0);
let values = [100.0, 101.0, 99.0, 100.5, 99.5, 100.0, 100.0, 99.0, 1000.0];
let mut last = None;
for v in values {
last = robust.update(v);
}
let band = last.unwrap();
assert!((band.median - 100.0).abs() < 1.0);
assert!(band.contains(100.5));
assert!(
!band.contains(1000.0),
"outlier must fall outside the robust band"
);
}
#[test]
fn test_robust_threshold_none_until_window_full() {
let mut robust = RollingRobustThreshold::new(5, 2.0);
for v in [1.0, 2.0, 3.0, 4.0] {
assert!(robust.update(v).is_none());
}
assert!(robust.update(5.0).is_some());
}
}