Skip to main content

rill_ml/preprocessing/
constant_imputer.rs

1//! Constant value imputer for missing data.
2//!
3//! Replaces `NaN` values with a fixed fill value. Non-NaN values are
4//! passed through unchanged. This imputer has no learnable state.
5
6use crate::error::{RillError, checked_increment, ensure_finite};
7#[cfg(feature = "serde")]
8use crate::persistence::ValidateState;
9use crate::traits::Transformer;
10
11/// Configuration for [`ConstantImputer`].
12#[derive(Debug, Clone)]
13#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
14#[non_exhaustive]
15pub struct ConstantImputerConfig {
16    /// The value to replace NaN with.
17    pub fill_value: f64,
18}
19
20impl Default for ConstantImputerConfig {
21    fn default() -> Self {
22        Self { fill_value: 0.0 }
23    }
24}
25
26/// Replaces `NaN` values with a constant.
27///
28/// This transformer accepts `NaN` in its input, unlike most other
29/// transformers. The fill value must be finite.
30#[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    /// Create a new imputer for `feature_count` features with default config.
40    ///
41    /// # Errors
42    /// Returns [`RillError::EmptyFeatures`] if `feature_count` is `0`.
43    pub fn new(feature_count: usize) -> Result<Self, RillError> {
44        Self::with_config(feature_count, ConstantImputerConfig::default())
45    }
46
47    /// Create a new imputer with a custom configuration.
48    ///
49    /// # Errors
50    /// Returns [`RillError::EmptyFeatures`] if `feature_count` is `0`.
51    /// Returns [`RillError::NonFiniteValue`] if `fill_value` is not finite.
52    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    /// The fill value used to replace NaN.
68    pub fn fill_value(&self) -> f64 {
69        self.config.fill_value
70    }
71
72    /// Validate only the dimension, allowing NaN values.
73    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}