use ndarray::{Array1, Array2};
use robust_rs_core::error::RobustError;
use robust_rs_core::scale::{Qn, ScaleEstimator};
use super::correction::hard_reweight;
use super::linalg::symmetric_eigen;
use super::{distances_from, ScatterFit};
use crate::util::median;
#[derive(Debug, Clone, Copy)]
pub struct Ogk<S = Qn> {
scale: S,
n_iter: usize,
reweight: bool,
reweight_quantile: f64,
}
impl Default for Ogk<Qn> {
fn default() -> Self {
Self::new(Qn::default())
}
}
impl<S: ScaleEstimator> Ogk<S> {
pub fn new(scale: S) -> Self {
Self {
scale,
n_iter: 2,
reweight: true,
reweight_quantile: 0.9,
}
}
pub fn n_iter(mut self, n: usize) -> Self {
self.n_iter = n;
self
}
pub fn reweight(mut self, on: bool) -> Self {
self.reweight = on;
self
}
pub fn fit(&self, x: &Array2<f64>) -> Result<ScatterFit, RobustError> {
let (n, p) = x.dim();
if p == 0 {
return Err(RobustError::SingularDesign);
}
if n < p + 1 {
return Err(RobustError::InsufficientData {
needed: p + 1,
got: n,
});
}
let (raw_location, raw_scatter) = self.raw(x)?;
let reweighted = if self.reweight {
hard_reweight(x, &raw_location, &raw_scatter, self.reweight_quantile, true)?
} else {
None
};
let (location, scatter, distances, weights) = match reweighted {
Some(rw) => (rw.location, rw.scatter, rw.distances, rw.weights),
None => {
let d = distances_from(x, &raw_location, &raw_scatter)?;
(raw_location, raw_scatter, d, Array1::ones(n))
}
};
Ok(ScatterFit {
location,
scatter,
distances,
weights,
})
}
fn raw(&self, x: &Array2<f64>) -> Result<(Array1<f64>, Array2<f64>), RobustError> {
let (n, p) = x.dim();
let mut w = x.clone();
let mut m = Array2::<f64>::eye(p);
for _ in 0..self.n_iter.max(1) {
let mut sigma = Array1::<f64>::zeros(p);
for j in 0..p {
sigma[j] = self.scale.scale(&w.column(j).to_vec())?.get();
}
let y = Array2::from_shape_fn((n, p), |(i, j)| w[[i, j]] / sigma[j]);
let mut u = Array2::<f64>::eye(p);
for j in 0..p {
for k in (j + 1)..p {
let sum: Vec<f64> = (0..n).map(|i| y[[i, j]] + y[[i, k]]).collect();
let dif: Vec<f64> = (0..n).map(|i| y[[i, j]] - y[[i, k]]).collect();
let sp = self.scale.scale(&sum)?.get();
let sm = self.scale.scale(&dif)?.get();
let ujk = 0.25 * (sp * sp - sm * sm);
u[[j, k]] = ujk;
u[[k, j]] = ujk;
}
}
let (_vals, e) = symmetric_eigen(&u)?;
let z = y.dot(&e); let b = Array2::from_shape_fn((p, p), |(i, j)| sigma[i] * e[[i, j]]); m = m.dot(&b);
w = z;
}
let mut nu = Array1::<f64>::zeros(p);
let mut gamma2 = Array1::<f64>::zeros(p);
for l in 0..p {
let g = self.scale.scale(&w.column(l).to_vec())?.get();
gamma2[l] = g * g;
nu[l] = median(&mut w.column(l).to_vec());
}
let mu = m.dot(&nu);
let mg = Array2::from_shape_fn((p, p), |(i, j)| m[[i, j]] * gamma2[j]);
let sigma = mg.dot(&m.t());
Ok((mu, sigma))
}
}