use super::{
checked_reduce_reset, compress_reduce_outer_axes, ErasedReduceScalar, ReduceInnerAxis,
ReduceInnerCursor, ReduceOuterAxis, ReduceOuterCursor,
};
use crate::{Result, StridedError};
use core::ops::Range;
pub(super) const REDUCE_LANES: usize = 8;
pub(super) const AXIS_RUN_KERNEL_MIN_LEN: usize = 2 * REDUCE_LANES;
pub(super) const AXIS_OUTPUT_BLOCK_MIN_LEN: usize = 16;
pub(super) const AXIS_OUTPUT_BLOCK_BYTES: usize = 64 * 1024;
pub(super) const AXIS_OUTPUT_BLOCK_STEPS: usize = 4;
fn axis_output_block_len<T>() -> usize {
(AXIS_OUTPUT_BLOCK_BYTES / core::mem::size_of::<T>().max(1)).max(AXIS_OUTPUT_BLOCK_MIN_LEN)
}
pub(super) trait ReduceKernel<T: ErasedReduceScalar>: 'static {
fn identity() -> T;
fn map(value: T) -> T;
fn combine(lhs: T, rhs: T) -> T;
#[inline(always)]
fn contiguous(values: &[T]) -> T {
lanes_contiguous::<T, Self>(values)
}
}
pub(super) struct SumKernel;
pub(super) struct ProductKernel;
pub(super) struct SumSquaresKernel;
pub(super) struct MaxKernel;
pub(super) struct MinKernel;
impl<T: ErasedReduceScalar> ReduceKernel<T> for SumKernel {
#[inline(always)]
fn identity() -> T {
T::zero()
}
#[inline(always)]
fn map(value: T) -> T {
value
}
#[inline(always)]
fn combine(lhs: T, rhs: T) -> T {
T::reduce_sum(lhs, rhs)
}
#[inline(always)]
fn contiguous(values: &[T]) -> T {
T::try_simd_sum(values).unwrap_or_else(|| lanes_contiguous::<T, Self>(values))
}
}
impl<T: ErasedReduceScalar> ReduceKernel<T> for ProductKernel {
#[inline(always)]
fn identity() -> T {
T::one()
}
#[inline(always)]
fn map(value: T) -> T {
value
}
#[inline(always)]
fn combine(lhs: T, rhs: T) -> T {
T::reduce_product(lhs, rhs)
}
#[inline(always)]
fn contiguous(values: &[T]) -> T {
T::try_simd_product(values).unwrap_or_else(|| lanes_contiguous::<T, Self>(values))
}
}
impl<T: ErasedReduceScalar> ReduceKernel<T> for SumSquaresKernel {
#[inline(always)]
fn identity() -> T {
T::zero()
}
#[inline(always)]
fn map(value: T) -> T {
T::reduce_product(value, value)
}
#[inline(always)]
fn combine(lhs: T, rhs: T) -> T {
T::reduce_sum(lhs, rhs)
}
#[inline(always)]
fn contiguous(values: &[T]) -> T {
T::try_simd_sum_squares(values).unwrap_or_else(|| lanes_contiguous::<T, Self>(values))
}
}
impl<T: ErasedReduceScalar> ReduceKernel<T> for MaxKernel {
#[inline(always)]
fn identity() -> T {
T::max_identity()
}
#[inline(always)]
fn map(value: T) -> T {
value
}
#[inline(always)]
fn combine(lhs: T, rhs: T) -> T {
T::reduce_max(lhs, rhs)
}
}
impl<T: ErasedReduceScalar> ReduceKernel<T> for MinKernel {
#[inline(always)]
fn identity() -> T {
T::min_identity()
}
#[inline(always)]
fn map(value: T) -> T {
value
}
#[inline(always)]
fn combine(lhs: T, rhs: T) -> T {
T::reduce_min(lhs, rhs)
}
}
pub(super) const CONTIGUOUS_REDUCE_LANES: usize = 16;
#[inline(always)]
fn lanes_contiguous<T, K>(values: &[T]) -> T
where
T: ErasedReduceScalar,
K: ReduceKernel<T> + ?Sized,
{
crate::simd::dispatch_if_large(values.len(), || {
let mut lanes = [K::identity(); CONTIGUOUS_REDUCE_LANES];
let mut chunks = values.chunks_exact(CONTIGUOUS_REDUCE_LANES);
for chunk in chunks.by_ref() {
for (lane, &value) in lanes.iter_mut().zip(chunk) {
*lane = K::combine(*lane, K::map(value));
}
}
for (lane, &value) in lanes.iter_mut().zip(chunks.remainder()) {
*lane = K::combine(*lane, K::map(value));
}
lanes.into_iter().fold(K::identity(), K::combine)
})
}
#[inline(always)]
unsafe fn lanes_strided<T, K>(ptr: *const T, stride: isize, len: usize) -> T
where
T: ErasedReduceScalar,
K: ReduceKernel<T>,
{
let mut lanes = [K::identity(); REDUCE_LANES];
let mut cursor = ptr;
for _ in 0..len / REDUCE_LANES {
let mut lane_ptr = cursor;
for lane in &mut lanes {
*lane = K::combine(*lane, K::map(unsafe { *lane_ptr }));
lane_ptr = lane_ptr.wrapping_offset(stride);
}
cursor = lane_ptr;
}
for lane in lanes.iter_mut().take(len % REDUCE_LANES) {
*lane = K::combine(*lane, K::map(unsafe { *cursor }));
cursor = cursor.wrapping_offset(stride);
}
lanes.into_iter().fold(K::identity(), K::combine)
}
#[inline(always)]
unsafe fn reduce_run<T, K>(ptr: *const T, stride: isize, len: usize) -> T
where
T: ErasedReduceScalar,
K: ReduceKernel<T>,
{
if stride == 1 {
K::contiguous(unsafe { core::slice::from_raw_parts(ptr, len) })
} else {
unsafe { lanes_strided::<T, K>(ptr, stride, len) }
}
}
#[derive(Clone, Debug)]
pub(super) struct FullTraversal {
total: usize,
run_len: usize,
run_stride: isize,
run_axes: Vec<ReduceOuterAxis>,
}
impl FullTraversal {
pub(super) fn compile(dims: &[usize], strides: &[isize]) -> Result<Self> {
if dims.len() != strides.len() {
return Err(StridedError::StrideLengthMismatch);
}
let total = dims
.iter()
.try_fold(1usize, |total, &dim| total.checked_mul(dim))
.ok_or(StridedError::OffsetOverflow)?;
if total == 0 {
return Ok(Self {
total,
run_len: 0,
run_stride: 1,
run_axes: Vec::new(),
});
}
let mut axes: Vec<(usize, isize)> = dims
.iter()
.copied()
.zip(strides.iter().copied())
.filter(|&(extent, _)| extent != 1)
.collect();
axes.sort_by_key(|&(_, stride)| {
if stride == 0 {
usize::MAX
} else {
stride.unsigned_abs()
}
});
let mut fused = compress_reduce_outer_axes(
axes.into_iter()
.map(|(extent, stride)| {
Ok(ReduceOuterAxis {
extent,
source_step: stride,
source_reset: checked_reduce_reset(extent, stride)?,
dest_step: 0,
dest_reset: 0,
})
})
.collect::<Result<Vec<_>>>()?,
)?;
if fused.is_empty() {
return Ok(Self {
total,
run_len: 1,
run_stride: 1,
run_axes: fused,
});
}
let run = fused.remove(0);
Ok(Self {
total,
run_len: run.extent,
run_stride: run.source_step,
run_axes: fused,
})
}
#[inline]
pub(super) fn total(&self) -> usize {
self.total
}
}
pub(super) unsafe fn reduce_full_range<T, K>(
source: *const T,
source_base: isize,
traversal: &FullTraversal,
range: Range<usize>,
) -> Result<T>
where
T: ErasedReduceScalar,
K: ReduceKernel<T>,
{
debug_assert!(!range.is_empty() && range.end <= traversal.total);
let run_len = traversal.run_len;
let run_stride = traversal.run_stride;
let mut col = range.start % run_len;
let mut runs =
ReduceOuterCursor::decode(range.start / run_len, source_base, 0, &traversal.run_axes)?;
let mut remaining = range.len();
let mut acc = None;
loop {
let len = (run_len - col).min(remaining);
let start = runs.source_offset + col as isize * run_stride;
let partial = unsafe { reduce_run::<T, K>(source.offset(start), run_stride, len) };
acc = Some(match acc {
Some(acc) => K::combine(acc, partial),
None => partial,
});
remaining -= len;
if remaining == 0 {
break;
}
col = 0;
runs.advance();
}
Ok(acc.unwrap_or_else(K::identity))
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum AxesStrategy {
ContiguousRuns,
OutputBlocks,
Sequential,
}
fn axes_strategy(outer_axes: &[ReduceOuterAxis], inner_axes: &[ReduceInnerAxis]) -> AxesStrategy {
if inner_axes
.first()
.is_some_and(|axis| axis.source_step == 1 && axis.extent >= AXIS_RUN_KERNEL_MIN_LEN)
{
AxesStrategy::ContiguousRuns
} else if outer_axes
.first()
.is_some_and(|axis| axis.source_step == 1 && axis.extent >= AXIS_OUTPUT_BLOCK_MIN_LEN)
{
AxesStrategy::OutputBlocks
} else {
AxesStrategy::Sequential
}
}
pub(super) struct AxesRange<'a, T> {
pub(super) source: *const T,
pub(super) source_base: isize,
pub(super) dest: *mut T,
pub(super) dest_base: isize,
pub(super) outer_axes: &'a [ReduceOuterAxis],
pub(super) inner_axes: &'a [ReduceInnerAxis],
pub(super) reduce_total: usize,
}
pub(super) unsafe fn reduce_axes_range<T, K>(
parts: AxesRange<'_, T>,
range: Range<usize>,
) -> Result<()>
where
T: ErasedReduceScalar,
K: ReduceKernel<T>,
{
if range.is_empty() {
return Ok(());
}
let source = parts.source;
let dest = parts.dest;
let reduce_total = parts.reduce_total;
debug_assert!(reduce_total != 0);
let mut outer = ReduceOuterCursor::decode(
range.start,
parts.source_base,
parts.dest_base,
parts.outer_axes,
)?;
match axes_strategy(parts.outer_axes, parts.inner_axes) {
AxesStrategy::ContiguousRuns => {
let run_len = parts.inner_axes[0].extent;
let run_count = reduce_total / run_len;
let mut runs = ReduceInnerCursor::new(0, &parts.inner_axes[1..]);
for output in range.clone() {
runs.reset(outer.source_offset);
let mut acc = K::contiguous(unsafe {
core::slice::from_raw_parts(source.offset(runs.source_offset), run_len)
});
for _ in 1..run_count {
runs.advance();
let partial = K::contiguous(unsafe {
core::slice::from_raw_parts(source.offset(runs.source_offset), run_len)
});
acc = K::combine(acc, partial);
}
unsafe { dest.offset(outer.dest_offset).write(acc) };
if output + 1 < range.end {
outer.advance();
}
}
}
AxesStrategy::OutputBlocks => {
let block_axis = parts.outer_axes[0];
let mut inner = ReduceInnerCursor::new(0, parts.inner_axes);
let block = axis_output_block_len::<T>().min(range.len());
let mut acc = vec![K::identity(); block];
let mut output = range.start;
while output < range.end {
let block_room = block_axis.extent - outer.leading_coord();
let len = block_room.min(range.end - output).min(block);
let acc = &mut acc[..len];
acc.fill(K::identity());
inner.reset(outer.source_offset);
let mut offsets = [0isize; AXIS_OUTPUT_BLOCK_STEPS];
let mut value_index = 0;
while value_index + AXIS_OUTPUT_BLOCK_STEPS <= reduce_total {
for offset in &mut offsets {
*offset = inner.source_offset;
value_index += 1;
if value_index < reduce_total {
inner.advance();
}
}
let [v0, v1, v2, v3] = offsets.map(|offset| unsafe {
core::slice::from_raw_parts(source.offset(offset), len)
});
for (index, slot) in acc.iter_mut().enumerate() {
let mut value = *slot;
value = K::combine(value, K::map(v0[index]));
value = K::combine(value, K::map(v1[index]));
value = K::combine(value, K::map(v2[index]));
value = K::combine(value, K::map(v3[index]));
*slot = value;
}
}
while value_index < reduce_total {
let values = unsafe {
core::slice::from_raw_parts(source.offset(inner.source_offset), len)
};
for (slot, &value) in acc.iter_mut().zip(values) {
*slot = K::combine(*slot, K::map(value));
}
value_index += 1;
if value_index < reduce_total {
inner.advance();
}
}
if block_axis.dest_step == 1 {
unsafe {
core::ptr::copy_nonoverlapping(
acc.as_ptr(),
dest.offset(outer.dest_offset),
len,
)
};
} else {
for (index, &value) in acc.iter().enumerate() {
let offset = outer.dest_offset + index as isize * block_axis.dest_step;
unsafe { dest.offset(offset).write(value) };
}
}
for _ in 0..len {
output += 1;
if output < range.end {
outer.advance();
}
}
}
}
AxesStrategy::Sequential => {
let mut inner = ReduceInnerCursor::new(0, parts.inner_axes);
for output in range.clone() {
inner.reset(outer.source_offset);
let acc = unsafe { sequential_fold::<T, K>(source, &mut inner, reduce_total) };
unsafe { dest.offset(outer.dest_offset).write(acc) };
if output + 1 < range.end {
outer.advance();
}
}
}
}
Ok(())
}
#[inline(always)]
unsafe fn sequential_fold<T, K>(
source: *const T,
inner: &mut ReduceInnerCursor<'_>,
count: usize,
) -> T
where
T: ErasedReduceScalar,
K: ReduceKernel<T>,
{
let mut acc = K::identity();
for value_index in 0..count {
let value = unsafe { *source.offset(inner.source_offset) };
acc = K::combine(acc, K::map(value));
if value_index + 1 < count {
inner.advance();
}
}
acc
}