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

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