pub fn fuse_dims(dims: &[usize], all_strides: &[&[isize]]) -> Vec<usize> {
let n = dims.len();
if n <= 1 || all_strides.is_empty() {
return dims.to_vec();
}
let mut result = dims.to_vec();
for i in (1..n).rev() {
let mut can_merge = true;
for strides in all_strides {
if strides[i - 1] == 0 && strides[i] == 0 {
continue;
}
let expected = result[i - 1] as isize * strides[i - 1];
if strides[i] != expected {
can_merge = false;
break;
}
}
if can_merge {
result[i - 1] *= result[i];
result[i] = 1;
}
}
result
}
pub fn compress_dims(dims: &[usize], all_strides: &[Vec<isize>]) -> (Vec<usize>, Vec<Vec<isize>>) {
let kept: Vec<usize> = (0..dims.len()).filter(|&i| dims[i] != 1).collect();
if kept.is_empty() {
if dims.is_empty() {
return (vec![], all_strides.to_vec());
}
let new_strides = all_strides.iter().map(|s| vec![s[0]]).collect();
return (vec![1], new_strides);
}
let new_dims: Vec<usize> = kept.iter().map(|&i| dims[i]).collect();
let new_strides: Vec<Vec<isize>> = all_strides
.iter()
.map(|s| kept.iter().map(|&i| s[i]).collect())
.collect();
(new_dims, new_strides)
}
pub fn compute_importance(
dims: &[usize],
all_strides: &[&[isize]],
index_orders: &[Vec<usize>],
) -> Vec<u64> {
let n = dims.len();
let m = all_strides.len();
if n == 0 || m == 0 {
return vec![];
}
let g = (64 - (m as u64 + 1).leading_zeros()) as u64;
let mut importance = vec![0u64; n];
let output_weight = 1u64 << (g + 1);
for i in 0..n {
if all_strides[0][i] != 0 {
let shift = g * (n - index_orders[0][i]) as u64;
importance[i] = output_weight * (1u64 << shift);
}
}
#[allow(clippy::needless_range_loop)]
for k in 1..m {
for i in 0..n {
if all_strides[k][i] != 0 {
let shift = g * (n - index_orders[k][i]) as u64;
importance[i] += 1u64 << shift;
}
}
}
for i in 0..n {
if dims[i] <= 1 {
importance[i] = 0;
}
}
importance
}
pub fn sort_by_importance(importance: &[u64]) -> Vec<usize> {
let mut indices: Vec<usize> = (0..importance.len()).collect();
indices.sort_by(|&a, &b| importance[b].cmp(&importance[a]));
indices
}
pub fn compute_costs<S: AsRef<[isize]>>(all_strides: &[S]) -> Vec<isize> {
if all_strides.is_empty() {
return vec![];
}
let n = all_strides[0].as_ref().len();
let mut costs = vec![isize::MAX; n];
for strides in all_strides {
let strides = strides.as_ref();
for i in 0..n {
costs[i] = costs[i].min(strides[i].abs());
}
}
for cost in &mut costs {
if *cost == 0 {
*cost = 1;
} else {
*cost *= 2;
}
}
costs
}
#[cfg(test)]
mod tests {
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());
}
}