use num::{Float, FromPrimitive};
use std::ops::{AddAssign, SubAssign};
use crate::mean::Mean;
use crate::stats::{Revertable, RollableUnivariate, Univariate};
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, Serialize, Deserialize)]
pub struct Variance<F: Float + FromPrimitive + AddAssign + SubAssign> {
pub mean: Mean<F>,
pub ddof: u32,
pub state: F,
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> Variance<F> {
pub fn new(ddof: u32) -> Self {
Self {
mean: Mean::new(),
ddof,
state: F::from_f64(0.).unwrap(),
}
}
}
impl<F> Default for Variance<F>
where
F: Float + FromPrimitive + AddAssign + SubAssign,
{
fn default() -> Self {
Self {
mean: Mean::new(),
ddof: 1,
state: F::from_f64(0.).unwrap(),
}
}
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> Univariate<F> for Variance<F> {
fn update(&mut self, x: F) {
let mean_old = self.mean.get();
self.mean.update(x);
let mean_new = self.mean.get();
self.state += (x - mean_old) * (x - mean_new);
}
fn get(&self) -> F {
let mean_n = self.mean.n.get();
if mean_n > F::from_u32(self.ddof).unwrap() {
return self.state / (mean_n - F::from_u32(self.ddof).unwrap());
}
F::from_f64(0.).unwrap()
}
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> Revertable<F> for Variance<F> {
fn revert(&mut self, x: F) -> Result<(), &'static str> {
let mean_old = self.mean.get();
self.mean.revert(x)?;
let mean_new = self.mean.get();
self.state -= (x - mean_old) * (x - mean_new);
Ok(())
}
}
impl<F: Float + FromPrimitive + AddAssign + SubAssign> RollableUnivariate<F> for Variance<F> {}