use super::MLError;
use crate::DataFrame;
pub struct AnomalyDetector {
pub threshold: f64,
pub method: AnomalyMethod,
pub statistics: Option<DataStatistics>,
pub anomalies: Vec<AnomalyPoint>,
}
#[derive(Debug, Clone, PartialEq)]
pub enum AnomalyMethod {
Statistical,
IsolationForest,
LocalOutlierFactor,
OneClassSVM,
}
#[derive(Debug, Clone)]
pub struct DataStatistics {
pub mean: f64,
pub std_dev: f64,
pub median: f64,
pub q1: f64,
pub q3: f64,
pub iqr: f64,
pub min: f64,
pub max: f64,
}
#[derive(Debug, Clone)]
pub struct AnomalyPoint {
pub index: usize,
pub value: f64,
pub score: f64,
pub reason: AnomalyReason,
}
#[derive(Debug, Clone, PartialEq)]
pub enum AnomalyReason {
HighZScore,
OutsideIQR,
IsolationForest,
LocalOutlier,
OneClassSVM,
}
#[derive(Debug, Clone)]
pub struct AnomalyResults {
pub anomalies: Vec<AnomalyPoint>,
pub count: usize,
pub rate: f64,
pub statistics: DataStatistics,
}
impl AnomalyDetector {
pub fn new(threshold: f64, method: AnomalyMethod) -> Self {
Self {
threshold,
method,
statistics: None,
anomalies: Vec::new(),
}
}
pub fn train(&mut self, data: &DataFrame) -> Result<(), MLError> {
if data.height() < 5 {
return Err(MLError::InsufficientData(
"Need at least 5 data points for anomaly detection training".to_string(),
));
}
let values = self.extract_values(data)?;
self.statistics = Some(self.calculate_statistics(&values));
Ok(())
}
pub fn detect(&mut self, data: &DataFrame) -> Result<AnomalyResults, MLError> {
if self.statistics.is_none() {
return Err(MLError::ModelNotTrained);
}
let values = self.extract_values(data)?;
let statistics = self.statistics.as_ref().unwrap();
let anomalies = match self.method {
AnomalyMethod::Statistical => {
self.detect_statistical_anomalies(&values, statistics)?
}
AnomalyMethod::IsolationForest => {
self.detect_isolation_forest_anomalies(&values)?
}
AnomalyMethod::LocalOutlierFactor => {
self.detect_lof_anomalies(&values)?
}
AnomalyMethod::OneClassSVM => {
self.detect_one_class_svm_anomalies(&values)?
}
};
self.anomalies = anomalies.clone();
let count = anomalies.len();
let rate = (count as f64 / values.len() as f64) * 100.0;
Ok(AnomalyResults {
anomalies,
count,
rate,
statistics: statistics.clone(),
})
}
fn extract_values(&self, data: &DataFrame) -> Result<Vec<f64>, MLError> {
let columns = data.get_columns();
for column in columns {
if let Ok(series) = column.f64() {
return Ok(series.into_iter().filter_map(|v| v).collect());
}
}
Err(MLError::InvalidData(
"No numeric column found for anomaly detection".to_string(),
))
}
fn calculate_statistics(&self, values: &[f64]) -> DataStatistics {
let n = values.len() as f64;
let mean = values.iter().sum::<f64>() / n;
let variance = values
.iter()
.map(|x| (x - mean).powi(2))
.sum::<f64>()
/ (n - 1.0);
let std_dev = variance.sqrt();
let mut sorted_values = values.to_vec();
sorted_values.sort_by(|a, b| a.partial_cmp(b).unwrap());
let median = if sorted_values.len() % 2 == 0 {
let mid = sorted_values.len() / 2;
(sorted_values[mid - 1] + sorted_values[mid]) / 2.0
} else {
sorted_values[sorted_values.len() / 2]
};
let q1 = self.calculate_percentile(&sorted_values, 25.0);
let q3 = self.calculate_percentile(&sorted_values, 75.0);
let iqr = q3 - q1;
DataStatistics {
mean,
std_dev,
median,
q1,
q3,
iqr,
min: *sorted_values.first().unwrap(),
max: *sorted_values.last().unwrap(),
}
}
fn calculate_percentile(&self, sorted_values: &[f64], percentile: f64) -> f64 {
let n = sorted_values.len() as f64;
let index = (percentile / 100.0) * (n - 1.0);
let lower = index.floor() as usize;
let upper = index.ceil() as usize;
if lower == upper {
sorted_values[lower]
} else {
let weight = index - lower as f64;
sorted_values[lower] * (1.0 - weight) + sorted_values[upper] * weight
}
}
fn detect_statistical_anomalies(
&self,
values: &[f64],
statistics: &DataStatistics,
) -> Result<Vec<AnomalyPoint>, MLError> {
let mut anomalies = Vec::new();
for (i, &value) in values.iter().enumerate() {
let mut is_anomaly = false;
let mut reason = AnomalyReason::HighZScore;
let mut score = 0.0;
if statistics.std_dev > 0.0 {
let z_score = (value - statistics.mean).abs() / statistics.std_dev;
if z_score > self.threshold {
is_anomaly = true;
score = z_score;
reason = AnomalyReason::HighZScore;
}
}
if !is_anomaly {
let lower_bound = statistics.q1 - 1.5 * statistics.iqr;
let upper_bound = statistics.q3 + 1.5 * statistics.iqr;
if value < lower_bound || value > upper_bound {
is_anomaly = true;
score = if value < lower_bound {
(lower_bound - value) / statistics.iqr
} else {
(value - upper_bound) / statistics.iqr
};
reason = AnomalyReason::OutsideIQR;
}
}
if is_anomaly {
anomalies.push(AnomalyPoint {
index: i,
value,
score,
reason,
});
}
}
Ok(anomalies)
}
fn detect_isolation_forest_anomalies(&self, values: &[f64]) -> Result<Vec<AnomalyPoint>, MLError> {
let mut anomalies = Vec::new();
for (i, &value) in values.iter().enumerate() {
let median = self.calculate_median(values);
let mad = self.calculate_mad(values, median);
if mad > 0.0 {
let score = (value - median).abs() / mad;
if score > self.threshold {
anomalies.push(AnomalyPoint {
index: i,
value,
score,
reason: AnomalyReason::IsolationForest,
});
}
}
}
Ok(anomalies)
}
fn detect_lof_anomalies(&self, values: &[f64]) -> Result<Vec<AnomalyPoint>, MLError> {
let mut anomalies = Vec::new();
for (i, &value) in values.iter().enumerate() {
let mut distances = Vec::new();
for (j, &other_value) in values.iter().enumerate() {
if i != j {
distances.push((value - other_value).abs());
}
}
distances.sort_by(|a, b| a.partial_cmp(b).unwrap());
let k = 3.min(distances.len());
if k > 0 {
let avg_distance = distances[..k].iter().sum::<f64>() / k as f64;
let score = avg_distance;
if score > self.threshold {
anomalies.push(AnomalyPoint {
index: i,
value,
score,
reason: AnomalyReason::LocalOutlier,
});
}
}
}
Ok(anomalies)
}
fn detect_one_class_svm_anomalies(&self, values: &[f64]) -> Result<Vec<AnomalyPoint>, MLError> {
let mut anomalies = Vec::new();
let mean = values.iter().sum::<f64>() / values.len() as f64;
let std_dev = self.calculate_std_dev(values, mean);
for (i, &value) in values.iter().enumerate() {
if std_dev > 0.0 {
let distance = (value - mean).abs() / std_dev;
if distance > self.threshold {
anomalies.push(AnomalyPoint {
index: i,
value,
score: distance,
reason: AnomalyReason::OneClassSVM,
});
}
}
}
Ok(anomalies)
}
fn calculate_median(&self, values: &[f64]) -> f64 {
let mut sorted_values = values.to_vec();
sorted_values.sort_by(|a, b| a.partial_cmp(b).unwrap());
if sorted_values.len() % 2 == 0 {
let mid = sorted_values.len() / 2;
(sorted_values[mid - 1] + sorted_values[mid]) / 2.0
} else {
sorted_values[sorted_values.len() / 2]
}
}
fn calculate_mad(&self, values: &[f64], median: f64) -> f64 {
let mut deviations = values.iter().map(|&x| (x - median).abs()).collect::<Vec<_>>();
deviations.sort_by(|a, b| a.partial_cmp(b).unwrap());
if deviations.len() % 2 == 0 {
let mid = deviations.len() / 2;
(deviations[mid - 1] + deviations[mid]) / 2.0
} else {
deviations[deviations.len() / 2]
}
}
fn calculate_std_dev(&self, values: &[f64], mean: f64) -> f64 {
let n = values.len() as f64;
let variance = values
.iter()
.map(|x| (x - mean).powi(2))
.sum::<f64>()
/ (n - 1.0);
variance.sqrt()
}
}