robust_rs/regression/
m_estimator.rs1use ndarray::{Array1, Array2};
4use robust_rs_core::error::RobustError;
5use robust_rs_core::rho::RhoFunction;
6use robust_rs_core::scale::ScaleEstimator;
7use robust_rs_core::solver::Control;
8
9use crate::estimator::RegressionFit;
10
11pub struct MEstimator<R, S> {
14 pub rho: R,
16 pub scale: S,
18 pub control: Control,
20}
21
22impl<R, S> MEstimator<R, S>
23where
24 R: RhoFunction + Clone + 'static,
25 S: ScaleEstimator,
26{
27 pub fn new(rho: R, scale: S) -> Self {
29 Self {
30 rho,
31 scale,
32 control: Control::default(),
33 }
34 }
35
36 pub fn fit(&self, x: &Array2<f64>, y: &Array1<f64>) -> Result<RegressionFit, RobustError> {
39 use crate::wls::weighted_least_squares;
40
41 let n = x.nrows();
42
43 let ones = Array1::from_elem(n, 1.0);
45 let mut beta = weighted_least_squares(x, y, &ones)?;
46 let mut resid = y - &x.dot(&beta);
47
48 let scale_est = self.scale.scale(resid.as_slice().expect("contiguous"))?;
50 let s: f64 = scale_est.get();
51
52 let mut converged = false;
53 for _ in 0..self.control.max_iter {
54 let w = resid.mapv(|r| self.rho.weight(r / s));
55 let next = weighted_least_squares(x, y, &w)?;
56
57 let diff = &next - β
59 let rel = diff.dot(&diff).sqrt() / (next.dot(&next).sqrt() + 1e-12);
60
61 beta = next;
62 resid = y - &x.dot(&beta);
63 if rel <= self.control.tol {
64 converged = true;
65 break;
66 }
67 }
68 if !converged {
69 return Err(RobustError::NonConvergence {
70 iters: self.control.max_iter,
71 });
72 }
73
74 let weights = resid.mapv(|r| self.rho.weight(r / s));
76
77 Ok(RegressionFit {
78 coefficients: beta,
79 scale: scale_est,
80 residuals: resid,
81 weights,
82 rho: Box::new(self.rho.clone()), breakdown_point: 0.0, })
85 }
86}