Skip to main content

rill_ml/preprocessing/
missing_indicator.rs

1//! Missing value indicator transformer.
2//!
3//! For each input feature, appends a binary indicator (1.0 if NaN, 0.0
4//! otherwise). Output dimension = 2 * input dimension.
5//!
6//! This transformer accepts NaN values in its input, unlike most other
7//! transformers.
8
9use crate::error::{RillError, checked_increment, ensure_finite};
10#[cfg(feature = "serde")]
11use crate::persistence::ValidateState;
12use crate::traits::Transformer;
13
14/// Adds a missing-value indicator for each feature.
15///
16/// `transform([x1, x2, ...])` produces `[x1, is_nan(x1) as f64, x2, is_nan(x2) as f64, ...]`.
17/// The original values are preserved (including NaN), with the indicator
18/// appended after each value.
19#[derive(Debug, Clone)]
20#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
21pub struct MissingIndicator {
22    feature_count: usize,
23    samples_seen: u64,
24}
25
26impl MissingIndicator {
27    /// Create a new indicator for `feature_count` features.
28    ///
29    /// # Errors
30    /// Returns [`RillError::EmptyFeatures`] if `feature_count` is `0`.
31    pub fn new(feature_count: usize) -> Result<Self, RillError> {
32        if feature_count == 0 {
33            return Err(RillError::EmptyFeatures);
34        }
35        Ok(Self {
36            feature_count,
37            samples_seen: 0,
38        })
39    }
40
41    /// Validate only the dimension, allowing NaN values.
42    fn check_dimension(&self, features: &[f64]) -> Result<(), RillError> {
43        if features.is_empty() {
44            return Err(RillError::EmptyFeatures);
45        }
46        if features.len() != self.feature_count {
47            return Err(RillError::DimensionMismatch {
48                expected: self.feature_count,
49                actual: features.len(),
50            });
51        }
52        Ok(())
53    }
54}
55
56#[cfg(feature = "serde")]
57impl ValidateState for MissingIndicator {
58    fn validate_state(&self) -> Result<(), RillError> {
59        if self.feature_count == 0 {
60            return Err(RillError::EmptyFeatures);
61        }
62        Ok(())
63    }
64}
65
66impl Transformer for MissingIndicator {
67    fn input_dim(&self) -> usize {
68        self.feature_count
69    }
70
71    fn output_dim(&self) -> usize {
72        self.feature_count * 2
73    }
74
75    fn transform(&self, features: &[f64]) -> Result<Vec<f64>, RillError> {
76        self.check_dimension(features)?;
77        let mut out = Vec::with_capacity(features.len() * 2);
78        for &v in features {
79            if !v.is_nan() {
80                ensure_finite("feature", v)?;
81            }
82            out.push(v);
83            out.push(if v.is_nan() { 1.0 } else { 0.0 });
84        }
85        Ok(out)
86    }
87
88    fn update(&mut self, features: &[f64]) -> Result<(), RillError> {
89        self.check_dimension(features)?;
90        for &v in features {
91            if !v.is_nan() {
92                ensure_finite("feature", v)?;
93            }
94        }
95        self.samples_seen = checked_increment(self.samples_seen, "samples_seen")?;
96        Ok(())
97    }
98
99    fn samples_seen(&self) -> u64 {
100        self.samples_seen
101    }
102
103    fn reset(&mut self) {
104        self.samples_seen = 0;
105    }
106}
107
108#[cfg(test)]
109mod tests {
110    use super::*;
111
112    #[test]
113    fn no_missing_values() {
114        let mi = MissingIndicator::new(3).unwrap();
115        let out = mi.transform(&[1.0, 2.0, 3.0]).unwrap();
116        assert_eq!(out, vec![1.0, 0.0, 2.0, 0.0, 3.0, 0.0]);
117    }
118
119    #[test]
120    fn all_missing_values() {
121        let mi = MissingIndicator::new(2).unwrap();
122        let out = mi.transform(&[f64::NAN, f64::NAN]).unwrap();
123        assert!(out[0].is_nan());
124        assert_eq!(out[1], 1.0);
125        assert!(out[2].is_nan());
126        assert_eq!(out[3], 1.0);
127    }
128
129    #[test]
130    fn mixed_missing_and_present() {
131        let mi = MissingIndicator::new(3).unwrap();
132        let out = mi.transform(&[1.0, f64::NAN, 3.0]).unwrap();
133        assert_eq!(out[0], 1.0);
134        assert_eq!(out[1], 0.0);
135        assert!(out[2].is_nan());
136        assert_eq!(out[3], 1.0);
137        assert_eq!(out[4], 3.0);
138        assert_eq!(out[5], 0.0);
139    }
140
141    #[test]
142    fn dimension_mismatch_rejected() {
143        let mi = MissingIndicator::new(3).unwrap();
144        assert!(matches!(
145            mi.transform(&[1.0, 2.0]),
146            Err(RillError::DimensionMismatch { .. })
147        ));
148        let mut mi = mi;
149        assert!(matches!(
150            mi.update(&[1.0, 2.0, 3.0, 4.0]),
151            Err(RillError::DimensionMismatch { .. })
152        ));
153    }
154
155    #[test]
156    fn reset_clears_state() {
157        let mut mi = MissingIndicator::new(2).unwrap();
158        mi.update(&[1.0, 2.0]).unwrap();
159        mi.update(&[3.0, 4.0]).unwrap();
160        assert_eq!(mi.samples_seen(), 2);
161        mi.reset();
162        assert_eq!(mi.samples_seen(), 0);
163    }
164
165    #[test]
166    #[cfg(feature = "serde")]
167    fn serde_roundtrip() {
168        let mut mi = MissingIndicator::new(3).unwrap();
169        mi.update(&[1.0, 2.0, 3.0]).unwrap();
170        let json = serde_json::to_string(&mi).unwrap();
171        let restored: MissingIndicator = serde_json::from_str(&json).unwrap();
172        assert_eq!(restored.input_dim(), mi.input_dim());
173        assert_eq!(restored.output_dim(), mi.output_dim());
174        assert_eq!(restored.samples_seen(), mi.samples_seen());
175    }
176}