use super::{checked_reduce_reset, compress_reduce_outer_axes, ReduceOuterAxis, ReduceOuterCursor};
use crate::{ExecContext, Result, StridedError};
use core::ops::Range;
pub(super) const PANEL: usize = 64;
#[derive(Clone, Debug)]
pub(super) struct LineLayout {
pub(super) axis_len: usize,
pub(super) src_axis_stride: isize,
pub(super) dest_axis_stride: isize,
pub(super) dest_lane_stride: isize,
outer_axes: Vec<ReduceOuterAxis>,
units: usize,
panel_extent: Option<usize>,
#[cfg_attr(not(feature = "parallel"), allow(dead_code))]
total: usize,
}
impl LineLayout {
pub(super) fn compile(
src_dims: &[usize],
src_strides: &[isize],
dest_outer_strides: &[isize],
dest_axis_stride: isize,
axis: usize,
panel_needs_unit_dest: bool,
) -> Result<Self> {
let rank = src_dims.len();
if axis >= rank {
return Err(StridedError::InvalidAxis { axis, rank });
}
debug_assert_eq!(dest_outer_strides.len() + 1, rank);
let axis_len = src_dims[axis];
let src_axis_stride = src_strides[axis];
let total = src_dims
.iter()
.try_fold(1usize, |acc, &dim| acc.checked_mul(dim))
.ok_or(StridedError::OffsetOverflow)?;
let mut outer: Vec<(usize, isize, isize)> = (0..rank)
.filter(|&source_axis| source_axis != axis)
.zip(dest_outer_strides)
.map(|(source_axis, &dest_step)| {
(src_dims[source_axis], src_strides[source_axis], dest_step)
})
.collect();
let units_before_panel = outer
.iter()
.try_fold(1usize, |acc, &(extent, _, _)| acc.checked_mul(extent))
.ok_or(StridedError::OffsetOverflow)?;
outer.retain(|&(extent, _, _)| extent != 1);
outer.sort_by_key(|&(_, source_step, _)| match source_step.unsigned_abs() {
0 => usize::MAX,
step => step,
});
let mut panel_extent = None;
if src_axis_stride != 1 && axis_len > 1 {
if let Some(&(extent, source_step, dest_step)) = outer.first() {
if source_step == 1 && (!panel_needs_unit_dest || dest_step == 1) && extent > 1 {
panel_extent = Some(extent);
}
}
}
let axes = outer
.iter()
.enumerate()
.map(|(index, &(extent, source_step, dest_step))| {
let (extent, source_step, dest_step) = match panel_extent {
Some(_) if index == 0 => {
let panel =
isize::try_from(PANEL).map_err(|_| StridedError::OffsetOverflow)?;
(
extent.div_ceil(PANEL),
source_step
.checked_mul(panel)
.ok_or(StridedError::OffsetOverflow)?,
dest_step
.checked_mul(panel)
.ok_or(StridedError::OffsetOverflow)?,
)
}
_ => (extent, source_step, dest_step),
};
Ok(ReduceOuterAxis {
extent,
source_step,
source_reset: checked_reduce_reset(extent, source_step)?,
dest_step,
dest_reset: checked_reduce_reset(extent, dest_step)?,
})
})
.collect::<Result<Vec<_>>>()?;
let dest_lane_stride = match panel_extent {
Some(_) => outer[0].2,
None => 0,
};
let outer_axes = match panel_extent {
Some(_) => {
let mut axes = axes.into_iter();
let mut fused = vec![axes.next().expect("panel axis exists")];
fused.extend(compress_reduce_outer_axes(axes.collect())?);
fused
}
None => compress_reduce_outer_axes(axes)?,
};
let units = if units_before_panel == 0 {
0
} else {
outer_axes
.iter()
.try_fold(1usize, |acc, axis| acc.checked_mul(axis.extent))
.ok_or(StridedError::OffsetOverflow)?
};
Ok(Self {
axis_len,
src_axis_stride,
dest_axis_stride,
dest_lane_stride,
outer_axes,
units,
panel_extent,
total,
})
}
#[inline]
pub(super) fn is_empty(&self) -> bool {
self.units == 0 || self.axis_len == 0
}
#[inline(always)]
fn width(&self, leading_coord: usize) -> usize {
match self.panel_extent {
Some(extent) => (extent - leading_coord * PANEL).min(PANEL),
None => 1,
}
}
}
pub(super) trait UnitKernel: Copy + crate::MaybeSendSync {
unsafe fn unit(self, source_offset: isize, dest_offset: isize, width: usize);
}
pub(super) unsafe fn for_each_unit<K: UnitKernel>(
ctx: &ExecContext,
layout: &LineLayout,
source_base: isize,
dest_base: isize,
kernel: K,
) -> Result<()> {
if layout.is_empty() {
return Ok(());
}
let task = |range: Range<usize>| {
dispatch_range(RangeTask {
layout,
source_base,
dest_base,
range,
kernel,
})
};
if super::reduce_context_is_serial(ctx) {
return task(0..layout.units);
}
ctx.run(|| {
#[cfg(feature = "parallel")]
{
let nthreads =
crate::threading::parallel_threads_for_len(layout.total).min(layout.units);
if nthreads > 1 {
return crate::threading::parallel_map_reduce(
0..layout.units,
nthreads,
&task,
&|left, right| left.and(right),
);
}
}
task(0..layout.units)
})
}
struct RangeTask<'a, K> {
layout: &'a LineLayout,
source_base: isize,
dest_base: isize,
range: Range<usize>,
kernel: K,
}
#[cfg(feature = "simd")]
impl<K: UnitKernel> pulp::WithSimd for RangeTask<'_, K> {
type Output = Result<()>;
#[inline(always)]
fn with_simd<S: pulp::Simd>(self, _simd: S) -> Self::Output {
run_units(self)
}
}
#[cfg(feature = "simd")]
fn dispatch_range<K: UnitKernel>(task: RangeTask<'_, K>) -> Result<()> {
pulp::Arch::new().dispatch(task)
}
#[cfg(not(feature = "simd"))]
fn dispatch_range<K: UnitKernel>(task: RangeTask<'_, K>) -> Result<()> {
run_units(task)
}
#[inline(always)]
fn run_units<K: UnitKernel>(task: RangeTask<'_, K>) -> Result<()> {
let RangeTask {
layout,
source_base,
dest_base,
range,
kernel,
} = task;
let end = range.end;
let mut cursor =
ReduceOuterCursor::decode(range.start, source_base, dest_base, &layout.outer_axes)?;
for index in range {
unsafe {
kernel.unit(
cursor.source_offset,
cursor.dest_offset,
layout.width(cursor.leading_coord()),
)
};
if index + 1 < end {
cursor.advance();
}
}
Ok(())
}
#[derive(Debug)]
pub(super) struct UnitPtr<T>(pub(super) *mut T);
impl<T> Clone for UnitPtr<T> {
fn clone(&self) -> Self {
*self
}
}
impl<T> Copy for UnitPtr<T> {}
unsafe impl<T> Send for UnitPtr<T> {}
unsafe impl<T> Sync for UnitPtr<T> {}
impl<T> UnitPtr<T> {
#[inline(always)]
pub(super) fn get(self) -> *mut T {
self.0
}
}
#[cfg(all(test, feature = "parallel"))]
#[path = "line/tests/tests.rs"]
mod tests;