rill_ml/preprocessing/
constant_imputer.rs1use crate::error::{RillError, checked_increment, ensure_finite};
7#[cfg(feature = "serde")]
8use crate::persistence::ValidateState;
9use crate::traits::Transformer;
10
11#[derive(Debug, Clone)]
13#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
14#[non_exhaustive]
15pub struct ConstantImputerConfig {
16 pub fill_value: f64,
18}
19
20impl Default for ConstantImputerConfig {
21 fn default() -> Self {
22 Self { fill_value: 0.0 }
23 }
24}
25
26#[derive(Debug, Clone)]
31#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
32pub struct ConstantImputer {
33 feature_count: usize,
34 config: ConstantImputerConfig,
35 samples_seen: u64,
36}
37
38impl ConstantImputer {
39 pub fn new(feature_count: usize) -> Result<Self, RillError> {
44 Self::with_config(feature_count, ConstantImputerConfig::default())
45 }
46
47 pub fn with_config(
53 feature_count: usize,
54 config: ConstantImputerConfig,
55 ) -> Result<Self, RillError> {
56 if feature_count == 0 {
57 return Err(RillError::EmptyFeatures);
58 }
59 ensure_finite("fill_value", config.fill_value)?;
60 Ok(Self {
61 feature_count,
62 config,
63 samples_seen: 0,
64 })
65 }
66
67 pub fn fill_value(&self) -> f64 {
69 self.config.fill_value
70 }
71
72 fn check_dimension(&self, features: &[f64]) -> Result<(), RillError> {
74 if features.is_empty() {
75 return Err(RillError::EmptyFeatures);
76 }
77 if features.len() != self.feature_count {
78 return Err(RillError::DimensionMismatch {
79 expected: self.feature_count,
80 actual: features.len(),
81 });
82 }
83 Ok(())
84 }
85}
86
87#[cfg(feature = "serde")]
88impl ValidateState for ConstantImputer {
89 fn validate_state(&self) -> Result<(), RillError> {
90 if self.feature_count == 0 {
91 return Err(RillError::EmptyFeatures);
92 }
93 ensure_finite("fill_value", self.config.fill_value)?;
94 Ok(())
95 }
96}
97
98impl Transformer for ConstantImputer {
99 fn input_dim(&self) -> usize {
100 self.feature_count
101 }
102
103 fn output_dim(&self) -> usize {
104 self.feature_count
105 }
106
107 fn transform(&self, features: &[f64]) -> Result<Vec<f64>, RillError> {
108 self.check_dimension(features)?;
109 let mut out = Vec::with_capacity(features.len());
110 for &x in features {
111 if x.is_nan() {
112 out.push(self.config.fill_value);
113 } else {
114 ensure_finite("feature", x)?;
115 out.push(x);
116 }
117 }
118 Ok(out)
119 }
120
121 fn update(&mut self, features: &[f64]) -> Result<(), RillError> {
122 self.check_dimension(features)?;
123 for &x in features {
124 if !x.is_nan() {
125 ensure_finite("feature", x)?;
126 }
127 }
128 self.samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
129 Ok(())
130 }
131
132 fn samples_seen(&self) -> u64 {
133 self.samples_seen
134 }
135
136 fn reset(&mut self) {
137 self.samples_seen = 0;
138 }
139}
140
141#[cfg(test)]
142mod tests {
143 use super::*;
144
145 #[test]
146 fn nan_replaced_with_fill_value() {
147 let imp = ConstantImputer::new(3).unwrap();
148 let out = imp.transform(&[1.0, f64::NAN, 3.0]).unwrap();
149 assert_eq!(out, vec![1.0, 0.0, 3.0]);
150 }
151
152 #[test]
153 fn non_nan_passed_through() {
154 let imp = ConstantImputer::new(3).unwrap();
155 let out = imp.transform(&[1.5, -2.0, 3.0]).unwrap();
156 assert_eq!(out, vec![1.5, -2.0, 3.0]);
157 }
158
159 #[test]
160 fn custom_fill_value() {
161 let imp =
162 ConstantImputer::with_config(2, ConstantImputerConfig { fill_value: -1.0 }).unwrap();
163 let out = imp.transform(&[f64::NAN, 5.0]).unwrap();
164 assert_eq!(out, vec![-1.0, 5.0]);
165 }
166
167 #[test]
168 fn dimension_mismatch_rejected() {
169 let imp = ConstantImputer::new(3).unwrap();
170 assert!(matches!(
171 imp.transform(&[1.0, 2.0]),
172 Err(RillError::DimensionMismatch { .. })
173 ));
174 let mut imp = imp;
175 assert!(matches!(
176 imp.update(&[1.0, 2.0, 3.0, 4.0]),
177 Err(RillError::DimensionMismatch { .. })
178 ));
179 }
180
181 #[test]
182 fn reset_clears_state() {
183 let mut imp = ConstantImputer::new(2).unwrap();
184 imp.update(&[1.0, f64::NAN]).unwrap();
185 imp.update(&[f64::NAN, 2.0]).unwrap();
186 assert_eq!(imp.samples_seen(), 2);
187 imp.reset();
188 assert_eq!(imp.samples_seen(), 0);
189 }
190
191 #[test]
192 #[cfg(feature = "serde")]
193 fn serde_roundtrip() {
194 let mut imp =
195 ConstantImputer::with_config(3, ConstantImputerConfig { fill_value: 7.0 }).unwrap();
196 imp.update(&[1.0, f64::NAN, 3.0]).unwrap();
197 let json = serde_json::to_string(&imp).unwrap();
198 let restored: ConstantImputer = serde_json::from_str(&json).unwrap();
199 assert_eq!(restored.input_dim(), imp.input_dim());
200 assert_eq!(restored.output_dim(), imp.output_dim());
201 assert_eq!(restored.samples_seen(), imp.samples_seen());
202 assert_eq!(restored.fill_value(), imp.fill_value());
203 }
204}