use crate::io::IoResult;
pub(crate) fn compute_strides(dims: &[u64], element_size: u64) -> Vec<u64> {
let ndims = dims.len();
if ndims == 0 {
return vec![];
}
let mut strides = vec![0u64; ndims];
strides[ndims - 1] = element_size;
for d in (0..ndims - 1).rev() {
strides[d] = strides[d + 1] * dims[d + 1];
}
strides
}
pub(crate) fn for_each_contiguous_run(
dims: &[u64],
starts: &[u64],
counts: &[u64],
element_size: u64,
mut f: impl FnMut(u64, usize, usize) -> IoResult<()>,
) -> IoResult<()> {
let ndims = dims.len();
debug_assert_eq!(starts.len(), ndims);
debug_assert_eq!(counts.len(), ndims);
if ndims == 0 {
return Ok(());
}
let strides = compute_strides(dims, element_size);
let mut m = ndims - 1;
while m > 0 && counts[m] == dims[m] {
m -= 1;
}
let run_elems: u64 = counts[m..].iter().product();
let run_bytes = (run_elems * element_size) as usize;
let inner_base: u64 = (m..ndims).map(|d| starts[d] * strides[d]).sum();
let n_outer: u64 = counts[..m].iter().product(); let mut coords = vec![0u64; m];
for outer in 0..n_outer {
let mut src_off = inner_base;
for d in 0..m {
src_off += (starts[d] + coords[d]) * strides[d];
}
f(src_off, outer as usize * run_bytes, run_bytes)?;
for d in (0..m).rev() {
coords[d] += 1;
if coords[d] < counts[d] {
break;
}
coords[d] = 0;
}
}
Ok(())
}