use super::line::{for_each_unit, LineLayout, UnitKernel, UnitPtr, PANEL};
use super::{
check_reduce_layout_offset_arithmetic, checked_total_len, reduce_uninit_writer, reduce_writer,
ReduceWriter,
};
use crate::erased_common::{check_dtype, validate_uninit_no_overlap};
use crate::*;
use num_complex::{Complex32, Complex64};
#[non_exhaustive]
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ScanOp {
Sum,
Product,
}
#[non_exhaustive]
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct ScanOptions {
exclusive: bool,
reverse: bool,
}
impl ScanOptions {
#[inline]
pub const fn new() -> Self {
Self {
exclusive: false,
reverse: false,
}
}
#[inline]
pub const fn exclusive(mut self, exclusive: bool) -> Self {
self.exclusive = exclusive;
self
}
#[inline]
pub const fn reverse(mut self, reverse: bool) -> Self {
self.reverse = reverse;
self
}
#[inline]
pub const fn is_exclusive(&self) -> bool {
self.exclusive
}
#[inline]
pub const fn is_reverse(&self) -> bool {
self.reverse
}
}
#[derive(Clone, Debug)]
pub struct ErasedScanPlan {
dtype: KernelDType,
op: ScanOp,
options: ScanOptions,
dims: Vec<usize>,
src_strides: Vec<isize>,
dest_strides: Vec<isize>,
layout: LineLayout,
}
impl ErasedScanPlan {
pub fn compile(
dtype: KernelDType,
op: ScanOp,
dims: &[usize],
src_strides: &[isize],
dest_strides: &[isize],
axis: usize,
options: ScanOptions,
) -> Result<Self> {
check_scan_dtype(dtype)?;
if dims.len() != src_strides.len() || dims.len() != dest_strides.len() {
return Err(StridedError::StrideLengthMismatch);
}
if axis >= dims.len() {
return Err(StridedError::InvalidAxis {
axis,
rank: dims.len(),
});
}
checked_total_len(dims)?;
check_reduce_layout_offset_arithmetic(dims, src_strides)?;
check_reduce_layout_offset_arithmetic(dims, dest_strides)?;
if !crate::layout_check::is_injective_layout(dims, dest_strides) {
return Err(StridedError::NonInjectiveOutputLayout);
}
let dest_outer: Vec<isize> = (0..dims.len())
.filter(|&a| a != axis)
.map(|a| dest_strides[a])
.collect();
let layout = LineLayout::compile(
dims,
src_strides,
&dest_outer,
dest_strides[axis],
axis,
true,
)?;
Ok(Self {
dtype,
op,
options,
dims: dims.to_vec(),
src_strides: src_strides.to_vec(),
dest_strides: dest_strides.to_vec(),
layout,
})
}
#[inline]
pub fn dtype(&self) -> KernelDType {
self.dtype
}
#[inline]
pub fn op(&self) -> ScanOp {
self.op
}
#[inline]
pub fn options(&self) -> ScanOptions {
self.options
}
fn check_layouts(
&self,
dest_dims: &[usize],
dest_strides: &[isize],
src: &ErasedRawStridedRef<'_>,
) -> Result<()> {
if src.dims() != self.dims.as_slice()
|| src.strides() != self.src_strides.as_slice()
|| dest_dims != self.dims.as_slice()
|| dest_strides != self.dest_strides.as_slice()
{
return Err(StridedError::PlanLayoutMismatch);
}
Ok(())
}
pub fn execute(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedMut<'_>,
src: &ErasedRawStridedRef<'_>,
) -> Result<()> {
check_dtype(self.dtype, dest.dtype())?;
check_dtype(self.dtype, src.dtype())?;
self.check_layouts(dest.dims(), dest.strides(), src)?;
macro_rules! run {
($ty:ty) => {{
let mut writer = reduce_writer::<$ty>(dest)?;
self.dispatch::<$ty, _>(ctx, &mut writer, src)
}};
}
match self.dtype {
KernelDType::F32 => run!(f32),
KernelDType::F64 => run!(f64),
KernelDType::I32 => run!(i32),
KernelDType::I64 => run!(i64),
KernelDType::C32 => run!(Complex32),
KernelDType::C64 => run!(Complex64),
_ => Err(unsupported(self.dtype)),
}
}
pub fn execute_uninit(
&self,
ctx: &ExecContext,
dest: &mut ErasedRawStridedUninitMut<'_>,
src: &ErasedRawStridedPtr<'_>,
) -> Result<()> {
check_dtype(self.dtype, dest.dtype())?;
check_dtype(self.dtype, src.dtype())?;
validate_uninit_no_overlap(dest, src, 0)?;
let src = unsafe { src.try_as_ref_after_no_overlap() }?;
self.check_layouts(dest.dims(), dest.strides(), &src)?;
macro_rules! run {
($ty:ty) => {{
let mut writer = reduce_uninit_writer::<$ty>(dest)?;
self.dispatch::<$ty, _>(ctx, &mut writer, &src)
}};
}
match self.dtype {
KernelDType::F32 => run!(f32),
KernelDType::F64 => run!(f64),
KernelDType::I32 => run!(i32),
KernelDType::I64 => run!(i64),
KernelDType::C32 => run!(Complex32),
KernelDType::C64 => run!(Complex64),
_ => Err(unsupported(self.dtype)),
}
}
fn dispatch<T, W>(
&self,
ctx: &ExecContext,
dest: &mut W,
src: &ErasedRawStridedRef<'_>,
) -> Result<()>
where
T: ScanScalar,
W: ReduceWriter<T>,
{
match (self.op, self.options.exclusive, self.options.reverse) {
(ScanOp::Sum, false, false) => self.run::<T, W, SumScan, false, false>(ctx, dest, src),
(ScanOp::Sum, false, true) => self.run::<T, W, SumScan, false, true>(ctx, dest, src),
(ScanOp::Sum, true, false) => self.run::<T, W, SumScan, true, false>(ctx, dest, src),
(ScanOp::Sum, true, true) => self.run::<T, W, SumScan, true, true>(ctx, dest, src),
(ScanOp::Product, false, false) => {
self.run::<T, W, ProductScan, false, false>(ctx, dest, src)
}
(ScanOp::Product, false, true) => {
self.run::<T, W, ProductScan, false, true>(ctx, dest, src)
}
(ScanOp::Product, true, false) => {
self.run::<T, W, ProductScan, true, false>(ctx, dest, src)
}
(ScanOp::Product, true, true) => {
self.run::<T, W, ProductScan, true, true>(ctx, dest, src)
}
}
}
fn run<T, W, K, const EXCLUSIVE: bool, const REVERSE: bool>(
&self,
ctx: &ExecContext,
dest: &mut W,
src: &ErasedRawStridedRef<'_>,
) -> Result<()>
where
T: ScanScalar,
W: ReduceWriter<T>,
K: ScanKernel<T>,
{
let layout = &self.layout;
let source = UnitPtr(src.data_as::<T>()?.as_ptr() as *mut T);
let target = UnitPtr(unsafe { dest.ptr() });
let n = layout.axis_len;
let ss = layout.src_axis_stride;
let ds = layout.dest_axis_stride;
let kernel = ScanUnit::<T, K, EXCLUSIVE, REVERSE> {
source,
target,
n,
ss,
ds,
_kernel: core::marker::PhantomData,
};
unsafe { for_each_unit(ctx, layout, src.offset(), dest.offset(), kernel) }
}
}
struct ScanUnit<T, K, const EXCLUSIVE: bool, const REVERSE: bool> {
source: UnitPtr<T>,
target: UnitPtr<T>,
n: usize,
ss: isize,
ds: isize,
_kernel: core::marker::PhantomData<fn() -> K>,
}
impl<T, K, const EXCLUSIVE: bool, const REVERSE: bool> Clone
for ScanUnit<T, K, EXCLUSIVE, REVERSE>
{
fn clone(&self) -> Self {
*self
}
}
impl<T, K, const EXCLUSIVE: bool, const REVERSE: bool> Copy for ScanUnit<T, K, EXCLUSIVE, REVERSE> {}
impl<T, K, const EXCLUSIVE: bool, const REVERSE: bool> UnitKernel
for ScanUnit<T, K, EXCLUSIVE, REVERSE>
where
T: ScanScalar,
K: ScanKernel<T>,
{
#[inline(always)]
unsafe fn unit(self, so: isize, d_o: isize, width: usize) {
let Self {
source,
target,
n,
ss,
ds,
..
} = self;
unsafe {
if width == 1 {
scan_line::<T, K, EXCLUSIVE, REVERSE>(
source.get(),
so,
ss,
target.get(),
d_o,
ds,
n,
)
} else {
scan_panel::<T, K, EXCLUSIVE, REVERSE>(
source.get(),
so,
ss,
target.get(),
d_o,
ds,
n,
width,
)
}
}
}
}
fn unsupported(dtype: KernelDType) -> StridedError {
StridedError::UnsupportedDType {
dtype: dtype.label(),
}
}
fn check_scan_dtype(dtype: KernelDType) -> Result<()> {
match dtype {
KernelDType::F32
| KernelDType::F64
| KernelDType::I32
| KernelDType::I64
| KernelDType::C32
| KernelDType::C64 => Ok(()),
_ => Err(unsupported(dtype)),
}
}
pub(super) trait ScanScalar: KernelStorageElement + MaybeSendSync {
fn zero() -> Self;
fn one() -> Self;
fn scan_add(lhs: Self, rhs: Self) -> Self;
fn scan_mul(lhs: Self, rhs: Self) -> Self;
}
macro_rules! impl_scan_scalar {
($add:ident, $mul:ident; $($ty:ty => $zero:expr, $one:expr),* $(,)?) => {$(
impl ScanScalar for $ty {
#[inline(always)]
fn zero() -> Self { $zero }
#[inline(always)]
fn one() -> Self { $one }
#[inline(always)]
fn scan_add(lhs: Self, rhs: Self) -> Self { $add(lhs, rhs) }
#[inline(always)]
fn scan_mul(lhs: Self, rhs: Self) -> Self { $mul(lhs, rhs) }
}
)*};
}
#[inline(always)]
fn plain_add<T: core::ops::Add<Output = T>>(lhs: T, rhs: T) -> T {
lhs + rhs
}
#[inline(always)]
fn plain_mul<T: core::ops::Mul<Output = T>>(lhs: T, rhs: T) -> T {
lhs * rhs
}
trait WrappingArith {
fn wadd(self, rhs: Self) -> Self;
fn wmul(self, rhs: Self) -> Self;
}
impl WrappingArith for i32 {
#[inline(always)]
fn wadd(self, rhs: Self) -> Self {
self.wrapping_add(rhs)
}
#[inline(always)]
fn wmul(self, rhs: Self) -> Self {
self.wrapping_mul(rhs)
}
}
impl WrappingArith for i64 {
#[inline(always)]
fn wadd(self, rhs: Self) -> Self {
self.wrapping_add(rhs)
}
#[inline(always)]
fn wmul(self, rhs: Self) -> Self {
self.wrapping_mul(rhs)
}
}
#[inline(always)]
fn wrapping_add<T: WrappingArith>(lhs: T, rhs: T) -> T {
lhs.wadd(rhs)
}
#[inline(always)]
fn wrapping_mul<T: WrappingArith>(lhs: T, rhs: T) -> T {
lhs.wmul(rhs)
}
impl_scan_scalar!(plain_add, plain_mul;
f32 => 0.0, 1.0,
f64 => 0.0, 1.0,
Complex32 => Complex32::new(0.0, 0.0), Complex32::new(1.0, 0.0),
Complex64 => Complex64::new(0.0, 0.0), Complex64::new(1.0, 0.0),
);
impl_scan_scalar!(wrapping_add, wrapping_mul;
i32 => 0, 1,
i64 => 0, 1,
);
pub(super) trait ScanKernel<T>: 'static {
fn identity() -> T;
fn combine(acc: T, value: T) -> T;
}
pub(super) struct SumScan;
pub(super) struct ProductScan;
impl<T: ScanScalar> ScanKernel<T> for SumScan {
#[inline(always)]
fn identity() -> T {
T::zero()
}
#[inline(always)]
fn combine(acc: T, value: T) -> T {
T::scan_add(acc, value)
}
}
impl<T: ScanScalar> ScanKernel<T> for ProductScan {
#[inline(always)]
fn identity() -> T {
T::one()
}
#[inline(always)]
fn combine(acc: T, value: T) -> T {
T::scan_mul(acc, value)
}
}
#[inline(always)]
fn scan_order(base: isize, stride: isize, n: usize, reverse: bool) -> (isize, isize) {
if reverse {
(base + (n as isize - 1) * stride, -stride)
} else {
(base, stride)
}
}
#[inline(always)]
unsafe fn scan_line<T, K, const EXCLUSIVE: bool, const REVERSE: bool>(
src: *const T,
so: isize,
ss: isize,
dst: *mut T,
d_o: isize,
ds: isize,
n: usize,
) where
T: ScanScalar,
K: ScanKernel<T>,
{
let (mut s, ss) = scan_order(so, ss, n, REVERSE);
let (mut d, ds) = scan_order(d_o, ds, n, REVERSE);
let mut acc = K::identity();
for _ in 0..n {
unsafe {
let value = src.offset(s).read();
if EXCLUSIVE {
dst.offset(d).write(acc);
acc = K::combine(acc, value);
} else {
acc = K::combine(acc, value);
dst.offset(d).write(acc);
}
}
s += ss;
d += ds;
}
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
unsafe fn scan_panel<T, K, const EXCLUSIVE: bool, const REVERSE: bool>(
src: *const T,
so: isize,
ss: isize,
dst: *mut T,
d_o: isize,
ds: isize,
n: usize,
width: usize,
) where
T: ScanScalar,
K: ScanKernel<T>,
{
debug_assert!(width <= PANEL);
let (mut s, ss) = scan_order(so, ss, n, REVERSE);
let (mut d, ds) = scan_order(d_o, ds, n, REVERSE);
let mut acc = [K::identity(); PANEL];
let acc = &mut acc[..width];
for _ in 0..n {
unsafe {
let input = core::slice::from_raw_parts(src.offset(s), width);
let output = dst.offset(d);
for (lane, (acc, &value)) in acc.iter_mut().zip(input).enumerate() {
if EXCLUSIVE {
output.add(lane).write(*acc);
*acc = K::combine(*acc, value);
} else {
*acc = K::combine(*acc, value);
output.add(lane).write(*acc);
}
}
}
s += ss;
d += ds;
}
}