Skip to main content

rill_ml/preprocessing/
mean_imputer.rs

1//! Mean imputer for missing data.
2//!
3//! Replaces `NaN` values with the running mean of observed non-NaN
4//! values for each feature. Uses Welford's algorithm for numerical
5//! stability.
6
7use crate::error::{RillError, checked_finite_add, checked_increment, ensure_finite};
8#[cfg(feature = "serde")]
9use crate::persistence::ValidateState;
10use crate::traits::Transformer;
11
12/// Replaces `NaN` values with the per-feature running mean.
13///
14/// When a feature has seen zero non-NaN values, `NaN` is replaced
15/// with `0.0`. This transformer accepts `NaN` in its input.
16#[derive(Debug, Clone)]
17#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
18pub struct MeanImputer {
19    feature_count: usize,
20    counts: Vec<u64>,
21    means: Vec<f64>,
22    samples_seen: u64,
23}
24
25impl MeanImputer {
26    /// Create a new imputer for `feature_count` features.
27    ///
28    /// # Errors
29    /// Returns [`RillError::EmptyFeatures`] if `feature_count` is `0`.
30    pub fn new(feature_count: usize) -> Result<Self, RillError> {
31        if feature_count == 0 {
32            return Err(RillError::EmptyFeatures);
33        }
34        Ok(Self {
35            feature_count,
36            counts: vec![0; feature_count],
37            means: vec![0.0; feature_count],
38            samples_seen: 0,
39        })
40    }
41
42    /// The per-feature running means of observed non-NaN values.
43    pub fn means(&self) -> &[f64] {
44        &self.means
45    }
46
47    /// The per-feature counts of observed non-NaN values.
48    pub fn counts(&self) -> &[u64] {
49        &self.counts
50    }
51
52    /// Validate only the dimension, allowing NaN values.
53    fn check_dimension(&self, features: &[f64]) -> Result<(), RillError> {
54        if features.is_empty() {
55            return Err(RillError::EmptyFeatures);
56        }
57        if features.len() != self.feature_count {
58            return Err(RillError::DimensionMismatch {
59                expected: self.feature_count,
60                actual: features.len(),
61            });
62        }
63        Ok(())
64    }
65
66    /// Update the running mean for feature `idx` using Welford's algorithm.
67    fn update_mean(&mut self, idx: usize, value: f64) -> Result<(), RillError> {
68        let n = checked_increment(self.counts[idx], "feature count")?;
69        self.counts[idx] = n;
70        let delta = value - self.means[idx];
71        ensure_finite("mean delta", delta)?;
72        self.means[idx] = checked_finite_add(self.means[idx], delta / n as f64, "mean")?;
73        Ok(())
74    }
75}
76
77#[cfg(feature = "serde")]
78impl ValidateState for MeanImputer {
79    fn validate_state(&self) -> Result<(), RillError> {
80        if self.feature_count == 0 {
81            return Err(RillError::EmptyFeatures);
82        }
83        if self.counts.len() != self.feature_count {
84            return Err(RillError::InvalidState(format!(
85                "mean imputer counts length {} does not match feature_count {}",
86                self.counts.len(),
87                self.feature_count
88            )));
89        }
90        if self.means.len() != self.feature_count {
91            return Err(RillError::InvalidState(format!(
92                "mean imputer means length {} does not match feature_count {}",
93                self.means.len(),
94                self.feature_count
95            )));
96        }
97        for &m in &self.means {
98            ensure_finite("mean", m)?;
99        }
100        Ok(())
101    }
102}
103
104impl Transformer for MeanImputer {
105    fn input_dim(&self) -> usize {
106        self.feature_count
107    }
108
109    fn output_dim(&self) -> usize {
110        self.feature_count
111    }
112
113    fn transform(&self, features: &[f64]) -> Result<Vec<f64>, RillError> {
114        self.check_dimension(features)?;
115        let mut out = Vec::with_capacity(features.len());
116        for (i, &x) in features.iter().enumerate() {
117            if x.is_nan() {
118                out.push(if self.counts[i] == 0 {
119                    0.0
120                } else {
121                    self.means[i]
122                });
123            } else {
124                ensure_finite("feature", x)?;
125                out.push(x);
126            }
127        }
128        Ok(out)
129    }
130
131    fn update(&mut self, features: &[f64]) -> Result<(), RillError> {
132        self.check_dimension(features)?;
133        for (i, &x) in features.iter().enumerate() {
134            if x.is_nan() {
135                continue;
136            }
137            ensure_finite("feature", x)?;
138            self.update_mean(i, x)?;
139        }
140        self.samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
141        Ok(())
142    }
143
144    fn samples_seen(&self) -> u64 {
145        self.samples_seen
146    }
147
148    fn reset(&mut self) {
149        for c in &mut self.counts {
150            *c = 0;
151        }
152        for m in &mut self.means {
153            *m = 0.0;
154        }
155        self.samples_seen = 0;
156    }
157}
158
159#[cfg(test)]
160mod tests {
161    use super::*;
162
163    #[test]
164    fn nan_replaced_with_mean_after_update() {
165        let mut imp = MeanImputer::new(2).unwrap();
166        // feature 0: observed [2.0, 4.0] -> mean 3.0
167        // feature 1: observed [10.0] -> mean 10.0
168        imp.update(&[2.0, 10.0]).unwrap();
169        imp.update(&[4.0, f64::NAN]).unwrap();
170        let out = imp.transform(&[f64::NAN, f64::NAN]).unwrap();
171        assert!((out[0] - 3.0).abs() < 1e-12);
172        assert!((out[1] - 10.0).abs() < 1e-12);
173    }
174
175    #[test]
176    fn nan_replaced_with_zero_when_no_data() {
177        let imp = MeanImputer::new(2).unwrap();
178        let out = imp.transform(&[f64::NAN, f64::NAN]).unwrap();
179        assert_eq!(out, vec![0.0, 0.0]);
180    }
181
182    #[test]
183    fn non_nan_passed_through() {
184        let mut imp = MeanImputer::new(2).unwrap();
185        imp.update(&[5.0, 6.0]).unwrap();
186        let out = imp.transform(&[1.5, -2.0]).unwrap();
187        assert_eq!(out, vec![1.5, -2.0]);
188    }
189
190    #[test]
191    fn mean_updates_correctly() {
192        let mut imp = MeanImputer::new(1).unwrap();
193        imp.update(&[1.0]).unwrap();
194        imp.update(&[2.0]).unwrap();
195        imp.update(&[3.0]).unwrap();
196        assert!((imp.means()[0] - 2.0).abs() < 1e-12);
197        assert_eq!(imp.counts()[0], 3);
198    }
199
200    #[test]
201    fn nan_skipped_in_update() {
202        let mut imp = MeanImputer::new(2).unwrap();
203        // feature 0: [1.0, NaN, 3.0] -> mean 2.0 (NaN skipped)
204        // feature 1: [NaN, NaN, NaN] -> count 0, mean 0.0
205        imp.update(&[1.0, f64::NAN]).unwrap();
206        imp.update(&[f64::NAN, f64::NAN]).unwrap();
207        imp.update(&[3.0, f64::NAN]).unwrap();
208        assert!((imp.means()[0] - 2.0).abs() < 1e-12);
209        assert_eq!(imp.counts()[0], 2);
210        assert_eq!(imp.counts()[1], 0);
211        assert!((imp.means()[1] - 0.0).abs() < 1e-12);
212    }
213
214    #[test]
215    fn dimension_mismatch_rejected() {
216        let imp = MeanImputer::new(3).unwrap();
217        assert!(matches!(
218            imp.transform(&[1.0, 2.0]),
219            Err(RillError::DimensionMismatch { .. })
220        ));
221        let mut imp = imp;
222        assert!(matches!(
223            imp.update(&[1.0, 2.0, 3.0, 4.0]),
224            Err(RillError::DimensionMismatch { .. })
225        ));
226    }
227
228    #[test]
229    fn reset_clears_state() {
230        let mut imp = MeanImputer::new(2).unwrap();
231        imp.update(&[1.0, 2.0]).unwrap();
232        imp.update(&[3.0, 4.0]).unwrap();
233        assert_eq!(imp.samples_seen(), 2);
234        assert_eq!(imp.counts()[0], 2);
235        imp.reset();
236        assert_eq!(imp.samples_seen(), 0);
237        assert_eq!(imp.counts()[0], 0);
238        assert!((imp.means()[0] - 0.0).abs() < 1e-12);
239    }
240
241    #[test]
242    #[cfg(feature = "serde")]
243    fn serde_roundtrip() {
244        let mut imp = MeanImputer::new(2).unwrap();
245        imp.update(&[1.0, f64::NAN]).unwrap();
246        imp.update(&[3.0, 5.0]).unwrap();
247        let json = serde_json::to_string(&imp).unwrap();
248        let restored: MeanImputer = serde_json::from_str(&json).unwrap();
249        assert_eq!(restored.input_dim(), imp.input_dim());
250        assert_eq!(restored.output_dim(), imp.output_dim());
251        assert_eq!(restored.samples_seen(), imp.samples_seen());
252        assert_eq!(restored.counts(), imp.counts());
253        assert!((restored.means()[0] - imp.means()[0]).abs() < 1e-12);
254    }
255}