rill_ml/preprocessing/
missing_indicator.rs1use crate::error::{RillError, checked_increment, ensure_finite};
10#[cfg(feature = "serde")]
11use crate::persistence::ValidateState;
12use crate::traits::Transformer;
13
14#[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 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 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}