434 lines
13 KiB
Rust
434 lines
13 KiB
Rust
#![cfg(feature = "disabled_tests")]
|
|
#![allow(clippy::approx_constant)]
|
|
use approx::assert_abs_diff_eq;
|
|
use rtx_preprocessing::{InvertibleTransformer, PreprocessingError, StandardScaler, Transformer};
|
|
use rtx_tensor::{Device, Tensor};
|
|
|
|
/// Test fixture for StandardScaler tests
|
|
struct StandardScalerTestFixture {
|
|
simple_data: Tensor,
|
|
multi_feature_data: Tensor,
|
|
single_value_data: Tensor,
|
|
zero_variance_data: Tensor,
|
|
with_nan_data: Tensor,
|
|
large_data: Tensor,
|
|
}
|
|
|
|
impl StandardScalerTestFixture {
|
|
fn new() -> Self {
|
|
// Simple 1D data: [1, 2, 3, 4, 5] -> mean=3, std=sqrt(2.5)
|
|
let simple_data =
|
|
Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0, 5.0], &[5, 1], &Device::cpu()).unwrap();
|
|
|
|
// Multi-feature data: 2 features, 4 samples
|
|
let multi_feature_data = Tensor::from_slice(
|
|
&[1.0, 10.0, 2.0, 20.0, 3.0, 30.0, 4.0, 40.0],
|
|
&[4, 2],
|
|
&Device::cpu(),
|
|
)
|
|
.unwrap();
|
|
|
|
// Single value (edge case)
|
|
let single_value_data = Tensor::from_slice(&[5.0], &[1, 1], &Device::cpu()).unwrap();
|
|
|
|
// Zero variance data (all same values)
|
|
let zero_variance_data =
|
|
Tensor::from_slice(&[2.0, 2.0, 2.0, 2.0], &[4, 1], &Device::cpu()).unwrap();
|
|
|
|
// Data with NaN values
|
|
let with_nan_data = Tensor::from_slice(
|
|
&[1.0, f32::NAN, 3.0, 4.0, f32::NAN],
|
|
&[5, 1],
|
|
&Device::cpu(),
|
|
)
|
|
.unwrap();
|
|
|
|
// Large dataset for performance testing (1000 samples, 50 features)
|
|
let mut large_vec = Vec::new();
|
|
for i in 0..1000 {
|
|
for j in 0..50 {
|
|
large_vec.push((i * j) as f32);
|
|
}
|
|
}
|
|
let large_data = Tensor::from_slice(&large_vec, &[1000, 50], &Device::cpu()).unwrap();
|
|
|
|
Self {
|
|
simple_data,
|
|
multi_feature_data,
|
|
single_value_data,
|
|
zero_variance_data,
|
|
with_nan_data,
|
|
large_data,
|
|
}
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_creation() {
|
|
// Should create with default parameters
|
|
let scaler = StandardScaler::new();
|
|
assert!(!scaler.is_fitted());
|
|
|
|
// Should create with custom parameters
|
|
let scaler = StandardScaler::with_mean_std(true, true);
|
|
assert!(!scaler.is_fitted());
|
|
|
|
// Should create without mean centering
|
|
let scaler = StandardScaler::with_mean_std(false, true);
|
|
assert!(!scaler.is_fitted());
|
|
|
|
// Should create without scaling
|
|
let scaler = StandardScaler::with_mean_std(true, false);
|
|
assert!(!scaler.is_fitted());
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_fit_simple() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Should fit successfully on simple data
|
|
let result = scaler.fit(&fixture.simple_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Should have computed correct mean and std
|
|
assert_abs_diff_eq!(scaler.mean()[0], 3.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 1.5811388300841898, epsilon = 1e-5); // sqrt(2.5)
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_fit_multi_feature() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::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 means and scales for each feature
|
|
assert_eq!(scaler.mean().len(), 2);
|
|
assert_eq!(scaler.scale().len(), 2);
|
|
|
|
// Feature 0: [1, 2, 3, 4] -> mean=2.5, std=sqrt(1.25)
|
|
assert_abs_diff_eq!(scaler.mean()[0], 2.5, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 1.2909944487358056, epsilon = 1e-5);
|
|
|
|
// Feature 1: [10, 20, 30, 40] -> mean=25, std=sqrt(125)
|
|
assert_abs_diff_eq!(scaler.mean()[1], 25.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[1], 12.909944487358056, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_fit_empty_data() {
|
|
let empty_data = Tensor::zeros(&[0, 1], &Device::cpu()).unwrap();
|
|
let mut scaler = StandardScaler::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_standard_scaler_fit_single_value() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Should handle single value gracefully
|
|
let result = scaler.fit(&fixture.single_value_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Mean should be the value, scale should be 1.0 (to avoid division by zero)
|
|
assert_abs_diff_eq!(scaler.mean()[0], 5.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 1.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_fit_zero_variance() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Should handle zero variance data
|
|
let result = scaler.fit(&fixture.zero_variance_data);
|
|
assert!(result.is_ok());
|
|
assert!(scaler.is_fitted());
|
|
|
|
// Mean should be 2.0, scale should be 1.0 (to avoid division by zero)
|
|
assert_abs_diff_eq!(scaler.mean()[0], 2.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(scaler.scale()[0], 1.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_transform_not_fitted() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let scaler = StandardScaler::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_standard_scaler_transform_simple() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Fit first
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
|
|
// Transform should produce mean=0, std=1
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
assert_eq!(transformed.shape(), &[5, 1]);
|
|
|
|
// Check that transformed data has mean ≈ 0 and std ≈ 1
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let mean: f32 = values.iter().sum::<f32>() / values.len() as f32;
|
|
let variance: f32 =
|
|
values.iter().map(|&x| (x - mean).powi(2)).sum::<f32>() / (values.len() - 1) as f32;
|
|
let std = variance.sqrt();
|
|
|
|
assert_abs_diff_eq!(mean, 0.0, epsilon = 1e-5);
|
|
assert_abs_diff_eq!(std, 1.0, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_transform_dimension_mismatch() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::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_standard_scaler_inverse_transform() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::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_standard_scaler_fit_transform() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// fit_transform should be equivalent to fit + transform
|
|
let result1 = scaler.fit_transform(&fixture.simple_data).unwrap();
|
|
|
|
let mut scaler2 = StandardScaler::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_standard_scaler_with_mean_false() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::with_mean_std(false, true);
|
|
|
|
// Should scale but not center
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
let transformed = scaler.transform(&fixture.simple_data).unwrap();
|
|
|
|
// Mean should not be zero
|
|
let values = transformed
|
|
.to_cpu()
|
|
.unwrap()
|
|
.iter()
|
|
.copied()
|
|
.collect::<Vec<f32>>();
|
|
let mean: f32 = values.iter().sum::<f32>() / values.len() as f32;
|
|
assert!(mean.abs() > 1e-10); // Should not be zero
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_with_std_false() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::with_mean_std(true, false);
|
|
|
|
// 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 mean: f32 = values.iter().sum::<f32>() / values.len() as f32;
|
|
let variance: f32 =
|
|
values.iter().map(|&x| (x - mean).powi(2)).sum::<f32>() / (values.len() - 1) as f32;
|
|
let std = variance.sqrt();
|
|
|
|
assert_abs_diff_eq!(mean, 0.0, epsilon = 1e-5);
|
|
// Std should not be 1.0 (should be original std)
|
|
assert_abs_diff_eq!(std, 1.5811388300841898, epsilon = 1e-5);
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_reset() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::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_standard_scaler_nan_handling() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Should handle NaN values appropriately
|
|
let result = scaler.fit(&fixture.with_nan_data);
|
|
// This might be implementation-dependent - could skip NaNs or error
|
|
// For now, we'll test that it either succeeds or fails gracefully
|
|
match result {
|
|
Ok(_) => {
|
|
// If it succeeds, check that mean/scale are not NaN
|
|
assert!(!scaler.mean()[0].is_nan());
|
|
assert!(!scaler.scale()[0].is_nan());
|
|
}
|
|
Err(e) => {
|
|
// If it fails, should be a proper error
|
|
assert!(matches!(e, PreprocessingError::NumericalError { .. }));
|
|
}
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_performance_large_data() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Should handle large datasets efficiently
|
|
let start = std::time::Instant::now();
|
|
scaler.fit(&fixture.large_data).unwrap();
|
|
let fit_time = start.elapsed();
|
|
|
|
let start = std::time::Instant::now();
|
|
let _transformed = scaler.transform(&fixture.large_data).unwrap();
|
|
let transform_time = start.elapsed();
|
|
|
|
// Performance check - should complete in reasonable time
|
|
// These are generous bounds for now
|
|
assert!(fit_time.as_secs() < 10);
|
|
assert!(transform_time.as_secs() < 10);
|
|
}
|
|
|
|
#[test]
|
|
fn test_standard_scaler_serialization() {
|
|
let fixture = StandardScalerTestFixture::new();
|
|
let mut scaler = StandardScaler::new();
|
|
|
|
// Fit scaler
|
|
scaler.fit(&fixture.simple_data).unwrap();
|
|
|
|
// Should be serializable and deserializable
|
|
let serialized = bincode::serialize(&scaler).unwrap();
|
|
let deserialized: StandardScaler = bincode::deserialize(&serialized).unwrap();
|
|
|
|
// Should maintain fitted state and parameters
|
|
assert!(deserialized.is_fitted());
|
|
assert_eq!(scaler.mean(), deserialized.mean());
|
|
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);
|
|
}
|
|
}
|