Files
rustytorch/crates/training/rtx-preprocessing/tests/robust_scaler_tests.rs
T
2026-03-04 00:08:42 +00:00

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);
}
}