475 lines
15 KiB
Rust
475 lines
15 KiB
Rust
#![allow(clippy::approx_constant)]
|
|
use approx::assert_abs_diff_eq;
|
|
use rtx_preprocessing::{InvertibleTransformer, PreprocessingError, RobustScaler, Transformer};
|
|
use rtx_tensor::{Device, Tensor};
|
|
|
|
/// Test fixture for RobustScaler tests
|
|
struct RobustScalerTestFixture {
|
|
simple_data: Tensor,
|
|
outlier_data: Tensor,
|
|
multi_feature_data: Tensor,
|
|
single_value_data: Tensor,
|
|
constant_data: Tensor,
|
|
symmetric_data: Tensor,
|
|
}
|
|
|
|
impl RobustScalerTestFixture {
|
|
fn new() -> Self {
|
|
// Simple data: [1, 2, 3, 4, 5] -> median=3, IQR=2
|
|
let simple_data =
|
|
Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0, 5.0], &[5, 1], &Device::cpu()).unwrap();
|
|
|
|
// Data with outliers: [1, 2, 3, 4, 100] -> median=3, should be robust to outlier
|
|
let outlier_data =
|
|
Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0, 100.0], &[5, 1], &Device::cpu()).unwrap();
|
|
|
|
// Multi-feature data
|
|
let multi_feature_data = Tensor::from_slice(
|
|
&[1.0, 10.0, 2.0, 20.0, 3.0, 30.0, 4.0, 40.0, 5.0, 50.0],
|
|
&[5, 2],
|
|
&Device::cpu(),
|
|
)
|
|
.unwrap();
|
|
|
|
// Single value
|
|
let single_value_data = Tensor::from_slice(&[5.0], &[1, 1], &Device::cpu()).unwrap();
|
|
|
|
// Constant data
|
|
let constant_data =
|
|
Tensor::from_slice(&[3.0, 3.0, 3.0, 3.0], &[4, 1], &Device::cpu()).unwrap();
|
|
|
|
// Symmetric data around zero: [-4, -2, 0, 2, 4]
|
|
let symmetric_data =
|
|
Tensor::from_slice(&[-4.0, -2.0, 0.0, 2.0, 4.0], &[5, 1], &Device::cpu()).unwrap();
|
|
|
|
Self {
|
|
simple_data,
|
|
outlier_data,
|
|
multi_feature_data,
|
|
single_value_data,
|
|
constant_data,
|
|
symmetric_data,
|
|
}
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_creation() {
|
|
// Default parameters
|
|
let scaler = RobustScaler::new();
|
|
assert!(!scaler.is_fitted());
|
|
assert!(scaler.with_centering());
|
|
assert!(scaler.with_scaling());
|
|
assert_eq!(scaler.quantile_range(), (25.0, 75.0));
|
|
|
|
// Custom parameters
|
|
let scaler = RobustScaler::with_params(false, true, (10.0, 90.0));
|
|
assert!(!scaler.is_fitted());
|
|
assert!(!scaler.with_centering());
|
|
assert!(scaler.with_scaling());
|
|
assert_eq!(scaler.quantile_range(), (10.0, 90.0));
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_invalid_quantile_range() {
|
|
// Should panic with invalid quantile range
|
|
let result = std::panic::catch_unwind(|| RobustScaler::with_params(true, true, (75.0, 25.0)));
|
|
assert!(result.is_err());
|
|
|
|
let result = std::panic::catch_unwind(|| RobustScaler::with_params(true, true, (-5.0, 75.0)));
|
|
assert!(result.is_err());
|
|
|
|
let result = std::panic::catch_unwind(|| RobustScaler::with_params(true, true, (25.0, 105.0)));
|
|
assert!(result.is_err());
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_simple() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Should fit successfully
|
|
let result = scaler.fit(&fixture.simple_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Should have computed correct median and IQR
|
|
assert_abs_diff_eq!(scaler.center()[0], 3.0, epsilon = 1e-5); // median
|
|
assert_abs_diff_eq!(scaler.scale()[0], 2.0, epsilon = 1e-5); // IQR: Q3(4) - Q1(2) = 2
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_outlier_data() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Should be robust to outliers
|
|
scaler.fit(&fixture.outlier_data).unwrap();
|
|
|
|
// Median should still be 3 (not affected by outlier)
|
|
assert_abs_diff_eq!(scaler.center()[0], 3.0, epsilon = 1e-5);
|
|
// IQR should still be 2 (Q1=2, Q3=4, outlier doesn't affect)
|
|
assert_abs_diff_eq!(scaler.scale()[0], 2.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_multi_feature() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Should fit successfully on multi-feature data
|
|
let result = scaler.fit(&fixture.multi_feature_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Should have computed correct median and IQR for each feature
|
|
assert_eq!(scaler.center().len(), 2);
|
|
assert_eq!(scaler.scale().len(), 2);
|
|
|
|
// Feature 0: [1, 2, 3, 4, 5] -> median=3, IQR=2
|
|
assert_abs_diff_eq!(scaler.center()[0], 3.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 2.0, epsilon = 1e-5);
|
|
|
|
// Feature 1: [10, 20, 30, 40, 50] -> median=30, IQR=20
|
|
assert_abs_diff_eq!(scaler.center()[1], 30.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[1], 20.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_empty_data() {
|
|
let empty_data = Tensor::zeros(&[0, 1], &Device::cpu()).unwrap();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Should fail on empty data
|
|
let result = scaler.fit(&empty_data);
|
|
assert!(result.is_err());
|
|
assert!(matches!(
|
|
result.unwrap_err(),
|
|
PreprocessingError::EmptyDataset
|
|
));
|
|
assert!(!scaler.is_fitted());
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_single_value() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Should handle single value
|
|
let result = scaler.fit(&fixture.single_value_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Center should be the value, scale should be 1.0
|
|
assert_abs_diff_eq!(scaler.center()[0], 5.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 1.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_constant_data() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Should handle constant data
|
|
let result = scaler.fit(&fixture.constant_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Center should be the constant value, scale should be 1.0 (IQR=0)
|
|
assert_abs_diff_eq!(scaler.center()[0], 3.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 1.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_not_fitted() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let scaler = RobustScaler::new();
|
|
|
|
// Should fail when not fitted
|
|
let result = scaler.transform(&fixture.simple_data);
|
|
assert!(result.is_err());
|
|
assert!(matches!(result.unwrap_err(), PreprocessingError::NotFitted));
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_simple() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Fit first
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
|
|
// Transform: (x - median) / IQR = (x - 3) / 2
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
assert_eq!(transformed.shape(), &[5, 1]);
|
|
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let expected = vec![-1.0, -0.5, 0.0, 0.5, 1.0]; // (1-3)/2, (2-3)/2, etc.
|
|
|
|
for (actual, expected) in values.iter().zip(expected.iter()) {
|
|
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_outlier_robustness() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Fit on data with outliers
|
|
scaler.fit(&fixture.outlier_data).unwrap();
|
|
|
|
// Transform should handle outliers gracefully
|
|
let transformed = scaler.transform(&fixture.outlier_data).unwrap();
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
|
|
// First 4 values should be scaled normally, outlier will be far but not break scaling
|
|
assert_abs_diff_eq!(values[0], -1.0, epsilon = 1e-5); // (1-3)/2
|
|
assert_abs_diff_eq!(values[1], -0.5, epsilon = 1e-5); // (2-3)/2
|
|
assert_abs_diff_eq!(values[2], 0.0, epsilon = 1e-5); // (3-3)/2
|
|
assert_abs_diff_eq!(values[3], 0.5, epsilon = 1e-5); // (4-3)/2
|
|
assert_abs_diff_eq!(values[4], 48.5, epsilon = 1e-5); // (100-3)/2 = 48.5
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_without_centering() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::with_params(false, true, (25.0, 75.0));
|
|
|
|
// Should scale but not center
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let expected = vec![0.5, 1.0, 1.5, 2.0, 2.5]; // x / IQR = x / 2
|
|
|
|
for (actual, expected) in values.iter().zip(expected.iter()) {
|
|
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_without_scaling() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::with_params(true, false, (25.0, 75.0));
|
|
|
|
// Should center but not scale
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let expected = vec![-2.0, -1.0, 0.0, 1.0, 2.0]; // x - median = x - 3
|
|
|
|
for (actual, expected) in values.iter().zip(expected.iter()) {
|
|
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_custom_quantiles() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::with_params(true, true, (10.0, 90.0));
|
|
|
|
// Should use 10th and 90th percentiles instead of 25th and 75th
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
|
|
// For [1,2,3,4,5]: 10th percentile ≈ 1.4, 90th percentile ≈ 4.6
|
|
// IQR = 4.6 - 1.4 = 3.2, median = 3
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
|
|
// Values should be different from standard 25-75 percentile scaling
|
|
// Just check that transformation occurred
|
|
assert_ne!(values[0], -1.0); // Should be different from 25-75 scaling
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_transform_dimension_mismatch() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Fit on simple data (1 feature)
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
|
|
// Try to transform multi-feature data (2 features)
|
|
let result = scaler.transform(&fixture.multi_feature_data);
|
|
assert!(result.is_err());
|
|
assert!(matches!(
|
|
result.unwrap_err(),
|
|
PreprocessingError::DimensionMismatch {
|
|
expected: 1,
|
|
actual: 2
|
|
}
|
|
));
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_inverse_transform() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Fit and transform
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
|
|
// Inverse transform should recover original data
|
|
let recovered = scaler.inverse_transform(&transformed).unwrap();
|
|
assert_eq!(recovered.shape(), fixture.simple_data.shape());
|
|
|
|
let original_values = fixture
|
|
.simple_data
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let recovered_values = recovered
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
|
|
for (orig, rec) in original_values.iter().zip(recovered_values.iter()) {
|
|
assert_abs_diff_eq!(orig, rec, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_fit_transform() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// fit_transform should be equivalent to fit + transform
|
|
let result1 = scaler.fit_transform(&fixture.simple_data).unwrap();
|
|
|
|
let mut scaler2 = RobustScaler::new();
|
|
scaler2.fit(&fixture.simple_data).unwrap();
|
|
let result2 = scaler2.transform(&fixture.simple_data).unwrap();
|
|
|
|
assert_eq!(result1.shape(), result2.shape());
|
|
let values1 = result1
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let values2 = result2
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
|
|
for (v1, v2) in values1.iter().zip(values2.iter()) {
|
|
assert_abs_diff_eq!(v1, v2, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_reset() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Fit scaler
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Reset should make it unfitted
|
|
scaler.reset();
|
|
assert!(!scaler.is_fitted());
|
|
|
|
// Should fail to transform after reset
|
|
let result = scaler.transform(&fixture.simple_data);
|
|
assert!(result.is_err());
|
|
assert!(matches!(result.unwrap_err(), PreprocessingError::NotFitted));
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_symmetric_data() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::new();
|
|
|
|
// Test with symmetric data around zero
|
|
scaler.fit(&fixture.symmetric_data).unwrap();
|
|
let transformed = scaler.transform(&fixture.symmetric_data).unwrap();
|
|
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
|
|
// Should be centered around zero and scaled by IQR
|
|
let mean: f32 = values.iter().sum::<f32>() / values.len() as f32;
|
|
assert_abs_diff_eq!(mean, 0.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_robust_scaler_serialization() {
|
|
let fixture = RobustScalerTestFixture::new();
|
|
let mut scaler = RobustScaler::with_params(true, true, (10.0, 90.0));
|
|
|
|
// Fit scaler
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
|
|
// Should be serializable and deserializable
|
|
let serialized = bincode::serialize(&scaler).unwrap();
|
|
let deserialized: RobustScaler = bincode::deserialize(&serialized).unwrap();
|
|
|
|
// Should maintain fitted state and parameters
|
|
assert!(deserialized.is_fitted());
|
|
assert_eq!(scaler.with_centering(), deserialized.with_centering());
|
|
assert_eq!(scaler.with_scaling(), deserialized.with_scaling());
|
|
assert_eq!(scaler.quantile_range(), deserialized.quantile_range());
|
|
assert_eq!(scaler.center(), deserialized.center());
|
|
assert_eq!(scaler.scale(), deserialized.scale());
|
|
|
|
// Should produce same transform results
|
|
let original_result = scaler.transform(&fixture.simple_data).unwrap();
|
|
let deserialized_result = deserialized.transform(&fixture.simple_data).unwrap();
|
|
|
|
let orig_values = original_result
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let deser_values = deserialized_result
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
|
|
for (orig, deser) in orig_values.iter().zip(deser_values.iter()) {
|
|
assert_abs_diff_eq!(orig, deser, epsilon = 1e-5);
|
|
}
|
|
}
|