use crate::fuse::{compress_dims, fuse_dims};
use crate::{block, order, Result};
pub(crate) const SMALL_TENSOR_THRESHOLD: usize = 1024;
pub(crate) struct KernelPlan {
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) order: Vec<usize>, pub(crate) block: Vec<usize>,
}
#[cfg(test)]
pub(crate) fn build_plan(
dims: &[usize],
strides_list: &[&[isize]],
dest_index: Option<usize>,
elem_size: usize,
) -> KernelPlan {
let order = order::compute_order(dims, strides_list, dest_index);
let block = block::compute_block_sizes(dims, &order, strides_list, elem_size);
KernelPlan { order, block }
}
pub(crate) fn build_plan_fused(
dims: &[usize],
strides_list: &[&[isize]],
dest_index: Option<usize>,
elem_size: usize,
) -> (Vec<usize>, Vec<Vec<isize>>, KernelPlan) {
let order = order::compute_order(dims, strides_list, dest_index);
let ordered_dims: Vec<usize> = order.iter().map(|&d| dims[d]).collect();
let ordered_strides: Vec<Vec<isize>> = strides_list
.iter()
.map(|strides| order.iter().map(|&d| strides[d]).collect())
.collect();
let ordered_strides_refs: Vec<&[isize]> =
ordered_strides.iter().map(|s| s.as_slice()).collect();
let fused_dims = fuse_dims(&ordered_dims, &ordered_strides_refs);
let (compressed_dims, compressed_strides) = compress_dims(&fused_dims, &ordered_strides);
let compressed_strides_refs: Vec<&[isize]> =
compressed_strides.iter().map(|s| s.as_slice()).collect();
let identity: Vec<usize> = (0..compressed_dims.len()).collect();
let block = block::compute_block_sizes(
&compressed_dims,
&identity,
&compressed_strides_refs,
elem_size,
);
(
compressed_dims,
compressed_strides,
KernelPlan {
order: identity,
block,
},
)
}
pub(crate) fn build_plan_fused_small(
dims: &[usize],
strides_list: &[&[isize]],
) -> (Vec<usize>, Vec<Vec<isize>>, KernelPlan) {
let strides_owned: Vec<Vec<isize>> = strides_list.iter().map(|s| s.to_vec()).collect();
let fused = fuse_dims(dims, strides_list);
let (fused_dims, fused_strides) = compress_dims(&fused, &strides_owned);
let block = fused_dims.clone();
let identity: Vec<usize> = (0..fused_dims.len()).collect();
(
fused_dims,
fused_strides,
KernelPlan {
order: identity,
block,
},
)
}
#[cfg(test)]
#[inline]
pub(crate) fn for_each_inner_block<F>(
dims: &[usize],
plan: &KernelPlan,
strides_list: &[&[isize]],
mut f: F,
) -> Result<()>
where
F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
{
let rank = dims.len();
if rank == 0 {
let offsets = vec![0isize; strides_list.len()];
return f(&offsets, 1, &[]);
}
let ordered_dims: Vec<usize> = plan.order.iter().map(|&d| dims[d]).collect();
let ordered_blocks: Vec<usize> = plan.block.clone();
let num_arrays = strides_list.len();
let mut ordered_strides: Vec<Vec<isize>> = Vec::with_capacity(num_arrays);
for strides in strides_list {
let s: Vec<isize> = plan.order.iter().map(|&d| strides[d]).collect();
ordered_strides.push(s);
}
let mut offsets = vec![0isize; num_arrays];
match rank {
1 => kernel_1d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
2 => kernel_2d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
3 => kernel_3d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
4 => kernel_4d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
5 => kernel_5d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
6 => kernel_6d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
7 => kernel_7d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
8 => kernel_8d_inner(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
_ => kernel_nd_inner_iterative(
&ordered_dims,
&ordered_blocks,
&ordered_strides,
&mut offsets,
&mut f,
),
}
}
macro_rules! elem_loops {
($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident; $lv:literal) => {
for _ in 0..$blens[$lv] {
$f($offsets, $blens[0], &$is)?;
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o += s[$lv];
}
}
};
($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident;
$lv:literal, $next:literal $(, $rest:literal)*) => {
for _ in 0..$blens[$lv] {
elem_loops!($offsets, $strides, $f, $blens, $is; $next $(, $rest)*);
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= ($blens[$next] as isize) * s[$next];
*o += s[$lv];
}
}
};
}
macro_rules! block_loop {
($dims:ident, $blocks:ident, $strides:ident, $offsets:ident, $f:ident,
$blens:ident, $is:ident; elem=[$($el:literal),+]; $lv0:literal; top=$top:literal) => {{
let mut _j = 0usize;
while _j < $dims[$lv0] {
$blens[$lv0] = $blocks[$lv0].max(1).min($dims[$lv0]).min($dims[$lv0] - _j);
elem_loops!($offsets, $strides, $f, $blens, $is; $($el),+);
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= ($blens[$top] as isize) * s[$top];
*o += ($blens[$lv0] as isize) * s[$lv0];
}
_j += $blens[$lv0];
}
}};
($dims:ident, $blocks:ident, $strides:ident, $offsets:ident, $f:ident,
$blens:ident, $is:ident; elem=[$($el:literal),+];
$lv:literal, $next:literal $(, $rest:literal)*; top=$top:literal) => {{
let mut _j = 0usize;
while _j < $dims[$lv] {
$blens[$lv] = $blocks[$lv].max(1).min($dims[$lv]).min($dims[$lv] - _j);
block_loop!($dims, $blocks, $strides, $offsets, $f, $blens, $is;
elem=[$($el),+]; $next $(, $rest)*; top=$top);
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= ($dims[$next] as isize) * s[$next];
*o += ($blens[$lv] as isize) * s[$lv];
}
_j += $blens[$lv];
}
}};
}
macro_rules! make_kernel {
($name:ident, rank=1) => {
#[inline]
fn $name<F>(
dims: &[usize],
blocks: &[usize],
strides: &[Vec<isize>],
offsets: &mut [isize],
f: &mut F,
) -> Result<()>
where
F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
{
let d0 = dims[0];
let b0 = blocks[0].max(1).min(d0);
let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
let mut j0 = 0usize;
while j0 < d0 {
let blen0 = b0.min(d0 - j0);
f(offsets, blen0, &inner_strides)?;
for (o, s) in offsets.iter_mut().zip(strides.iter()) {
*o += (blen0 as isize) * s[0];
}
j0 += blen0;
}
for (o, s) in offsets.iter_mut().zip(strides.iter()) {
*o -= (d0 as isize) * s[0];
}
Ok(())
}
};
($name:ident, rank=$rank:literal,
block=[$($blk:literal),+], elem=[$($el:literal),+], top=$top:literal) => {
#[inline]
fn $name<F>(
dims: &[usize],
blocks: &[usize],
strides: &[Vec<isize>],
offsets: &mut [isize],
f: &mut F,
) -> Result<()>
where
F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
{
let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
let mut blens = [0usize; $rank];
block_loop!(dims, blocks, strides, offsets, f, blens, inner_strides;
elem=[$($el),+]; $($blk),+; top=$top);
for (o, s) in offsets.iter_mut().zip(strides.iter()) {
*o -= (dims[$top] as isize) * s[$top];
}
Ok(())
}
};
}
make_kernel!(kernel_1d_inner, rank = 1);
make_kernel!(
kernel_2d_inner,
rank = 2,
block = [1, 0],
elem = [1],
top = 1
);
make_kernel!(
kernel_3d_inner,
rank = 3,
block = [2, 1, 0],
elem = [2, 1],
top = 2
);
make_kernel!(
kernel_4d_inner,
rank = 4,
block = [3, 2, 1, 0],
elem = [3, 2, 1],
top = 3
);
make_kernel!(
kernel_5d_inner,
rank = 5,
block = [4, 3, 2, 1, 0],
elem = [4, 3, 2, 1],
top = 4
);
make_kernel!(
kernel_6d_inner,
rank = 6,
block = [5, 4, 3, 2, 1, 0],
elem = [5, 4, 3, 2, 1],
top = 5
);
make_kernel!(
kernel_7d_inner,
rank = 7,
block = [6, 5, 4, 3, 2, 1, 0],
elem = [6, 5, 4, 3, 2, 1],
top = 6
);
make_kernel!(
kernel_8d_inner,
rank = 8,
block = [7, 6, 5, 4, 3, 2, 1, 0],
elem = [7, 6, 5, 4, 3, 2, 1],
top = 7
);
#[inline]
fn kernel_nd_inner_iterative<F>(
dims: &[usize],
blocks: &[usize],
strides: &[Vec<isize>],
offsets: &mut [isize],
f: &mut F,
) -> Result<()>
where
F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
{
let rank = dims.len();
debug_assert!(rank >= 9);
let d0 = dims[0];
let b0 = blocks[0].max(1).min(d0);
let inner_strides: Vec<isize> = strides.iter().map(|s| s[0]).collect();
let mut idx = vec![0usize; rank];
loop {
let mut j0 = 0usize;
while j0 < d0 {
let blen0 = b0.min(d0 - j0);
f(offsets, blen0, &inner_strides)?;
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset += (blen0 as isize) * s[0];
}
j0 += blen0;
}
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset -= (d0 as isize) * s[0];
}
let mut level = 1usize;
loop {
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset += s[level];
}
idx[level] += 1;
if idx[level] < dims[level] {
break;
}
idx[level] = 0;
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset -= (dims[level] as isize) * s[level];
}
level += 1;
if level == rank {
return Ok(());
}
}
}
}
#[inline]
pub(crate) fn for_each_inner_block_preordered<F>(
dims: &[usize],
blocks: &[usize],
strides: &[Vec<isize>],
initial_offsets: &[isize],
mut f: F,
) -> Result<()>
where
F: FnMut(&[isize], usize, &[isize]) -> Result<()>,
{
let rank = dims.len();
if rank == 0 {
return f(initial_offsets, 1, &[]);
}
let mut offsets = initial_offsets.to_vec();
match rank {
1 => kernel_1d_inner(dims, blocks, strides, &mut offsets, &mut f),
2 => kernel_2d_inner(dims, blocks, strides, &mut offsets, &mut f),
3 => kernel_3d_inner(dims, blocks, strides, &mut offsets, &mut f),
4 => kernel_4d_inner(dims, blocks, strides, &mut offsets, &mut f),
5 => kernel_5d_inner(dims, blocks, strides, &mut offsets, &mut f),
6 => kernel_6d_inner(dims, blocks, strides, &mut offsets, &mut f),
7 => kernel_7d_inner(dims, blocks, strides, &mut offsets, &mut f),
8 => kernel_8d_inner(dims, blocks, strides, &mut offsets, &mut f),
_ => kernel_nd_inner_iterative(dims, blocks, strides, &mut offsets, &mut f),
}
}
pub(crate) fn ensure_same_shape(a: &[usize], b: &[usize]) -> Result<()> {
if a.len() != b.len() {
return Err(crate::StridedError::RankMismatch(a.len(), b.len()));
}
if a != b {
return Err(crate::StridedError::ShapeMismatch(a.to_vec(), b.to_vec()));
}
Ok(())
}
#[derive(Copy, Clone, Debug, Eq, PartialEq)]
pub(crate) enum ContiguousLayout {
RowMajor,
ColMajor,
}
pub(crate) fn contiguous_layout(dims: &[usize], strides: &[isize]) -> Option<ContiguousLayout> {
if dims.len() != strides.len() {
return None;
}
if dims.is_empty() {
return Some(ContiguousLayout::RowMajor);
}
let mut expected = 1isize;
let mut row_ok = true;
for (&dim, &stride) in dims.iter().rev().zip(strides.iter().rev()) {
if dim <= 1 {
continue;
}
if stride != expected {
row_ok = false;
break;
}
expected = expected.saturating_mul(dim as isize);
}
if row_ok {
return Some(ContiguousLayout::RowMajor);
}
let mut expected = 1isize;
for (&dim, &stride) in dims.iter().zip(strides.iter()) {
if dim <= 1 {
continue;
}
if stride != expected {
return None;
}
expected = expected.saturating_mul(dim as isize);
}
Some(ContiguousLayout::ColMajor)
}
pub(crate) fn total_len(dims: &[usize]) -> usize {
if dims.is_empty() {
return 1;
}
dims.iter().product()
}
#[inline]
pub(crate) fn use_sequential_fast_path(total: usize) -> bool {
#[cfg(feature = "parallel")]
{
total <= crate::threading::MINTHREADLENGTH || rayon::current_num_threads() <= 1
}
#[cfg(not(feature = "parallel"))]
{
let _ = total;
true
}
}
#[inline]
pub(crate) fn same_contiguous_layout(
dims: &[usize],
strides_list: &[&[isize]],
) -> Option<ContiguousLayout> {
let first = contiguous_layout(dims, strides_list.first()?)?;
for strides in &strides_list[1..] {
if contiguous_layout(dims, strides)? != first {
return None;
}
}
Some(first)
}
#[inline]
pub(crate) fn sequential_contiguous_layout(
dims: &[usize],
strides_list: &[&[isize]],
) -> Option<ContiguousLayout> {
if !use_sequential_fast_path(total_len(dims)) {
return None;
}
same_contiguous_layout(dims, strides_list)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_kernel_inner_block() {
let dims = vec![2, 4];
let strides1 = vec![4isize, 1];
let strides2 = vec![4isize, 1];
let strides_list: Vec<&[isize]> = vec![&strides1, &strides2];
let plan = build_plan(&dims, &strides_list, Some(0), 8);
let mut total_elements = 0usize;
for_each_inner_block(&dims, &plan, &strides_list, |_offsets, len, _strides| {
total_elements += len;
Ok(())
})
.unwrap();
assert_eq!(total_elements, 8);
}
#[test]
fn test_contiguous_layout_row_vs_col() {
let dims = [3usize, 4];
let row = [4isize, 1];
let col = [1isize, 3];
assert_eq!(
contiguous_layout(&dims, &row),
Some(ContiguousLayout::RowMajor)
);
assert_eq!(
contiguous_layout(&dims, &col),
Some(ContiguousLayout::ColMajor)
);
assert!(contiguous_layout(&dims, &row).is_some());
assert!(contiguous_layout(&dims, &col).is_some());
}
#[test]
fn test_contiguous_layout_ignores_dim1_axes() {
let dims = [2usize, 1, 3];
let strides = [3isize, 999, 1];
assert_eq!(
contiguous_layout(&dims, &strides),
Some(ContiguousLayout::RowMajor)
);
}
#[test]
fn test_same_contiguous_layout_all_row_major() {
let dims = [3usize, 4];
let s1 = [4isize, 1];
let s2 = [4isize, 1];
assert_eq!(
same_contiguous_layout(&dims, &[&s1, &s2]),
Some(ContiguousLayout::RowMajor)
);
}
#[test]
fn test_same_contiguous_layout_all_col_major() {
let dims = [3usize, 4];
let s1 = [1isize, 3];
let s2 = [1isize, 3];
assert_eq!(
same_contiguous_layout(&dims, &[&s1, &s2]),
Some(ContiguousLayout::ColMajor)
);
}
#[test]
fn test_same_contiguous_layout_mixed_layouts() {
let dims = [3usize, 4];
let row = [4isize, 1];
let col = [1isize, 3];
assert_eq!(same_contiguous_layout(&dims, &[&row, &col]), None);
}
#[test]
fn test_same_contiguous_layout_one_noncontiguous() {
let dims = [3usize, 4];
let row = [4isize, 1];
let bad = [8isize, 2];
assert_eq!(same_contiguous_layout(&dims, &[&row, &bad]), None);
}
#[test]
fn test_same_contiguous_layout_empty_strides_list() {
let dims = [3usize, 4];
let empty: &[&[isize]] = &[];
assert_eq!(same_contiguous_layout(&dims, empty), None);
}
#[test]
fn test_same_contiguous_layout_single_array() {
let dims = [3usize, 4];
let s = [4isize, 1];
assert_eq!(
same_contiguous_layout(&dims, &[&s[..]]),
Some(ContiguousLayout::RowMajor)
);
}
#[test]
fn test_same_contiguous_layout_many_arrays() {
let dims = [2usize, 3];
let s = [3isize, 1];
assert_eq!(
same_contiguous_layout(&dims, &[&s[..], &s[..], &s[..], &s[..], &s[..]]),
Some(ContiguousLayout::RowMajor)
);
}
#[test]
fn test_same_contiguous_layout_empty_dims() {
let dims: [usize; 0] = [];
let s: [isize; 0] = [];
assert_eq!(
same_contiguous_layout(&dims, &[&s[..], &s[..]]),
Some(ContiguousLayout::RowMajor)
);
}
#[test]
fn test_sequential_contiguous_layout_small_array() {
let dims = [3usize, 4];
let s1 = [4isize, 1];
let s2 = [4isize, 1];
assert_eq!(
sequential_contiguous_layout(&dims, &[&s1, &s2]),
Some(ContiguousLayout::RowMajor)
);
}
#[test]
fn test_sequential_contiguous_layout_noncontiguous() {
let dims = [3usize, 4];
let s1 = [4isize, 1];
let s2 = [8isize, 2];
assert_eq!(sequential_contiguous_layout(&dims, &[&s1, &s2]), None);
}
#[test]
fn test_sequential_contiguous_layout_col_major() {
let dims = [3usize, 4];
let col = [1isize, 3];
assert_eq!(
sequential_contiguous_layout(&dims, &[&col]),
Some(ContiguousLayout::ColMajor)
);
}
#[test]
fn test_build_plan_fused_compresses() {
let dims = [2usize, 3];
let strides = [1isize, 2];
let strides_list: Vec<&[isize]> = vec![&strides];
let (fused_dims, fused_strides, plan) = build_plan_fused(&dims, &strides_list, Some(0), 8);
assert_eq!(fused_dims, vec![6]);
assert_eq!(fused_strides.len(), 1);
assert_eq!(fused_strides[0], vec![1]);
assert_eq!(plan.block.len(), 1);
}
#[test]
fn test_build_plan_fused_keeps_broadcast_compact_inner_loop() {
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 strides_list: Vec<&[isize]> = vec![&out, &lhs, &rhs];
let (fused_dims, fused_strides, _) = build_plan_fused(&dims, &strides_list, Some(0), 8);
assert_eq!(fused_dims[0], 256);
assert_eq!(fused_strides[0][0], 1);
assert_eq!(fused_strides[1][0], 1);
assert_eq!(fused_strides[2][0], 0);
}
#[test]
fn test_kernel_nd_iterative_total_elements_match() {
let dims = vec![2usize, 2, 2, 2, 2, 2, 2, 2, 3];
let blocks = vec![2usize, 1, 1, 1, 1, 1, 1, 1, 1];
let mut stride_val = 1isize;
let mut sv = Vec::new();
for &d in &dims {
sv.push(stride_val);
stride_val *= d as isize;
}
let strides = vec![sv.clone(), sv];
let mut offsets = vec![0isize, 0isize];
let mut total = 0usize;
kernel_nd_inner_iterative(
&dims,
&blocks,
&strides,
&mut offsets,
&mut |_off, len, _s| {
total += len;
Ok(())
},
)
.unwrap();
assert_eq!(total, dims.iter().product::<usize>());
assert_eq!(offsets, vec![0isize, 0isize]);
}
fn verify_kernel_total_and_offsets(rank: usize) {
let mut dims = vec![3usize];
for _ in 1..rank {
dims.push(2);
}
let blocks = vec![2usize; rank];
let mut stride_val = 1isize;
let mut sv = Vec::new();
for &d in &dims {
sv.push(stride_val);
stride_val *= d as isize;
}
let strides = vec![sv.clone(), sv];
let mut offsets = vec![0isize, 0isize];
let mut total = 0usize;
let expected: usize = dims.iter().product();
let result = match rank {
1 => kernel_1d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
2 => kernel_2d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
3 => kernel_3d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
4 => kernel_4d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
5 => kernel_5d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
6 => kernel_6d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
7 => kernel_7d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
8 => kernel_8d_inner(&dims, &blocks, &strides, &mut offsets, &mut |_o, l, _s| {
total += l;
Ok(())
}),
_ => panic!("unsupported rank"),
};
result.unwrap();
assert_eq!(total, expected, "rank={rank}: total mismatch");
assert_eq!(offsets, vec![0, 0], "rank={rank}: offsets not reset");
}
#[test]
fn test_macro_kernels_total_elements_1d() {
verify_kernel_total_and_offsets(1);
}
#[test]
fn test_macro_kernels_total_elements_2d() {
verify_kernel_total_and_offsets(2);
}
#[test]
fn test_macro_kernels_total_elements_3d() {
verify_kernel_total_and_offsets(3);
}
#[test]
fn test_macro_kernels_total_elements_4d() {
verify_kernel_total_and_offsets(4);
}
#[test]
fn test_macro_kernels_total_elements_5d() {
verify_kernel_total_and_offsets(5);
}
#[test]
fn test_macro_kernels_total_elements_6d() {
verify_kernel_total_and_offsets(6);
}
#[test]
fn test_macro_kernels_total_elements_7d() {
verify_kernel_total_and_offsets(7);
}
#[test]
fn test_macro_kernels_total_elements_8d() {
verify_kernel_total_and_offsets(8);
}
fn verify_kernel_visits_all_elements(rank: usize) {
assert!(rank >= 2 && rank <= 8);
let mut dims = vec![3usize];
for _ in 1..rank {
dims.push(2);
}
let blocks = vec![2usize; rank];
let mut stride_val = 1isize;
let mut sv = Vec::new();
for &d in &dims {
sv.push(stride_val);
stride_val *= d as isize;
}
let strides = vec![sv];
let mut visited = std::collections::HashSet::new();
let mut offsets = vec![0isize];
let result = match rank {
2 => kernel_2d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
3 => kernel_3d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
4 => kernel_4d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
5 => kernel_5d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
6 => kernel_6d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
7 => kernel_7d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
8 => kernel_8d_inner(&dims, &blocks, &strides, &mut offsets, &mut |o, len, s| {
for i in 0..len {
visited.insert(o[0] + (i as isize) * s[0]);
}
Ok(())
}),
_ => unreachable!(),
};
result.unwrap();
let total: usize = dims.iter().product();
let expected: std::collections::HashSet<isize> = (0..total as isize).collect();
assert_eq!(
visited, expected,
"rank={rank}: not all elements visited exactly once"
);
}
#[test]
fn test_macro_kernel_5d_visits_all_elements() {
verify_kernel_visits_all_elements(5);
}
#[test]
fn test_macro_kernel_6d_visits_all_elements() {
verify_kernel_visits_all_elements(6);
}
#[test]
fn test_macro_kernel_7d_visits_all_elements() {
verify_kernel_visits_all_elements(7);
}
#[test]
fn test_macro_kernel_8d_visits_all_elements() {
verify_kernel_visits_all_elements(8);
}
}