226 lines
6.9 KiB
Rust
226 lines
6.9 KiB
Rust
//! 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::<f32>().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));
|
|
}
|