Skip to main content

rill_ml/preprocessing/
forward_fill.rs

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