use crate::kernel::SimdKernel;
use crate::ops::elementwise::Clamp;
use crate::scalar::{NumericElement, Scalar};
pub trait UnaryOp<T: Scalar>: crate::private::Sealed + Copy + 'static {
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector;
fn apply_scalar(self, a: T) -> T;
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Abs;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Neg;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Sqrt;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RecipSqrt;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Popcount;
impl crate::private::Sealed for Abs {}
impl crate::private::Sealed for Neg {}
impl crate::private::Sealed for Sqrt {}
impl crate::private::Sealed for RecipSqrt {}
impl crate::private::Sealed for Popcount {}
impl<T: Scalar> UnaryOp<T> for Abs {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector {
Arch::abs(v)
}
#[inline(always)]
fn apply_scalar(self, a: T) -> T {
a.abs()
}
}
impl<T: Scalar> UnaryOp<T> for Neg {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector {
Arch::neg(v)
}
#[inline(always)]
fn apply_scalar(self, a: T) -> T {
T::ZERO - a
}
}
impl<T: Scalar> UnaryOp<T> for Sqrt {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector {
Arch::sqrt(v)
}
#[inline(always)]
fn apply_scalar(self, a: T) -> T {
a.sqrt()
}
}
impl<T: Scalar> UnaryOp<T> for RecipSqrt {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector {
Arch::recip_sqrt(v)
}
#[inline(always)]
fn apply_scalar(self, a: T) -> T {
T::ONE / a.sqrt()
}
}
impl<T: Scalar + PartialOrd + NumericElement> UnaryOp<T> for Clamp<T> {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector {
let lo_vec = Arch::splat(self.lo);
let hi_vec = Arch::splat(self.hi);
let clamped_hi = Arch::min(v, hi_vec);
Arch::max(clamped_hi, lo_vec)
}
#[inline(always)]
fn apply_scalar(self, a: T) -> T {
a.min_scalar(self.hi).max_scalar(self.lo)
}
}
impl<T: Scalar> UnaryOp<T> for Popcount {
#[inline(always)]
unsafe fn apply<Arch: SimdKernel<T>>(self, v: Arch::Vector) -> Arch::Vector {
Arch::popcount(v)
}
#[inline(always)]
fn apply_scalar(self, a: T) -> T {
T::cast_from(a.count_ones() as i32)
}
}