use crate::kernel::SimdKernel;
use crate::scalar::Scalar;
pub trait ElementOp<T: Scalar>: crate::private::Sealed + Copy + 'static {
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector;
fn apply_scalar(self, a: T, b: T) -> T;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Mul;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Add;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Sub;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Div;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BitAnd;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BitOr;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BitXor;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FmaAdd;
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Clamp<T: Copy> {
pub lo: T,
pub hi: T,
}
impl<T: Copy> Clamp<T> {
#[inline(always)]
pub fn new(lo: T, hi: T) -> Self {
Self { lo, hi }
}
}
impl crate::private::Sealed for Mul {}
impl crate::private::Sealed for Add {}
impl crate::private::Sealed for Sub {}
impl crate::private::Sealed for Div {}
impl crate::private::Sealed for BitAnd {}
impl crate::private::Sealed for BitOr {}
impl crate::private::Sealed for BitXor {}
impl crate::private::Sealed for FmaAdd {}
impl<T: Copy + 'static> crate::private::Sealed for Clamp<T> {}
impl<T: Scalar> ElementOp<T> for Mul {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::mul(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a * b
}
}
impl<T: Scalar> ElementOp<T> for Add {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::add(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a + b
}
}
impl<T: Scalar> ElementOp<T> for Sub {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::sub(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a - b
}
}
impl<T: Scalar> ElementOp<T> for Div {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::div(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a / b
}
}
impl<T: Scalar> ElementOp<T> for BitAnd {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::bitand(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a.bitand(b)
}
}
impl<T: Scalar> ElementOp<T> for BitOr {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::bitor(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a.bitor(b)
}
}
impl<T: Scalar> ElementOp<T> for BitXor {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
Arch::bitxor(a, b)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a.bitxor(b)
}
}
impl<T: Scalar> ElementOp<T> for FmaAdd {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, b: Arch::Vector) -> Arch::Vector {
let zero = Arch::zero();
Arch::fmadd(a, b, zero)
}
#[inline(always)]
fn apply_scalar(self, a: T, b: T) -> T {
a.scalar_fmadd(b, T::ZERO)
}
}
impl<T: Scalar + Copy> ElementOp<T> for Clamp<T> {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, a: Arch::Vector, _b: Arch::Vector) -> Arch::Vector {
let lo_vec = Arch::splat(self.lo);
let hi_vec = Arch::splat(self.hi);
Arch::min(Arch::max(a, lo_vec), hi_vec)
}
#[inline(always)]
fn apply_scalar(self, a: T, _b: T) -> T {
a.max_scalar(self.lo).min_scalar(self.hi)
}
}