rill_ml/preprocessing/
mean_imputer.rs1use crate::error::{RillError, checked_finite_add, 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 MeanImputer {
19 feature_count: usize,
20 counts: Vec<u64>,
21 means: Vec<f64>,
22 samples_seen: u64,
23}
24
25impl MeanImputer {
26 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 pub fn means(&self) -> &[f64] {
44 &self.means
45 }
46
47 pub fn counts(&self) -> &[u64] {
49 &self.counts
50 }
51
52 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 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 self.counts.fill(0);
150 self.means.fill(0.0);
151 self.samples_seen = 0;
152 }
153}
154
155#[cfg(test)]
156mod tests {
157 use super::*;
158
159 #[test]
160 fn nan_replaced_with_mean_after_update() {
161 let mut imp = MeanImputer::new(2).unwrap();
162 imp.update(&[2.0, 10.0]).unwrap();
165 imp.update(&[4.0, f64::NAN]).unwrap();
166 let out = imp.transform(&[f64::NAN, f64::NAN]).unwrap();
167 assert!((out[0] - 3.0).abs() < 1e-12);
168 assert!((out[1] - 10.0).abs() < 1e-12);
169 }
170
171 #[test]
172 fn nan_replaced_with_zero_when_no_data() {
173 let imp = MeanImputer::new(2).unwrap();
174 let out = imp.transform(&[f64::NAN, f64::NAN]).unwrap();
175 assert_eq!(out, vec![0.0, 0.0]);
176 }
177
178 #[test]
179 fn non_nan_passed_through() {
180 let mut imp = MeanImputer::new(2).unwrap();
181 imp.update(&[5.0, 6.0]).unwrap();
182 let out = imp.transform(&[1.5, -2.0]).unwrap();
183 assert_eq!(out, vec![1.5, -2.0]);
184 }
185
186 #[test]
187 fn mean_updates_correctly() {
188 let mut imp = MeanImputer::new(1).unwrap();
189 imp.update(&[1.0]).unwrap();
190 imp.update(&[2.0]).unwrap();
191 imp.update(&[3.0]).unwrap();
192 assert!((imp.means()[0] - 2.0).abs() < 1e-12);
193 assert_eq!(imp.counts()[0], 3);
194 }
195
196 #[test]
197 fn nan_skipped_in_update() {
198 let mut imp = MeanImputer::new(2).unwrap();
199 imp.update(&[1.0, f64::NAN]).unwrap();
202 imp.update(&[f64::NAN, f64::NAN]).unwrap();
203 imp.update(&[3.0, f64::NAN]).unwrap();
204 assert!((imp.means()[0] - 2.0).abs() < 1e-12);
205 assert_eq!(imp.counts()[0], 2);
206 assert_eq!(imp.counts()[1], 0);
207 assert!((imp.means()[1] - 0.0).abs() < 1e-12);
208 }
209
210 #[test]
211 fn dimension_mismatch_rejected() {
212 let imp = MeanImputer::new(3).unwrap();
213 assert!(matches!(
214 imp.transform(&[1.0, 2.0]),
215 Err(RillError::DimensionMismatch { .. })
216 ));
217 let mut imp = imp;
218 assert!(matches!(
219 imp.update(&[1.0, 2.0, 3.0, 4.0]),
220 Err(RillError::DimensionMismatch { .. })
221 ));
222 }
223
224 #[test]
225 fn reset_clears_state() {
226 let mut imp = MeanImputer::new(2).unwrap();
227 imp.update(&[1.0, 2.0]).unwrap();
228 imp.update(&[3.0, 4.0]).unwrap();
229 assert_eq!(imp.samples_seen(), 2);
230 assert_eq!(imp.counts()[0], 2);
231 imp.reset();
232 assert_eq!(imp.samples_seen(), 0);
233 assert_eq!(imp.counts()[0], 0);
234 assert!((imp.means()[0] - 0.0).abs() < 1e-12);
235 }
236
237 #[test]
238 #[cfg(feature = "serde")]
239 fn serde_roundtrip() {
240 let mut imp = MeanImputer::new(2).unwrap();
241 imp.update(&[1.0, f64::NAN]).unwrap();
242 imp.update(&[3.0, 5.0]).unwrap();
243 let json = serde_json::to_string(&imp).unwrap();
244 let restored: MeanImputer = serde_json::from_str(&json).unwrap();
245 assert_eq!(restored.input_dim(), imp.input_dim());
246 assert_eq!(restored.output_dim(), imp.output_dim());
247 assert_eq!(restored.samples_seen(), imp.samples_seen());
248 assert_eq!(restored.counts(), imp.counts());
249 assert!((restored.means()[0] - imp.means()[0]).abs() < 1e-12);
250 }
251}