use crate::error::RobustError;
use crate::rho::RhoFunction;
use crate::scale::ScaleEstimator;
use crate::solver::Control;
use crate::theory::gauss_hermite;
use crate::types::Scale;
#[derive(Debug, Clone, Copy)]
pub struct SScale<R: RhoFunction> {
rho: R,
delta: f64,
control: Control,
}
impl<R: RhoFunction> SScale<R> {
pub fn new(rho: R, delta: f64) -> Result<Self, RobustError> {
let sup = rho.rho_sup().ok_or(RobustError::UnboundedLoss)?;
if delta.is_finite() && delta > 0.0 && delta < sup {
Ok(Self {
rho,
delta,
control: Control::default(),
})
} else {
Err(RobustError::InvalidTuning { value: delta })
}
}
pub fn fisher_consistent(rho: R, quad_points: usize) -> Result<Self, RobustError> {
let (nodes, weights) = gauss_hermite(quad_points);
let delta = nodes
.iter()
.zip(&weights)
.map(|(&x, &w)| w * rho.rho(x))
.sum();
Self::new(rho, delta)
}
pub fn with_control(mut self, control: Control) -> Self {
self.control = control;
self
}
pub fn delta(&self) -> f64 {
self.delta
}
}
impl<R: RhoFunction> ScaleEstimator for SScale<R> {
fn scale(&self, residuals: &[f64]) -> Result<Scale, RobustError> {
let n = residuals.len();
if n == 0 {
return Err(RobustError::InsufficientData { needed: 1, got: 0 });
}
let mut buf = residuals.to_vec();
let med = median(&mut buf);
for (b, &r) in buf.iter_mut().zip(residuals) {
*b = (r - med).abs();
}
let mut s = 1.482_602_218_505_602 * median(&mut buf);
if !(s.is_finite() && s > 0.0) {
return Err(RobustError::DegenerateScale); }
for _ in 0..self.control.max_iter {
let mean_rho = residuals.iter().map(|&r| self.rho.rho(r / s)).sum::<f64>() / n as f64;
let s_next = s * (mean_rho / self.delta).sqrt();
if !(s_next.is_finite() && s_next > 0.0) {
return Err(RobustError::DegenerateScale);
}
if (s_next - s).abs() <= self.control.tol * s {
return Scale::new(s_next);
}
s = s_next;
}
Err(RobustError::NonConvergence {
iters: self.control.max_iter,
})
}
}
fn median(v: &mut [f64]) -> f64 {
v.sort_unstable_by(f64::total_cmp);
let n = v.len();
let mid = n / 2;
if n % 2 == 1 {
v[mid]
} else {
0.5 * (v[mid - 1] + v[mid])
}
}