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 = isize::try_from(result[i - 1])
.ok()
.and_then(|dim| dim.checked_mul(strides[i - 1]));
if expected != Some(strides[i]) {
can_merge = false;
break;
}
}
if can_merge {
if let Some(merged) = result[i - 1].checked_mul(result[i]) {
result[i - 1] = merged;
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(crate) 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].checked_abs().unwrap_or(isize::MAX));
}
}
for cost in &mut costs {
if *cost == 0 {
*cost = 1;
} else {
*cost = cost.saturating_mul(2);
}
}
costs
}
#[cfg(test)]
#[path = "fuse/tests/tests.rs"]
mod tests;