use legume_numeric::matrix::sparse_stat::SparseRunningStatistics;
use legume_numeric::matrix::traits::RunningStatOps;
const MIN_MEAN_FOR_FIT: f32 = 1e-4;
const PHI_FLOOR: f32 = 0.0;
const PHI_CEIL: f32 = 100.0;
#[derive(Clone, Debug)]
pub struct DispersionTrend {
a: f32,
b: f32,
pub num_fit: usize,
}
impl DispersionTrend {
pub fn fit(means: &[f32], vars: &[f32]) -> Self {
assert_eq!(means.len(), vars.len(), "means and vars length mismatch");
let mut x: Vec<f64> = Vec::with_capacity(means.len());
let mut y: Vec<f64> = Vec::with_capacity(means.len());
let mut w: Vec<f64> = Vec::with_capacity(means.len());
for (&mu, &var) in means.iter().zip(vars.iter()) {
if !mu.is_finite() || !var.is_finite() || mu < MIN_MEAN_FOR_FIT {
continue;
}
let phi_hat = ((var - mu) / (mu * mu)) as f64;
if phi_hat <= 0.0 {
continue;
}
x.push((mu as f64).ln());
y.push(phi_hat.ln());
w.push(mu as f64);
}
let num_fit = x.len();
if num_fit < 2 {
return Self {
a: f32::NEG_INFINITY,
b: 0.0,
num_fit,
};
}
let w_sum: f64 = w.iter().sum();
let x_mean: f64 = x.iter().zip(&w).map(|(xi, wi)| xi * wi).sum::<f64>() / w_sum;
let y_mean: f64 = y.iter().zip(&w).map(|(yi, wi)| yi * wi).sum::<f64>() / w_sum;
let mut sxx = 0.0f64;
let mut sxy = 0.0f64;
for ((xi, yi), wi) in x.iter().zip(&y).zip(&w) {
let dx = xi - x_mean;
sxx += wi * dx * dx;
sxy += wi * dx * (yi - y_mean);
}
if sxx <= 0.0 {
return Self {
a: y_mean as f32,
b: 0.0,
num_fit,
};
}
let b = sxy / sxx;
let a = y_mean - b * x_mean;
Self {
a: a as f32,
b: b as f32,
num_fit,
}
}
pub fn from_sparse_stats(stats: &SparseRunningStatistics<f32>) -> Self {
let means = stats.mean();
let vars = stats.variance();
Self::fit(&means, &vars)
}
pub fn phi_at(&self, mu: f32) -> f32 {
if !mu.is_finite() || mu <= 0.0 {
return PHI_FLOOR;
}
let log_mu = mu.ln();
let log_phi = self.a + self.b * log_mu;
log_phi.exp().clamp(PHI_FLOOR, PHI_CEIL)
}
pub fn excess(&self, mu: f32, var: f32) -> f32 {
if !mu.is_finite() || mu <= 0.0 || !var.is_finite() {
return f32::NEG_INFINITY;
}
let phi_hat = (var - mu) / (mu * mu);
phi_hat - self.phi_at(mu)
}
pub fn fisher_weight(&self, pi: f32, avg_s: f32, mu: f32) -> f32 {
let phi = self.phi_at(mu);
1.0 / (1.0 + pi * avg_s * phi)
}
}
#[cfg(test)]
#[path = "nb_dispersion_tests.rs"]
mod tests;