131 lines
4.3 KiB
Rust
131 lines
4.3 KiB
Rust
#![allow(clippy::approx_constant)]
|
|
use approx::assert_abs_diff_eq;
|
|
use rtx_preprocessing::{InvertibleTransformer, LabelEncoder, PreprocessingError, Transformer};
|
|
use rtx_tensor::{Device, Tensor};
|
|
|
|
#[test]
|
|
fn test_label_encoder_creation() {
|
|
let encoder = LabelEncoder::new();
|
|
assert!(!encoder.is_fitted());
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_fit_simple() {
|
|
let data = Tensor::from_slice(&[1.0, 3.0, 2.0, 1.0, 3.0], &[5, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
let result = encoder.fit(&data);
|
|
assert!(result.is_ok());
|
|
assert!(encoder.is_fitted());
|
|
assert_eq!(encoder.classes().len(), 3); // [1, 2, 3]
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_transform_simple() {
|
|
let data = Tensor::from_slice(&[1.0f32, 3.0, 2.0, 1.0, 3.0], &[5, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
encoder.fit(&data).unwrap();
|
|
let encoded = encoder.transform(&data).unwrap();
|
|
|
|
let values: Vec<f32> = encoded.to_cpu().unwrap().iter().copied().collect();
|
|
// Classes should be sorted: [1, 2, 3] -> encoded as [0, 1, 2]
|
|
// Original [1, 3, 2, 1, 3] -> [0, 2, 1, 0, 2]
|
|
let expected = vec![0.0f32, 2.0, 1.0, 0.0, 2.0];
|
|
for (actual, expected) in values.iter().zip(expected.iter()) {
|
|
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_string_labels() {
|
|
// This would require string support in tensors, might be implementation dependent
|
|
// For now, just test numeric labels
|
|
let data = Tensor::from_slice(&[10.0f32, 30.0, 20.0], &[3, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
encoder.fit(&data).unwrap();
|
|
let encoded = encoder.transform(&data).unwrap();
|
|
|
|
let values: Vec<f32> = encoded.to_cpu().unwrap().iter().copied().collect();
|
|
// Should be [0, 2, 1] for sorted classes [10, 20, 30]
|
|
let expected = vec![0.0f32, 2.0, 1.0];
|
|
for (actual, expected) in values.iter().zip(expected.iter()) {
|
|
assert_abs_diff_eq!(actual, expected, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_unknown_label() {
|
|
let data = Tensor::from_slice(&[1.0f32, 2.0, 3.0], &[3, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
encoder.fit(&data).unwrap();
|
|
|
|
let unknown_data = Tensor::from_slice(&[4.0f32], &[1, 1], &Device::cpu()).unwrap();
|
|
let result = encoder.transform(&unknown_data);
|
|
|
|
// Should fail on unknown label
|
|
assert!(result.is_err());
|
|
assert!(matches!(
|
|
result.unwrap_err(),
|
|
PreprocessingError::InvalidInput { .. }
|
|
));
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_inverse_transform() {
|
|
let data = Tensor::from_slice(&[1.0f32, 3.0, 2.0, 1.0], &[4, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
encoder.fit(&data).unwrap();
|
|
let encoded = encoder.transform(&data).unwrap();
|
|
let decoded = encoder.inverse_transform(&encoded).unwrap();
|
|
|
|
let original_values: Vec<f32> = data.to_cpu().unwrap().iter().copied().collect();
|
|
let decoded_values: Vec<f32> = decoded.to_cpu().unwrap().iter().copied().collect();
|
|
|
|
for (orig, dec) in original_values.iter().zip(decoded_values.iter()) {
|
|
assert_abs_diff_eq!(orig, dec, epsilon = 1e-5);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_empty_data() {
|
|
let empty_data = Tensor::zeros(&[0, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
let result = encoder.fit(&empty_data);
|
|
assert!(result.is_err());
|
|
assert!(matches!(
|
|
result.unwrap_err(),
|
|
PreprocessingError::EmptyDataset
|
|
));
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_multi_column() {
|
|
let data = Tensor::from_slice(&[1.0f32, 2.0, 3.0, 1.0], &[2, 2], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
// Label encoder typically works on 1D data
|
|
let result = encoder.fit(&data);
|
|
// Might succeed and flatten, or might error - implementation dependent
|
|
match result {
|
|
Ok(_) => (),
|
|
Err(e) => assert!(matches!(e, PreprocessingError::InvalidInput { .. })),
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_label_encoder_reset() {
|
|
let data = Tensor::from_slice(&[1.0f32, 2.0, 3.0], &[3, 1], &Device::cpu()).unwrap();
|
|
let mut encoder = LabelEncoder::new();
|
|
|
|
encoder.fit(&data).unwrap();
|
|
assert!(encoder.is_fitted());
|
|
|
|
encoder.reset();
|
|
assert!(!encoder.is_fitted());
|
|
}
|