datarust/imputer/
simple.rs1use crate::error::{DatarustError, Result};
2use crate::matrix::Matrix;
3use crate::stats;
4use crate::traits::{default_input_names, FeatureNames};
5use crate::Transformer;
6
7#[derive(Debug, Clone, Default, PartialEq)]
9#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
10pub enum ImputeStrategy {
11 #[default]
13 Mean,
14 Median,
16 MostFrequent,
18 Constant(f64),
20}
21
22#[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 pub fn new(strategy: ImputeStrategy) -> Self {
36 Self {
37 strategy,
38 fill_values: vec![],
39 fitted: false,
40 }
41 }
42
43 pub fn strategy(&self) -> &ImputeStrategy {
45 &self.strategy
46 }
47
48 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 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 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 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 assert!((imp.fill_values()[0] - 2.0).abs() < 1e-9);
205 assert!((out.get(0, 0) - 2.0).abs() < 1e-9);
206 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 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 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}