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 src_starts = vec![0u64; counts.len()];
for_each_dual_run(
dims,
starts,
counts,
&src_starts,
counts,
element_size,
|dst_off, src_off, len| f(dst_off, src_off as usize, len),
)
}
pub(crate) fn for_each_dual_run(
dst_dims: &[u64],
dst_starts: &[u64],
src_dims: &[u64],
src_starts: &[u64],
counts: &[u64],
element_size: u64,
mut f: impl FnMut(u64, u64, usize) -> IoResult<()>,
) -> IoResult<()> {
let ndims = counts.len();
debug_assert_eq!(dst_dims.len(), ndims);
debug_assert_eq!(dst_starts.len(), ndims);
debug_assert_eq!(src_dims.len(), ndims);
debug_assert_eq!(src_starts.len(), ndims);
if ndims == 0 || counts.contains(&0) {
return Ok(());
}
let dst_strides = compute_strides(dst_dims, element_size);
let src_strides = compute_strides(src_dims, element_size);
let mut m = ndims - 1;
while m > 0 && counts[m] == dst_dims[m] && counts[m] == src_dims[m] {
m -= 1;
}
let run_elems: u64 = counts[m..].iter().product();
let run_bytes = (run_elems * element_size) as usize;
let dst_base: u64 = (m..ndims).map(|d| dst_starts[d] * dst_strides[d]).sum();
let src_base: u64 = (m..ndims).map(|d| src_starts[d] * src_strides[d]).sum();
let n_outer: u64 = counts[..m].iter().product(); let mut coords = vec![0u64; m];
for _ in 0..n_outer {
let mut dst_off = dst_base;
let mut src_off = src_base;
for d in 0..m {
dst_off += (dst_starts[d] + coords[d]) * dst_strides[d];
src_off += (src_starts[d] + coords[d]) * src_strides[d];
}
f(dst_off, src_off, run_bytes)?;
for d in (0..m).rev() {
coords[d] += 1;
if coords[d] < counts[d] {
break;
}
coords[d] = 0;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
fn dual(
dst_dims: &[u64],
dst_starts: &[u64],
src_dims: &[u64],
src_starts: &[u64],
counts: &[u64],
es: u64,
) -> Vec<(u64, u64, usize)> {
let mut runs = Vec::new();
for_each_dual_run(
dst_dims,
dst_starts,
src_dims,
src_starts,
counts,
es,
|d, s, l| {
runs.push((d, s, l));
Ok(())
},
)
.unwrap();
runs
}
#[test]
fn dual_run_coalesces_when_trailing_dim_is_full_on_both_sides() {
let runs = dual(&[4, 6], &[2, 0], &[2, 6], &[0, 0], &[2, 6], 4);
assert_eq!(runs, vec![(48, 0, 48)]);
}
#[test]
fn dual_run_splits_when_trailing_dim_is_partial_on_the_source() {
let runs = dual(&[4, 6], &[1, 0], &[4, 12], &[1, 3], &[2, 6], 4);
assert_eq!(runs, vec![(24, 60, 24), (48, 108, 24)]);
}
#[test]
fn dual_run_splits_when_trailing_dim_is_partial_on_the_destination() {
let runs = dual(&[4, 12], &[1, 3], &[4, 6], &[1, 0], &[2, 6], 4);
assert_eq!(runs, vec![(60, 24, 24), (108, 48, 24)]);
}
#[test]
fn dual_run_walks_three_dimensions() {
let runs = dual(
&[2, 4, 3],
&[0, 1, 0],
&[2, 2, 3],
&[0, 0, 0],
&[2, 2, 3],
2,
);
assert_eq!(runs, vec![(6, 0, 12), (30, 12, 12)]);
}
#[test]
fn empty_selection_visits_no_runs() {
assert!(dual(&[4, 6], &[0, 0], &[4, 6], &[0, 0], &[0, 6], 4).is_empty());
assert!(dual(&[4, 6], &[0, 0], &[4, 6], &[0, 0], &[2, 0], 4).is_empty());
}
#[test]
fn contiguous_run_is_the_dual_walk_against_the_selection_itself() {
for (dims, starts, counts) in [
(vec![4u64, 6], vec![1u64, 2], vec![2u64, 3]),
(vec![4, 6], vec![2, 0], vec![2, 6]),
(vec![5, 3, 2], vec![1, 0, 0], vec![3, 3, 2]),
(vec![7], vec![2], vec![4]),
] {
let mut got = Vec::new();
for_each_contiguous_run(&dims, &starts, &counts, 4, |dst, src, len| {
got.push((dst, src as u64, len));
Ok(())
})
.unwrap();
let zeros = vec![0u64; counts.len()];
assert_eq!(got, dual(&dims, &starts, &counts, &zeros, &counts, 4));
let total: usize = got.iter().map(|r| r.2).sum();
assert_eq!(total as u64, counts.iter().product::<u64>() * 4);
}
}
}