use serde::{Deserialize, Serialize};
#[derive(Serialize, Deserialize, Debug, Clone)]
pub struct FeatureStatistics {
pub count: u64,
pub mean: f64,
pub m2: f64,
pub min: f64,
pub max: f64,
pub null_count: u64,
}
impl FeatureStatistics {
pub fn new() -> Self {
Self {
count: 0,
mean: 0.0,
m2: 0.0,
min: f64::INFINITY,
max: f64::NEG_INFINITY,
null_count: 0,
}
}
pub fn update(&mut self, value: f64) {
self.count += 1;
let delta = value - self.mean;
self.mean += delta / self.count as f64;
let delta2 = value - self.mean;
self.m2 += delta * delta2;
if value < self.min {
self.min = value;
}
if value > self.max {
self.max = value;
}
}
pub fn update_null(&mut self) {
self.null_count += 1;
}
pub fn variance(&self) -> f64 {
if self.count < 2 {
0.0
} else {
self.m2 / self.count as f64
}
}
pub fn sample_variance(&self) -> f64 {
if self.count < 2 {
0.0
} else {
self.m2 / (self.count - 1) as f64
}
}
pub fn stddev(&self) -> f64 {
self.variance().sqrt()
}
pub fn null_rate(&self) -> f64 {
let total = self.count + self.null_count;
if total == 0 {
0.0
} else {
self.null_count as f64 / total as f64
}
}
}
impl Default for FeatureStatistics {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone)]
pub struct ColumnStatistics {
pub mean: f64,
pub stddev: f64,
pub null_rate: f64,
pub min: f64,
pub max: f64,
pub count: u64,
}
pub fn compute_psi(expected: &[f64], actual: &[f64]) -> f64 {
if expected.len() != actual.len() || expected.is_empty() {
return 0.0;
}
let epsilon = 1e-10; let mut psi = 0.0;
for (e, a) in expected.iter().zip(actual.iter()) {
let e_safe = e.max(epsilon);
let a_safe = a.max(epsilon);
psi += (a_safe - e_safe) * (a_safe / e_safe).ln();
}
psi
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_welford_basic() {
let mut stats = FeatureStatistics::new();
stats.update(10.0);
stats.update(20.0);
stats.update(30.0);
assert_eq!(stats.count, 3);
assert!((stats.mean - 20.0).abs() < 1e-10);
assert_eq!(stats.min, 10.0);
assert_eq!(stats.max, 30.0);
}
#[test]
fn test_welford_variance() {
let mut stats = FeatureStatistics::new();
for v in &[2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0] {
stats.update(*v);
}
assert!((stats.mean - 5.0).abs() < 1e-10);
assert!((stats.variance() - 4.0).abs() < 1e-10);
assert!((stats.sample_variance() - 32.0 / 7.0).abs() < 1e-10);
}
#[test]
fn test_null_rate() {
let mut stats = FeatureStatistics::new();
stats.update(1.0);
stats.update(2.0);
stats.update_null();
stats.update_null();
assert!((stats.null_rate() - 0.5).abs() < 1e-10);
}
#[test]
fn test_psi_identical() {
let dist = vec![0.25, 0.25, 0.25, 0.25];
let psi = compute_psi(&dist, &dist);
assert!(psi.abs() < 1e-10);
}
#[test]
fn test_psi_shifted() {
let expected = vec![0.25, 0.25, 0.25, 0.25];
let actual = vec![0.1, 0.1, 0.4, 0.4];
let psi = compute_psi(&expected, &actual);
assert!(psi > 0.1);
}
#[test]
fn test_single_value() {
let mut stats = FeatureStatistics::new();
stats.update(42.0);
assert_eq!(stats.count, 1);
assert_eq!(stats.mean, 42.0);
assert_eq!(stats.variance(), 0.0);
assert_eq!(stats.min, 42.0);
assert_eq!(stats.max, 42.0);
}
}