Skip to main content

sklears_datasets/validation/
basic.rs

1//! Basic validation functions for dataset statistical properties
2//!
3//! This module provides fundamental validation functions that check
4//! basic statistical properties of datasets including size, consistency,
5//! distribution properties, correlations, and outliers.
6
7use super::types::{ValidationConfig, ValidationReport, ValidationResult};
8use std::collections::HashMap;
9
10/// Validate basic statistical properties of a dataset
11pub fn validate_basic_statistics(data: &[Vec<f64>], config: &ValidationConfig) -> ValidationReport {
12    let mut report = ValidationReport::new();
13
14    if data.is_empty() {
15        report.add_result(ValidationResult {
16            property: "Dataset Size".to_string(),
17            passed: false,
18            expected: config.min_samples as f64,
19            actual: 0.0,
20            tolerance: 0.0,
21            message: "Dataset is empty".to_string(),
22        });
23        return report;
24    }
25
26    let n_samples = data.len();
27    let n_features = data[0].len();
28
29    // Check minimum sample size
30    let min_samples_passed = n_samples >= config.min_samples;
31    report.add_result(ValidationResult {
32        property: "Minimum Samples".to_string(),
33        passed: min_samples_passed,
34        expected: config.min_samples as f64,
35        actual: n_samples as f64,
36        tolerance: 0.0,
37        message: if min_samples_passed {
38            format!(
39                "Dataset has {} samples (≥ {})",
40                n_samples, config.min_samples
41            )
42        } else {
43            format!(
44                "Dataset has {} samples (< {})",
45                n_samples, config.min_samples
46            )
47        },
48    });
49
50    // Validate feature consistency
51    let consistent_features = data.iter().all(|row| row.len() == n_features);
52    report.add_result(ValidationResult {
53        property: "Feature Consistency".to_string(),
54        passed: consistent_features,
55        expected: n_features as f64,
56        actual: if consistent_features {
57            n_features as f64
58        } else {
59            -1.0
60        },
61        tolerance: 0.0,
62        message: if consistent_features {
63            format!("All samples have {} features", n_features)
64        } else {
65            "Inconsistent number of features across samples".to_string()
66        },
67    });
68
69    // Check for NaN and infinite values
70    let mut has_nan = false;
71    let mut has_inf = false;
72
73    for row in data {
74        for &value in row {
75            if value.is_nan() {
76                has_nan = true;
77            }
78            if value.is_infinite() {
79                has_inf = true;
80            }
81        }
82    }
83
84    report.add_result(ValidationResult {
85        property: "No NaN Values".to_string(),
86        passed: !has_nan,
87        expected: 0.0,
88        actual: if has_nan { 1.0 } else { 0.0 },
89        tolerance: 0.0,
90        message: if has_nan {
91            "Dataset contains NaN values".to_string()
92        } else {
93            "No NaN values found".to_string()
94        },
95    });
96
97    report.add_result(ValidationResult {
98        property: "No Infinite Values".to_string(),
99        passed: !has_inf,
100        expected: 0.0,
101        actual: if has_inf { 1.0 } else { 0.0 },
102        tolerance: 0.0,
103        message: if has_inf {
104            "Dataset contains infinite values".to_string()
105        } else {
106            "No infinite values found".to_string()
107        },
108    });
109
110    report
111}
112
113/// Validate distribution properties (mean, variance, etc.)
114pub fn validate_distribution_properties(
115    data: &[Vec<f64>],
116    expected_mean: Option<f64>,
117    expected_std: Option<f64>,
118    config: &ValidationConfig,
119) -> ValidationReport {
120    let mut report = ValidationReport::new();
121
122    if data.is_empty() {
123        return report;
124    }
125
126    let n_samples = data.len();
127    let n_features = data[0].len();
128
129    // Calculate statistics for each feature
130    for feature_idx in 0..n_features {
131        let values: Vec<f64> = data.iter().map(|row| row[feature_idx]).collect();
132
133        // Calculate mean
134        let mean = values.iter().sum::<f64>() / n_samples as f64;
135
136        // Calculate standard deviation
137        let variance = values.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / n_samples as f64;
138        let std_dev = variance.sqrt();
139
140        // Validate mean if expected
141        if let Some(expected_mean) = expected_mean {
142            let mean_diff = (mean - expected_mean).abs();
143            let mean_passed = mean_diff <= config.tolerance;
144
145            report.add_result(ValidationResult {
146                property: format!("Feature {} Mean", feature_idx),
147                passed: mean_passed,
148                expected: expected_mean,
149                actual: mean,
150                tolerance: config.tolerance,
151                message: if mean_passed {
152                    format!(
153                        "Mean {:.4} is within tolerance of {:.4}",
154                        mean, expected_mean
155                    )
156                } else {
157                    format!(
158                        "Mean {:.4} differs from expected {:.4} by {:.4}",
159                        mean, expected_mean, mean_diff
160                    )
161                },
162            });
163        }
164
165        // Validate standard deviation if expected
166        if let Some(expected_std) = expected_std {
167            let std_diff = (std_dev - expected_std).abs();
168            let std_passed = std_diff <= config.tolerance;
169
170            report.add_result(ValidationResult {
171                property: format!("Feature {} Std Dev", feature_idx),
172                passed: std_passed,
173                expected: expected_std,
174                actual: std_dev,
175                tolerance: config.tolerance,
176                message: if std_passed {
177                    format!(
178                        "Std dev {:.4} is within tolerance of {:.4}",
179                        std_dev, expected_std
180                    )
181                } else {
182                    format!(
183                        "Std dev {:.4} differs from expected {:.4} by {:.4}",
184                        std_dev, expected_std, std_diff
185                    )
186                },
187            });
188        }
189    }
190
191    report
192}
193
194/// Validate correlation structure
195pub fn validate_correlation_structure(
196    data: &[Vec<f64>],
197    expected_correlations: Option<&HashMap<(usize, usize), f64>>,
198    config: &ValidationConfig,
199) -> ValidationReport {
200    let mut report = ValidationReport::new();
201
202    if data.is_empty() {
203        return report;
204    }
205
206    let n_samples = data.len();
207    let n_features = data[0].len();
208
209    if n_features < 2 {
210        return report;
211    }
212
213    // Calculate correlation matrix
214    let mut correlations = HashMap::new();
215
216    for i in 0..n_features {
217        for j in (i + 1)..n_features {
218            let values_i: Vec<f64> = data.iter().map(|row| row[i]).collect();
219            let values_j: Vec<f64> = data.iter().map(|row| row[j]).collect();
220
221            let mean_i = values_i.iter().sum::<f64>() / n_samples as f64;
222            let mean_j = values_j.iter().sum::<f64>() / n_samples as f64;
223
224            let numerator: f64 = values_i
225                .iter()
226                .zip(values_j.iter())
227                .map(|(&x, &y)| (x - mean_i) * (y - mean_j))
228                .sum();
229
230            let sum_sq_i: f64 = values_i.iter().map(|&x| (x - mean_i).powi(2)).sum();
231            let sum_sq_j: f64 = values_j.iter().map(|&x| (x - mean_j).powi(2)).sum();
232
233            let correlation = if sum_sq_i > 0.0 && sum_sq_j > 0.0 {
234                numerator / (sum_sq_i.sqrt() * sum_sq_j.sqrt())
235            } else {
236                0.0
237            };
238
239            correlations.insert((i, j), correlation);
240        }
241    }
242
243    // Validate expected correlations
244    if let Some(expected_correlations) = expected_correlations {
245        for (&(i, j), &expected_corr) in expected_correlations {
246            if let Some(&actual_corr) = correlations.get(&(i, j)) {
247                let corr_diff = (actual_corr - expected_corr).abs();
248                let corr_passed = corr_diff <= config.tolerance;
249
250                report.add_result(ValidationResult {
251                    property: format!("Correlation [{}, {}]", i, j),
252                    passed: corr_passed,
253                    expected: expected_corr,
254                    actual: actual_corr,
255                    tolerance: config.tolerance,
256                    message: if corr_passed {
257                        format!(
258                            "Correlation {:.4} is within tolerance of {:.4}",
259                            actual_corr, expected_corr
260                        )
261                    } else {
262                        format!(
263                            "Correlation {:.4} differs from expected {:.4} by {:.4}",
264                            actual_corr, expected_corr, corr_diff
265                        )
266                    },
267                });
268            }
269        }
270    }
271
272    report
273}
274
275/// Validate normality using Shapiro-Wilk approximation
276pub fn validate_normality(data: &[Vec<f64>], _config: &ValidationConfig) -> ValidationReport {
277    let mut report = ValidationReport::new();
278
279    if data.is_empty() {
280        return report;
281    }
282
283    let n_features = data[0].len();
284
285    for feature_idx in 0..n_features {
286        let mut values: Vec<f64> = data.iter().map(|row| row[feature_idx]).collect();
287        values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
288
289        let n = values.len();
290        if n < 3 {
291            continue;
292        }
293
294        // Simple normality test: check if data is approximately normal
295        // by comparing percentiles with expected normal distribution
296        let mean = values.iter().sum::<f64>() / n as f64;
297        let variance = values.iter().map(|&x| (x - mean).powi(2)).sum::<f64>() / n as f64;
298        let std_dev = variance.sqrt();
299
300        if std_dev == 0.0 {
301            continue;
302        }
303
304        // Check if 68% of data is within 1 standard deviation
305        let within_1_std = values
306            .iter()
307            .filter(|&&x| (x - mean).abs() <= std_dev)
308            .count();
309        let expected_within_1_std = (n as f64 * 0.68) as usize;
310        let tolerance_1_std = (n as f64 * 0.1) as usize; // 10% tolerance
311
312        let normality_passed =
313            (within_1_std as i32 - expected_within_1_std as i32).abs() <= tolerance_1_std as i32;
314
315        report.add_result(ValidationResult {
316            property: format!("Feature {} Normality", feature_idx),
317            passed: normality_passed,
318            expected: expected_within_1_std as f64,
319            actual: within_1_std as f64,
320            tolerance: tolerance_1_std as f64,
321            message: if normality_passed {
322                format!("Feature {} appears normally distributed", feature_idx)
323            } else {
324                format!("Feature {} may not be normally distributed", feature_idx)
325            },
326        });
327    }
328
329    report
330}
331
332/// Validate outlier detection
333pub fn validate_outliers(
334    data: &[Vec<f64>],
335    expected_outlier_ratio: Option<f64>,
336    config: &ValidationConfig,
337) -> ValidationReport {
338    let mut report = ValidationReport::new();
339
340    if data.is_empty() {
341        return report;
342    }
343
344    let n_samples = data.len();
345    let n_features = data[0].len();
346
347    for feature_idx in 0..n_features {
348        let values: Vec<f64> = data.iter().map(|row| row[feature_idx]).collect();
349
350        // Calculate Q1, Q3, and IQR
351        let mut sorted_values = values.clone();
352        sorted_values.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
353
354        let q1_idx = n_samples / 4;
355        let q3_idx = 3 * n_samples / 4;
356        let q1 = sorted_values[q1_idx];
357        let q3 = sorted_values[q3_idx];
358        let iqr = q3 - q1;
359
360        // Count outliers (values beyond 1.5 * IQR from quartiles)
361        let outlier_threshold = 1.5 * iqr;
362        let outliers = values
363            .iter()
364            .filter(|&&x| x < q1 - outlier_threshold || x > q3 + outlier_threshold)
365            .count();
366
367        let outlier_ratio = outliers as f64 / n_samples as f64;
368
369        if let Some(expected_ratio) = expected_outlier_ratio {
370            let ratio_diff = (outlier_ratio - expected_ratio).abs();
371            let outlier_passed = ratio_diff <= config.tolerance;
372
373            report.add_result(ValidationResult {
374                property: format!("Feature {} Outlier Ratio", feature_idx),
375                passed: outlier_passed,
376                expected: expected_ratio,
377                actual: outlier_ratio,
378                tolerance: config.tolerance,
379                message: if outlier_passed {
380                    format!(
381                        "Outlier ratio {:.4} is within tolerance of {:.4}",
382                        outlier_ratio, expected_ratio
383                    )
384                } else {
385                    format!(
386                        "Outlier ratio {:.4} differs from expected {:.4} by {:.4}",
387                        outlier_ratio, expected_ratio, ratio_diff
388                    )
389                },
390            });
391        }
392    }
393
394    report
395}