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

493 lines
15 KiB
Rust

#![allow(clippy::approx_constant)]
use approx::assert_abs_diff_eq;
use rtx_preprocessing::{InvertibleTransformer, MinMaxScaler, PreprocessingError, Transformer};
use rtx_tensor::{Device, Tensor};
/// Test fixture for MinMaxScaler tests
struct MinMaxScalerTestFixture {
simple_data: Tensor,
multi_feature_data: Tensor,
single_value_data: Tensor,
constant_data: Tensor,
negative_data: Tensor,
large_range_data: Tensor,
}
impl MinMaxScalerTestFixture {
fn new() -> Self {
// Simple data: [1, 2, 3, 4, 5] -> should scale to [0, 0.25, 0.5, 0.75, 1.0]
let simple_data =
Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0, 5.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],
&[4, 2],
&Device::cpu(),
)
.unwrap();
// Single value
let single_value_data = Tensor::from_slice(&[5.0], &[1, 1], &Device::cpu()).unwrap();
// Constant data (all same values)
let constant_data =
Tensor::from_slice(&[3.0, 3.0, 3.0, 3.0], &[4, 1], &Device::cpu()).unwrap();
// Data with negative values: [-2, -1, 0, 1, 2]
let negative_data =
Tensor::from_slice(&[-2.0, -1.0, 0.0, 1.0, 2.0], &[5, 1], &Device::cpu()).unwrap();
// Large range data: [0, 1000]
let large_range_data =
Tensor::from_slice(&[0.0, 250.0, 500.0, 750.0, 1000.0], &[5, 1], &Device::cpu())
.unwrap();
Self {
simple_data,
multi_feature_data,
single_value_data,
constant_data,
negative_data,
large_range_data,
}
}
}
#[test]
fn test_minmax_scaler_creation() {
// Default range [0, 1]
let scaler = MinMaxScaler::new();
assert!(!scaler.is_fitted());
assert_eq!(scaler.feature_range(), (0.0, 1.0));
// Custom range [-1, 1]
let scaler = MinMaxScaler::with_range(-1.0, 1.0);
assert!(!scaler.is_fitted());
assert_eq!(scaler.feature_range(), (-1.0, 1.0));
// Custom range [0, 10]
let scaler = MinMaxScaler::with_range(0.0, 10.0);
assert!(!scaler.is_fitted());
assert_eq!(scaler.feature_range(), (0.0, 10.0));
}
#[test]
fn test_minmax_scaler_invalid_range() {
// Should fail with invalid range (min >= max)
let result = std::panic::catch_unwind(|| MinMaxScaler::with_range(1.0, 1.0));
assert!(result.is_err());
let result = std::panic::catch_unwind(|| MinMaxScaler::with_range(2.0, 1.0));
assert!(result.is_err());
}
#[test]
fn test_minmax_scaler_fit_simple() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Should fit successfully
let result = scaler.fit(&fixture.simple_data);
assert!(result.is_ok());
assert!(scaler.is_fitted());
// Should have computed correct min and max
assert_abs_diff_eq!(scaler.data_min()[0], 1.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_max()[0], 5.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_range()[0], 4.0, epsilon = 1e-5);
}
#[test]
fn test_minmax_scaler_fit_multi_feature() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::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 min/max for each feature
assert_eq!(scaler.data_min().len(), 2);
assert_eq!(scaler.data_max().len(), 2);
assert_eq!(scaler.data_range().len(), 2);
// Feature 0: [1, 2, 3, 4]
assert_abs_diff_eq!(scaler.data_min()[0], 1.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_max()[0], 4.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_range()[0], 3.0, epsilon = 1e-5);
// Feature 1: [10, 20, 30, 40]
assert_abs_diff_eq!(scaler.data_min()[1], 10.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_max()[1], 40.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_range()[1], 30.0, epsilon = 1e-5);
}
#[test]
fn test_minmax_scaler_fit_empty_data() {
let empty_data = Tensor::zeros(&[0, 1], &Device::cpu()).unwrap();
let mut scaler = MinMaxScaler::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_minmax_scaler_fit_single_value() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Should handle single value
let result = scaler.fit(&fixture.single_value_data);
assert!(result.is_ok());
assert!(scaler.is_fitted());
// Min and max should be the same value
assert_abs_diff_eq!(scaler.data_min()[0], 5.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_max()[0], 5.0, epsilon = 1e-5);
assert_abs_diff_eq!(scaler.data_range()[0], 0.0, epsilon = 1e-5);
}
#[test]
fn test_minmax_scaler_fit_constant_data() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Should handle constant data
let result = scaler.fit(&fixture.constant_data);
assert!(result.is_ok());
assert!(scaler.is_fitted());
// Range should be zero
assert_abs_diff_eq!(scaler.data_range()[0], 0.0, epsilon = 1e-5);
}
#[test]
fn test_minmax_scaler_transform_not_fitted() {
let fixture = MinMaxScalerTestFixture::new();
let scaler = MinMaxScaler::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_minmax_scaler_transform_simple() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Fit first
scaler.fit(&fixture.simple_data).unwrap();
// Transform should scale to [0, 1] range
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![0.0, 0.25, 0.5, 0.75, 1.0];
for (actual, expected) in values.iter().zip(expected.iter()) {
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
}
}
#[test]
fn test_minmax_scaler_transform_custom_range() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::with_range(-1.0, 1.0);
// Fit and transform
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![-1.0, -0.5, 0.0, 0.5, 1.0];
for (actual, expected) in values.iter().zip(expected.iter()) {
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
}
}
#[test]
fn test_minmax_scaler_transform_negative_data() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Fit and transform negative data
scaler.fit(&fixture.negative_data).unwrap();
let transformed = scaler.transform(&fixture.negative_data).unwrap();
let values = transformed
.to_cpu()
.unwrap()
.iter()
.copied()
.collect::<Vec<f32>>();
let expected = vec![0.0, 0.25, 0.5, 0.75, 1.0];
for (actual, expected) in values.iter().zip(expected.iter()) {
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
}
}
#[test]
fn test_minmax_scaler_transform_constant_data() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Fit and transform constant data
scaler.fit(&fixture.constant_data).unwrap();
let transformed = scaler.transform(&fixture.constant_data).unwrap();
let values = transformed
.to_cpu()
.unwrap()
.iter()
.copied()
.collect::<Vec<f32>>();
// All values should be 0.0 (min of range) when data has no variance
for value in values {
assert_abs_diff_eq!(value, 0.0, epsilon = 1e-5);
}
}
#[test]
fn test_minmax_scaler_transform_dimension_mismatch() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::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_minmax_scaler_inverse_transform() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::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_minmax_scaler_inverse_transform_custom_range() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::with_range(-5.0, 5.0);
// 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();
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_minmax_scaler_fit_transform() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// fit_transform should be equivalent to fit + transform
let result1 = scaler.fit_transform(&fixture.simple_data).unwrap();
let mut scaler2 = MinMaxScaler::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_minmax_scaler_reset() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::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_minmax_scaler_out_of_range_data() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Fit on simple data [1, 2, 3, 4, 5]
scaler.fit(&fixture.simple_data).unwrap();
// Transform data outside the fitted range
let out_of_range_data = Tensor::from_slice(&[0.0, 6.0], &[2, 1], &Device::cpu()).unwrap();
let transformed = scaler.transform(&out_of_range_data).unwrap();
let values = transformed
.to_cpu()
.unwrap()
.iter()
.copied()
.collect::<Vec<f32>>();
// Should extrapolate: 0 -> -0.25, 6 -> 1.25
assert_abs_diff_eq!(values[0], -0.25, epsilon = 1e-5);
assert_abs_diff_eq!(values[1], 1.25, epsilon = 1e-5);
}
#[test]
fn test_minmax_scaler_large_range_data() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::new();
// Should handle large range data efficiently
scaler.fit(&fixture.large_range_data).unwrap();
let transformed = scaler.transform(&fixture.large_range_data).unwrap();
let values = transformed
.to_cpu()
.unwrap()
.iter()
.copied()
.collect::<Vec<f32>>();
let expected = vec![0.0, 0.25, 0.5, 0.75, 1.0];
for (actual, expected) in values.iter().zip(expected.iter()) {
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
}
}
#[test]
fn test_minmax_scaler_serialization() {
let fixture = MinMaxScalerTestFixture::new();
let mut scaler = MinMaxScaler::with_range(-2.0, 2.0);
// Fit scaler
scaler.fit(&fixture.simple_data).unwrap();
// Should be serializable and deserializable
let serialized = bincode::serialize(&scaler).unwrap();
let deserialized: MinMaxScaler = bincode::deserialize(&serialized).unwrap();
// Should maintain fitted state and parameters
assert!(deserialized.is_fitted());
assert_eq!(scaler.feature_range(), deserialized.feature_range());
assert_eq!(scaler.data_min(), deserialized.data_min());
assert_eq!(scaler.data_max(), deserialized.data_max());
assert_eq!(scaler.data_range(), deserialized.data_range());
// 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);
}
}