use approx::assert_relative_eq;
use num_complex::Complex64;
use strided_kernel::{
add, axpy, batched_outer_product_into, broadcast_mul_into, copy_conj, copy_into, copy_scale,
copy_transpose_scale_into, dot, fma, map_into, mul, mul_into, reduce, reduce_axis, sum,
symmetrize_conj_into, symmetrize_into, zip_map2_into, zip_map3_into, zip_map4_into,
StridedArray, StridedError,
};
fn make_tensor(rows: usize, cols: usize) -> StridedArray<f64> {
StridedArray::from_fn_row_major(&[rows, cols], |idx| (idx[0] * cols + idx[1]) as f64)
}
#[test]
fn test_map_into_transposed() {
let a = make_tensor(8, 5);
let a_view = a.view();
let a_t = a_view.permute(&[1, 0]).unwrap();
let mut out = StridedArray::<f64>::row_major(&[5, 8]);
map_into(&mut out.view_mut(), &a_t, |x| x * 2.0).unwrap();
for i in 0..5 {
for j in 0..8 {
let expected = a.get(&[j, i]) * 2.0;
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_zip_map2_mixed_strides() {
let a = make_tensor(6, 4);
let b = make_tensor(6, 4);
let a_view = a.view();
let b_view = b.view();
let a_t = a_view.permute(&[1, 0]).unwrap();
let b_t = b_view.permute(&[1, 0]).unwrap();
let mut out = StridedArray::<f64>::row_major(&[4, 6]);
zip_map2_into(&mut out.view_mut(), &a_t, &b_t, |x, y| x + y).unwrap();
for i in 0..4 {
for j in 0..6 {
let expected = a.get(&[j, i]) + b.get(&[j, i]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_mul_into_contiguous() {
let a = make_tensor(4, 5);
let b = StridedArray::<f64>::from_fn_row_major(&[4, 5], |idx| (1 + idx[0] + 2 * idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[4, 5]);
mul_into(&mut out.view_mut(), &a.view(), &b.view()).unwrap();
for i in 0..4 {
for j in 0..5 {
assert_relative_eq!(
out.get(&[i, j]),
a.get(&[i, j]) * b.get(&[i, j]),
epsilon = 1e-10
);
}
}
}
#[test]
fn test_mul_into_broadcast_stride_zero() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 1], |idx| (1 + idx[0]) as f64);
let b = StridedArray::<f64>::from_fn_row_major(&[1, 4], |idx| (10 + idx[1]) as f64);
let a_broadcast = a.view().broadcast(&[3, 4]).unwrap();
let b_broadcast = b.view().broadcast(&[3, 4]).unwrap();
let mut out = StridedArray::<f64>::row_major(&[3, 4]);
mul_into(&mut out.view_mut(), &a_broadcast, &b_broadcast).unwrap();
for i in 0..3 {
for j in 0..4 {
assert_relative_eq!(
out.get(&[i, j]),
a.get(&[i, 0]) * b.get(&[0, j]),
epsilon = 1e-10
);
}
}
}
#[test]
fn test_mul_into_permuted_noncompact() {
let a =
StridedArray::<f64>::from_fn_row_major(&[5, 4], |idx| (1 + idx[0] * 10 + idx[1]) as f64);
let b = StridedArray::<f64>::from_fn_row_major(&[5, 4], |idx| (2 + idx[0] * 7 + idx[1]) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap();
let b_t = b.view().permute(&[1, 0]).unwrap();
let mut out_base = StridedArray::<f64>::row_major(&[5, 4]);
let mut out_t = out_base.view_mut().permute(&[1, 0]).unwrap();
mul_into(&mut out_t, &a_t, &b_t).unwrap();
for i in 0..4 {
for j in 0..5 {
assert_relative_eq!(
out_t.get(&[i, j]),
a_t.get(&[i, j]) * b_t.get(&[i, j]),
epsilon = 1e-10
);
}
}
}
#[test]
fn test_mul_into_rejects_shape_mismatch() {
let a = StridedArray::<f64>::row_major(&[2, 3]);
let b = StridedArray::<f64>::row_major(&[3, 2]);
let mut out = StridedArray::<f64>::row_major(&[2, 3]);
let err = mul_into(&mut out.view_mut(), &a.view(), &b.view()).unwrap_err();
assert!(matches!(err, StridedError::ShapeMismatch(_, _)));
}
#[test]
fn test_broadcast_mul_into_maps_source_axes_without_copy() {
let lhs =
StridedArray::<f64>::from_fn_col_major(&[2, 3], |idx| (1 + idx[0] + 10 * idx[1]) as f64);
let rhs =
StridedArray::<f64>::from_fn_col_major(&[4, 3], |idx| (2 + idx[0] + 20 * idx[1]) as f64);
let mut out = StridedArray::<f64>::col_major(&[2, 4, 3]);
broadcast_mul_into(
&mut out.view_mut(),
&lhs.view(),
&[0, 2],
&rhs.view(),
&[1, 2],
)
.unwrap();
for j in 0..2 {
for o in 0..4 {
for t in 0..3 {
assert_relative_eq!(
out.get(&[j, o, t]),
lhs.get(&[j, t]) * rhs.get(&[o, t]),
epsilon = 1e-10
);
}
}
}
}
#[test]
fn test_broadcast_mul_into_broadcasts_mapped_size_one_axes() {
let lhs = StridedArray::<f64>::from_fn_row_major(&[3, 1], |idx| (1 + idx[0]) as f64);
let rhs = StridedArray::<f64>::from_fn_row_major(&[1, 4], |idx| (10 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[3, 4]);
broadcast_mul_into(
&mut out.view_mut(),
&lhs.view(),
&[0, 1],
&rhs.view(),
&[0, 1],
)
.unwrap();
for i in 0..3 {
for j in 0..4 {
assert_relative_eq!(
out.get(&[i, j]),
lhs.get(&[i, 0]) * rhs.get(&[0, j]),
epsilon = 1e-10
);
}
}
}
#[test]
fn test_broadcast_mul_into_rejects_invalid_axis_map() {
let lhs = StridedArray::<f64>::row_major(&[2, 3]);
let rhs = StridedArray::<f64>::row_major(&[4]);
let mut out = StridedArray::<f64>::row_major(&[2, 4, 3]);
let err = broadcast_mul_into(&mut out.view_mut(), &lhs.view(), &[0, 3], &rhs.view(), &[1])
.unwrap_err();
assert!(matches!(
err,
StridedError::InvalidAxis { axis: 3, rank: 3 }
));
}
#[test]
fn test_batched_outer_product_into_compact() {
let lhs =
StridedArray::<f64>::from_fn_col_major(&[2, 3], |idx| (1 + idx[0] + 10 * idx[1]) as f64);
let rhs =
StridedArray::<f64>::from_fn_col_major(&[4, 3], |idx| (2 + idx[0] + 20 * idx[1]) as f64);
let mut out = StridedArray::<f64>::col_major(&[2, 4, 3]);
batched_outer_product_into(&mut out.view_mut(), &lhs.view(), &rhs.view(), 1, 1).unwrap();
for j in 0..2 {
for o in 0..4 {
for t in 0..3 {
assert_relative_eq!(
out.get(&[j, o, t]),
lhs.get(&[j, t]) * rhs.get(&[o, t]),
epsilon = 1e-10
);
}
}
}
}
#[test]
fn test_batched_outer_product_into_mixed_types() {
let lhs = StridedArray::<Complex64>::from_fn_col_major(&[2, 3], |idx| {
Complex64::new((1 + idx[0] + 10 * idx[1]) as f64, (2 + idx[0]) as f64)
});
let rhs =
StridedArray::<f64>::from_fn_col_major(&[4, 3], |idx| (2 + idx[0] + 20 * idx[1]) as f64);
let mut out = StridedArray::<Complex64>::col_major(&[2, 4, 3]);
batched_outer_product_into(&mut out.view_mut(), &lhs.view(), &rhs.view(), 1, 1).unwrap();
for j in 0..2 {
for o in 0..4 {
for t in 0..3 {
let expected = lhs.get(&[j, t]) * rhs.get(&[o, t]);
assert_relative_eq!(out.get(&[j, o, t]).re, expected.re, epsilon = 1e-10);
assert_relative_eq!(out.get(&[j, o, t]).im, expected.im, epsilon = 1e-10);
}
}
}
}
#[test]
fn test_batched_outer_product_into_noncompact() {
let lhs_base =
StridedArray::<f64>::from_fn_col_major(&[3, 2], |idx| (1 + idx[0] + 10 * idx[1]) as f64);
let rhs_base =
StridedArray::<f64>::from_fn_col_major(&[3, 4], |idx| (2 + idx[0] + 20 * idx[1]) as f64);
let lhs = lhs_base.view().permute(&[1, 0]).unwrap(); let rhs = rhs_base.view().permute(&[1, 0]).unwrap(); let mut out_base = StridedArray::<f64>::col_major(&[3, 4, 2]);
let mut out = out_base.view_mut().permute(&[2, 1, 0]).unwrap();
batched_outer_product_into(&mut out, &lhs, &rhs, 1, 1).unwrap();
for j in 0..2 {
for o in 0..4 {
for t in 0..3 {
assert_relative_eq!(
out.get(&[j, o, t]),
lhs.get(&[j, t]) * rhs.get(&[o, t]),
epsilon = 1e-10
);
}
}
}
}
#[test]
fn test_batched_outer_product_into_matches_broadcast_mul_into() {
let lhs_base = StridedArray::<f64>::from_fn_col_major(&[3, 2, 5], |idx| {
(1 + idx[0] + 10 * idx[1] + 100 * idx[2]) as f64
});
let lhs = lhs_base.view().permute(&[1, 0, 2]).unwrap();
let rhs =
StridedArray::<f64>::from_fn_col_major(&[4, 5], |idx| (2 + idx[0] + 20 * idx[1]) as f64);
let mut outer_out = StridedArray::<f64>::col_major(&[2, 3, 4, 5]);
let mut broadcast_out = StridedArray::<f64>::col_major(&[2, 3, 4, 5]);
batched_outer_product_into(&mut outer_out.view_mut(), &lhs, &rhs.view(), 2, 1).unwrap();
broadcast_mul_into(
&mut broadcast_out.view_mut(),
&lhs,
&[0, 1, 3],
&rhs.view(),
&[2, 3],
)
.unwrap();
for j in 0..2 {
for k in 0..3 {
for o in 0..4 {
for t in 0..5 {
assert_relative_eq!(
outer_out.get(&[j, k, o, t]),
broadcast_out.get(&[j, k, o, t]),
epsilon = 1e-10
);
}
}
}
}
}
#[test]
fn test_batched_outer_product_into_rejects_mismatched_batch_dims() {
let lhs = StridedArray::<f64>::col_major(&[2, 3]);
let rhs = StridedArray::<f64>::col_major(&[4, 5]);
let mut out = StridedArray::<f64>::col_major(&[2, 4, 3]);
let err = batched_outer_product_into(&mut out.view_mut(), &lhs.view(), &rhs.view(), 1, 1)
.unwrap_err();
assert!(matches!(
err,
StridedError::ShapeMismatch(actual, expected)
if actual == vec![3] && expected == vec![5]
));
}
#[test]
fn test_reduce_sum() {
let a = make_tensor(10, 12);
let result = reduce(&a.view(), |x| x, |a, b| a + b, 0.0).unwrap();
let expected: f64 = a.iter().copied().sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_axis_sum() {
let a = StridedArray::<f64>::from_fn_row_major(&[4, 3, 2], |idx| {
(idx[0] + 2 * idx[1] + 3 * idx[2]) as f64
});
let result = reduce_axis(&a.view(), 1, |x| x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[4, 2]);
for i in 0..4 {
for k in 0..2 {
let mut expected = 0.0;
for j in 0..3 {
expected += (i + 2 * j + 3 * k) as f64;
}
assert_relative_eq!(result.get(&[i, k]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_dot() {
let a = make_tensor(7, 3);
let b = make_tensor(7, 3);
let result = dot(&a.view(), &b.view()).unwrap();
let expected: f64 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_copy_into() {
let a = make_tensor(4, 5);
let mut out = StridedArray::<f64>::row_major(&[4, 5]);
copy_into(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..4 {
for j in 0..5 {
assert_relative_eq!(out.get(&[i, j]), a.get(&[i, j]), epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_into_mixed_layouts() {
let a = StridedArray::<f64>::from_fn_col_major(&[4, 5], |idx| (idx[0] * 10 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[4, 5]);
copy_into(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..4 {
for j in 0..5 {
assert_relative_eq!(out.get(&[i, j]), a.get(&[i, j]), epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_transpose_scale_into() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 10 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[3, 2]);
copy_transpose_scale_into(&mut out.view_mut(), &a.view(), 3.0).unwrap();
for i in 0..2 {
for j in 0..3 {
assert_relative_eq!(out.get(&[j, i]), 3.0 * a.get(&[i, j]), epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_transpose_scale_into_zero_does_not_read_source_values() {
let a = StridedArray::<f64>::from_fn_col_major(&[2, 3], |_| f64::NAN);
let mut out = StridedArray::<f64>::from_fn_col_major(&[3, 2], |_| 1.0);
copy_transpose_scale_into(&mut out.view_mut(), &a.view(), 0.0).unwrap();
for i in 0..3 {
for j in 0..2 {
assert_eq!(out.get(&[i, j]), 0.0);
}
}
}
#[test]
fn test_symmetrize_into() {
let n = 4;
let a = StridedArray::<f64>::from_fn_row_major(&[n, n], |idx| (idx[0] * 10 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[n, n]);
symmetrize_into(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..n {
for j in 0..n {
let expected = (a.get(&[i, j]) + a.get(&[j, i])) * 0.5;
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_zip_map4_into_contiguous() {
let a = make_tensor(4, 5);
let b = make_tensor(4, 5);
let c = make_tensor(4, 5);
let d = make_tensor(4, 5);
let mut out = StridedArray::<f64>::row_major(&[4, 5]);
zip_map4_into(
&mut out.view_mut(),
&a.view(),
&b.view(),
&c.view(),
&d.view(),
|a, b, c, d| a + b + c + d,
)
.unwrap();
for i in 0..4 {
for j in 0..5 {
let expected = a.get(&[i, j]) + b.get(&[i, j]) + c.get(&[i, j]) + d.get(&[i, j]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_zip_map4_into_permuted() {
let size = 8usize;
let a = StridedArray::<f64>::from_fn_row_major(&[size, size, size, size], |idx| {
(idx[0] + 2 * idx[1] + 3 * idx[2] + 4 * idx[3]) as f64
});
let av = a.view();
let p1 = av.permute(&[0, 1, 2, 3]).unwrap();
let p2 = av.permute(&[1, 2, 3, 0]).unwrap();
let p3 = av.permute(&[2, 3, 0, 1]).unwrap();
let p4 = av.permute(&[3, 0, 1, 2]).unwrap();
let mut out = StridedArray::<f64>::row_major(&[size, size, size, size]);
zip_map4_into(&mut out.view_mut(), &p1, &p2, &p3, &p4, |a, b, c, d| {
a + b + c + d
})
.unwrap();
for i in 0..size {
for j in 0..size {
for k in 0..size {
for l in 0..size {
let expected = p1.get(&[i, j, k, l])
+ p2.get(&[i, j, k, l])
+ p3.get(&[i, j, k, l])
+ p4.get(&[i, j, k, l]);
assert_relative_eq!(out.get(&[i, j, k, l]), expected, epsilon = 1e-10);
}
}
}
}
}
#[test]
fn test_col_major_tensor() {
let t = StridedArray::<f64>::from_fn_col_major(&[3, 4], |idx| (idx[0] * 10 + idx[1]) as f64);
assert_eq!(t.strides(), &[1, 3]);
assert_eq!(t.get(&[0, 0]), 0.0);
assert_eq!(t.get(&[1, 0]), 10.0);
assert_eq!(t.get(&[2, 3]), 23.0);
let v = t.view();
let vt = v.transpose_2d().unwrap();
assert_eq!(vt.dims(), &[4, 3]);
assert_eq!(vt.get(&[0, 0]), 0.0);
assert_eq!(vt.get(&[0, 1]), 10.0);
assert_eq!(vt.get(&[3, 2]), 23.0);
}
#[test]
fn test_strided_view_broadcast_and_copy() {
let data = vec![1.0, 2.0, 3.0];
let row = strided_kernel::StridedView::<f64>::new(&data, &[1, 3], &[3, 1], 0).unwrap();
let broad = row.broadcast(&[4, 3]).unwrap();
let mut dest = StridedArray::<f64>::row_major(&[4, 3]);
copy_into(&mut dest.view_mut(), &broad).unwrap();
for i in 0..4 {
assert_relative_eq!(dest.get(&[i, 0]), 1.0, epsilon = 1e-10);
assert_relative_eq!(dest.get(&[i, 1]), 2.0, epsilon = 1e-10);
assert_relative_eq!(dest.get(&[i, 2]), 3.0, epsilon = 1e-10);
}
}
fn make_large(n: usize) -> StridedArray<f64> {
StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] * n + idx[1]) as f64)
}
#[test]
fn test_large_map_into() {
let n = 200;
let a = make_large(n);
let mut out = StridedArray::<f64>::col_major(&[n, n]);
map_into(&mut out.view_mut(), &a.view(), |x| x * 3.0).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
assert_relative_eq!(out.get(&[i, j]), a.get(&[i, j]) * 3.0, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_map_into_permuted() {
let n = 200;
let a = make_large(n);
let a_t = a.view().permute(&[1, 0]).unwrap();
let mut out = StridedArray::<f64>::col_major(&[n, n]);
map_into(&mut out.view_mut(), &a_t, |x| x * 2.0).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
assert_relative_eq!(out.get(&[i, j]), a.get(&[j, i]) * 2.0, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_zip_map2() {
let n = 200;
let a = make_large(n);
let b = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] + idx[1]) as f64);
let mut out = StridedArray::<f64>::col_major(&[n, n]);
zip_map2_into(&mut out.view_mut(), &a.view(), &b.view(), |x, y| x + y).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
assert_relative_eq!(
out.get(&[i, j]),
a.get(&[i, j]) + b.get(&[i, j]),
epsilon = 1e-10
);
}
}
}
#[test]
fn test_large_zip_map4_permuted() {
let n = 14; let a = StridedArray::from_fn_col_major(&[n, n, n, n], |idx| {
(idx[0] + 2 * idx[1] + 3 * idx[2] + 4 * idx[3]) as f64
});
let av = a.view();
let p1 = av.permute(&[0, 1, 2, 3]).unwrap();
let p2 = av.permute(&[1, 2, 3, 0]).unwrap();
let p3 = av.permute(&[2, 3, 0, 1]).unwrap();
let p4 = av.permute(&[3, 0, 1, 2]).unwrap();
let mut out = StridedArray::<f64>::col_major(&[n, n, n, n]);
zip_map4_into(&mut out.view_mut(), &p1, &p2, &p3, &p4, |a, b, c, d| {
a + b + c + d
})
.unwrap();
for i in (0..n).step_by(3) {
for j in (0..n).step_by(3) {
for k in (0..n).step_by(3) {
for l in (0..n).step_by(3) {
let expected = p1.get(&[i, j, k, l])
+ p2.get(&[i, j, k, l])
+ p3.get(&[i, j, k, l])
+ p4.get(&[i, j, k, l]);
assert_relative_eq!(out.get(&[i, j, k, l]), expected, epsilon = 1e-10);
}
}
}
}
}
#[test]
fn test_large_reduce() {
let n = 200;
let a = make_large(n);
let result = sum(&a.view()).unwrap();
let expected: f64 = a.iter().copied().sum();
assert_relative_eq!(result, expected, epsilon = 1e-6);
}
#[test]
fn test_large_add() {
let n = 200;
let b = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] + idx[1]) as f64);
let mut dest = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] * 3 + idx[1]) as f64);
let expected_base = dest.iter().cloned().collect::<Vec<_>>();
add(&mut dest.view_mut(), &b.view()).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
let flat = j * n + i; let expected = expected_base[flat] + b.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_mul() {
let n = 200;
let a = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] * 2 + idx[1] + 1) as f64);
let expected_base = dest.iter().cloned().collect::<Vec<_>>();
mul(&mut dest.view_mut(), &a.view()).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
let flat = j * n + i;
let expected = expected_base[flat] * a.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_axpy() {
let n = 200;
let x = make_large(n);
let mut y = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] + idx[1]) as f64);
let y_orig = y.iter().cloned().collect::<Vec<_>>();
let alpha = 2.5;
axpy(&mut y.view_mut(), &x.view(), alpha).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
let flat = j * n + i;
let expected = alpha * x.get(&[i, j]) + y_orig[flat];
assert_relative_eq!(y.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_fma() {
let n = 200;
let a = make_large(n);
let b = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] * 3 + idx[1]) as f64);
let dest_orig = dest.iter().cloned().collect::<Vec<_>>();
fma(&mut dest.view_mut(), &a.view(), &b.view()).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
let flat = j * n + i;
let expected = dest_orig[flat] + a.get(&[i, j]) * b.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_dot() {
let n = 200;
let a = make_large(n);
let b = StridedArray::from_fn_col_major(&[n, n], |idx| (idx[0] + idx[1] + 1) as f64);
let result = dot(&a.view(), &b.view()).unwrap();
let expected: f64 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
assert_relative_eq!(result, expected, epsilon = 1e-4);
}
#[test]
fn test_large_symmetrize() {
let n = 200;
let a = make_large(n);
let mut out = StridedArray::<f64>::col_major(&[n, n]);
symmetrize_into(&mut out.view_mut(), &a.view()).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
let expected = (a.get(&[i, j]) + a.get(&[j, i])) * 0.5;
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_large_copy_into_permuted() {
let n = 200;
let a = make_large(n);
let a_t = a.view().permute(&[1, 0]).unwrap();
let mut out = StridedArray::<f64>::col_major(&[n, n]);
copy_into(&mut out.view_mut(), &a_t).unwrap();
for i in (0..n).step_by(17) {
for j in (0..n).step_by(19) {
assert_relative_eq!(out.get(&[i, j]), a.get(&[j, i]), epsilon = 1e-10);
}
}
}
#[test]
fn test_mul_permuted() {
let rows = 6;
let cols = 5;
let src =
StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest = StridedArray::<f64>::from_fn_col_major(&[cols, rows], |idx| {
(idx[0] * 2 + idx[1] + 1) as f64
});
let mut dest_orig = vec![0.0; cols * rows];
for i in 0..cols {
for j in 0..rows {
dest_orig[i * rows + j] = dest.get(&[i, j]);
}
}
let src_t = src.view().permute(&[1, 0]).unwrap(); mul(&mut dest.view_mut(), &src_t).unwrap();
for i in 0..cols {
for j in 0..rows {
let expected = dest_orig[i * rows + j] * src.get(&[j, i]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_axpy_permuted() {
let rows = 6;
let cols = 5;
let x = StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| {
(idx[0] * cols + idx[1]) as f64
});
let mut y =
StridedArray::<f64>::from_fn_col_major(&[cols, rows], |idx| (idx[0] + idx[1]) as f64);
let mut y_orig = vec![0.0; cols * rows];
for i in 0..cols {
for j in 0..rows {
y_orig[i * rows + j] = y.get(&[i, j]);
}
}
let alpha = 3.5;
let x_t = x.view().permute(&[1, 0]).unwrap(); axpy(&mut y.view_mut(), &x_t, alpha).unwrap();
for i in 0..cols {
for j in 0..rows {
let expected = alpha * x.get(&[j, i]) + y_orig[i * rows + j];
assert_relative_eq!(y.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_fma_permuted() {
let rows = 6;
let cols = 5;
let a = StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| {
(idx[0] * cols + idx[1]) as f64
});
let b =
StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest =
StridedArray::<f64>::from_fn_col_major(&[cols, rows], |idx| (idx[0] * 3 + idx[1]) as f64);
let mut dest_orig = vec![0.0; cols * rows];
for i in 0..cols {
for j in 0..rows {
dest_orig[i * rows + j] = dest.get(&[i, j]);
}
}
let a_t = a.view().permute(&[1, 0]).unwrap();
let b_t = b.view().permute(&[1, 0]).unwrap();
fma(&mut dest.view_mut(), &a_t, &b_t).unwrap();
for i in 0..cols {
for j in 0..rows {
let expected = dest_orig[i * rows + j] + a.get(&[j, i]) * b.get(&[j, i]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_dot_permuted() {
let rows = 7;
let cols = 5;
let a = StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| {
(idx[0] * cols + idx[1]) as f64
});
let b =
StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| (idx[0] + idx[1] + 1) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap(); let b_t = b.view().permute(&[1, 0]).unwrap();
let result = dot(&a_t, &b_t).unwrap();
let mut expected = 0.0;
for i in 0..cols {
for j in 0..rows {
expected += a.get(&[j, i]) * b.get(&[j, i]);
}
}
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_sum_permuted() {
let rows = 8;
let cols = 6;
let a = StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| {
(idx[0] * cols + idx[1]) as f64
});
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = sum(&a_t).unwrap();
let expected: f64 = a.iter().copied().sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_copy_into_conj_complex() {
let a = StridedArray::<Complex64>::from_fn_row_major(&[3, 4], |idx| {
Complex64::new((idx[0] * 4 + idx[1]) as f64, (idx[0] + idx[1]) as f64)
});
let mut out = StridedArray::<Complex64>::row_major(&[3, 4]);
let a_conj = a.view().conj();
copy_into(&mut out.view_mut(), &a_conj).unwrap();
for i in 0..3 {
for j in 0..4 {
let expected = a.get(&[i, j]).conj();
assert_relative_eq!(out.get(&[i, j]).re, expected.re, epsilon = 1e-10);
assert_relative_eq!(out.get(&[i, j]).im, expected.im, epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_into_contiguous_identity() {
let a = StridedArray::<f64>::from_fn_row_major(&[10, 12], |idx| (idx[0] * 12 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[10, 12]);
copy_into(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..10 {
for j in 0..12 {
assert_relative_eq!(out.get(&[i, j]), a.get(&[i, j]), epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_scale() {
let a = StridedArray::<f64>::from_fn_row_major(&[5, 6], |idx| (idx[0] * 6 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[5, 6]);
let scale = 2.5;
copy_scale(&mut out.view_mut(), &a.view(), scale).unwrap();
for i in 0..5 {
for j in 0..6 {
let expected = scale * a.get(&[i, j]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_conj_complex() {
let a = StridedArray::<Complex64>::from_fn_row_major(&[4, 3], |idx| {
Complex64::new((idx[0] * 3 + idx[1]) as f64, idx[0] as f64 - idx[1] as f64)
});
let mut out = StridedArray::<Complex64>::row_major(&[4, 3]);
copy_conj(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..4 {
for j in 0..3 {
let expected = a.get(&[i, j]).conj();
assert_relative_eq!(out.get(&[i, j]).re, expected.re, epsilon = 1e-10);
assert_relative_eq!(out.get(&[i, j]).im, expected.im, epsilon = 1e-10);
}
}
}
#[test]
fn test_symmetrize_conj_into_complex() {
let n = 4;
let a = StridedArray::<Complex64>::from_fn_row_major(&[n, n], |idx| {
Complex64::new((idx[0] * n + idx[1]) as f64, idx[0] as f64 - idx[1] as f64)
});
let mut out = StridedArray::<Complex64>::row_major(&[n, n]);
symmetrize_conj_into(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..n {
for j in 0..n {
let expected = (a.get(&[i, j]) + a.get(&[j, i]).conj()) * 0.5;
assert_relative_eq!(out.get(&[i, j]).re, expected.re, epsilon = 1e-10);
assert_relative_eq!(out.get(&[i, j]).im, expected.im, epsilon = 1e-10);
}
}
}
#[test]
fn test_symmetrize_into_error_non_2d() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 3, 3], |idx| {
(idx[0] * 9 + idx[1] * 3 + idx[2]) as f64
});
let mut out = StridedArray::<f64>::row_major(&[3, 3, 3]);
let result = symmetrize_into(&mut out.view_mut(), &a.view());
assert!(result.is_err());
match result.unwrap_err() {
StridedError::RankMismatch(ndim, 2) => assert_eq!(ndim, 3),
e => panic!("expected RankMismatch, got: {:?}", e),
}
}
#[test]
fn test_symmetrize_into_error_non_square() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 5], |idx| (idx[0] * 5 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[3, 5]);
let result = symmetrize_into(&mut out.view_mut(), &a.view());
assert!(result.is_err());
match result.unwrap_err() {
StridedError::NonSquare { rows, cols } => {
assert_eq!(rows, 3);
assert_eq!(cols, 5);
}
e => panic!("expected NonSquare, got: {:?}", e),
}
}
#[test]
fn test_copy_transpose_scale_into_error_non_2d() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3, 4], |idx| {
(idx[0] * 12 + idx[1] * 4 + idx[2]) as f64
});
let mut out = StridedArray::<f64>::row_major(&[2, 3, 4]);
let result = copy_transpose_scale_into(&mut out.view_mut(), &a.view(), 2.0);
assert!(result.is_err());
match result.unwrap_err() {
StridedError::RankMismatch(ndim, 2) => assert_eq!(ndim, 3),
e => panic!("expected RankMismatch, got: {:?}", e),
}
}
#[test]
fn test_add_small_contiguous() {
let src = StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] + idx[1]) as f64);
let mut dest =
StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] * 4 + idx[1]) as f64);
let dest_orig: Vec<f64> = dest.iter().copied().collect();
add(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..3 {
for j in 0..4 {
let flat = i * 4 + j;
let expected = dest_orig[flat] + src.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_mul_small_contiguous() {
let src = StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest =
StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] * 4 + idx[1] + 1) as f64);
let dest_orig: Vec<f64> = dest.iter().copied().collect();
mul(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..3 {
for j in 0..4 {
let flat = i * 4 + j;
let expected = dest_orig[flat] * src.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_copy_scale_permuted() {
let a = StridedArray::<f64>::from_fn_row_major(&[5, 7], |idx| (idx[0] * 7 + idx[1]) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap(); let mut out = StridedArray::<f64>::row_major(&[7, 5]);
let scale = -1.5;
copy_scale(&mut out.view_mut(), &a_t, scale).unwrap();
for i in 0..7 {
for j in 0..5 {
let expected = scale * a.get(&[j, i]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_add_mixed_layout() {
let m = 5;
let n = 6;
let src = StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] + idx[1]) as f64);
let mut dest =
StridedArray::<f64>::from_fn_row_major(&[m, n], |idx| (idx[0] * n + idx[1]) as f64);
let dest_orig: Vec<f64> = dest.iter().copied().collect();
add(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..m {
for j in 0..n {
let flat = i * n + j;
let expected = dest_orig[flat] + src.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_mul_mixed_layout() {
let m = 5;
let n = 6;
let src = StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest =
StridedArray::<f64>::from_fn_row_major(&[m, n], |idx| (idx[0] * n + idx[1] + 1) as f64);
let dest_orig: Vec<f64> = dest.iter().copied().collect();
mul(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..m {
for j in 0..n {
let flat = i * n + j;
let expected = dest_orig[flat] * src.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_axpy_mixed_layout() {
let m = 5;
let n = 6;
let x = StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] * n + idx[1]) as f64);
let mut y = StridedArray::<f64>::from_fn_row_major(&[m, n], |idx| (idx[0] + idx[1]) as f64);
let y_orig: Vec<f64> = y.iter().copied().collect();
let alpha = 2.5;
axpy(&mut y.view_mut(), &x.view(), alpha).unwrap();
for i in 0..m {
for j in 0..n {
let flat = i * n + j;
let expected = alpha * x.get(&[i, j]) + y_orig[flat];
assert_relative_eq!(y.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_fma_mixed_layout() {
let m = 5;
let n = 6;
let a = StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] * n + idx[1]) as f64);
let b = StridedArray::<f64>::from_fn_row_major(&[m, n], |idx| (idx[0] + idx[1] + 1) as f64);
let mut dest =
StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] * 3 + idx[1]) as f64);
let dest_orig: Vec<f64> = dest.iter().copied().collect();
fma(&mut dest.view_mut(), &a.view(), &b.view()).unwrap();
for i in 0..m {
for j in 0..n {
let flat = j * m + i; let expected = dest_orig[flat] + a.get(&[i, j]) * b.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_dot_mixed_layout() {
let m = 7;
let n = 5;
let a = StridedArray::<f64>::from_fn_row_major(&[m, n], |idx| (idx[0] * n + idx[1]) as f64);
let b = StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] + idx[1] + 1) as f64);
let result = dot(&a.view(), &b.view()).unwrap();
let mut expected = 0.0;
for i in 0..m {
for j in 0..n {
expected += a.get(&[i, j]) * b.get(&[i, j]);
}
}
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_sum_col_major_small() {
let a = StridedArray::<f64>::from_fn_col_major(&[3, 4], |idx| (idx[0] * 4 + idx[1]) as f64);
let result = sum(&a.view()).unwrap();
let expected: f64 = (0..3)
.flat_map(|i| (0..4).map(move |j| (i * 4 + j) as f64))
.sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_sum_f32_contiguous() {
let a = StridedArray::<f32>::from_fn_row_major(&[100], |idx| (idx[0] + 1) as f32);
let result = sum(&a.view()).unwrap();
let expected: f32 = (1..=100).map(|x| x as f32).sum();
assert_relative_eq!(result, expected, epsilon = 1e-3);
}
#[test]
fn test_dot_f32_contiguous() {
let a = StridedArray::<f32>::from_fn_row_major(&[100], |idx| (idx[0] + 1) as f32);
let b = StridedArray::<f32>::from_fn_row_major(&[100], |idx| (idx[0] * 2 + 1) as f32);
let result = dot(&a.view(), &b.view()).unwrap();
let expected: f32 = (0..100).map(|i| (i + 1) as f32 * (i * 2 + 1) as f32).sum();
assert_relative_eq!(result, expected, epsilon = 1e-1);
}
#[test]
fn test_copy_into_conj_mixed_layout() {
let a = StridedArray::<Complex64>::from_fn_col_major(&[3, 4], |idx| {
Complex64::new((idx[0] * 4 + idx[1]) as f64, (idx[0] + idx[1]) as f64)
});
let mut out = StridedArray::<Complex64>::row_major(&[3, 4]);
let a_conj = a.view().conj();
copy_into(&mut out.view_mut(), &a_conj).unwrap();
for i in 0..3 {
for j in 0..4 {
let expected = a.get(&[i, j]).conj();
assert_relative_eq!(out.get(&[i, j]).re, expected.re, epsilon = 1e-10);
assert_relative_eq!(out.get(&[i, j]).im, expected.im, epsilon = 1e-10);
}
}
}
#[test]
fn test_zip_map3_into_contiguous() {
let a = StridedArray::<f64>::from_fn_row_major(&[4, 5], |idx| (idx[0] * 5 + idx[1]) as f64);
let b = StridedArray::<f64>::from_fn_row_major(&[4, 5], |idx| (idx[0] + idx[1] + 1) as f64);
let c = StridedArray::<f64>::from_fn_row_major(&[4, 5], |idx| (idx[0] * 2 + idx[1] * 3) as f64);
let mut out = StridedArray::<f64>::row_major(&[4, 5]);
zip_map3_into(
&mut out.view_mut(),
&a.view(),
&b.view(),
&c.view(),
|x, y, z| x + y * z,
)
.unwrap();
for i in 0..4 {
for j in 0..5 {
let expected = a.get(&[i, j]) + b.get(&[i, j]) * c.get(&[i, j]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_zip_map3_into_permuted() {
let rows = 6;
let cols = 5;
let a = StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| {
(idx[0] * cols + idx[1]) as f64
});
let b =
StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| (idx[0] + idx[1] + 1) as f64);
let c = StridedArray::<f64>::from_fn_row_major(&[rows, cols], |idx| {
(idx[0] * 2 + idx[1] * 3) as f64
});
let a_t = a.view().permute(&[1, 0]).unwrap(); let b_t = b.view().permute(&[1, 0]).unwrap();
let c_t = c.view().permute(&[1, 0]).unwrap();
let mut out = StridedArray::<f64>::row_major(&[cols, rows]);
zip_map3_into(&mut out.view_mut(), &a_t, &b_t, &c_t, |x, y, z| x * y + z).unwrap();
for i in 0..cols {
for j in 0..rows {
let expected = a.get(&[j, i]) * b.get(&[j, i]) + c.get(&[j, i]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_zip_map3_into_mixed_strides() {
let a = StridedArray::<f64>::from_fn_row_major(&[5, 4], |idx| (idx[0] * 4 + idx[1]) as f64);
let b = StridedArray::<f64>::from_fn_col_major(&[5, 4], |idx| (idx[0] + idx[1]) as f64);
let c = StridedArray::<f64>::from_fn_row_major(&[4, 5], |idx| (idx[0] * 5 + idx[1] + 1) as f64);
let c_t = c.view().permute(&[1, 0]).unwrap();
let mut out = StridedArray::<f64>::col_major(&[5, 4]);
zip_map3_into(
&mut out.view_mut(),
&a.view(),
&b.view(),
&c_t,
|x, y, z| x + y + z,
)
.unwrap();
for i in 0..5 {
for j in 0..4 {
let expected = a.get(&[i, j]) + b.get(&[i, j]) + c.get(&[j, i]);
assert_relative_eq!(out.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_reduce_mixed_layout() {
let m = 8;
let n = 6;
let a = StridedArray::<f64>::from_fn_col_major(&[m, n], |idx| (idx[0] * n + idx[1] + 1) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = reduce(&a_t, |x| x, |a, b| a + b, 0.0).unwrap();
let expected: f64 = (0..m)
.flat_map(|i| (0..n).map(move |j| (i * n + j + 1) as f64))
.sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_product_permuted() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| {
1.0 + 0.01 * (idx[0] * 4 + idx[1]) as f64
});
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = reduce(&a_t, |x| x, |a, b| a * b, 1.0).unwrap();
let expected: f64 = a.iter().copied().product();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_axis_invalid_axis() {
let a =
StridedArray::<f64>::from_fn_row_major(&[3, 4, 2], |idx| (idx[0] + idx[1] + idx[2]) as f64);
let result = reduce_axis(&a.view(), 3, |x| x, |a, b| a + b, 0.0);
assert!(result.is_err());
match result.unwrap_err() {
StridedError::InvalidAxis { axis, rank } => {
assert_eq!(axis, 3);
assert_eq!(rank, 3);
}
e => panic!("expected InvalidAxis, got: {:?}", e),
}
}
#[test]
fn test_reduce_axis_1d_to_scalar() {
let a = StridedArray::<f64>::from_fn_row_major(&[5], |idx| (idx[0] + 1) as f64);
let result = reduce_axis(&a.view(), 0, |x| x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[1]);
assert_relative_eq!(result.get(&[0]), 15.0, epsilon = 1e-10);
}
#[test]
fn test_reduce_axis_permuted() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] * 4 + idx[1] + 1) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = reduce_axis(&a_t, 0, |x| x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[3]);
for j in 0..3 {
let mut expected = 0.0;
for i in 0..4 {
expected += a.get(&[j, i]); }
assert_relative_eq!(result.get(&[j]), expected, epsilon = 1e-10);
}
}
#[test]
fn test_small_contiguous_map_into() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[2, 3]);
map_into(&mut out.view_mut(), &a.view(), |x| x * 5.0).unwrap();
for i in 0..2 {
for j in 0..3 {
assert_relative_eq!(out.get(&[i, j]), a.get(&[i, j]) * 5.0, epsilon = 1e-10);
}
}
}
#[test]
fn test_small_contiguous_sum() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1]) as f64);
let result = sum(&a.view()).unwrap();
assert_relative_eq!(result, 15.0, epsilon = 1e-10);
}
#[test]
fn test_small_contiguous_dot() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
let b = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] + idx[1]) as f64);
let result = dot(&a.view(), &b.view()).unwrap();
let expected: f64 = a.iter().zip(b.iter()).map(|(x, y)| x * y).sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_small_contiguous_add() {
let src = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] + idx[1]) as f64);
let mut dest =
StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1]) as f64);
let dest_orig: Vec<f64> = dest.iter().copied().collect();
add(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..2 {
for j in 0..3 {
let flat = i * 3 + j;
let expected = dest_orig[flat] + src.get(&[i, j]);
assert_relative_eq!(dest.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_small_contiguous_copy_into() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 2], |idx| (idx[0] * 10 + idx[1]) as f64);
let mut out = StridedArray::<f64>::row_major(&[2, 2]);
copy_into(&mut out.view_mut(), &a.view()).unwrap();
for i in 0..2 {
for j in 0..2 {
assert_relative_eq!(out.get(&[i, j]), a.get(&[i, j]), epsilon = 1e-10);
}
}
}
#[test]
fn test_reduce_3d_permuted_blocked() {
let a = StridedArray::<f64>::from_fn_row_major(&[4, 5, 6], |idx| {
(idx[0] * 30 + idx[1] * 6 + idx[2] + 1) as f64
});
let a_perm = a.view().permute(&[2, 0, 1]).unwrap();
let result = reduce(&a_perm, |x| x, |a, b| a + b, 0.0).unwrap();
let expected: f64 = a.iter().copied().sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_sum_of_squares_permuted() {
let a = StridedArray::<f64>::from_fn_row_major(&[5, 7], |idx| (idx[0] * 7 + idx[1] + 1) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = reduce(&a_t, |x| x * x, |a, b| a + b, 0.0).unwrap();
let expected: f64 = a.iter().copied().map(|x| x * x).sum();
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_max_permuted() {
let a = StridedArray::<f64>::from_fn_row_major(&[6, 8], |idx| (idx[0] * 8 + idx[1]) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = reduce(
&a_t,
|x| x,
|a, b| if a > b { a } else { b },
f64::NEG_INFINITY,
)
.unwrap();
let expected: f64 = a.iter().copied().fold(f64::NEG_INFINITY, f64::max);
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_col_major_3d() {
let a = StridedArray::<f64>::from_fn_col_major(&[3, 5, 4], |idx| {
(idx[0] * 20 + idx[1] * 4 + idx[2] + 1) as f64
});
let a_perm = a.view().permute(&[1, 2, 0]).unwrap();
let result = reduce(&a_perm, |x| x, |a, b| a + b, 0.0).unwrap();
let mut expected = 0.0;
for i in 0..3 {
for j in 0..5 {
for k in 0..4 {
expected += a.get(&[i, j, k]);
}
}
}
assert_relative_eq!(result, expected, epsilon = 1e-10);
}
#[test]
fn test_reduce_axis_3d_permuted_blocked() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 4, 5], |idx| {
(idx[0] * 20 + idx[1] * 5 + idx[2] + 1) as f64
});
let a_perm = a.view().permute(&[2, 0, 1]).unwrap();
let result = reduce_axis(&a_perm, 1, |x| x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[5, 4]);
for i in 0..5 {
for j in 0..4 {
let mut expected = 0.0;
for k in 0..3 {
expected += a.get(&[k, j, i]);
}
assert_relative_eq!(result.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_reduce_axis_col_major_mixed() {
let a = StridedArray::<f64>::from_fn_col_major(&[4, 3, 2], |idx| {
(idx[0] * 6 + idx[1] * 2 + idx[2] + 1) as f64
});
let result = reduce_axis(&a.view(), 2, |x| x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[4, 3]);
for i in 0..4 {
for j in 0..3 {
let mut expected = 0.0;
for k in 0..2 {
expected += a.get(&[i, j, k]);
}
assert_relative_eq!(result.get(&[i, j]), expected, epsilon = 1e-10);
}
}
}
#[test]
fn test_reduce_axis_map_fn_permuted() {
let a = StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] * 4 + idx[1] + 1) as f64);
let a_t = a.view().permute(&[1, 0]).unwrap();
let result = reduce_axis(&a_t, 1, |x| x * x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[4]);
for i in 0..4 {
let mut expected = 0.0;
for j in 0..3 {
let val = a.get(&[j, i]); expected += val * val;
}
assert_relative_eq!(result.get(&[i]), expected, epsilon = 1e-10);
}
}
#[test]
fn test_reduce_axis_4d_blocked() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3, 4, 5], |idx| {
(idx[0] * 60 + idx[1] * 20 + idx[2] * 5 + idx[3] + 1) as f64
});
let a_perm = a.view().permute(&[3, 1, 0, 2]).unwrap();
let result = reduce_axis(&a_perm, 2, |x| x, |a, b| a + b, 0.0).unwrap();
assert_eq!(result.dims(), &[5, 3, 4]);
for i in 0..5 {
for j in 0..3 {
for k in 0..4 {
let mut expected = 0.0;
for l in 0..2 {
expected += a.get(&[l, j, k, i]);
}
assert_relative_eq!(result.get(&[i, j, k]), expected, epsilon = 1e-10);
}
}
}
}
fn c(re: f64, im: f64) -> Complex64 {
Complex64::new(re, im)
}
#[test]
fn test_map_into_mixed_f64_to_c64() {
let src = StridedArray::<f64>::from_fn_row_major(&[3, 4], |idx| (idx[0] * 4 + idx[1]) as f64);
let mut dest = StridedArray::<Complex64>::row_major(&[3, 4]);
map_into(&mut dest.view_mut(), &src.view(), |x| {
Complex64::new(x, x * 2.0)
})
.unwrap();
for i in 0..3 {
for j in 0..4 {
let v = (i * 4 + j) as f64;
assert_eq!(dest.get(&[i, j]), c(v, v * 2.0));
}
}
}
#[test]
fn test_zip_map2_into_mixed_f64_c64() {
let a = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
let b = StridedArray::<Complex64>::from_fn_row_major(&[2, 3], |idx| {
c((idx[0] + 1) as f64, (idx[1] + 1) as f64)
});
let mut dest = StridedArray::<Complex64>::row_major(&[2, 3]);
zip_map2_into(&mut dest.view_mut(), &a.view(), &b.view(), |x, y| {
Complex64::new(x, 0.0) + y
})
.unwrap();
for i in 0..2 {
for j in 0..3 {
let x = (i * 3 + j + 1) as f64;
let y = c((i + 1) as f64, (j + 1) as f64);
assert_eq!(dest.get(&[i, j]), Complex64::new(x, 0.0) + y);
}
}
}
#[test]
fn test_add_mixed_c64_plus_f64() {
let mut dest = StridedArray::<Complex64>::from_fn_row_major(&[3, 3], |idx| {
c(idx[0] as f64, idx[1] as f64)
});
let src =
StridedArray::<f64>::from_fn_row_major(&[3, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
let orig: Vec<Complex64> = dest.iter().copied().collect();
add(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..3 {
for j in 0..3 {
let flat = i * 3 + j;
let expected = orig[flat] + (flat + 1) as f64;
assert_eq!(dest.get(&[i, j]), expected);
}
}
}
#[test]
fn test_mul_mixed_c64_times_f64() {
let mut dest = StridedArray::<Complex64>::from_fn_row_major(&[2, 3], |idx| {
c(idx[0] as f64, idx[1] as f64)
});
let src = StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] + idx[1] + 1) as f64);
let orig: Vec<Complex64> = dest.iter().copied().collect();
mul(&mut dest.view_mut(), &src.view()).unwrap();
for i in 0..2 {
for j in 0..3 {
let flat = i * 3 + j;
let scale = (i + j + 1) as f64;
assert_eq!(dest.get(&[i, j]), orig[flat] * scale);
}
}
}
#[test]
fn test_axpy_mixed() {
let mut dest = StridedArray::<Complex64>::from_fn_row_major(&[2, 4], |idx| {
c((idx[0] * 4 + idx[1]) as f64, 1.0)
});
let src = StridedArray::<f64>::from_fn_row_major(&[2, 4], |idx| (idx[0] + idx[1] + 1) as f64);
let alpha = 2.5_f64;
let orig: Vec<Complex64> = dest.iter().copied().collect();
let alpha_c = Complex64::new(alpha, 0.0);
axpy(&mut dest.view_mut(), &src.view(), alpha_c).unwrap();
for i in 0..2 {
for j in 0..4 {
let flat = i * 4 + j;
let s = (i + j + 1) as f64;
let expected = orig[flat] + alpha_c * s;
assert_eq!(dest.get(&[i, j]), expected);
}
}
}
#[test]
fn test_fma_mixed_f64_c64() {
let mut dest = StridedArray::<Complex64>::from_fn_row_major(&[3, 2], |_| c(0.0, 0.0));
let a = StridedArray::<f64>::from_fn_row_major(&[3, 2], |idx| (idx[0] * 2 + idx[1] + 1) as f64);
let b = StridedArray::<Complex64>::from_fn_row_major(&[3, 2], |idx| {
c(idx[0] as f64, idx[1] as f64)
});
fma(&mut dest.view_mut(), &a.view(), &b.view()).unwrap();
for i in 0..3 {
for j in 0..2 {
let av = (i * 2 + j + 1) as f64;
let bv = c(i as f64, j as f64);
assert_eq!(dest.get(&[i, j]), av * bv);
}
}
}
#[test]
fn test_dot_mixed_f64_c64() {
let a = StridedArray::<f64>::from_fn_row_major(&[4], |idx| (idx[0] + 1) as f64);
let b = StridedArray::<Complex64>::from_fn_row_major(&[4], |idx| c(idx[0] as f64, 1.0));
let result: Complex64 = dot(&a.view(), &b.view()).unwrap();
assert_eq!(result, c(20.0, 10.0));
}
#[test]
fn test_dot_same_type_simd_regression() {
let n = 1000;
let a = StridedArray::<f64>::from_fn_row_major(&[n], |idx| (idx[0] + 1) as f64);
let b = StridedArray::<f64>::from_fn_row_major(&[n], |idx| (idx[0] + 1) as f64);
let result: f64 = dot(&a.view(), &b.view()).unwrap();
let expected = (n * (n + 1) * (2 * n + 1)) as f64 / 6.0;
assert_relative_eq!(result, expected, epsilon = 1e-6);
}
#[test]
fn test_copy_scale_mixed() {
let src =
StridedArray::<f64>::from_fn_row_major(&[2, 3], |idx| (idx[0] * 3 + idx[1] + 1) as f64);
let mut dest = StridedArray::<Complex64>::row_major(&[2, 3]);
let scale = c(0.0, 1.0);
copy_scale(&mut dest.view_mut(), &src.view(), scale).unwrap();
for i in 0..2 {
for j in 0..3 {
let v = (i * 3 + j + 1) as f64;
assert_eq!(dest.get(&[i, j]), c(0.0, v));
}
}
}
#[test]
fn test_custom_type_map_into_without_element_op_apply() {
#[derive(Debug, Clone, Copy, PartialEq, Default)]
struct Wrapper(f64);
let src_data = vec![Wrapper(1.0), Wrapper(2.0), Wrapper(3.0), Wrapper(4.0)];
let src = StridedArray::from_parts(src_data, &[2, 2], &[2, 1], 0).unwrap();
let mut dest = StridedArray::<Wrapper>::col_major(&[2, 2]);
map_into(&mut dest.view_mut(), &src.view(), |Wrapper(x)| {
Wrapper(x * 2.0)
})
.unwrap();
assert_eq!(dest.get(&[0, 0]), Wrapper(2.0));
assert_eq!(dest.get(&[0, 1]), Wrapper(4.0));
assert_eq!(dest.get(&[1, 0]), Wrapper(6.0));
assert_eq!(dest.get(&[1, 1]), Wrapper(8.0));
}