#![allow(clippy::approx_constant)] use approx::assert_abs_diff_eq; use rtx_preprocessing::{PolynomialFeatures, PreprocessingError, Transformer}; use rtx_tensor::{Device, Tensor}; #[test] fn test_polynomial_features_creation() { let transformer = PolynomialFeatures::new(); assert!(!transformer.is_fitted()); assert_eq!(transformer.degree(), 2); assert!(transformer.include_bias()); assert!(transformer.interaction_only()); let transformer = PolynomialFeatures::with_params(3, false, false); assert_eq!(transformer.degree(), 3); assert!(!transformer.include_bias()); assert!(!transformer.interaction_only()); } #[test] fn test_polynomial_features_invalid_degree() { let result = std::panic::catch_unwind(|| PolynomialFeatures::with_params(0, true, true)); assert!(result.is_err()); } #[test] fn test_polynomial_features_fit_simple() { let data = Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0], &[2, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::new(); let result = transformer.fit(&data); assert!(result.is_ok()); assert!(transformer.is_fitted()); } #[test] fn test_polynomial_features_transform_degree2() { let data = Tensor::from_slice(&[2.0, 3.0], &[1, 2], &Device::cpu()).unwrap(); // [x1=2, x2=3] let mut transformer = PolynomialFeatures::with_params(2, true, false); transformer.fit(&data).unwrap(); let poly_features = transformer.transform(&data).unwrap(); let values = poly_features .to_cpu() .unwrap() .iter() .copied() .collect::>(); // Expected features: [1, x1, x2, x1^2, x1*x2, x2^2] = [1, 2, 3, 4, 6, 9] let expected = vec![1.0, 2.0, 3.0, 4.0, 6.0, 9.0]; assert_eq!(values.len(), expected.len()); for (actual, expected) in values.iter().zip(expected.iter()) { assert_abs_diff_eq!(actual, expected, epsilon = 1e-5); } } #[test] fn test_polynomial_features_transform_no_bias() { let data = Tensor::from_slice(&[2.0, 3.0], &[1, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::with_params(2, false, false); transformer.fit(&data).unwrap(); let poly_features = transformer.transform(&data).unwrap(); let values = poly_features .to_cpu() .unwrap() .iter() .copied() .collect::>(); // Expected features: [x1, x2, x1^2, x1*x2, x2^2] = [2, 3, 4, 6, 9] let expected = vec![2.0, 3.0, 4.0, 6.0, 9.0]; assert_eq!(values.len(), expected.len()); for (actual, expected) in values.iter().zip(expected.iter()) { assert_abs_diff_eq!(actual, expected, epsilon = 1e-5); } } #[test] fn test_polynomial_features_interaction_only() { let data = Tensor::from_slice(&[2.0, 3.0], &[1, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::with_params(2, true, true); transformer.fit(&data).unwrap(); let poly_features = transformer.transform(&data).unwrap(); let values = poly_features .to_cpu() .unwrap() .iter() .copied() .collect::>(); // Expected features: [1, x1, x2, x1*x2] = [1, 2, 3, 6] (no x1^2, x2^2) let expected = vec![1.0, 2.0, 3.0, 6.0]; assert_eq!(values.len(), expected.len()); for (actual, expected) in values.iter().zip(expected.iter()) { assert_abs_diff_eq!(actual, expected, epsilon = 1e-5); } } #[test] fn test_polynomial_features_degree3() { let data = Tensor::from_slice(&[2.0, 3.0], &[1, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::with_params(3, false, false); transformer.fit(&data).unwrap(); let poly_features = transformer.transform(&data).unwrap(); let values = poly_features .to_cpu() .unwrap() .iter() .copied() .collect::>(); // Should include terms up to degree 3: x1, x2, x1^2, x1*x2, x2^2, x1^3, x1^2*x2, x1*x2^2, x2^3 assert!(values.len() >= 9); // At least 9 features for degree 3 with 2 input features // Check some specific values assert_abs_diff_eq!(values[0], 2.0, epsilon = 1e-5); // x1 assert_abs_diff_eq!(values[1], 3.0, epsilon = 1e-5); // x2 // Check for x1^3 = 8 and x2^3 = 27 somewhere in the features assert!(values.iter().any(|&v| (v - 8.0).abs() < 1e-10)); assert!(values.iter().any(|&v| (v - 27.0).abs() < 1e-10)); } #[test] fn test_polynomial_features_multiple_samples() { let data = Tensor::from_slice(&[1.0, 2.0, 3.0, 4.0], &[2, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::with_params(2, true, false); transformer.fit(&data).unwrap(); let poly_features = transformer.transform(&data).unwrap(); assert_eq!(poly_features.shape()[0], 2); // Still 2 samples assert!(poly_features.shape()[1] > 2); // More features than input } #[test] fn test_polynomial_features_single_feature() { let data = Tensor::from_slice(&[2.0, 3.0], &[2, 1], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::with_params(3, true, false); transformer.fit(&data).unwrap(); let poly_features = transformer.transform(&data).unwrap(); assert_eq!(poly_features.shape()[0], 2); // 2 samples assert_eq!(poly_features.shape()[1], 4); // [1, x, x^2, x^3] let values = poly_features .to_cpu() .unwrap() .iter() .copied() .collect::>(); // First sample: x=2 -> [1, 2, 4, 8] assert_abs_diff_eq!(values[0], 1.0, epsilon = 1e-5); assert_abs_diff_eq!(values[1], 2.0, epsilon = 1e-5); assert_abs_diff_eq!(values[2], 4.0, epsilon = 1e-5); assert_abs_diff_eq!(values[3], 8.0, epsilon = 1e-5); // Second sample: x=3 -> [1, 3, 9, 27] assert_abs_diff_eq!(values[4], 1.0, epsilon = 1e-5); assert_abs_diff_eq!(values[5], 3.0, epsilon = 1e-5); assert_abs_diff_eq!(values[6], 9.0, epsilon = 1e-5); assert_abs_diff_eq!(values[7], 27.0, epsilon = 1e-5); } #[test] fn test_polynomial_features_get_feature_names() { let data = Tensor::from_slice(&[1.0, 2.0], &[1, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::with_params(2, true, false); transformer.fit(&data).unwrap(); let feature_names = transformer.get_feature_names(); // Should have names like ["1", "x0", "x1", "x0^2", "x0 x1", "x1^2"] assert!(feature_names.len() >= 6); assert!(feature_names.contains(&"1".to_string())); assert!(feature_names.contains(&"x0".to_string())); assert!(feature_names.contains(&"x1".to_string())); } #[test] fn test_polynomial_features_reset() { let data = Tensor::from_slice(&[1.0, 2.0], &[1, 2], &Device::cpu()).unwrap(); let mut transformer = PolynomialFeatures::new(); transformer.fit(&data).unwrap(); assert!(transformer.is_fitted()); transformer.reset(); assert!(!transformer.is_fitted()); }