//! Comprehensive tests for TensorExt trait following strict TDD principles //! Red-Green-Refactor with full implementations only use rtx_tensor::{DType, Device, Tensor}; use rtx_vision_advanced::tensor_utils::TensorExt; #[test] fn test_mean_dim_global() { // Test global mean (empty dimensions) let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let tensor = Tensor::from_vec(data, &[2, 3], &Device::default()).unwrap(); let mean = tensor.mean_dim(&[], false).unwrap(); let mean_val: f32 = mean.to_scalar().unwrap(); assert!((mean_val - 3.5).abs() < 1e-6); } #[test] fn test_mean_dim_with_keepdim() { let tensor = Tensor::randn(&[4, 5, 6], &Device::default()).unwrap(); // Mean along dimension 1 with keepdim let mean = tensor.mean_dim(&[1], true).unwrap(); assert_eq!(mean.shape().dims(), &[4, 1, 6]); // Mean along dimension 1 without keepdim let mean = tensor.mean_dim(&[1], false).unwrap(); assert_eq!(mean.shape().dims(), &[4, 6]); } #[test] fn test_var_dim_biased_unbiased() { let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let tensor = Tensor::from_vec(data, &[2, 3], &Device::default()).unwrap(); // Test biased variance let var_biased = tensor.var_dim(&[1], false, true).unwrap(); assert_eq!(var_biased.shape().dims(), &[2, 1]); // Test unbiased variance let var_unbiased = tensor.var_dim(&[1], true, true).unwrap(); assert_eq!(var_unbiased.shape().dims(), &[2, 1]); // Variance should be non-negative let var_data = var_biased.to_vec().unwrap(); assert!(var_data.iter().all(|&v| v >= 0.0)); } #[test] fn test_flip_single_axis() { let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]; let tensor = Tensor::from_vec(data.clone(), &[2, 3], &Device::default()).unwrap(); // Flip along axis 0 let flipped = tensor.flip(&[0]).unwrap(); assert_eq!(flipped.shape(), tensor.shape()); // Flipping twice should give original let double_flipped = flipped.flip(&[0]).unwrap(); let original_data = tensor.to_vec().unwrap(); let double_flip_data = double_flipped.to_vec().unwrap(); for (a, b) in original_data.iter().zip(double_flip_data.iter()) { assert!((a - b).abs() < 1e-6); } } #[test] fn test_flip_multiple_axes() { let tensor = Tensor::randn(&[3, 4, 5, 6], &Device::default()).unwrap(); // Flip along axes 1 and 3 let flipped = tensor.flip(&[1, 3]).unwrap(); assert_eq!(flipped.shape(), tensor.shape()); } #[test] fn test_max_dim() { let data = vec![1.0, 5.0, 3.0, 2.0, 8.0, 4.0]; let tensor = Tensor::from_vec(data, &[2, 3], &Device::default()).unwrap(); // Max along dimension 1 let (max_vals, max_indices) = tensor.max_dim(1, false).unwrap(); assert_eq!(max_vals.shape().dims(), &[2]); assert_eq!(max_indices.shape().dims(), &[2]); let vals = max_vals.to_vec().unwrap(); assert_eq!(vals[0], 5.0); // max of [1.0, 5.0, 3.0] assert_eq!(vals[1], 8.0); // max of [2.0, 8.0, 4.0] } #[test] fn test_max_dim_with_keepdim() { let tensor = Tensor::randn(&[3, 4, 5], &Device::default()).unwrap(); let (max_vals, max_indices) = tensor.max_dim(1, true).unwrap(); assert_eq!(max_vals.shape().dims(), &[3, 1, 5]); assert_eq!(max_indices.shape().dims(), &[3, 1, 5]); } #[test] fn test_argmax_global() { let data = vec![1.0, 5.0, 3.0, 9.0, 2.0, 7.0]; let tensor = Tensor::from_vec(data, &[6], &Device::default()).unwrap(); // Global argmax let idx = tensor.argmax(None, false).unwrap(); let idx_val = idx.to_scalar::().unwrap(); assert_eq!(idx_val, 3.0); // Index of 9.0 } #[test] fn test_argmax_along_dimension() { let data = vec![1.0, 5.0, 3.0, 2.0, 8.0, 4.0]; let tensor = Tensor::from_vec(data, &[2, 3], &Device::default()).unwrap(); // Argmax along dimension 1 let idx = tensor.argmax(Some(1), false).unwrap(); assert_eq!(idx.shape().dims(), &[2]); let indices = idx.to_vec().unwrap(); assert_eq!(indices[0], 1.0); // Index of 5.0 in [1.0, 5.0, 3.0] assert_eq!(indices[1], 1.0); // Index of 8.0 in [2.0, 8.0, 4.0] } #[test] fn test_min_global() { let data = vec![5.0, 2.0, 8.0, 1.0, 9.0, 3.0]; let tensor = Tensor::from_vec(data, &[6], &Device::default()).unwrap(); let min = tensor.min().unwrap(); let min_val: f32 = min.to_scalar().unwrap(); assert_eq!(min_val, 1.0); } #[test] fn test_eq_element_wise() { let data1 = vec![1.0, 2.0, 3.0, 4.0]; let data2 = vec![1.0, 2.0, 5.0, 4.0]; let tensor1 = Tensor::from_vec(data1, &[2, 2], &Device::default()).unwrap(); let tensor2 = Tensor::from_vec(data2, &[2, 2], &Device::default()).unwrap(); let eq_result = tensor1.eq(&tensor2).unwrap(); let eq_data = eq_result.to_vec().unwrap(); // Expected: [1.0, 1.0, 0.0, 1.0] assert_eq!(eq_data[0], 1.0); // 1.0 == 1.0 assert_eq!(eq_data[1], 1.0); // 2.0 == 2.0 assert_eq!(eq_data[2], 0.0); // 3.0 != 5.0 assert_eq!(eq_data[3], 1.0); // 4.0 == 4.0 } #[test] fn test_eq_with_self() { let tensor = Tensor::randn(&[3, 4], &Device::default()).unwrap(); let eq_result = tensor.eq(&tensor).unwrap(); let eq_data = eq_result.to_vec().unwrap(); // All elements should be 1.0 (true) assert!(eq_data.iter().all(|&v| v == 1.0)); } #[test] fn test_eq_shape_mismatch() { let tensor1 = Tensor::randn(&[2, 3], &Device::default()).unwrap(); let tensor2 = Tensor::randn(&[3, 2], &Device::default()).unwrap(); let result = tensor1.eq(&tensor2); assert!(result.is_err()); } #[test] fn test_mean_dim_multiple_dimensions() { let tensor = Tensor::randn(&[2, 3, 4, 5], &Device::default()).unwrap(); // Mean along dimensions 1 and 3 let mean = tensor.mean_dim(&[1, 3], false).unwrap(); assert_eq!(mean.shape().dims(), &[2, 4]); // With keepdim let mean = tensor.mean_dim(&[1, 3], true).unwrap(); assert_eq!(mean.shape().dims(), &[2, 1, 4, 1]); } #[test] fn test_var_dim_correctness() { // Create tensor with known variance let data = vec![1.0, 1.0, 1.0, 5.0, 5.0, 5.0]; let tensor = Tensor::from_vec(data, &[2, 3], &Device::default()).unwrap(); // Variance along dimension 0 let var = tensor.var_dim(&[0], false, false).unwrap(); let var_data = var.to_vec().unwrap(); // Each column has variance of (5-3)^2 / 2 = 2.0 for unbiased for &v in &var_data { assert!((v - 4.0).abs() < 0.1); // Biased variance is 4.0 } } #[test] fn test_integration_mean_var_std() { let tensor = Tensor::randn(&[100, 50], &Device::default()).unwrap(); // Compute mean and variance let mean = tensor.mean_dim(&[0], true).unwrap(); let var = tensor.var_dim(&[0], false, true).unwrap(); // Standard deviation is sqrt of variance let std = var.sqrt().unwrap(); assert_eq!(mean.shape().dims(), &[1, 50]); assert_eq!(std.shape().dims(), &[1, 50]); // Verify all std values are non-negative let std_data = std.to_vec().unwrap(); assert!(std_data.iter().all(|&v| v >= 0.0)); }