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

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