use crate::kernel::SimdKernel;
use crate::scalar::Scalar;
pub trait ReductionOp<T: Scalar>: crate::private::Sealed + Copy + 'static {
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector;
#[inline(always)]
unsafe fn fma_pair_accumulate<Arch: SimdKernel<T>>(
acc: Arch::Vector,
a: Arch::Vector,
b: Arch::Vector,
) -> Arch::Vector {
Self::accumulate::<Arch>(acc, Arch::mul(a, b))
}
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T;
fn identity_scalar() -> T;
fn scalar_combine(a: T, b: T) -> T;
#[inline(always)]
fn scalar_accumulate(acc: T, elem: T) -> T {
Self::scalar_combine(acc, elem)
}
#[inline(always)]
unsafe fn identity_vector<Arch: SimdKernel<T>>() -> Arch::Vector {
Arch::splat(Self::identity_scalar())
}
#[inline(always)]
unsafe fn transform_vector<Arch: SimdKernel<T>>(v: Arch::Vector) -> Arch::Vector {
v
}
#[inline(always)]
unsafe fn combine_vectors<Arch: SimdKernel<T>>(
a: Arch::Vector,
b: Arch::Vector,
) -> Arch::Vector {
Self::accumulate::<Arch>(a, b)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Sum;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Dot;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Min;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Max;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AbsSum;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AbsMax;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Product;
impl crate::private::Sealed for Sum {}
impl crate::private::Sealed for Dot {}
impl crate::private::Sealed for Min {}
impl crate::private::Sealed for Max {}
impl crate::private::Sealed for AbsSum {}
impl crate::private::Sealed for AbsMax {}
impl crate::private::Sealed for Product {}
impl<T: Scalar> ReductionOp<T> for Sum {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::add(acc, v)
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
Arch::sum_reduce(acc)
}
#[inline(always)]
fn identity_scalar() -> T {
T::ZERO
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a + b
}
}
impl<T: Scalar> ReductionOp<T> for Dot {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::add(acc, v)
}
#[inline(always)]
unsafe fn fma_pair_accumulate<Arch: SimdKernel<T>>(
acc: Arch::Vector,
a: Arch::Vector,
b: Arch::Vector,
) -> Arch::Vector {
Arch::fmadd(a, b, acc)
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
Arch::sum_reduce(acc)
}
#[inline(always)]
fn identity_scalar() -> T {
T::ZERO
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a + b
}
}
impl<T: Scalar> ReductionOp<T> for Min {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::min(acc, v)
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
Arch::min_reduce(acc)
}
#[inline(always)]
fn identity_scalar() -> T {
T::MAX_VALUE
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a.min_scalar(b)
}
}
impl<T: Scalar> ReductionOp<T> for Max {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::max(acc, v)
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
Arch::max_reduce(acc)
}
#[inline(always)]
fn identity_scalar() -> T {
T::MIN_VALUE
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a.max_scalar(b)
}
}
impl<T: Scalar> ReductionOp<T> for AbsSum {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::add(acc, Arch::abs(v))
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
Arch::sum_reduce(acc)
}
#[inline(always)]
fn identity_scalar() -> T {
T::ZERO
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a + b
}
#[inline(always)]
fn scalar_accumulate(acc: T, elem: T) -> T {
acc + elem.abs()
}
#[inline(always)]
unsafe fn transform_vector<Arch: SimdKernel<T>>(v: Arch::Vector) -> Arch::Vector {
Arch::abs(v)
}
#[inline(always)]
unsafe fn combine_vectors<Arch: SimdKernel<T>>(
a: Arch::Vector,
b: Arch::Vector,
) -> Arch::Vector {
Arch::add(a, b)
}
}
impl<T: Scalar> ReductionOp<T> for AbsMax {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::max(acc, Arch::abs(v))
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
Arch::max_reduce(acc)
}
#[inline(always)]
fn identity_scalar() -> T {
T::ZERO
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a.max_scalar(b)
}
#[inline(always)]
fn scalar_accumulate(acc: T, elem: T) -> T {
acc.max_scalar(elem.abs())
}
#[inline(always)]
unsafe fn transform_vector<Arch: SimdKernel<T>>(v: Arch::Vector) -> Arch::Vector {
Arch::abs(v)
}
#[inline(always)]
unsafe fn combine_vectors<Arch: SimdKernel<T>>(
a: Arch::Vector,
b: Arch::Vector,
) -> Arch::Vector {
Arch::max(a, b)
}
}
impl<T: Scalar> ReductionOp<T> for Product {
#[inline(always)]
unsafe fn accumulate<Arch: SimdKernel<T>>(acc: Arch::Vector, v: Arch::Vector) -> Arch::Vector {
Arch::mul(acc, v)
}
#[inline(always)]
unsafe fn finalize<Arch: SimdKernel<T>>(acc: Arch::Vector) -> T {
const { <Arch as SimdKernel<T>>::LANE_BOUND_CHECK };
let mut buf = [T::ZERO; crate::kernel::MAX_SIMD_LANES];
Arch::store_unaligned(buf.as_mut_ptr(), acc);
let mut result = T::ONE;
for i in 0..Arch::LANE_COUNT {
result = result * buf[i];
}
result
}
#[inline(always)]
fn identity_scalar() -> T {
T::ONE
}
#[inline(always)]
fn scalar_combine(a: T, b: T) -> T {
a * b
}
}