Skip to main content

sklears_utils/
preprocessing.rs

1//! Data preprocessing utilities for machine learning
2//!
3//! This module provides utilities for data cleaning, outlier detection,
4//! data transformation, and quality assessment.
5
6use crate::{UtilsError, UtilsResult};
7use scirs2_core::ndarray::{Array1, Array2, ArrayView1, Axis};
8use scirs2_core::numeric::Float;
9use std::cmp::Ordering;
10use std::collections::HashMap;
11
12/// Helper function to safely compare floats, treating NaN as greater than all other values
13#[inline]
14fn compare_floats<T: Float>(a: &T, b: &T) -> Ordering {
15    match a.partial_cmp(b) {
16        Some(ord) => ord,
17        None => {
18            // Handle NaN cases: NaN is treated as greater than any number
19            if a.is_nan() && b.is_nan() {
20                Ordering::Equal
21            } else if a.is_nan() {
22                Ordering::Greater
23            } else {
24                Ordering::Less
25            }
26        }
27    }
28}
29
30/// Helper function to safely convert usize to Float type
31#[inline]
32fn usize_to_float<T: Float>(value: usize) -> UtilsResult<T> {
33    T::from(value).ok_or_else(|| {
34        UtilsError::InvalidParameter(format!("Failed to convert usize {} to float type", value))
35    })
36}
37
38/// Helper function to safely convert constant to Float type
39#[inline]
40fn const_to_float<T: Float>(value: f64) -> UtilsResult<T> {
41    T::from(value).ok_or_else(|| {
42        UtilsError::InvalidParameter(format!(
43            "Failed to convert constant {} to float type",
44            value
45        ))
46    })
47}
48
49/// Data cleaning utilities
50pub struct DataCleaner;
51
52impl DataCleaner {
53    /// Remove rows with missing values (NaN)
54    pub fn drop_missing_rows<T>(data: &Array2<T>) -> UtilsResult<Array2<T>>
55    where
56        T: Float + Clone + std::iter::Sum,
57    {
58        let mut valid_rows = Vec::new();
59
60        for (row_idx, row) in data.axis_iter(Axis(0)).enumerate() {
61            if !row.iter().any(|&x| x.is_nan()) {
62                valid_rows.push(row_idx);
63            }
64        }
65
66        if valid_rows.is_empty() {
67            return Err(UtilsError::EmptyInput);
68        }
69
70        let mut result = Array2::zeros((valid_rows.len(), data.ncols()));
71        for (new_idx, &old_idx) in valid_rows.iter().enumerate() {
72            result.row_mut(new_idx).assign(&data.row(old_idx));
73        }
74
75        Ok(result)
76    }
77
78    /// Fill missing values with specified value
79    pub fn fill_missing<T>(data: &mut Array2<T>, fill_value: T)
80    where
81        T: Float + Clone + std::iter::Sum,
82    {
83        data.mapv_inplace(|x| if x.is_nan() { fill_value } else { x });
84    }
85
86    /// Fill missing values with column means
87    pub fn fill_with_mean<T>(data: &mut Array2<T>) -> UtilsResult<()>
88    where
89        T: Float + Clone + std::iter::Sum,
90    {
91        for col_idx in 0..data.ncols() {
92            let col = data.column(col_idx);
93            let valid_values: Vec<T> = col.iter().cloned().filter(|x| !x.is_nan()).collect();
94
95            if !valid_values.is_empty() {
96                let mean = valid_values.iter().cloned().sum::<T>()
97                    / usize_to_float::<T>(valid_values.len())?;
98
99                for row_idx in 0..data.nrows() {
100                    if data[[row_idx, col_idx]].is_nan() {
101                        data[[row_idx, col_idx]] = mean;
102                    }
103                }
104            }
105        }
106        Ok(())
107    }
108
109    /// Fill missing values with column medians
110    pub fn fill_with_median<T>(data: &mut Array2<T>) -> UtilsResult<()>
111    where
112        T: Float + Clone + PartialOrd,
113    {
114        for col_idx in 0..data.ncols() {
115            let col = data.column(col_idx);
116            let mut valid_values: Vec<T> = col.iter().cloned().filter(|x| !x.is_nan()).collect();
117
118            if !valid_values.is_empty() {
119                valid_values.sort_by(compare_floats);
120                let median = if valid_values.len().is_multiple_of(2) {
121                    let mid = valid_values.len() / 2;
122                    (valid_values[mid - 1] + valid_values[mid]) / const_to_float::<T>(2.0)?
123                } else {
124                    valid_values[valid_values.len() / 2]
125                };
126
127                for row_idx in 0..data.nrows() {
128                    if data[[row_idx, col_idx]].is_nan() {
129                        data[[row_idx, col_idx]] = median;
130                    }
131                }
132            }
133        }
134        Ok(())
135    }
136}
137
138/// Outlier detection methods
139pub struct OutlierDetector;
140
141impl OutlierDetector {
142    /// Detect outliers using Z-score method
143    ///
144    /// Returns an empty vector if conversion fails or if standard deviation is zero.
145    pub fn zscore_outliers<T>(data: &ArrayView1<T>, threshold: T) -> Vec<usize>
146    where
147        T: Float + Clone + std::iter::Sum,
148    {
149        // Safe conversion - return empty on failure
150        let Ok(len_float) = usize_to_float::<T>(data.len()) else {
151            return Vec::new();
152        };
153
154        let mean = data.iter().cloned().sum::<T>() / len_float;
155        let variance = data.iter().map(|&x| (x - mean).powi(2)).sum::<T>() / len_float;
156        let std_dev = variance.sqrt();
157
158        if std_dev == T::zero() {
159            return Vec::new();
160        }
161
162        data.iter()
163            .enumerate()
164            .filter_map(|(idx, &value)| {
165                let z_score = (value - mean).abs() / std_dev;
166                if z_score > threshold {
167                    Some(idx)
168                } else {
169                    None
170                }
171            })
172            .collect()
173    }
174
175    /// Detect outliers using IQR (Interquartile Range) method
176    pub fn iqr_outliers<T>(data: &ArrayView1<T>, multiplier: T) -> Vec<usize>
177    where
178        T: Float + Clone + PartialOrd,
179    {
180        let mut sorted_data: Vec<T> = data.iter().cloned().collect();
181        sorted_data.sort_by(compare_floats);
182
183        let n = sorted_data.len();
184        if n < 4 {
185            return Vec::new();
186        }
187
188        let q1_idx = n / 4;
189        let q3_idx = 3 * n / 4;
190        let q1 = sorted_data[q1_idx];
191        let q3 = sorted_data[q3_idx];
192        let iqr = q3 - q1;
193
194        let lower_bound = q1 - multiplier * iqr;
195        let upper_bound = q3 + multiplier * iqr;
196
197        data.iter()
198            .enumerate()
199            .filter_map(|(idx, &value)| {
200                if value < lower_bound || value > upper_bound {
201                    Some(idx)
202                } else {
203                    None
204                }
205            })
206            .collect()
207    }
208
209    /// Detect outliers using modified Z-score method (using median)
210    ///
211    /// Returns an empty vector if conversion fails or if MAD is zero.
212    pub fn modified_zscore_outliers<T>(data: &ArrayView1<T>, threshold: T) -> Vec<usize>
213    where
214        T: Float + Clone + PartialOrd,
215    {
216        let mut sorted_data: Vec<T> = data.iter().cloned().collect();
217        sorted_data.sort_by(compare_floats);
218
219        let n = sorted_data.len();
220        if n == 0 {
221            return Vec::new();
222        }
223
224        // Safe conversion - return empty on failure
225        let Ok(two) = const_to_float::<T>(2.0) else {
226            return Vec::new();
227        };
228
229        let median = if n.is_multiple_of(2) {
230            (sorted_data[n / 2 - 1] + sorted_data[n / 2]) / two
231        } else {
232            sorted_data[n / 2]
233        };
234
235        // Calculate MAD (Median Absolute Deviation)
236        let mut deviations: Vec<T> = data.iter().map(|&x| (x - median).abs()).collect();
237        deviations.sort_by(compare_floats);
238
239        let mad = if deviations.len().is_multiple_of(2) {
240            let mid = deviations.len() / 2;
241            (deviations[mid - 1] + deviations[mid]) / two
242        } else {
243            deviations[deviations.len() / 2]
244        };
245
246        if mad == T::zero() {
247            return Vec::new();
248        }
249
250        // Safe conversion for scale factors
251        let Ok(scale_factor) = const_to_float::<T>(1.4826) else {
252            return Vec::new();
253        };
254        let Ok(modified_z_factor) = const_to_float::<T>(0.6745) else {
255            return Vec::new();
256        };
257
258        let mad_scaled = mad * scale_factor;
259
260        data.iter()
261            .enumerate()
262            .filter_map(|(idx, &value)| {
263                let modified_z = modified_z_factor * (value - median).abs() / mad_scaled;
264                if modified_z > threshold {
265                    Some(idx)
266                } else {
267                    None
268                }
269            })
270            .collect()
271    }
272}
273
274/// Feature scaling utilities
275pub struct FeatureScaler;
276
277impl FeatureScaler {
278    /// Standard scaling (z-score normalization)
279    pub fn standard_scale<T>(data: &Array2<T>) -> UtilsResult<(Array2<T>, Array1<T>, Array1<T>)>
280    where
281        T: Float + Clone + std::iter::Sum,
282    {
283        let mut scaled_data = data.clone();
284        let mut means = Array1::zeros(data.ncols());
285        let mut stds = Array1::zeros(data.ncols());
286
287        for col_idx in 0..data.ncols() {
288            let col = data.column(col_idx);
289            let col_len = usize_to_float::<T>(col.len())?;
290            let mean = col.iter().cloned().sum::<T>() / col_len;
291            let variance = col.iter().map(|&x| (x - mean).powi(2)).sum::<T>() / col_len;
292            let std_dev = variance.sqrt();
293
294            means[col_idx] = mean;
295            stds[col_idx] = std_dev;
296
297            if std_dev != T::zero() {
298                for row_idx in 0..data.nrows() {
299                    scaled_data[[row_idx, col_idx]] = (data[[row_idx, col_idx]] - mean) / std_dev;
300                }
301            }
302        }
303
304        Ok((scaled_data, means, stds))
305    }
306
307    /// Min-max scaling to [0, 1] range
308    pub fn minmax_scale<T>(data: &Array2<T>) -> UtilsResult<(Array2<T>, Array1<T>, Array1<T>)>
309    where
310        T: Float + Clone + PartialOrd,
311    {
312        let mut scaled_data = data.clone();
313        let mut mins = Array1::zeros(data.ncols());
314        let mut maxs = Array1::zeros(data.ncols());
315
316        for col_idx in 0..data.ncols() {
317            let col = data.column(col_idx);
318            let min_val = col
319                .iter()
320                .cloned()
321                .fold(col[0], |acc, x| if x < acc { x } else { acc });
322            let max_val = col
323                .iter()
324                .cloned()
325                .fold(col[0], |acc, x| if x > acc { x } else { acc });
326
327            mins[col_idx] = min_val;
328            maxs[col_idx] = max_val;
329
330            let range = max_val - min_val;
331            if range != T::zero() {
332                for row_idx in 0..data.nrows() {
333                    scaled_data[[row_idx, col_idx]] = (data[[row_idx, col_idx]] - min_val) / range;
334                }
335            }
336        }
337
338        Ok((scaled_data, mins, maxs))
339    }
340
341    /// Robust scaling using median and IQR
342    pub fn robust_scale<T>(data: &Array2<T>) -> UtilsResult<(Array2<T>, Array1<T>, Array1<T>)>
343    where
344        T: Float + Clone + PartialOrd,
345    {
346        let mut scaled_data = data.clone();
347        let mut medians = Array1::zeros(data.ncols());
348        let mut iqrs = Array1::zeros(data.ncols());
349
350        for col_idx in 0..data.ncols() {
351            let col = data.column(col_idx);
352            let mut sorted_col: Vec<T> = col.iter().cloned().collect();
353            sorted_col.sort_by(compare_floats);
354
355            let n = sorted_col.len();
356            let median = if n.is_multiple_of(2) {
357                (sorted_col[n / 2 - 1] + sorted_col[n / 2]) / const_to_float::<T>(2.0)?
358            } else {
359                sorted_col[n / 2]
360            };
361
362            let q1_idx = n / 4;
363            let q3_idx = 3 * n / 4;
364            let q1 = sorted_col[q1_idx];
365            let q3 = sorted_col[q3_idx];
366            let iqr = q3 - q1;
367
368            medians[col_idx] = median;
369            iqrs[col_idx] = iqr;
370
371            if iqr != T::zero() {
372                for row_idx in 0..data.nrows() {
373                    scaled_data[[row_idx, col_idx]] = (data[[row_idx, col_idx]] - median) / iqr;
374                }
375            }
376        }
377
378        Ok((scaled_data, medians, iqrs))
379    }
380}
381
382/// Data quality assessment utilities
383pub struct DataQualityAssessor;
384
385impl DataQualityAssessor {
386    /// Calculate missing value statistics
387    pub fn missing_value_stats<T>(data: &Array2<T>) -> HashMap<String, f64>
388    where
389        T: Float,
390    {
391        let total_cells = data.len() as f64;
392        let mut missing_count = 0;
393        let mut missing_per_column = Vec::new();
394        let mut missing_per_row = Vec::new();
395
396        // Count missing values per column
397        for col_idx in 0..data.ncols() {
398            let col_missing = data.column(col_idx).iter().filter(|&&x| x.is_nan()).count();
399            missing_per_column.push(col_missing as f64 / data.nrows() as f64);
400            missing_count += col_missing;
401        }
402
403        // Count missing values per row
404        for row_idx in 0..data.nrows() {
405            let row_missing = data.row(row_idx).iter().filter(|&&x| x.is_nan()).count();
406            missing_per_row.push(row_missing as f64 / data.ncols() as f64);
407        }
408
409        let mut stats = HashMap::new();
410        stats.insert(
411            "total_missing_ratio".to_string(),
412            missing_count as f64 / total_cells,
413        );
414        stats.insert(
415            "max_column_missing_ratio".to_string(),
416            missing_per_column.iter().cloned().fold(0.0, f64::max),
417        );
418        stats.insert(
419            "max_row_missing_ratio".to_string(),
420            missing_per_row.iter().cloned().fold(0.0, f64::max),
421        );
422        stats.insert(
423            "columns_with_missing".to_string(),
424            missing_per_column.iter().filter(|&&x| x > 0.0).count() as f64,
425        );
426        stats.insert(
427            "rows_with_missing".to_string(),
428            missing_per_row.iter().filter(|&&x| x > 0.0).count() as f64,
429        );
430
431        stats
432    }
433
434    /// Calculate basic data quality metrics
435    pub fn quality_metrics<T>(data: &Array2<T>) -> HashMap<String, f64>
436    where
437        T: Float + PartialOrd + std::iter::Sum + std::fmt::Display,
438    {
439        let mut metrics = HashMap::new();
440
441        // Calculate completeness (non-missing ratio)
442        let total_cells = data.len() as f64;
443        let missing_count = data.iter().filter(|&&x| x.is_nan()).count() as f64;
444        metrics.insert(
445            "completeness".to_string(),
446            1.0 - (missing_count / total_cells),
447        );
448
449        // Calculate uniformity (check for repeated values)
450        let mut unique_counts = Vec::new();
451        for col_idx in 0..data.ncols() {
452            let col = data.column(col_idx);
453            let mut unique_values = std::collections::HashSet::new();
454            for &value in col.iter() {
455                if !value.is_nan() {
456                    // Convert to string for hashing (approximation)
457                    unique_values.insert(format!("{value:.6}"));
458                }
459            }
460            let uniqueness = unique_values.len() as f64 / col.len() as f64;
461            unique_counts.push(uniqueness);
462        }
463
464        let avg_uniqueness = unique_counts.iter().sum::<f64>() / unique_counts.len() as f64;
465        metrics.insert("uniqueness".to_string(), avg_uniqueness);
466
467        // Calculate consistency (low coefficient of variation)
468        let mut cv_values = Vec::new();
469        for col_idx in 0..data.ncols() {
470            let col = data.column(col_idx);
471            let valid_values: Vec<T> = col.iter().cloned().filter(|x| !x.is_nan()).collect();
472
473            if valid_values.len() > 1 {
474                // Safe conversion - skip column if conversion fails
475                if let Ok(valid_len) = usize_to_float::<T>(valid_values.len()) {
476                    let mean = valid_values.iter().cloned().sum::<T>() / valid_len;
477                    let variance =
478                        valid_values.iter().map(|&x| (x - mean).powi(2)).sum::<T>() / valid_len;
479                    let std_dev = variance.sqrt();
480
481                    if mean != T::zero() {
482                        // Safe conversion to f64 - skip if fails
483                        if let Some(cv_f64) = (std_dev / mean.abs()).to_f64() {
484                            cv_values.push(cv_f64);
485                        }
486                    }
487                }
488            }
489        }
490
491        if !cv_values.is_empty() {
492            let avg_cv = cv_values.iter().sum::<f64>() / cv_values.len() as f64;
493            metrics.insert("consistency".to_string(), 1.0 / (1.0 + avg_cv)); // Higher is better
494        }
495
496        metrics
497    }
498}
499
500#[allow(non_snake_case)]
501#[cfg(test)]
502mod tests {
503    use super::*;
504    use approx::assert_abs_diff_eq;
505    use scirs2_core::ndarray::array;
506
507    #[test]
508    fn test_drop_missing_rows() {
509        let data = array![
510            [1.0, 2.0, 3.0],
511            [4.0, f64::NAN, 6.0],
512            [7.0, 8.0, 9.0],
513            [f64::NAN, 11.0, 12.0]
514        ];
515
516        let cleaned = DataCleaner::drop_missing_rows(&data)
517            .expect("drop_missing_rows should succeed with valid data");
518        assert_eq!(cleaned.nrows(), 2);
519        assert_eq!(cleaned.row(0), array![1.0, 2.0, 3.0]);
520        assert_eq!(cleaned.row(1), array![7.0, 8.0, 9.0]);
521    }
522
523    #[test]
524    fn test_fill_missing_with_value() {
525        let mut data = array![[1.0, 2.0], [f64::NAN, 4.0], [5.0, f64::NAN]];
526
527        DataCleaner::fill_missing(&mut data, 0.0);
528
529        assert_eq!(data, array![[1.0, 2.0], [0.0, 4.0], [5.0, 0.0]]);
530    }
531
532    #[test]
533    fn test_fill_with_mean() {
534        let mut data = array![[1.0, 2.0], [f64::NAN, 4.0], [5.0, f64::NAN]];
535
536        DataCleaner::fill_with_mean(&mut data)
537            .expect("fill_with_mean should succeed with valid data");
538
539        // Mean of first column (1, 5) = 3, mean of second column (2, 4) = 3
540        assert_abs_diff_eq!(data[[1, 0]], 3.0, epsilon = 1e-10);
541        assert_abs_diff_eq!(data[[2, 1]], 3.0, epsilon = 1e-10);
542    }
543
544    #[test]
545    fn test_zscore_outliers() {
546        let data = array![1.0, 2.0, 3.0, 4.0, 100.0]; // 100 is clearly an outlier
547        let outliers = OutlierDetector::zscore_outliers(&data.view(), 1.5);
548        assert_eq!(outliers, vec![4]);
549    }
550
551    #[test]
552    fn test_iqr_outliers() {
553        let data = array![1.0, 2.0, 3.0, 4.0, 5.0, 100.0]; // 100 is an outlier
554        let outliers = OutlierDetector::iqr_outliers(&data.view(), 1.5);
555        assert_eq!(outliers, vec![5]);
556    }
557
558    #[test]
559    fn test_standard_scaling() {
560        let data = array![[1.0, 10.0], [2.0, 20.0], [3.0, 30.0]];
561
562        let (scaled, _means, _stds) = FeatureScaler::standard_scale(&data)
563            .expect("standard_scale should succeed with valid data");
564
565        // Check that scaled data has mean ~0 and std ~1
566        for col_idx in 0..scaled.ncols() {
567            let col = scaled.column(col_idx);
568            let mean = col.iter().sum::<f64>() / col.len() as f64;
569            assert_abs_diff_eq!(mean, 0.0, epsilon = 1e-10);
570        }
571    }
572
573    #[test]
574    fn test_minmax_scaling() {
575        let data = array![[1.0, 10.0], [2.0, 20.0], [3.0, 30.0]];
576
577        let (scaled, _mins, _maxs) = FeatureScaler::minmax_scale(&data)
578            .expect("minmax_scale should succeed with valid data");
579
580        // Check that scaled data is in [0, 1] range
581        for col_idx in 0..scaled.ncols() {
582            let col = scaled.column(col_idx);
583            let min_val = col.iter().cloned().fold(col[0], f64::min);
584            let max_val = col.iter().cloned().fold(col[0], f64::max);
585
586            assert_abs_diff_eq!(min_val, 0.0, epsilon = 1e-10);
587            assert_abs_diff_eq!(max_val, 1.0, epsilon = 1e-10);
588        }
589    }
590
591    #[test]
592    fn test_missing_value_stats() {
593        let data = array![[1.0, 2.0, 3.0], [f64::NAN, 5.0, 6.0], [7.0, f64::NAN, 9.0]];
594
595        let stats = DataQualityAssessor::missing_value_stats(&data);
596
597        assert_abs_diff_eq!(stats["total_missing_ratio"], 2.0 / 9.0, epsilon = 1e-10);
598        assert_eq!(stats["columns_with_missing"], 2.0);
599        assert_eq!(stats["rows_with_missing"], 2.0);
600    }
601
602    #[test]
603    fn test_quality_metrics() {
604        let data = array![[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]];
605
606        let metrics = DataQualityAssessor::quality_metrics(&data);
607
608        // All data is present, so completeness should be 1.0
609        assert_abs_diff_eq!(metrics["completeness"], 1.0, epsilon = 1e-10);
610        assert!(metrics.contains_key("uniqueness"));
611        assert!(metrics.contains_key("consistency"));
612    }
613}