mod backend;
mod error;
mod tensor;
pub use backend::Backend;
pub use error::{Result, TensorError};
pub use tensor::{ops::TensorOps, shape::Shape, Tensor};
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_tensor_zeros() {
let tensor = Tensor::zeros(vec![2, 3]).unwrap();
assert_eq!(tensor.shape().dims(), &[2, 3]);
assert_eq!(tensor.numel(), 6);
let data = tensor.to_vec().unwrap();
assert!(data.iter().all(|&x| x == 0.0));
}
#[test]
fn test_tensor_ones() {
let tensor = Tensor::ones(vec![3, 2]).unwrap();
assert_eq!(tensor.shape().dims(), &[3, 2]);
assert_eq!(tensor.numel(), 6);
let data = tensor.to_vec().unwrap();
assert!(data.iter().all(|&x| x == 1.0));
}
#[test]
fn test_tensor_from_vec() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let tensor = Tensor::from_vec(data.clone(), vec![2, 3]).unwrap();
assert_eq!(tensor.shape().dims(), &[2, 3]);
assert_eq!(tensor.to_vec().unwrap(), data);
}
#[test]
fn test_tensor_1d() {
let tensor = Tensor::ones(vec![5]).unwrap();
assert_eq!(tensor.shape().dims(), &[5]);
assert_eq!(tensor.numel(), 5);
}
#[test]
fn test_tensor_3d() {
let tensor = Tensor::zeros(vec![2, 3, 4]).unwrap();
assert_eq!(tensor.shape().dims(), &[2, 3, 4]);
assert_eq!(tensor.numel(), 24);
}
#[test]
fn test_tensor_scalar() {
let tensor = Tensor::from_vec(vec![42.0], vec![]).unwrap();
assert_eq!(tensor.shape().dims(), &[] as &[usize]);
assert_eq!(tensor.numel(), 1);
assert_eq!(tensor.to_vec().unwrap(), vec![42.0]);
}
#[test]
fn test_tensor_addition() {
let a = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![5.0, 6.0, 7.0, 8.0], vec![2, 2]).unwrap();
let c = (a + b).unwrap();
assert_eq!(c.to_vec().unwrap(), vec![6.0, 8.0, 10.0, 12.0]);
}
#[test]
fn test_tensor_subtraction() {
let a = Tensor::from_vec(vec![5.0, 6.0, 7.0, 8.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let c = (a - b).unwrap();
assert_eq!(c.to_vec().unwrap(), vec![4.0, 4.0, 4.0, 4.0]);
}
#[test]
fn test_tensor_multiplication() {
let a = Tensor::from_vec(vec![2.0, 3.0, 4.0, 5.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let c = (a * b).unwrap();
assert_eq!(c.to_vec().unwrap(), vec![2.0, 6.0, 12.0, 20.0]);
}
#[test]
fn test_tensor_division() {
let a = Tensor::from_vec(vec![8.0, 12.0, 16.0, 20.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![2.0, 3.0, 4.0, 5.0], vec![2, 2]).unwrap();
let c = (a / b).unwrap();
assert_eq!(c.to_vec().unwrap(), vec![4.0, 4.0, 4.0, 4.0]);
}
#[test]
fn test_tensor_chain_operations() {
let a = Tensor::ones(vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![2.0, 2.0, 2.0, 2.0], vec![2, 2]).unwrap();
let c = Tensor::from_vec(vec![3.0, 3.0, 3.0, 3.0], vec![2, 2]).unwrap();
let result = ((a + b).unwrap() * c).unwrap();
assert_eq!(result.to_vec().unwrap(), vec![9.0, 9.0, 9.0, 9.0]);
}
#[test]
fn test_broadcast_2d_1d() {
let a = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![10.0, 20.0], vec![2]).unwrap();
assert_eq!(a.shape().dims(), &[2, 2]);
assert_eq!(b.shape().dims(), &[2]);
}
#[test]
fn test_broadcast_same_shape() {
let a = Tensor::ones(vec![2, 3]).unwrap();
let b = Tensor::ones(vec![2, 3]).unwrap();
let c = (a + b).unwrap();
assert_eq!(c.shape().dims(), &[2, 3]);
let data = c.to_vec().unwrap();
assert!(data.iter().all(|&x| x == 2.0));
}
#[test]
fn test_broadcast_compatible_shapes() {
let a = Tensor::ones(vec![2, 1]).unwrap();
let b = Tensor::ones(vec![1, 3]).unwrap();
let c = (a + b).unwrap();
assert_eq!(c.shape().dims(), &[2, 3]);
let data = c.to_vec().unwrap();
assert!(data.iter().all(|&x| x == 2.0));
}
#[test]
fn test_broadcast_scalar() {
let a = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![5.0], vec![]).unwrap(); assert_eq!(a.shape().dims(), &[2, 2]);
assert_eq!(b.shape().dims(), &[] as &[usize]);
}
#[test]
fn test_tensor_sum() {
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let sum = tensor.sum(None).unwrap();
assert_eq!(sum.to_vec().unwrap(), vec![10.0]);
}
#[test]
fn test_tensor_mean() {
let tensor = Tensor::from_vec(vec![2.0, 4.0, 6.0, 8.0], vec![2, 2]).unwrap();
let mean = tensor.mean(None).unwrap();
assert_eq!(mean.to_vec().unwrap(), vec![5.0]);
}
#[test]
fn test_sum_ones() {
let tensor = Tensor::ones(vec![3, 3]).unwrap();
let sum = tensor.sum(None).unwrap();
assert_eq!(sum.to_vec().unwrap(), vec![9.0]);
}
#[test]
fn test_mean_zeros() {
let tensor = Tensor::zeros(vec![2, 5]).unwrap();
let mean = tensor.mean(None).unwrap();
assert_eq!(mean.to_vec().unwrap(), vec![0.0]);
}
#[test]
fn test_tensor_reshape() {
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]).unwrap();
let reshaped = tensor.reshape(vec![3, 2]).unwrap();
assert_eq!(reshaped.shape().dims(), &[3, 2]);
assert_eq!(
reshaped.to_vec().unwrap(),
vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]
);
}
#[test]
fn test_tensor_reshape_1d() {
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let reshaped = tensor.reshape(vec![4]).unwrap();
assert_eq!(reshaped.shape().dims(), &[4]);
assert_eq!(reshaped.to_vec().unwrap(), vec![1.0, 2.0, 3.0, 4.0]);
}
#[test]
fn test_tensor_transpose_2d() {
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]).unwrap();
let transposed = tensor.transpose().unwrap();
assert_eq!(transposed.shape().dims(), &[3, 2]);
assert_eq!(
transposed.to_vec().unwrap(),
vec![1.0, 4.0, 2.0, 5.0, 3.0, 6.0]
);
}
#[test]
fn test_shape_mismatch_from_vec() {
let data = vec![1.0, 2.0, 3.0];
let result = Tensor::from_vec(data, vec![2, 2]); assert!(result.is_err());
if let Err(TensorError::ShapeMismatch { expected, got }) = result {
assert_eq!(expected, vec![4]);
assert_eq!(got, vec![3]);
}
}
#[test]
fn test_incompatible_shapes_addition() {
let a = Tensor::ones(vec![2, 3]).unwrap();
let b = Tensor::ones(vec![3, 4]).unwrap();
let result = a + b;
match result {
Ok(_) => {} Err(_) => {} }
}
#[test]
fn test_invalid_reshape() {
let tensor = Tensor::ones(vec![2, 3]).unwrap(); let result = tensor.reshape(vec![2, 2]); assert!(result.is_err());
}
#[test]
fn test_transpose_1d() {
let tensor = Tensor::ones(vec![5]).unwrap();
let result = tensor.transpose();
assert!(result.is_ok() || result.is_err());
}
#[test]
fn test_empty_tensor() {
let tensor = Tensor::zeros(vec![0]).unwrap();
assert_eq!(tensor.numel(), 0);
assert_eq!(tensor.to_vec().unwrap(), Vec::<f32>::new());
}
#[test]
fn test_large_tensor() {
let tensor = Tensor::zeros(vec![100, 100]).unwrap();
assert_eq!(tensor.numel(), 10000);
assert_eq!(tensor.shape().dims(), &[100, 100]);
}
#[test]
fn test_operations_with_negative_numbers() {
let a = Tensor::from_vec(vec![-1.0, -2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let b = Tensor::from_vec(vec![1.0, 2.0, -3.0, -4.0], vec![2, 2]).unwrap();
let sum = (a.clone() + b.clone()).unwrap();
assert_eq!(sum.to_vec().unwrap(), vec![0.0, 0.0, 0.0, 0.0]);
let product = (a * b).unwrap();
assert_eq!(product.to_vec().unwrap(), vec![-1.0, -4.0, -9.0, -16.0]);
}
#[test]
fn test_operations_with_zero() {
let a = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let zeros = Tensor::zeros(vec![2, 2]).unwrap();
let sum = (a.clone() + zeros.clone()).unwrap();
assert_eq!(sum.to_vec().unwrap(), vec![1.0, 2.0, 3.0, 4.0]);
let product = (a * zeros).unwrap();
assert_eq!(product.to_vec().unwrap(), vec![0.0, 0.0, 0.0, 0.0]);
}
#[test]
fn test_display_formatting() {
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0], vec![2, 2]).unwrap();
let display_str = format!("{}", tensor);
assert!(!display_str.is_empty());
}
#[test]
fn test_shape_validation() {
use crate::tensor::shape::Shape;
assert!(Shape::new(vec![0]).is_ok());
assert!(Shape::new(vec![2, 0]).is_ok());
assert!(Shape::new(vec![0, 3]).is_ok());
assert!(Shape::new(vec![2, 0, 3]).is_ok());
assert!(Shape::new(vec![1]).is_ok());
assert!(Shape::new(vec![2, 3]).is_ok());
assert!(Shape::new(vec![]).is_ok());
let empty = Shape::new(vec![0]).unwrap();
assert_eq!(empty.numel(), 0);
let empty2 = Shape::new(vec![2, 0, 3]).unwrap();
assert_eq!(empty2.numel(), 0);
}
#[test]
fn test_overflow_protection() {
use crate::tensor::shape::Shape;
let huge_dims = vec![usize::MAX, 2];
assert!(Shape::new(huge_dims).is_err());
let large_dims = vec![1000000, 1000000, 1000000];
let result = Shape::new(large_dims);
if result.is_ok() {
let huge_dims = vec![usize::MAX / 2, usize::MAX / 2];
assert!(Shape::new(huge_dims).is_err());
} else {
assert!(result.is_err());
}
}
#[test]
fn test_tensor_creation_with_mismatched_data() {
let result = Tensor::from_vec_with_shape(vec![1.0, 2.0], vec![3, 2]);
assert!(result.is_err());
let result2 = Tensor::from_vec_with_shape(Vec::new(), vec![0]);
assert!(result2.is_ok());
let result3 = Tensor::from_vec_with_shape(vec![1.0, 2.0], vec![1, 2]);
assert!(result3.is_ok());
}
#[test]
fn test_division_by_zero_handling() {
let numerator = Tensor::from_vec(vec![1.0, -1.0, 0.0, 5.0], vec![4]).unwrap();
let denominator = Tensor::from_vec(vec![0.0, 0.0, 0.0, 2.0], vec![4]).unwrap();
let result = (numerator / denominator).unwrap();
let values = result.to_vec().unwrap();
assert!(values[0].is_infinite() && values[0].is_sign_positive()); assert!(values[1].is_infinite() && values[1].is_sign_negative()); assert!(values[2].is_nan()); assert_eq!(values[3], 2.5); }
#[test]
fn test_division_by_near_zero() {
let numerator = Tensor::from_vec(vec![1.0, 2.0], vec![2]).unwrap();
let denominator = Tensor::from_vec(vec![1e-10, 1e-20], vec![2]).unwrap();
let result = (numerator / denominator).unwrap();
let values = result.to_vec().unwrap();
assert!(values[0].is_finite());
assert!(values[1].is_finite());
assert!(values[0] > 1e9); assert!(values[1] > 1e19); }
#[test]
fn test_axis_specific_sum() {
use crate::tensor::ops::TensorOps;
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]).unwrap();
let sum_axis_0 = tensor.sum(Some(0)).unwrap();
let result_0 = sum_axis_0.to_vec().unwrap();
assert_eq!(result_0, vec![5.0, 7.0, 9.0]);
assert_eq!(sum_axis_0.shape().dims(), &[3]);
let sum_axis_1 = tensor.sum(Some(1)).unwrap();
let result_1 = sum_axis_1.to_vec().unwrap();
assert_eq!(result_1, vec![6.0, 15.0]);
assert_eq!(sum_axis_1.shape().dims(), &[2]);
let sum_all = tensor.sum(None).unwrap();
let result_all = sum_all.to_vec().unwrap();
assert_eq!(result_all, vec![21.0]);
assert_eq!(sum_all.shape().dims(), &[] as &[usize]);
}
#[test]
fn test_axis_specific_mean() {
use crate::tensor::ops::TensorOps;
let tensor = Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], vec![2, 3]).unwrap();
let mean_axis_0 = tensor.mean(Some(0)).unwrap();
let result_0 = mean_axis_0.to_vec().unwrap();
assert_eq!(result_0, vec![2.5, 3.5, 4.5]);
assert_eq!(mean_axis_0.shape().dims(), &[3]);
let mean_axis_1 = tensor.mean(Some(1)).unwrap();
let result_1 = mean_axis_1.to_vec().unwrap();
assert_eq!(result_1, vec![2.0, 5.0]);
assert_eq!(mean_axis_1.shape().dims(), &[2]);
let mean_all = tensor.mean(None).unwrap();
let result_all = mean_all.to_vec().unwrap();
assert_eq!(result_all, vec![3.5]);
assert_eq!(mean_all.shape().dims(), &[] as &[usize]);
}
#[test]
fn test_axis_sum_3d_tensor() {
use crate::tensor::ops::TensorOps;
let tensor =
Tensor::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0], vec![2, 2, 2]).unwrap();
let sum_axis_0 = tensor.sum(Some(0)).unwrap();
assert_eq!(sum_axis_0.shape().dims(), &[2, 2]);
let result_0 = sum_axis_0.to_vec().unwrap();
assert_eq!(result_0, vec![6.0, 8.0, 10.0, 12.0]);
let sum_axis_1 = tensor.sum(Some(1)).unwrap();
assert_eq!(sum_axis_1.shape().dims(), &[2, 2]);
let result_1 = sum_axis_1.to_vec().unwrap();
assert_eq!(result_1, vec![4.0, 6.0, 12.0, 14.0]);
let sum_axis_2 = tensor.sum(Some(2)).unwrap();
assert_eq!(sum_axis_2.shape().dims(), &[2, 2]);
let result_2 = sum_axis_2.to_vec().unwrap();
assert_eq!(result_2, vec![3.0, 7.0, 11.0, 15.0]); }
}