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 for c in &mut self.counts {
150 *c = 0;
151 }
152 for m in &mut self.means {
153 *m = 0.0;
154 }
155 self.samples_seen = 0;
156 }
157}
158
159#[cfg(test)]
160mod tests {
161 use super::*;
162
163 #[test]
164 fn nan_replaced_with_mean_after_update() {
165 let mut imp = MeanImputer::new(2).unwrap();
166 imp.update(&[2.0, 10.0]).unwrap();
169 imp.update(&[4.0, f64::NAN]).unwrap();
170 let out = imp.transform(&[f64::NAN, f64::NAN]).unwrap();
171 assert!((out[0] - 3.0).abs() < 1e-12);
172 assert!((out[1] - 10.0).abs() < 1e-12);
173 }
174
175 #[test]
176 fn nan_replaced_with_zero_when_no_data() {
177 let imp = MeanImputer::new(2).unwrap();
178 let out = imp.transform(&[f64::NAN, f64::NAN]).unwrap();
179 assert_eq!(out, vec![0.0, 0.0]);
180 }
181
182 #[test]
183 fn non_nan_passed_through() {
184 let mut imp = MeanImputer::new(2).unwrap();
185 imp.update(&[5.0, 6.0]).unwrap();
186 let out = imp.transform(&[1.5, -2.0]).unwrap();
187 assert_eq!(out, vec![1.5, -2.0]);
188 }
189
190 #[test]
191 fn mean_updates_correctly() {
192 let mut imp = MeanImputer::new(1).unwrap();
193 imp.update(&[1.0]).unwrap();
194 imp.update(&[2.0]).unwrap();
195 imp.update(&[3.0]).unwrap();
196 assert!((imp.means()[0] - 2.0).abs() < 1e-12);
197 assert_eq!(imp.counts()[0], 3);
198 }
199
200 #[test]
201 fn nan_skipped_in_update() {
202 let mut imp = MeanImputer::new(2).unwrap();
203 imp.update(&[1.0, f64::NAN]).unwrap();
206 imp.update(&[f64::NAN, f64::NAN]).unwrap();
207 imp.update(&[3.0, f64::NAN]).unwrap();
208 assert!((imp.means()[0] - 2.0).abs() < 1e-12);
209 assert_eq!(imp.counts()[0], 2);
210 assert_eq!(imp.counts()[1], 0);
211 assert!((imp.means()[1] - 0.0).abs() < 1e-12);
212 }
213
214 #[test]
215 fn dimension_mismatch_rejected() {
216 let imp = MeanImputer::new(3).unwrap();
217 assert!(matches!(
218 imp.transform(&[1.0, 2.0]),
219 Err(RillError::DimensionMismatch { .. })
220 ));
221 let mut imp = imp;
222 assert!(matches!(
223 imp.update(&[1.0, 2.0, 3.0, 4.0]),
224 Err(RillError::DimensionMismatch { .. })
225 ));
226 }
227
228 #[test]
229 fn reset_clears_state() {
230 let mut imp = MeanImputer::new(2).unwrap();
231 imp.update(&[1.0, 2.0]).unwrap();
232 imp.update(&[3.0, 4.0]).unwrap();
233 assert_eq!(imp.samples_seen(), 2);
234 assert_eq!(imp.counts()[0], 2);
235 imp.reset();
236 assert_eq!(imp.samples_seen(), 0);
237 assert_eq!(imp.counts()[0], 0);
238 assert!((imp.means()[0] - 0.0).abs() < 1e-12);
239 }
240
241 #[test]
242 #[cfg(feature = "serde")]
243 fn serde_roundtrip() {
244 let mut imp = MeanImputer::new(2).unwrap();
245 imp.update(&[1.0, f64::NAN]).unwrap();
246 imp.update(&[3.0, 5.0]).unwrap();
247 let json = serde_json::to_string(&imp).unwrap();
248 let restored: MeanImputer = serde_json::from_str(&json).unwrap();
249 assert_eq!(restored.input_dim(), imp.input_dim());
250 assert_eq!(restored.output_dim(), imp.output_dim());
251 assert_eq!(restored.samples_seen(), imp.samples_seen());
252 assert_eq!(restored.counts(), imp.counts());
253 assert!((restored.means()[0] - imp.means()[0]).abs() < 1e-12);
254 }
255}