use super::*;
use strided_view::auxiliary::index_order;
#[test]
fn test_fuse_dims_contiguous() {
let dims = [3, 4];
let strides1 = [1isize, 3];
let strides2 = [1isize, 3];
let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let fused = fuse_dims(&dims, &all_strides);
assert_eq!(fused, vec![12, 1]);
}
#[test]
fn test_fuse_dims_allows_broadcast_operand_across_fused_axes() {
let dims = [16usize, 16, 64, 64];
let out = [1isize, 16, 256, 16_384];
let lhs = [1isize, 16, 0, 256];
let rhs = [0isize, 0, 1, 64];
let all_strides: Vec<&[isize]> = vec![&out, &lhs, &rhs];
let fused = fuse_dims(&dims, &all_strides);
assert_eq!(fused, vec![256, 1, 64, 64]);
}
#[test]
fn test_fuse_dims_non_contiguous() {
let dims = [3, 4];
let strides1 = [1isize, 10]; let all_strides: Vec<&[isize]> = vec![&strides1];
let fused = fuse_dims(&dims, &all_strides);
assert_eq!(fused, vec![3, 4]); }
#[test]
fn test_fuse_dims_partial() {
let dims = [2, 3, 4];
let strides = [1isize, 2, 100]; let all_strides: Vec<&[isize]> = vec![&strides];
let fused = fuse_dims(&dims, &all_strides);
assert_eq!(fused, vec![6, 1, 4]); }
#[test]
fn test_fuse_dims_multiple_arrays() {
let dims = [3, 4];
let strides1 = [1isize, 3]; let strides2 = [1isize, 10]; let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let fused = fuse_dims(&dims, &all_strides);
assert_eq!(fused, vec![3, 4]); }
#[test]
fn test_compute_importance_2_arrays() {
let dims = [4usize, 5];
let strides1 = [1isize, 4]; let strides2 = [5isize, 1]; let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let order1 = index_order(&strides1);
let order2 = index_order(&strides2);
let index_orders = vec![order1, order2];
let importance = compute_importance(&dims, &all_strides, &index_orders);
assert!(importance[0] > importance[1]);
}
#[test]
fn test_sort_by_importance() {
let importance = vec![100u64, 50, 200, 10];
let perm = sort_by_importance(&importance);
assert_eq!(perm, vec![2, 0, 1, 3]); }
#[test]
fn test_compute_costs() {
let strides1 = [1isize, 4, 0];
let strides2 = [2isize, 1, 0];
let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let costs = compute_costs(&all_strides);
assert_eq!(costs, vec![2, 2, 1]);
}
#[test]
fn test_compute_importance_with_zero_stride() {
let dims = [4usize, 5];
let strides1 = [0isize, 1]; let all_strides: Vec<&[isize]> = vec![&strides1];
let order1 = index_order(&strides1);
assert_eq!(order1, vec![1, 1]);
let index_orders = vec![order1];
let importance = compute_importance(&dims, &all_strides, &index_orders);
assert_eq!(importance[0], 0);
assert!(importance[1] > 0);
}
#[test]
fn test_compute_importance_size_one_dim() {
let dims = [4usize, 1, 5];
let strides1 = [1isize, 4, 4];
let all_strides: Vec<&[isize]> = vec![&strides1];
let order1 = index_order(&strides1);
let index_orders = vec![order1];
let importance = compute_importance(&dims, &all_strides, &index_orders);
assert_eq!(importance[1], 0);
assert!(importance[0] > 0);
assert!(importance[2] > 0);
}
#[test]
fn test_compute_importance_output_weight() {
let dims = [4usize, 5];
let out_strides = [1isize, 4]; let in_strides = [5isize, 1]; let all_strides: Vec<&[isize]> = vec![&out_strides, &in_strides];
let order_out = index_order(&out_strides); let order_in = index_order(&in_strides); let index_orders = vec![order_out, order_in];
let importance = compute_importance(&dims, &all_strides, &index_orders);
assert!(importance[0] > importance[1]);
}
#[test]
fn test_compute_costs_owned_vecs() {
let strides_list: Vec<Vec<isize>> = vec![vec![1, 0, 3], vec![2, 0, 4]];
let costs = compute_costs(&strides_list);
assert_eq!(costs, vec![2, 1, 6]);
}
#[test]
fn test_compute_costs_with_zero() {
let strides1 = [0isize, 2, -3];
let strides2 = [1isize, 0, 2];
let all_strides: Vec<&[isize]> = vec![&strides1, &strides2];
let costs = compute_costs(&all_strides);
assert_eq!(costs, vec![1, 1, 4]);
}
#[test]
fn test_compress_dims_removes_fused() {
let dims = vec![12usize, 1];
let strides = vec![vec![1isize, 3]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![12]);
assert_eq!(cs, vec![vec![1]]);
}
#[test]
fn test_compress_dims_removes_multiple() {
let dims = vec![6usize, 1, 4];
let strides = vec![vec![1isize, 2, 100]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![6, 4]);
assert_eq!(cs, vec![vec![1, 100]]);
}
#[test]
fn test_compress_dims_no_removal() {
let dims = vec![3usize, 4];
let strides = vec![vec![1isize, 3]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![3, 4]);
assert_eq!(cs, vec![vec![1, 3]]);
}
#[test]
fn test_compress_dims_all_ones() {
let dims = vec![1usize, 1, 1];
let strides = vec![vec![1isize, 1, 1]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![1]);
assert_eq!(cs, vec![vec![1]]);
}
#[test]
fn test_compress_dims_multi_arrays() {
let dims = vec![6usize, 1, 4];
let strides = vec![vec![1isize, 6, 6], vec![4isize, 24, 1]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![6, 4]);
assert_eq!(cs, vec![vec![1, 6], vec![4, 1]]);
}
#[test]
fn test_compress_dims_single_dim() {
let dims = vec![5usize];
let strides = vec![vec![1isize]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![5]);
assert_eq!(cs, vec![vec![1]]);
}
#[test]
fn test_compress_dims_single_dim_one() {
let dims = vec![1usize];
let strides = vec![vec![1isize]];
let (cd, cs) = compress_dims(&dims, &strides);
assert_eq!(cd, vec![1]);
assert_eq!(cs, vec![vec![1]]);
}
#[test]
fn test_compress_dims_empty() {
let dims: Vec<usize> = vec![];
let strides: Vec<Vec<isize>> = vec![];
let (cd, cs) = compress_dims(&dims, &strides);
assert!(cd.is_empty());
assert!(cs.is_empty());
}