rill_ml/preprocessing/
forward_fill.rs1use crate::error::{RillError, checked_increment, ensure_finite};
8#[cfg(feature = "serde")]
9use crate::persistence::ValidateState;
10use crate::traits::Transformer;
11
12#[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 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 pub fn last_values(&self) -> &[Option<f64>] {
42 &self.last_values
43 }
44
45 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}