use crate::fuse::{compress_dims, fuse_dims};
use crate::{block, order, Result, StridedError};
pub const SMALL_TENSOR_THRESHOLD: usize = 1024;
#[derive(Debug)]
pub struct KernelPlan {
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) order: Vec<usize>, pub 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 _i in 0..$blens[$lv] {
$f($offsets, $blens[0], &$is)?;
if _i + 1 < $blens[$lv] {
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o += s[$lv];
}
}
}
let _back = $blens[$lv].saturating_sub(1) as isize;
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= _back * s[$lv];
}
};
($offsets:ident, $strides:ident, $f:ident, $blens:ident, $is:ident;
$lv:literal, $next:literal $(, $rest:literal)*) => {
for _i in 0..$blens[$lv] {
elem_loops!($offsets, $strides, $f, $blens, $is; $next $(, $rest)*);
if _i + 1 < $blens[$lv] {
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o += s[$lv];
}
}
}
let _back = $blens[$lv].saturating_sub(1) as isize;
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= _back * 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;
let mut _advanced = 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),+);
_j += $blens[$lv0];
if _j < $dims[$lv0] {
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o += ($blens[$lv0] as isize) * s[$lv0];
}
_advanced = _j;
}
}
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= (_advanced as isize) * s[$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;
let mut _advanced = 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);
_j += $blens[$lv];
if _j < $dims[$lv] {
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o += ($blens[$lv] as isize) * s[$lv];
}
_advanced = _j;
}
}
for (o, s) in $offsets.iter_mut().zip($strides.iter()) {
*o -= (_advanced as isize) * s[$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;
let mut advanced = 0usize;
while j0 < d0 {
let blen0 = b0.min(d0 - j0);
f(offsets, blen0, &inner_strides)?;
j0 += blen0;
if j0 < d0 {
for (o, s) in offsets.iter_mut().zip(strides.iter()) {
*o += (blen0 as isize) * s[0];
}
advanced = j0;
}
}
for (o, s) in offsets.iter_mut().zip(strides.iter()) {
*o -= (advanced 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);
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];
if dims.contains(&0) {
return Ok(());
}
loop {
let mut j0 = 0usize;
let mut advanced = 0usize;
while j0 < d0 {
let blen0 = b0.min(d0 - j0);
f(offsets, blen0, &inner_strides)?;
j0 += blen0;
if j0 < d0 {
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset += (blen0 as isize) * s[0];
}
advanced = j0;
}
}
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset -= (advanced as isize) * s[0];
}
let mut level = 1usize;
loop {
if idx[level] + 1 < dims[level] {
idx[level] += 1;
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset += s[level];
}
break;
}
let last = (dims[level] - 1) as isize;
idx[level] = 0;
for (offset, s) in offsets.iter_mut().zip(strides.iter()) {
*offset -= last * 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<()>,
{
preordered_dyn(dims, blocks, strides, initial_offsets, &mut f)
}
type BlockFn<'a> = &'a mut dyn FnMut(&[isize], usize, &[isize]) -> Result<()>;
#[inline(never)]
fn preordered_dyn(
dims: &[usize],
blocks: &[usize],
strides: &[Vec<isize>],
initial_offsets: &[isize],
mut f: BlockFn<'_>,
) -> 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 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]) -> Result<usize> {
if dims.contains(&0) {
return Ok(0);
}
dims.iter()
.try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
.ok_or(StridedError::OffsetOverflow)
}
#[inline]
pub(crate) fn use_sequential_fast_path(total: usize) -> bool {
#[cfg(feature = "parallel")]
{
total <= crate::threading::MINTHREADLENGTH || crate::execution_policy::rayon_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]],
) -> Result<Option<ContiguousLayout>> {
if !use_sequential_fast_path(total_len(dims)?) {
return Ok(None);
}
Ok(same_contiguous_layout(dims, strides_list))
}
#[cfg(test)]
#[path = "kernel/tests/tests.rs"]
mod tests;