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 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}