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        self.last_values.fill(None);
122        self.samples_seen = 0;
123    }
124}
125
126#[cfg(test)]
127mod tests {
128    use super::*;
129
130    #[test]
131    fn nan_replaced_with_last_value() {
132        let mut ff = ForwardFill::new(2).unwrap();
133        ff.update(&[1.0, 2.0]).unwrap();
134        ff.update(&[3.0, 4.0]).unwrap();
135        let out = ff.transform(&[f64::NAN, f64::NAN]).unwrap();
136        assert_eq!(out, vec![3.0, 4.0]);
137    }
138
139    #[test]
140    fn nan_replaced_with_zero_when_no_data() {
141        let ff = ForwardFill::new(2).unwrap();
142        let out = ff.transform(&[f64::NAN, f64::NAN]).unwrap();
143        assert_eq!(out, vec![0.0, 0.0]);
144    }
145
146    #[test]
147    fn non_nan_passed_through() {
148        let mut ff = ForwardFill::new(2).unwrap();
149        ff.update(&[5.0, 6.0]).unwrap();
150        let out = ff.transform(&[1.5, -2.0]).unwrap();
151        assert_eq!(out, vec![1.5, -2.0]);
152    }
153
154    #[test]
155    fn last_value_updates_correctly() {
156        let mut ff = ForwardFill::new(2).unwrap();
157        ff.update(&[1.0, 10.0]).unwrap();
158        assert_eq!(ff.last_values()[0], Some(1.0));
159        assert_eq!(ff.last_values()[1], Some(10.0));
160        ff.update(&[2.0, f64::NAN]).unwrap();
161        assert_eq!(ff.last_values()[0], Some(2.0));
162        assert_eq!(ff.last_values()[1], Some(10.0));
163    }
164
165    #[test]
166    fn nan_skipped_in_update() {
167        let mut ff = ForwardFill::new(2).unwrap();
168        ff.update(&[f64::NAN, 5.0]).unwrap();
169        assert_eq!(ff.last_values()[0], None);
170        assert_eq!(ff.last_values()[1], Some(5.0));
171        ff.update(&[7.0, f64::NAN]).unwrap();
172        assert_eq!(ff.last_values()[0], Some(7.0));
173        assert_eq!(ff.last_values()[1], Some(5.0));
174    }
175
176    #[test]
177    fn dimension_mismatch_rejected() {
178        let ff = ForwardFill::new(3).unwrap();
179        assert!(matches!(
180            ff.transform(&[1.0, 2.0]),
181            Err(RillError::DimensionMismatch { .. })
182        ));
183        let mut ff = ff;
184        assert!(matches!(
185            ff.update(&[1.0, 2.0, 3.0, 4.0]),
186            Err(RillError::DimensionMismatch { .. })
187        ));
188    }
189
190    #[test]
191    fn reset_clears_state() {
192        let mut ff = ForwardFill::new(2).unwrap();
193        ff.update(&[1.0, 2.0]).unwrap();
194        ff.update(&[3.0, 4.0]).unwrap();
195        assert_eq!(ff.samples_seen(), 2);
196        ff.reset();
197        assert_eq!(ff.samples_seen(), 0);
198        assert_eq!(ff.last_values()[0], None);
199        assert_eq!(ff.last_values()[1], None);
200    }
201
202    #[test]
203    #[cfg(feature = "serde")]
204    fn serde_roundtrip() {
205        let mut ff = ForwardFill::new(2).unwrap();
206        ff.update(&[1.0, f64::NAN]).unwrap();
207        ff.update(&[3.0, 5.0]).unwrap();
208        let json = serde_json::to_string(&ff).unwrap();
209        let restored: ForwardFill = serde_json::from_str(&json).unwrap();
210        assert_eq!(restored.input_dim(), ff.input_dim());
211        assert_eq!(restored.output_dim(), ff.output_dim());
212        assert_eq!(restored.samples_seen(), ff.samples_seen());
213        assert_eq!(restored.last_values(), ff.last_values());
214    }
215}