use core::mem::MaybeUninit;
use crate::{CopyPlan, MaybeSendSync, RawStridedMut, RawStridedRef, Result, StridedError};
#[cfg(feature = "parallel")]
type AxisVec<T> = smallvec::SmallVec<[T; crate::RAW_FUSED_RANK_LIMIT]>;
#[cfg(not(feature = "parallel"))]
type AxisVec<T> = Vec<T>;
#[derive(Clone, Debug)]
pub struct SlicePlan {
operand_dims: AxisVec<usize>,
operand_strides: AxisVec<isize>,
dest_dims: AxisVec<usize>,
dest_strides: AxisVec<isize>,
source_strides: AxisVec<isize>,
source_offset_delta: isize,
copy_plan: CopyPlan,
}
#[derive(Clone, Debug)]
pub struct ReversePlan {
operand_dims: AxisVec<usize>,
operand_strides: AxisVec<isize>,
dest_strides: AxisVec<isize>,
source_strides: AxisVec<isize>,
source_offset_delta: isize,
copy_plan: CopyPlan,
}
#[derive(Clone, Debug)]
pub struct PadPlan {
operand_dims: AxisVec<usize>,
operand_strides: AxisVec<isize>,
dest_dims: AxisVec<usize>,
dest_strides: AxisVec<isize>,
edge_padding_low: AxisVec<i64>,
interior_step: AxisVec<i64>,
operand_total: usize,
dest_total: usize,
contiguous_dest_fill: bool,
contiguous_axis0_run: Option<ContiguousPadAxis0Run>,
generic_fill: PadFillCursor,
generic_copy: PadCopyCursor,
}
#[derive(Clone, Debug)]
struct PadFillCursor {
steps: AxisVec<isize>,
resets: AxisVec<isize>,
}
#[derive(Clone, Debug)]
struct PadCopyCursor {
shape: AxisVec<usize>,
source_base_delta: isize,
dest_base_delta: isize,
source_steps: AxisVec<isize>,
source_resets: AxisVec<isize>,
dest_steps: AxisVec<isize>,
dest_resets: AxisVec<isize>,
total: usize,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
struct ContiguousPadAxis0Run {
operand_start: usize,
dest_start: usize,
len: usize,
}
#[derive(Clone, Debug)]
pub struct ConcatenatePlan {
input_dims: Vec<AxisVec<usize>>,
input_strides: Vec<AxisVec<isize>>,
dest_dims: AxisVec<usize>,
dest_strides: AxisVec<isize>,
dest_offset_deltas: Vec<isize>,
#[cfg_attr(not(feature = "parallel"), allow(dead_code))]
segment_starts: Vec<usize>,
copy_plans: Vec<CopyPlan>,
}
impl SlicePlan {
#[allow(clippy::too_many_arguments)]
pub fn compile(
operand_dims: &[usize],
operand_strides: &[isize],
dest_dims: &[usize],
dest_strides: &[isize],
starts: &[usize],
limits: &[usize],
slice_strides: &[usize],
) -> Result<Self> {
let rank = operand_dims.len();
if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
return Err(StridedError::StrideLengthMismatch);
}
if starts.len() != rank {
return Err(StridedError::RankMismatch(starts.len(), rank));
}
if limits.len() != rank {
return Err(StridedError::RankMismatch(limits.len(), rank));
}
if slice_strides.len() != rank {
return Err(StridedError::RankMismatch(slice_strides.len(), rank));
}
checked_total_len(operand_dims)?;
checked_total_len(dest_dims)?;
let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
let mut source_offset_delta = 0isize;
for axis in 0..rank {
let start = starts[axis];
let limit = limits[axis];
let stride = slice_strides[axis];
if start > limit || limit > operand_dims[axis] || stride == 0 {
return Err(StridedError::InvalidAxis { axis, rank });
}
let span = limit - start;
expected_dest_dims.push(span.div_ceil(stride));
source_strides.push(checked_stride_mul(operand_strides[axis], stride)?);
source_offset_delta =
checked_offset_add(source_offset_delta, operand_strides[axis], start)?;
}
if dest_dims != &expected_dest_dims[..] {
return Err(StridedError::ShapeMismatch(
dest_dims.to_vec(),
expected_dest_dims.to_vec(),
));
}
let copy_plan = CopyPlan::compile(dest_dims, dest_strides, &source_strides)?;
Ok(Self {
operand_dims: operand_dims.into(),
operand_strides: operand_strides.into(),
dest_dims: dest_dims.into(),
dest_strides: dest_strides.into(),
source_strides,
source_offset_delta,
copy_plan,
})
}
pub fn execute<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_call(dest, operand)?;
let source_offset = operand
.offset()
.checked_add(self.source_offset_delta)
.ok_or(StridedError::OffsetOverflow)?;
let source = unsafe {
RawStridedRef::new_unchecked(
operand.data(),
&self.dest_dims,
&self.source_strides,
source_offset,
)
};
self.copy_plan.execute(dest, &source)
}
pub fn execute_uninit<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_call(dest, operand)?;
let source_offset = operand
.offset()
.checked_add(self.source_offset_delta)
.ok_or(StridedError::OffsetOverflow)?;
let source = unsafe {
RawStridedRef::new_unchecked(
operand.data(),
&self.dest_dims,
&self.source_strides,
source_offset,
)
};
self.copy_plan.execute_uninit(dest, &source)
}
fn check_call<D, T>(
&self,
dest: &RawStridedMut<'_, D>,
operand: &RawStridedRef<'_, T>,
) -> Result<()> {
if operand.dims() != &self.operand_dims[..]
|| operand.strides() != &self.operand_strides[..]
|| dest.dims() != &self.dest_dims[..]
|| dest.strides() != &self.dest_strides[..]
{
return Err(StridedError::PlanLayoutMismatch);
}
Ok(())
}
}
impl PadPlan {
#[allow(clippy::too_many_arguments)]
pub fn compile(
operand_dims: &[usize],
operand_strides: &[isize],
dest_dims: &[usize],
dest_strides: &[isize],
edge_padding_low: &[i64],
edge_padding_high: &[i64],
interior_padding: &[i64],
) -> Result<Self> {
let rank = operand_dims.len();
if operand_strides.len() != rank || dest_dims.len() != rank || dest_strides.len() != rank {
return Err(StridedError::StrideLengthMismatch);
}
if edge_padding_low.len() != rank {
return Err(StridedError::RankMismatch(edge_padding_low.len(), rank));
}
if edge_padding_high.len() != rank {
return Err(StridedError::RankMismatch(edge_padding_high.len(), rank));
}
if interior_padding.len() != rank {
return Err(StridedError::RankMismatch(interior_padding.len(), rank));
}
let operand_total = checked_total_len(operand_dims)?;
let dest_total = checked_total_len(dest_dims)?;
if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
return Err(StridedError::NonInjectiveOutputLayout);
}
let mut expected_dest_dims: AxisVec<usize> = AxisVec::with_capacity(rank);
let mut interior_step: AxisVec<i64> = AxisVec::with_capacity(rank);
for axis in 0..rank {
if interior_padding[axis] < 0 {
return Err(StridedError::InvalidAxis { axis, rank });
}
let step = interior_padding[axis]
.checked_add(1)
.ok_or(StridedError::OffsetOverflow)?;
interior_step.push(step);
expected_dest_dims.push(checked_pad_output_dim(
operand_dims[axis],
edge_padding_low[axis],
edge_padding_high[axis],
step,
axis,
rank,
)?);
}
if dest_dims != &expected_dest_dims[..] {
return Err(StridedError::ShapeMismatch(
dest_dims.to_vec(),
expected_dest_dims.to_vec(),
));
}
let contiguous_dest_fill = is_dense_col_major(dest_dims, dest_strides);
let contiguous_axis0_run = compile_contiguous_pad_axis0_run(
operand_dims,
operand_strides,
dest_dims,
dest_strides,
edge_padding_low,
&interior_step,
);
let generic_fill = compile_pad_fill_cursor(dest_dims, dest_strides)?;
let generic_copy = compile_pad_copy_cursor(
operand_dims,
operand_strides,
dest_dims,
dest_strides,
edge_padding_low,
&interior_step,
)?;
Ok(Self {
operand_dims: operand_dims.into(),
operand_strides: operand_strides.into(),
dest_dims: dest_dims.into(),
dest_strides: dest_strides.into(),
edge_padding_low: edge_padding_low.into(),
interior_step,
operand_total,
dest_total,
contiguous_dest_fill,
contiguous_axis0_run,
generic_fill,
generic_copy,
})
}
pub fn execute<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
fill: T,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_call(dest, operand)?;
self.fill_dest(dest, fill)?;
if self.operand_total == 0 {
return Ok(());
}
if let Some(run) = self.contiguous_axis0_run {
return self.copy_operand_axis0_runs(dest, operand, run);
}
self.copy_operand(dest, operand)
}
pub fn execute_uninit<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
fill: T,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_call(dest, operand)?;
self.fill_dest(dest, MaybeUninit::new(fill))?;
if self.operand_total == 0 {
return Ok(());
}
if let Some(run) = self.contiguous_axis0_run {
return self.copy_operand_axis0_runs_uninit(dest, operand, run);
}
self.copy_operand_uninit(dest, operand)
}
fn fill_dest<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
where
T: Copy + MaybeSendSync,
{
if self.dest_total == 0 {
return Ok(());
}
if self.contiguous_dest_fill {
let dest_offset =
usize::try_from(dest.offset()).map_err(|_| StridedError::OffsetOverflow)?;
let dest_end = dest_offset
.checked_add(self.dest_total)
.ok_or(StridedError::OffsetOverflow)?;
let dest_data = dest.data_mut();
let dest_slice = dest_data
.get_mut(dest_offset..dest_end)
.ok_or(StridedError::OffsetOverflow)?;
crate::threading::fill_contiguous(dest_slice, fill);
return Ok(());
}
#[cfg(feature = "parallel")]
{
let nthreads = crate::threading::parallel_threads_for_len(self.dest_total);
if nthreads > 1 {
return self.fill_dest_parallel(dest, fill, nthreads);
}
}
self.fill_dest_serial(dest, fill)
}
fn copy_operand_axis0_runs<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
run: ContiguousPadAxis0Run,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let dest_offset = dest.offset();
let dest_ptr = dest.data_mut().as_mut_ptr();
unsafe { self.copy_axis0_runs_raw(dest_ptr, dest_offset, operand, run) }
}
fn copy_operand_axis0_runs_uninit<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
run: ContiguousPadAxis0Run,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let dest_offset = dest.offset();
let dest_ptr = dest.data_mut().as_mut_ptr().cast::<T>();
unsafe { self.copy_axis0_runs_raw(dest_ptr, dest_offset, operand, run) }
}
unsafe fn copy_axis0_runs_raw<T>(
&self,
dest_ptr: *mut T,
dest_offset: isize,
operand: &RawStridedRef<'_, T>,
run: ContiguousPadAxis0Run,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
if run.len == 0 {
return Ok(());
}
let outer_dims = &self.operand_dims[1..];
let outer_total = checked_total_len(outer_dims)?;
#[cfg(feature = "parallel")]
{
let copied = outer_total.saturating_mul(run.len);
let nthreads = crate::threading::parallel_threads_for_len(copied);
if nthreads > 1 && outer_total >= nthreads {
let dest_ptr = crate::threading::SendPtr(dest_ptr);
return crate::threading::parallel_map_reduce(
0..outer_total,
nthreads,
&|range| {
let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
let outer_idx = outer_idx_storage.as_mut_slice();
fill_col_major_index(range.start, outer_dims, outer_idx);
unsafe {
self.copy_axis0_run_range(
dest_ptr.as_ptr(),
dest_offset,
operand,
run,
outer_idx,
range.len(),
)
}
},
&|left, right| left.and(right),
);
}
}
let mut outer_idx_storage = CoordScratch::new(outer_dims.len());
let outer_idx = outer_idx_storage.as_mut_slice();
unsafe {
self.copy_axis0_run_range(dest_ptr, dest_offset, operand, run, outer_idx, outer_total)
}
}
unsafe fn copy_axis0_run_range<T>(
&self,
dest_ptr: *mut T,
dest_offset: isize,
operand: &RawStridedRef<'_, T>,
run: ContiguousPadAxis0Run,
outer_idx: &mut [usize],
count: usize,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let outer_dims = &self.operand_dims[1..];
let operand_ptr = operand.data().as_ptr();
for _ in 0..count {
let mut operand_offset =
checked_offset_add(operand.offset(), self.operand_strides[0], run.operand_start)?;
let mut dest_run_offset =
checked_offset_add(dest_offset, self.dest_strides[0], run.dest_start)?;
let mut in_bounds = true;
for (outer_axis, &coord) in outer_idx.iter().enumerate() {
let axis = outer_axis + 1;
let out_pos = i128::from(self.edge_padding_low[axis])
+ coord as i128 * i128::from(self.interior_step[axis]);
if out_pos < 0 || out_pos >= self.dest_dims[axis] as i128 {
in_bounds = false;
break;
}
operand_offset =
checked_offset_add(operand_offset, self.operand_strides[axis], coord)?;
dest_run_offset =
checked_offset_add(dest_run_offset, self.dest_strides[axis], out_pos as usize)?;
}
if in_bounds {
unsafe {
crate::threading::copy_contiguous(
operand_ptr.offset(operand_offset),
dest_ptr.offset(dest_run_offset),
run.len,
);
}
}
advance_col_major_index(outer_idx, outer_dims);
}
Ok(())
}
#[cfg(test)]
fn contiguous_axis0_run(&self) -> Option<(usize, usize, usize)> {
self.contiguous_axis0_run
.map(|run| (run.operand_start, run.dest_start, run.len))
}
#[cfg(test)]
fn has_contiguous_dest_fill(&self) -> bool {
self.contiguous_dest_fill
}
fn fill_dest_serial<T>(&self, dest: &mut RawStridedMut<'_, T>, fill: T) -> Result<()>
where
T: Copy,
{
let dest_ptr = dest.data_mut().as_mut_ptr();
let mut cursor = PadFillState::new(dest.offset(), &self.dest_dims, &self.generic_fill);
for _ in 0..self.dest_total {
unsafe {
*dest_ptr.offset(cursor.offset) = fill;
}
cursor.advance(&self.dest_dims, &self.generic_fill);
}
Ok(())
}
#[cfg(feature = "parallel")]
fn fill_dest_parallel<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
fill: T,
nthreads: usize,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let dest_offset_base = dest.offset();
let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
crate::threading::parallel_map_reduce(
0..self.dest_total,
nthreads,
&|range| {
let mut cursor = PadFillState::decode(
range.start,
dest_offset_base,
&self.dest_dims,
&self.generic_fill,
)?;
let dest_ptr = dest_ptr.as_ptr();
for _ in range {
unsafe {
*dest_ptr.offset(cursor.offset) = fill;
}
cursor.advance(&self.dest_dims, &self.generic_fill);
}
Ok(())
},
&|left, right| left.and(right),
)
}
fn copy_operand<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
#[cfg(feature = "parallel")]
{
let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
if nthreads > 1 {
return self.copy_operand_parallel(dest, operand, nthreads);
}
}
self.copy_operand_serial(dest, operand)
}
fn copy_operand_uninit<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
#[cfg(feature = "parallel")]
{
let nthreads = crate::threading::parallel_threads_for_len(self.generic_copy.total);
if nthreads > 1 {
return self.copy_operand_uninit_parallel(dest, operand, nthreads);
}
}
self.copy_operand_uninit_serial(dest, operand)
}
fn copy_operand_uninit_serial<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy,
{
if self.generic_copy.total == 0 {
return Ok(());
}
let operand_ptr = operand.data().as_ptr();
let dest_ptr = dest.data_mut().as_mut_ptr();
let source_base = operand
.offset()
.checked_add(self.generic_copy.source_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let dest_base = dest
.offset()
.checked_add(self.generic_copy.dest_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
for _ in 0..self.generic_copy.total {
unsafe {
(*dest_ptr.offset(cursor.dest_offset))
.write(*operand_ptr.offset(cursor.source_offset));
}
cursor.advance(&self.generic_copy);
}
Ok(())
}
#[cfg(feature = "parallel")]
fn copy_operand_uninit_parallel<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
nthreads: usize,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let operand_base = operand
.offset()
.checked_add(self.generic_copy.source_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
let dest_base = dest
.offset()
.checked_add(self.generic_copy.dest_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
crate::threading::parallel_map_reduce(
0..self.generic_copy.total,
nthreads,
&|range| {
let mut cursor =
PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
let operand_ptr = operand_ptr.as_const();
let dest_ptr = dest_ptr.as_ptr();
for _ in range {
unsafe {
(*dest_ptr.offset(cursor.dest_offset))
.write(*operand_ptr.offset(cursor.source_offset));
}
cursor.advance(&self.generic_copy);
}
Ok(())
},
&|left, right| left.and(right),
)
}
fn copy_operand_serial<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy,
{
if self.generic_copy.total == 0 {
return Ok(());
}
let operand_ptr = operand.data().as_ptr();
let dest_ptr = dest.data_mut().as_mut_ptr();
let source_base = operand
.offset()
.checked_add(self.generic_copy.source_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let dest_base = dest
.offset()
.checked_add(self.generic_copy.dest_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let mut cursor = PadCopyState::new(source_base, dest_base, &self.generic_copy);
for _ in 0..self.generic_copy.total {
unsafe {
*dest_ptr.offset(cursor.dest_offset) = *operand_ptr.offset(cursor.source_offset);
}
cursor.advance(&self.generic_copy);
}
Ok(())
}
#[cfg(feature = "parallel")]
fn copy_operand_parallel<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
nthreads: usize,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let operand_base = operand
.offset()
.checked_add(self.generic_copy.source_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let operand_ptr = crate::threading::SendPtr(operand.data().as_ptr() as *mut T);
let dest_base = dest
.offset()
.checked_add(self.generic_copy.dest_base_delta)
.ok_or(StridedError::OffsetOverflow)?;
let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
crate::threading::parallel_map_reduce(
0..self.generic_copy.total,
nthreads,
&|range| {
let mut cursor =
PadCopyState::decode(range.start, operand_base, dest_base, &self.generic_copy)?;
let operand_ptr = operand_ptr.as_const();
let dest_ptr = dest_ptr.as_ptr();
for _ in range {
unsafe {
*dest_ptr.offset(cursor.dest_offset) =
*operand_ptr.offset(cursor.source_offset);
}
cursor.advance(&self.generic_copy);
}
Ok(())
},
&|left, right| left.and(right),
)
}
fn check_call<D, T>(
&self,
dest: &RawStridedMut<'_, D>,
operand: &RawStridedRef<'_, T>,
) -> Result<()> {
if operand.dims() != &self.operand_dims[..]
|| operand.strides() != &self.operand_strides[..]
|| dest.dims() != &self.dest_dims[..]
|| dest.strides() != &self.dest_strides[..]
{
return Err(StridedError::PlanLayoutMismatch);
}
Ok(())
}
}
#[cfg(feature = "parallel")]
struct ConcatSegment<'a, T> {
layout: &'a crate::raw_ops::FusedPairLayout,
dest_offset: isize,
src_ptr: crate::threading::SendPtr<T>,
src_offset: isize,
}
impl ConcatenatePlan {
pub fn compile(
input_dims: &[&[usize]],
input_strides: &[&[isize]],
dest_dims: &[usize],
dest_strides: &[isize],
axis: usize,
) -> Result<Self> {
if input_dims.is_empty() {
return Err(StridedError::UnsupportedArity {
arity: 0,
max: usize::MAX,
});
}
if input_dims.len() != input_strides.len() {
return Err(StridedError::RankMismatch(
input_strides.len(),
input_dims.len(),
));
}
let rank = input_dims[0].len();
if dest_dims.len() != rank || dest_strides.len() != rank {
return Err(StridedError::StrideLengthMismatch);
}
if axis >= rank {
return Err(StridedError::InvalidAxis { axis, rank });
}
checked_total_len(dest_dims)?;
if !crate::layout_check::is_injective_layout(dest_dims, dest_strides) {
return Err(StridedError::NonInjectiveOutputLayout);
}
let mut expected_dest_dims: AxisVec<usize> = input_dims[0].into();
expected_dest_dims[axis] = 0;
let mut stored_input_dims = Vec::with_capacity(input_dims.len());
let mut stored_input_strides = Vec::with_capacity(input_dims.len());
let mut dest_offset_deltas = Vec::with_capacity(input_dims.len());
let mut copy_plans = Vec::with_capacity(input_dims.len());
let mut segment_starts = Vec::with_capacity(input_dims.len() + 1);
segment_starts.push(0usize);
let mut axis_base = 0usize;
for (dims, strides) in input_dims.iter().zip(input_strides.iter()) {
if dims.len() != rank {
return Err(StridedError::RankMismatch(dims.len(), rank));
}
if strides.len() != rank {
return Err(StridedError::StrideLengthMismatch);
}
let segment_total = checked_total_len(dims)?;
let segment_end = segment_starts[segment_starts.len() - 1]
.checked_add(segment_total)
.ok_or(StridedError::OffsetOverflow)?;
segment_starts.push(segment_end);
for dim in 0..rank {
if dim == axis {
expected_dest_dims[axis] = expected_dest_dims[axis]
.checked_add(dims[axis])
.ok_or(StridedError::OffsetOverflow)?;
} else if dims[dim] != input_dims[0][dim] {
return Err(StridedError::ShapeMismatch(
dims.to_vec(),
input_dims[0].to_vec(),
));
}
}
dest_offset_deltas.push(checked_offset_add(0, dest_strides[axis], axis_base)?);
axis_base = axis_base
.checked_add(dims[axis])
.ok_or(StridedError::OffsetOverflow)?;
copy_plans.push(CopyPlan::compile(dims, dest_strides, strides)?);
stored_input_dims.push((*dims).into());
stored_input_strides.push((*strides).into());
}
if dest_dims != &expected_dest_dims[..] {
return Err(StridedError::ShapeMismatch(
dest_dims.to_vec(),
expected_dest_dims.to_vec(),
));
}
Ok(Self {
input_dims: stored_input_dims,
input_strides: stored_input_strides,
dest_dims: dest_dims.into(),
dest_strides: dest_strides.into(),
dest_offset_deltas,
segment_starts,
copy_plans,
})
}
pub fn execute<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
inputs: &[RawStridedRef<'_, T>],
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_dest_layout(dest)?;
if inputs.len() != self.input_dims.len() {
return Err(StridedError::RankMismatch(
inputs.len(),
self.input_dims.len(),
));
}
for (position, input) in inputs.iter().enumerate() {
self.check_input_layout(position, input)?;
self.segment_offset(position, dest.offset())?;
}
#[cfg(feature = "parallel")]
if self.try_execute_parallel(dest, inputs, |dst: &mut T, value| *dst = value)? {
return Ok(());
}
for (position, input) in inputs.iter().enumerate() {
self.execute_segment(position, dest, input)?;
}
Ok(())
}
pub fn execute_uninit<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
inputs: &[RawStridedRef<'_, T>],
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_dest_layout(dest)?;
if inputs.len() != self.input_dims.len() {
return Err(StridedError::RankMismatch(
inputs.len(),
self.input_dims.len(),
));
}
for (position, input) in inputs.iter().enumerate() {
self.check_input_layout(position, input)?;
self.segment_offset(position, dest.offset())?;
}
#[cfg(feature = "parallel")]
if self.try_execute_parallel(dest, inputs, |dst: &mut MaybeUninit<T>, value| {
dst.write(value);
})? {
return Ok(());
}
for (position, input) in inputs.iter().enumerate() {
self.execute_segment_uninit(position, dest, input)?;
}
Ok(())
}
#[cfg(feature = "parallel")]
fn try_execute_parallel<D, T, Apply>(
&self,
dest: &mut RawStridedMut<'_, D>,
inputs: &[RawStridedRef<'_, T>],
apply: Apply,
) -> Result<bool>
where
D: Copy + MaybeSendSync,
T: Copy + MaybeSendSync,
Apply: Fn(&mut D, T) + MaybeSendSync,
{
let total = self.segment_starts[self.segment_starts.len() - 1];
let nthreads = crate::threading::parallel_threads_for_len(total);
if nthreads <= 1 {
return Ok(false);
}
let mut segments = Vec::with_capacity(inputs.len());
for (position, input) in inputs.iter().enumerate() {
let Some(layout) = self.copy_plans[position].fused_layout() else {
return Ok(false);
};
segments.push(ConcatSegment {
layout,
dest_offset: self.segment_offset(position, dest.offset())?,
src_ptr: crate::threading::SendPtr(input.data().as_ptr() as *mut T),
src_offset: input.offset(),
});
}
let dest_ptr = crate::threading::SendPtr(dest.data_mut().as_mut_ptr());
let starts = &self.segment_starts;
let segments = &segments;
crate::threading::parallel_for_each(0..total, nthreads, &|range| {
let mut position = starts[1..].partition_point(|&end| end <= range.start);
while position < segments.len() && starts[position] < range.end {
let segment = &segments[position];
let local_start = range.start.max(starts[position]) - starts[position];
let local_end = range.end.min(starts[position + 1]) - starts[position];
unsafe {
crate::raw_ops::apply_fused_range(
dest_ptr.as_ptr(),
segment.dest_offset,
segment.src_ptr.as_const(),
segment.src_offset,
segment.layout,
local_start,
local_end - local_start,
&apply,
&|value| value,
);
}
position += 1;
}
});
Ok(true)
}
pub(crate) fn check_dest_layout<T>(&self, dest: &RawStridedMut<'_, T>) -> Result<()> {
if dest.dims() != &self.dest_dims[..] || dest.strides() != &self.dest_strides[..] {
return Err(StridedError::PlanLayoutMismatch);
}
Ok(())
}
pub(crate) fn check_input_layout<T>(
&self,
position: usize,
input: &RawStridedRef<'_, T>,
) -> Result<()> {
if position >= self.input_dims.len()
|| input.dims() != &self.input_dims[position][..]
|| input.strides() != &self.input_strides[position][..]
{
return Err(StridedError::PlanLayoutMismatch);
}
Ok(())
}
pub(crate) fn input_count(&self) -> usize {
self.input_dims.len()
}
pub(crate) fn prefers_whole_plan(&self) -> bool {
#[cfg(feature = "parallel")]
{
let total = self.segment_starts[self.segment_starts.len() - 1];
crate::threading::parallel_threads_for_len(total) > 1
}
#[cfg(not(feature = "parallel"))]
{
false
}
}
pub(crate) fn segment_offset(&self, position: usize, dest_offset: isize) -> Result<isize> {
dest_offset
.checked_add(self.dest_offset_deltas[position])
.ok_or(StridedError::OffsetOverflow)
}
pub(crate) fn execute_segment<T>(
&self,
position: usize,
dest: &mut RawStridedMut<'_, T>,
input: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let segment_offset = self.segment_offset(position, dest.offset())?;
let dest_data = dest.data_mut();
let mut segment = unsafe {
RawStridedMut::new_unchecked(
dest_data,
&self.input_dims[position],
&self.dest_strides,
segment_offset,
)
};
self.copy_plans[position].execute(&mut segment, input)
}
pub(crate) fn execute_segment_uninit<T>(
&self,
position: usize,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
input: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
let segment_offset = self.segment_offset(position, dest.offset())?;
let dest_data = dest.data_mut();
let mut segment = unsafe {
RawStridedMut::new_unchecked(
dest_data,
&self.input_dims[position],
&self.dest_strides,
segment_offset,
)
};
self.copy_plans[position].execute_uninit(&mut segment, input)
}
}
impl ReversePlan {
pub fn compile(
operand_dims: &[usize],
operand_strides: &[isize],
dest_strides: &[isize],
axes: &[usize],
) -> Result<Self> {
let rank = operand_dims.len();
if operand_strides.len() != rank || dest_strides.len() != rank {
return Err(StridedError::StrideLengthMismatch);
}
checked_total_len(operand_dims)?;
let mut reverse_axis: AxisVec<bool> = (0..rank).map(|_| false).collect();
for &axis in axes {
if axis >= rank {
return Err(StridedError::InvalidAxis { axis, rank });
}
reverse_axis[axis] = true;
}
let mut source_strides: AxisVec<isize> = AxisVec::with_capacity(rank);
let mut source_offset_delta = 0isize;
for axis in 0..rank {
if reverse_axis[axis] {
source_strides.push(
operand_strides[axis]
.checked_neg()
.ok_or(StridedError::OffsetOverflow)?,
);
if operand_dims[axis] > 0 {
source_offset_delta = checked_offset_add(
source_offset_delta,
operand_strides[axis],
operand_dims[axis] - 1,
)?;
}
} else {
source_strides.push(operand_strides[axis]);
}
}
let copy_plan = CopyPlan::compile(operand_dims, dest_strides, &source_strides)?;
Ok(Self {
operand_dims: operand_dims.into(),
operand_strides: operand_strides.into(),
dest_strides: dest_strides.into(),
source_strides,
source_offset_delta,
copy_plan,
})
}
pub fn execute<T>(
&self,
dest: &mut RawStridedMut<'_, T>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_call(dest, operand)?;
let source_offset = operand
.offset()
.checked_add(self.source_offset_delta)
.ok_or(StridedError::OffsetOverflow)?;
let source = unsafe {
RawStridedRef::new_unchecked(
operand.data(),
&self.operand_dims,
&self.source_strides,
source_offset,
)
};
self.copy_plan.execute(dest, &source)
}
pub fn execute_uninit<T>(
&self,
dest: &mut RawStridedMut<'_, MaybeUninit<T>>,
operand: &RawStridedRef<'_, T>,
) -> Result<()>
where
T: Copy + MaybeSendSync,
{
self.check_call(dest, operand)?;
let source_offset = operand
.offset()
.checked_add(self.source_offset_delta)
.ok_or(StridedError::OffsetOverflow)?;
let source = unsafe {
RawStridedRef::new_unchecked(
operand.data(),
&self.operand_dims,
&self.source_strides,
source_offset,
)
};
self.copy_plan.execute_uninit(dest, &source)
}
fn check_call<D, T>(
&self,
dest: &RawStridedMut<'_, D>,
operand: &RawStridedRef<'_, T>,
) -> Result<()> {
if operand.dims() != &self.operand_dims[..]
|| operand.strides() != &self.operand_strides[..]
|| dest.dims() != &self.operand_dims[..]
|| dest.strides() != &self.dest_strides[..]
{
return Err(StridedError::PlanLayoutMismatch);
}
Ok(())
}
}
fn checked_total_len(dims: &[usize]) -> Result<usize> {
if dims.is_empty() {
return Ok(1);
}
dims.iter()
.try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
.ok_or(StridedError::OffsetOverflow)
}
fn checked_stride_mul(stride: isize, factor: usize) -> Result<isize> {
let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
stride
.checked_mul(factor)
.ok_or(StridedError::OffsetOverflow)
}
fn checked_pad_output_dim(
input_extent: usize,
edge_low: i64,
edge_high: i64,
interior_step: i64,
axis: usize,
rank: usize,
) -> Result<usize> {
let base = if input_extent == 0 {
0i128
} else {
(input_extent as i128 - 1)
.checked_mul(i128::from(interior_step))
.and_then(|value| value.checked_add(1))
.ok_or(StridedError::OffsetOverflow)?
};
let dim = i128::from(edge_low)
.checked_add(i128::from(edge_high))
.and_then(|value| value.checked_add(base))
.ok_or(StridedError::OffsetOverflow)?;
usize::try_from(dim).map_err(|_| StridedError::InvalidAxis { axis, rank })
}
fn compile_pad_fill_cursor(dims: &[usize], strides: &[isize]) -> Result<PadFillCursor> {
let mut steps = AxisVec::with_capacity(dims.len());
let mut resets = AxisVec::with_capacity(dims.len());
for (&dim, &stride) in dims.iter().zip(strides) {
steps.push(stride);
resets.push(checked_cursor_reset(stride, dim)?);
}
check_offset_span(dims, &steps)?;
Ok(PadFillCursor { steps, resets })
}
fn compile_pad_copy_cursor(
operand_dims: &[usize],
operand_strides: &[isize],
dest_dims: &[usize],
dest_strides: &[isize],
edge_padding_low: &[i64],
interior_step: &[i64],
) -> Result<PadCopyCursor> {
let mut shape = AxisVec::with_capacity(operand_dims.len());
let mut source_steps = AxisVec::with_capacity(operand_dims.len());
let mut source_resets = AxisVec::with_capacity(operand_dims.len());
let mut dest_steps = AxisVec::with_capacity(operand_dims.len());
let mut dest_resets = AxisVec::with_capacity(operand_dims.len());
let mut source_base_delta = 0isize;
let mut dest_base_delta = 0isize;
let mut copy_empty = false;
for axis in 0..operand_dims.len() {
let (start, end) = checked_pad_valid_interval(
operand_dims[axis],
dest_dims[axis],
edge_padding_low[axis],
interior_step[axis],
)?;
let extent = end - start;
shape.push(extent);
if !copy_empty && extent != 0 {
source_base_delta =
checked_offset_add(source_base_delta, operand_strides[axis], start)?;
let output_start = i128::from(edge_padding_low[axis])
.checked_add(
i128::try_from(start)
.map_err(|_| StridedError::OffsetOverflow)?
.checked_mul(i128::from(interior_step[axis]))
.ok_or(StridedError::OffsetOverflow)?,
)
.ok_or(StridedError::OffsetOverflow)?;
let output_start =
usize::try_from(output_start).map_err(|_| StridedError::OffsetOverflow)?;
dest_base_delta =
checked_offset_add(dest_base_delta, dest_strides[axis], output_start)?;
}
copy_empty |= extent == 0;
let source_step = operand_strides[axis];
let dest_step = checked_stride_mul_i64(dest_strides[axis], interior_step[axis])?;
source_steps.push(source_step);
source_resets.push(checked_cursor_reset(source_step, extent)?);
dest_steps.push(dest_step);
dest_resets.push(checked_cursor_reset(dest_step, extent)?);
}
check_offset_span(&shape, &source_steps)?;
check_offset_span(&shape, &dest_steps)?;
let total = if shape.iter().any(|&extent| extent == 0) {
0
} else {
checked_total_len(&shape)?
};
Ok(PadCopyCursor {
shape,
source_base_delta,
dest_base_delta,
source_steps,
source_resets,
dest_steps,
dest_resets,
total,
})
}
fn checked_pad_valid_interval(
input_extent: usize,
dest_extent: usize,
edge_low: i64,
step: i64,
) -> Result<(usize, usize)> {
if input_extent == 0 || dest_extent == 0 {
return Ok((0, 0));
}
let step = i128::from(step);
let lower = ceil_div_positive(-i128::from(edge_low), step)?;
let dest_last = i128::try_from(dest_extent)
.map_err(|_| StridedError::OffsetOverflow)?
.checked_sub(1)
.ok_or(StridedError::OffsetOverflow)?;
let upper = floor_div_positive(
dest_last
.checked_sub(i128::from(edge_low))
.ok_or(StridedError::OffsetOverflow)?,
step,
)?
.checked_add(1)
.ok_or(StridedError::OffsetOverflow)?;
let input_extent = i128::try_from(input_extent).map_err(|_| StridedError::OffsetOverflow)?;
let lower = lower.clamp(0, input_extent);
let upper = upper.clamp(0, input_extent);
if lower >= upper {
return Ok((0, 0));
}
Ok((
usize::try_from(lower).map_err(|_| StridedError::OffsetOverflow)?,
usize::try_from(upper).map_err(|_| StridedError::OffsetOverflow)?,
))
}
fn ceil_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
if denominator <= 0 {
return Err(StridedError::OffsetOverflow);
}
let quotient = numerator.div_euclid(denominator);
let remainder = numerator.rem_euclid(denominator);
quotient
.checked_add(i128::from(remainder != 0))
.ok_or(StridedError::OffsetOverflow)
}
fn floor_div_positive(numerator: i128, denominator: i128) -> Result<i128> {
if denominator <= 0 {
return Err(StridedError::OffsetOverflow);
}
Ok(numerator.div_euclid(denominator))
}
fn checked_stride_mul_i64(stride: isize, factor: i64) -> Result<isize> {
let factor = isize::try_from(factor).map_err(|_| StridedError::OffsetOverflow)?;
stride
.checked_mul(factor)
.ok_or(StridedError::OffsetOverflow)
}
fn checked_cursor_reset(step: isize, extent: usize) -> Result<isize> {
if extent == 0 {
return Ok(0);
}
let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
step.checked_mul(last)
.and_then(isize::checked_neg)
.ok_or(StridedError::OffsetOverflow)
}
fn check_offset_span(shape: &[usize], strides: &[isize]) -> Result<()> {
let mut min = 0isize;
let mut max = 0isize;
for (&extent, &stride) in shape.iter().zip(strides) {
if extent <= 1 {
continue;
}
let last = isize::try_from(extent - 1).map_err(|_| StridedError::OffsetOverflow)?;
let delta = stride
.checked_mul(last)
.ok_or(StridedError::OffsetOverflow)?;
if delta < 0 {
min = min.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
} else {
max = max.checked_add(delta).ok_or(StridedError::OffsetOverflow)?;
}
}
let _ = (min, max);
Ok(())
}
fn is_dense_col_major(dims: &[usize], strides: &[isize]) -> bool {
let mut expected = 1isize;
for (&dim, &stride) in dims.iter().zip(strides.iter()) {
if stride != expected {
return false;
}
let Ok(dim) = isize::try_from(dim) else {
return false;
};
let Some(next) = expected.checked_mul(dim) else {
return false;
};
expected = next;
}
true
}
fn compile_contiguous_pad_axis0_run(
operand_dims: &[usize],
operand_strides: &[isize],
dest_dims: &[usize],
dest_strides: &[isize],
edge_padding_low: &[i64],
interior_step: &[i64],
) -> Option<ContiguousPadAxis0Run> {
if operand_dims.is_empty()
|| operand_strides[0] != 1
|| dest_strides[0] != 1
|| interior_step[0] != 1
{
return None;
}
let operand_extent = operand_dims[0] as i128;
let dest_extent = dest_dims[0] as i128;
let edge_low = i128::from(edge_padding_low[0]);
let operand_start = (-edge_low).clamp(0, operand_extent);
let dest_start = edge_low.clamp(0, dest_extent);
let len = (operand_extent - operand_start).min(dest_extent - dest_start);
Some(ContiguousPadAxis0Run {
operand_start: usize::try_from(operand_start).ok()?,
dest_start: usize::try_from(dest_start).ok()?,
len: usize::try_from(len).ok()?,
})
}
fn checked_offset_add(base: isize, stride: isize, coord: usize) -> Result<isize> {
let coord = isize::try_from(coord).map_err(|_| StridedError::OffsetOverflow)?;
let scaled = stride
.checked_mul(coord)
.ok_or(StridedError::OffsetOverflow)?;
base.checked_add(scaled).ok_or(StridedError::OffsetOverflow)
}
fn advance_col_major_index(index: &mut [usize], shape: &[usize]) {
for axis in 0..index.len() {
index[axis] += 1;
if index[axis] < shape[axis] {
return;
}
index[axis] = 0;
}
}
#[cfg(feature = "parallel")]
fn fill_col_major_index(mut linear: usize, shape: &[usize], out: &mut [usize]) {
for (axis, coord) in out.iter_mut().enumerate() {
let dim = shape[axis];
*coord = linear % dim;
linear /= dim;
}
}
struct PadFillState {
coords: CoordScratch,
offset: isize,
}
impl PadFillState {
fn new(base: isize, shape: &[usize], _cursor: &PadFillCursor) -> Self {
Self {
coords: CoordScratch::new(shape.len()),
offset: base,
}
}
#[cfg(feature = "parallel")]
fn decode(linear: usize, base: isize, shape: &[usize], cursor: &PadFillCursor) -> Result<Self> {
let mut state = Self::new(base, shape, cursor);
fill_col_major_index(linear, shape, state.coords.as_mut_slice());
for (&coord, &step) in state.coords.as_mut_slice().iter().zip(&cursor.steps) {
state.offset = checked_offset_add(state.offset, step, coord)?;
}
Ok(state)
}
#[inline]
fn advance(&mut self, shape: &[usize], cursor: &PadFillCursor) {
for axis in 0..shape.len() {
let next = self.coords.as_mut_slice()[axis] + 1;
if next < shape[axis] {
self.coords.as_mut_slice()[axis] = next;
self.offset += cursor.steps[axis];
return;
}
self.coords.as_mut_slice()[axis] = 0;
self.offset += cursor.resets[axis];
}
}
}
struct PadCopyState {
coords: CoordScratch,
source_offset: isize,
dest_offset: isize,
}
impl PadCopyState {
fn new(source_base: isize, dest_base: isize, cursor: &PadCopyCursor) -> Self {
Self {
coords: CoordScratch::new(cursor.shape.len()),
source_offset: source_base,
dest_offset: dest_base,
}
}
#[cfg(feature = "parallel")]
fn decode(
linear: usize,
source_base: isize,
dest_base: isize,
cursor: &PadCopyCursor,
) -> Result<Self> {
let mut state = Self::new(source_base, dest_base, cursor);
fill_col_major_index(linear, &cursor.shape, state.coords.as_mut_slice());
for axis in 0..cursor.shape.len() {
let coord = state.coords.as_mut_slice()[axis];
state.source_offset =
checked_offset_add(state.source_offset, cursor.source_steps[axis], coord)?;
state.dest_offset =
checked_offset_add(state.dest_offset, cursor.dest_steps[axis], coord)?;
}
Ok(state)
}
#[inline]
fn advance(&mut self, cursor: &PadCopyCursor) {
for axis in 0..cursor.shape.len() {
let next = self.coords.as_mut_slice()[axis] + 1;
if next < cursor.shape[axis] {
self.coords.as_mut_slice()[axis] = next;
self.source_offset += cursor.source_steps[axis];
self.dest_offset += cursor.dest_steps[axis];
return;
}
self.coords.as_mut_slice()[axis] = 0;
self.source_offset += cursor.source_resets[axis];
self.dest_offset += cursor.dest_resets[axis];
}
}
}
struct CoordScratch {
inline: [usize; crate::RAW_FUSED_RANK_LIMIT],
heap: Option<Vec<usize>>,
len: usize,
}
impl CoordScratch {
fn new(len: usize) -> Self {
if len <= crate::RAW_FUSED_RANK_LIMIT {
Self {
inline: [0; crate::RAW_FUSED_RANK_LIMIT],
heap: None,
len,
}
} else {
Self {
inline: [0; crate::RAW_FUSED_RANK_LIMIT],
heap: Some(vec![0; len]),
len,
}
}
}
fn as_mut_slice(&mut self) -> &mut [usize] {
match &mut self.heap {
Some(heap) => heap,
None => &mut self.inline[..self.len],
}
}
}
#[cfg(test)]
#[path = "static_indexing_plan/tests/tests.rs"]
mod tests;
#[cfg(all(test, feature = "parallel"))]
#[path = "static_indexing_plan/tests/parallel_tests.rs"]
mod parallel_tests;