Skip to main content

datarust/imputer/
simple.rs

1use crate::error::{DatarustError, Result};
2use crate::matrix::Matrix;
3use crate::stats;
4use crate::traits::{default_input_names, FeatureNames};
5use crate::Transformer;
6
7/// Imputation strategy, mirroring `sklearn.impute.SimpleImputer`.
8#[derive(Debug, Clone, Default, PartialEq)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10pub enum ImputeStrategy {
11    /// Fill missing values with the column mean.
12    #[default]
13    Mean,
14    /// Fill missing values with the column median.
15    Median,
16    /// Fill missing values with the column most frequent value.
17    MostFrequent,
18    /// Fill missing values with the given constant.
19    Constant(f64),
20}
21
22/// Impute missing values (represented as `f64::NAN`) using a per-column statistic.
23///
24/// Mirrors `sklearn.impute.SimpleImputer`.
25#[derive(Debug, Clone)]
26#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
27pub struct SimpleImputer {
28    strategy: ImputeStrategy,
29    fill_values: Vec<f64>,
30    fitted: bool,
31}
32
33impl SimpleImputer {
34    /// Creates a new simple imputer with the given strategy.
35    pub fn new(strategy: ImputeStrategy) -> Self {
36        Self {
37            strategy,
38            fill_values: vec![],
39            fitted: false,
40        }
41    }
42
43    /// Returns the imputation strategy.
44    pub fn strategy(&self) -> &ImputeStrategy {
45        &self.strategy
46    }
47
48    /// Returns the learned per-column fill values.
49    pub fn fill_values(&self) -> &[f64] {
50        &self.fill_values
51    }
52
53    #[allow(clippy::needless_range_loop)]
54    fn compute_fill(x: &Matrix, strategy: &ImputeStrategy) -> Result<Vec<f64>> {
55        let data = x.rows_ref();
56        let cols = x.ncols();
57        let mut fills = Vec::with_capacity(cols);
58        for j in 0..cols {
59            let col: Vec<f64> = (0..x.nrows())
60                .filter_map(|i| {
61                    let v = data[i][j];
62                    if v.is_nan() {
63                        None
64                    } else {
65                        Some(v)
66                    }
67                })
68                .collect();
69            let fill = match strategy {
70                ImputeStrategy::Mean => {
71                    if col.is_empty() {
72                        return Err(DatarustError::AllMissing(format!("column {}", j)));
73                    }
74                    let s: f64 = col.iter().sum();
75                    s / col.len() as f64
76                }
77                ImputeStrategy::Median => {
78                    if col.is_empty() {
79                        return Err(DatarustError::AllMissing(format!("column {}", j)));
80                    }
81                    let mut c = col.clone();
82                    c.sort_by(|a, b| a.total_cmp(b));
83                    stats::median_sorted(&c).expect("column non-empty (checked above)")
84                }
85                ImputeStrategy::MostFrequent => {
86                    if col.is_empty() {
87                        return Err(DatarustError::AllMissing(format!("column {}", j)));
88                    }
89                    let mut c = col.clone();
90                    c.sort_by(|a, b| a.total_cmp(b));
91                    let single: Vec<Vec<f64>> = c.into_iter().map(|v| vec![v]).collect();
92                    stats::mode_column(&single)[0]
93                }
94                ImputeStrategy::Constant(v) => *v,
95            };
96            fills.push(fill);
97        }
98        Ok(fills)
99    }
100}
101
102impl Default for SimpleImputer {
103    fn default() -> Self {
104        Self::new(ImputeStrategy::Mean)
105    }
106}
107
108impl FeatureNames for SimpleImputer {
109    fn feature_names_out(&self, input_features: Option<&[String]>) -> Vec<String> {
110        match input_features {
111            Some(fs) => fs.to_vec(),
112            None => default_input_names(self.fill_values.len()),
113        }
114    }
115}
116
117impl Transformer for SimpleImputer {
118    fn name(&self) -> &'static str {
119        "SimpleImputer"
120    }
121
122    fn fit(&mut self, x: &Matrix) -> Result<()> {
123        self.fill_values = Self::compute_fill(x, &self.strategy)?;
124        self.fitted = true;
125        Ok(())
126    }
127
128    fn transform(&self, x: &Matrix) -> Result<Matrix> {
129        if !self.fitted {
130            return Err(DatarustError::NotFitted("SimpleImputer".into()));
131        }
132        if self.fill_values.len() != x.ncols() {
133            return Err(DatarustError::ShapeMismatch {
134                expected: format!("{} features", self.fill_values.len()),
135                actual: format!("{} features", x.ncols()),
136            });
137        }
138        let mut out = x.clone();
139        for i in 0..out.nrows() {
140            for j in 0..out.ncols() {
141                if out.get(i, j).is_nan() {
142                    out.set(i, j, self.fill_values[j]);
143                }
144            }
145        }
146        Ok(out)
147    }
148
149    fn is_fitted(&self) -> bool {
150        self.fitted
151    }
152}
153
154#[cfg(test)]
155mod tests {
156    use super::*;
157
158    fn nan() -> f64 {
159        f64::NAN
160    }
161
162    fn m_missing() -> Matrix {
163        Matrix::new(vec![
164            vec![1.0, 10.0, nan()],
165            vec![2.0, nan(), 5.0],
166            vec![3.0, 30.0, 5.0],
167            vec![4.0, 40.0, 5.0],
168        ])
169        .unwrap()
170    }
171
172    #[test]
173    fn mean_strategy() {
174        let mut imp = SimpleImputer::new(ImputeStrategy::Mean);
175        let out = imp.fit_transform(&m_missing()).unwrap();
176        // col1 mean of (10,30,40) = 26.666
177        assert!((imp.fill_values()[1] - (80.0 / 3.0)).abs() < 1e-9);
178        assert!((out.get(1, 1) - (80.0 / 3.0)).abs() < 1e-9);
179        // col2 mean of (5,5,5) = 5
180        assert!((out.get(0, 2) - 5.0).abs() < 1e-9);
181    }
182
183    #[test]
184    fn median_strategy() {
185        let mut imp = SimpleImputer::new(ImputeStrategy::Median);
186        let out = imp.fit_transform(&m_missing()).unwrap();
187        // col1: 10,30,40 sorted -> median 30
188        assert!((imp.fill_values()[1] - 30.0).abs() < 1e-9);
189        assert!((out.get(1, 1) - 30.0).abs() < 1e-9);
190    }
191
192    #[test]
193    fn most_frequent_strategy() {
194        let x = Matrix::new(vec![
195            vec![nan(), 5.0],
196            vec![1.0, 5.0],
197            vec![2.0, 9.0],
198            vec![2.0, 5.0],
199        ])
200        .unwrap();
201        let mut imp = SimpleImputer::new(ImputeStrategy::MostFrequent);
202        let out = imp.fit_transform(&x).unwrap();
203        // col0: 1,2,2 -> mode 2
204        assert!((imp.fill_values()[0] - 2.0).abs() < 1e-9);
205        assert!((out.get(0, 0) - 2.0).abs() < 1e-9);
206        // col1: 5,5,9,5 -> mode 5
207        assert!((imp.fill_values()[1] - 5.0).abs() < 1e-9);
208    }
209
210    #[test]
211    fn most_frequent_tie_smallest() {
212        let x = Matrix::new(vec![vec![nan()], vec![1.0], vec![2.0]]).unwrap();
213        let mut imp = SimpleImputer::new(ImputeStrategy::MostFrequent);
214        imp.fit(&x).unwrap();
215        // tie between 1 and 2 -> smallest wins
216        assert!((imp.fill_values()[0] - 1.0).abs() < 1e-9);
217    }
218
219    #[test]
220    fn constant_strategy() {
221        let mut imp = SimpleImputer::new(ImputeStrategy::Constant(-99.0));
222        let out = imp.fit_transform(&m_missing()).unwrap();
223        assert!((out.get(0, 2) - (-99.0)).abs() < 1e-9);
224        assert!((out.get(1, 1) - (-99.0)).abs() < 1e-9);
225        assert!((imp.fill_values()[0] - (-99.0)).abs() < 1e-9);
226    }
227
228    #[test]
229    fn no_missing_unchanged() {
230        let x = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0]]).unwrap();
231        let mut imp = SimpleImputer::new(ImputeStrategy::Mean);
232        let out = imp.fit_transform(&x).unwrap();
233        assert_eq!(out.rows_ref(), x.rows_ref());
234    }
235
236    #[test]
237    fn all_missing_column_errors() {
238        let x = Matrix::new(vec![vec![nan(), 1.0], vec![nan(), 2.0]]).unwrap();
239        let mut imp = SimpleImputer::new(ImputeStrategy::Mean);
240        let err = imp.fit(&x).unwrap_err();
241        assert!(matches!(err, DatarustError::AllMissing(_)));
242    }
243
244    #[test]
245    fn all_missing_median_errors() {
246        let x = Matrix::new(vec![vec![nan()], vec![nan()]]).unwrap();
247        let mut imp = SimpleImputer::new(ImputeStrategy::Median);
248        assert!(imp.fit(&x).is_err());
249    }
250
251    #[test]
252    fn constant_works_with_all_missing() {
253        // constant strategy fills even all-missing columns
254        let x = Matrix::new(vec![vec![nan()], vec![nan()]]).unwrap();
255        let mut imp = SimpleImputer::new(ImputeStrategy::Constant(0.0));
256        let out = imp.fit_transform(&x).unwrap();
257        for i in 0..2 {
258            assert!((out.get(i, 0) - 0.0).abs() < 1e-9);
259        }
260    }
261
262    #[test]
263    fn transform_before_fit_errors() {
264        let imp = SimpleImputer::new(ImputeStrategy::Mean);
265        assert!(matches!(
266            imp.transform(&m_missing()),
267            Err(DatarustError::NotFitted(_))
268        ));
269    }
270
271    #[test]
272    fn transform_new_data_uses_fitted() {
273        let mut imp = SimpleImputer::new(ImputeStrategy::Mean);
274        imp.fit(&m_missing()).unwrap();
275        let new = Matrix::new(vec![vec![nan(), nan(), nan()]]).unwrap();
276        let out = imp.transform(&new).unwrap();
277        assert!((out.get(0, 1) - (80.0 / 3.0)).abs() < 1e-9);
278    }
279}