mod common;
use candela::{Dimension, FloatLikeTensorElement, Layout, Tensor, ones};
use common::{assert_approx_eq, tensor_of};
use rstest::rstest;
#[test]
fn broadcast_layout_zero_stride_on_expanded_dim() {
let l = Layout::new(&[4]);
let b = l.broadcast(&[3, 4]).unwrap();
assert_eq!(b.shape(), &[3, 4]);
assert_eq!(b.stride()[0], 0);
assert_eq!(b.stride()[1], 1);
assert_eq!(b.len(), 12);
}
#[rstest]
#[case::f64(1.0f64, 2.0f64)]
#[case::f32(1.0f32, 2.0f32)]
fn broadcast_col_plus_row<T: FloatLikeTensorElement>(#[case] input1: T, #[case] input2: T) {
let col = Tensor::from_scalar(input1, &[3, 1]);
let row = Tensor::from_scalar(input2, &[1, 4]);
let result = (col + row).materialize();
assert_eq!(result.shape(), &[3, 4]);
assert_approx_eq(result.data(), &[3.0; 12]);
}
#[rstest]
#[case::f64(0.0f64)]
#[case::f32(0.0f32)]
fn broadcast_1d_against_2d<T: FloatLikeTensorElement>(#[case] _t: T) {
let v = tensor_of::<T>(&[0.0, 1.0, 2.0, 3.0], &[4]);
let m = Tensor::from_scalar(T::from_f64(1.0), &[3, 4]);
let result = (v + m).materialize();
assert_eq!(result.shape(), &[3, 4]);
let expected: Vec<f64> = [1.0, 2.0, 3.0, 4.0]
.iter()
.cycle()
.take(12)
.copied()
.collect();
assert_approx_eq(result.data(), &expected);
}
#[rstest]
#[case::f64(0.0f64)]
#[case::f32(0.0f32)]
fn broadcast_scalar_tensor_against_matrix<T: FloatLikeTensorElement>(#[case] _t: T) {
let s = Tensor::from_scalar(T::from_f64(5.0), &[1]);
let m = Tensor::from_scalar(T::from_f64(1.0), &[3, 3]);
let result = (s * m).materialize();
assert_approx_eq(result.data(), &[5.0; 9]);
}
#[rstest]
#[case::f64(0.0f64)]
#[case::f32(0.0f32)]
fn broadcast_size_one_dim_in_both<T: FloatLikeTensorElement>(#[case] _t: T) {
let col = tensor_of::<T>(&[0.0, 1.0, 2.0], &[3, 1]);
let row = tensor_of::<T>(&[0.0, 1.0, 2.0, 3.0], &[1, 4]);
let result = (col + row).materialize();
assert_eq!(result.shape(), &[3, 4]);
let expected = vec![
0.0, 1.0, 2.0, 3.0, 1.0, 2.0, 3.0, 4.0, 2.0, 3.0, 4.0, 5.0, ];
assert_approx_eq(result.data(), &expected);
}
#[rstest]
#[case::f64(0.0f64)]
#[case::f32(0.0f32)]
fn broadcast_mul_with_row_vector<T: FloatLikeTensorElement>(#[case] _t: T) {
let col = tensor_of::<T>(&[1.0, 2.0, 3.0], &[3, 1]);
let row = tensor_of::<T>(&[1.0, 2.0, 3.0], &[1, 3]);
let result = (col * row).materialize();
assert_eq!(result.shape(), &[3, 3]);
let expected = vec![
1.0, 2.0, 3.0, 2.0, 4.0, 6.0, 3.0, 6.0, 9.0, ];
assert_approx_eq(result.data(), &expected);
}
#[test]
#[should_panic]
fn broadcast_incompatible_shapes_panics() {
let a = ones!(&[3, 4]);
let b = ones!(&[2, 4]);
let _ = (a + b).materialize();
}