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

237 lines
7.0 KiB
Rust

//! Comprehensive tests for data imputation functionality
#![cfg(feature = "disabled_tests")]
use rtx_preprocessing::{
PreprocessingError,
transformers::{ImputationStrategy, Imputer},
};
use rtx_tensor::{Device, Tensor};
#[cfg(test)]
mod imputer_tests {
use super::*;
fn create_test_device() -> Device {
Device::cpu()
}
fn create_data_with_missing() -> Tensor {
// NaN represents missing values
Tensor::from_data(
vec![1.0, f32::NAN, 3.0, 4.0, f32::NAN, 6.0],
vec![6],
&create_test_device(),
)
.unwrap()
}
fn create_2d_data_with_missing() -> Tensor {
Tensor::from_data(
vec![1.0, 2.0, f32::NAN, 4.0, f32::NAN, 6.0, 7.0, 8.0, 9.0],
vec![3, 3],
&create_test_device(),
)
.unwrap()
}
#[test]
fn test_imputer_creation() {
let imputer = Imputer::new(ImputationStrategy::Mean);
assert_eq!(imputer.strategy(), ImputationStrategy::Mean);
}
#[test]
fn test_mean_imputation() {
let mut imputer = Imputer::new(ImputationStrategy::Mean);
let data = create_data_with_missing();
// Fit the imputer
let result = imputer.fit(&data);
assert!(result.is_ok());
// Transform the data
let transformed = imputer.transform(&data);
assert!(transformed.is_ok());
let output = transformed.unwrap();
let output_data = output.to_cpu().unwrap();
// Mean of [1, 3, 4, 6] = 3.5
assert!((output_data[1] - 3.5).abs() < 0.01);
assert!((output_data[4] - 3.5).abs() < 0.01);
}
#[test]
fn test_median_imputation() {
let mut imputer = Imputer::new(ImputationStrategy::Median);
let data = create_data_with_missing();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Median of [1, 3, 4, 6] = 3.5
assert!((output_data[1] - 3.5).abs() < 0.01);
assert!((output_data[4] - 3.5).abs() < 0.01);
}
#[test]
fn test_mode_imputation() {
let mut imputer = Imputer::new(ImputationStrategy::Mode);
// Create data with clear mode
let data = Tensor::from_data(
vec![1.0, 2.0, 2.0, f32::NAN, 2.0, 3.0, f32::NAN],
vec![7],
&create_test_device(),
)
.unwrap();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Mode is 2.0
assert_eq!(output_data[3], 2.0);
assert_eq!(output_data[6], 2.0);
}
#[test]
fn test_constant_imputation() {
let mut imputer = Imputer::with_constant(0.0);
let data = create_data_with_missing();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Should be filled with 0.0
assert_eq!(output_data[1], 0.0);
assert_eq!(output_data[4], 0.0);
}
#[test]
fn test_forward_fill_imputation() {
let mut imputer = Imputer::new(ImputationStrategy::ForwardFill);
let data = create_data_with_missing();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Forward fill: NaN at index 1 gets 1.0, NaN at index 4 gets 4.0
assert_eq!(output_data[1], 1.0);
assert_eq!(output_data[4], 4.0);
}
#[test]
fn test_backward_fill_imputation() {
let mut imputer = Imputer::new(ImputationStrategy::BackwardFill);
let data = create_data_with_missing();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Backward fill: NaN at index 1 gets 3.0, NaN at index 4 gets 6.0
assert_eq!(output_data[1], 3.0);
assert_eq!(output_data[4], 6.0);
}
#[test]
fn test_2d_imputation() {
let mut imputer = Imputer::new(ImputationStrategy::Mean);
let data = create_2d_data_with_missing();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Check that missing values are filled
assert!(!output_data[2].is_nan());
assert!(!output_data[4].is_nan());
}
#[test]
fn test_fit_transform() {
let mut imputer = Imputer::new(ImputationStrategy::Mean);
let data = create_data_with_missing();
// Fit and transform in one call
let transformed = imputer.fit_transform(&data);
assert!(transformed.is_ok());
let output = transformed.unwrap();
let output_data = output.to_cpu().unwrap();
// Check that missing values are filled
assert!(!output_data[1].is_nan());
assert!(!output_data[4].is_nan());
}
#[test]
fn test_no_missing_values() {
let mut imputer = Imputer::new(ImputationStrategy::Mean);
let data =
Tensor::from_data(vec![1.0, 2.0, 3.0, 4.0], vec![4], &create_test_device()).unwrap();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Data should remain unchanged
assert_eq!(output_data, vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn test_all_missing_values() {
let mut imputer = Imputer::with_constant(0.0);
let data = Tensor::from_data(
vec![f32::NAN, f32::NAN, f32::NAN],
vec![3],
&create_test_device(),
)
.unwrap();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Should all be filled with constant value
assert_eq!(output_data, vec![0.0, 0.0, 0.0]);
}
#[test]
fn test_interpolation() {
let mut imputer = Imputer::new(ImputationStrategy::Interpolate);
let data = Tensor::from_data(
vec![1.0, f32::NAN, 3.0, f32::NAN, f32::NAN, 6.0],
vec![6],
&create_test_device(),
)
.unwrap();
imputer.fit(&data).unwrap();
let transformed = imputer.transform(&data).unwrap();
let output_data = transformed.to_cpu().unwrap();
// Linear interpolation
assert!((output_data[1] - 2.0).abs() < 0.01); // Between 1 and 3
assert!((output_data[3] - 4.0).abs() < 0.01); // Between 3 and 6
assert!((output_data[4] - 5.0).abs() < 0.01); // Between 3 and 6
}
#[test]
fn test_not_fitted_error() {
let imputer = Imputer::new(ImputationStrategy::Mean);
let data = create_data_with_missing();
// Should error if not fitted
let result = imputer.transform(&data);
assert!(result.is_err());
}
}