use core::arch::aarch64::*;
#[cfg(feature = "half")]
use half::{bf16, f16};
#[cfg(any(feature = "half", feature = "int8"))]
use super::KernelSimd;
use super::{Simd, SimdOps};
#[cfg(feature = "half")]
use crate::scalar::NarrowFloat;
#[derive(Copy, Clone, Default)]
pub struct Neon;
impl Simd for Neon {
#[inline(always)]
unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
#[target_feature(enable = "neon")]
unsafe fn inner<R>(f: impl FnOnce() -> R) -> R {
f()
}
unsafe { inner(f) }
}
}
impl SimdOps<f32> for Neon {
type Reg = float32x4_t;
const LANES: usize = 4;
const LANE_FMA: bool = true;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { vdupq_n_f32(0.0) }
}
#[inline(always)]
unsafe fn splat(self, v: f32) -> Self::Reg {
unsafe { vdupq_n_f32(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const f32) -> Self::Reg {
unsafe { vld1q_f32(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f32, v: Self::Reg) {
unsafe { vst1q_f32(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vmulq_f32(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vaddq_f32(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { vfmaq_f32(c, a, b) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { vfmsq_f32(c, a, b) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vmaxnmq_f32(a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vminnmq_f32(a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f32 {
unsafe { vaddvq_f32(v) }
}
#[inline(always)]
unsafe fn fma_bvec<const MR_REG: usize>(
self,
a_regs: &[Self::Reg; MR_REG],
bvec: Self::Reg,
acc: &mut [[Self::Reg; MR_REG]],
) {
debug_assert_eq!(acc.len(), 4);
unsafe {
for i in 0..MR_REG {
acc[0][i] = vfmaq_laneq_f32::<0>(acc[0][i], a_regs[i], bvec);
}
for i in 0..MR_REG {
acc[1][i] = vfmaq_laneq_f32::<1>(acc[1][i], a_regs[i], bvec);
}
for i in 0..MR_REG {
acc[2][i] = vfmaq_laneq_f32::<2>(acc[2][i], a_regs[i], bvec);
}
for i in 0..MR_REG {
acc[3][i] = vfmaq_laneq_f32::<3>(acc[3][i], a_regs[i], bvec);
}
}
}
}
impl SimdOps<f64> for Neon {
type Reg = float64x2_t;
const LANES: usize = 2;
const LANE_FMA: bool = true;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { vdupq_n_f64(0.0) }
}
#[inline(always)]
unsafe fn splat(self, v: f64) -> Self::Reg {
unsafe { vdupq_n_f64(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const f64) -> Self::Reg {
unsafe { vld1q_f64(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f64, v: Self::Reg) {
unsafe { vst1q_f64(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vmulq_f64(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vaddq_f64(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { vfmaq_f64(c, a, b) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { vfmsq_f64(c, a, b) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vmaxnmq_f64(a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { vminnmq_f64(a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f64 {
unsafe { vaddvq_f64(v) }
}
#[inline(always)]
unsafe fn fma_bvec<const MR_REG: usize>(
self,
a_regs: &[Self::Reg; MR_REG],
bvec: Self::Reg,
acc: &mut [[Self::Reg; MR_REG]],
) {
debug_assert_eq!(acc.len(), 2);
unsafe {
for i in 0..MR_REG {
acc[0][i] = vfmaq_laneq_f64::<0>(acc[0][i], a_regs[i], bvec);
}
for i in 0..MR_REG {
acc[1][i] = vfmaq_laneq_f64::<1>(acc[1][i], a_regs[i], bvec);
}
}
}
}
#[cfg(feature = "half")]
impl KernelSimd<f16, f16, f32, f16> for Neon {
#[inline(always)]
unsafe fn load_lhs(self, p: *const f16) -> float32x4_t {
unsafe {
let a = [
(*p).widen(),
(*p.add(1)).widen(),
(*p.add(2)).widen(),
(*p.add(3)).widen(),
];
vld1q_f32(a.as_ptr())
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: f16) -> float32x4_t {
unsafe { vdupq_n_f32(v.widen()) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const f16) -> float32x4_t {
unsafe { <Self as KernelSimd<f16, f16, f32, f16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut f16, v: float32x4_t) {
unsafe {
let mut t = [0.0f32; 4];
vst1q_f32(t.as_mut_ptr(), v);
for (i, &x) in t.iter().enumerate() {
*p.add(i) = f16::narrow(x);
}
}
}
}
#[cfg(feature = "half")]
impl KernelSimd<bf16, bf16, f32, bf16> for Neon {
#[inline(always)]
unsafe fn load_lhs(self, p: *const bf16) -> float32x4_t {
unsafe {
let a = [
(*p).widen(),
(*p.add(1)).widen(),
(*p.add(2)).widen(),
(*p.add(3)).widen(),
];
vld1q_f32(a.as_ptr())
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: bf16) -> float32x4_t {
unsafe { vdupq_n_f32(v.widen()) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const bf16) -> float32x4_t {
unsafe { <Self as KernelSimd<bf16, bf16, f32, bf16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut bf16, v: float32x4_t) {
unsafe {
let mut t = [0.0f32; 4];
vst1q_f32(t.as_mut_ptr(), v);
for (i, &x) in t.iter().enumerate() {
*p.add(i) = bf16::narrow(x);
}
}
}
}
#[cfg(feature = "int8")]
impl SimdOps<i32> for Neon {
type Reg = int32x4_t;
const LANES: usize = 4;
#[inline(always)]
unsafe fn zero(self) -> int32x4_t {
unsafe { vdupq_n_s32(0) }
}
#[inline(always)]
unsafe fn splat(self, v: i32) -> int32x4_t {
unsafe { vdupq_n_s32(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const i32) -> int32x4_t {
unsafe { vld1q_s32(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut i32, v: int32x4_t) {
unsafe { vst1q_s32(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: int32x4_t, b: int32x4_t) -> int32x4_t {
unsafe { vmulq_s32(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: int32x4_t, b: int32x4_t) -> int32x4_t {
unsafe { vaddq_s32(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: int32x4_t, b: int32x4_t, c: int32x4_t) -> int32x4_t {
unsafe { vmlaq_s32(c, a, b) }
}
#[inline(always)]
unsafe fn fnma(self, a: int32x4_t, b: int32x4_t, c: int32x4_t) -> int32x4_t {
unsafe { vmlsq_s32(c, a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: int32x4_t) -> i32 {
unsafe { vaddvq_s32(v) }
}
}
#[cfg(feature = "int8")]
#[inline(always)]
unsafe fn requant_pair_neon(
x: int64x2_t,
scale_v: float64x2_t,
zp_v: float64x2_t,
lo_v: float64x2_t,
hi_v: float64x2_t,
) -> int64x2_t {
unsafe {
let t = vcvtq_f64_s64(x);
let t = vmulq_f64(t, scale_v);
let t = vrndnq_f64(t);
let u = vaddq_f64(t, zp_v);
let u = vmaxq_f64(u, lo_v);
let u = vminq_f64(u, hi_v);
vcvtq_s64_f64(u)
}
}
#[cfg(feature = "int8")]
#[inline(always)]
unsafe fn requant_store_neon(dst: *mut i8, v: int32x4_t, scale: f64, zp: i32, lo: i32, hi: i32) {
unsafe {
let scale_v = vdupq_n_f64(scale);
let zp_v = vdupq_n_f64(zp as f64);
let lo_v = vdupq_n_f64(lo as f64);
let hi_v = vdupq_n_f64(hi as f64);
let i_lo = requant_pair_neon(vmovl_s32(vget_low_s32(v)), scale_v, zp_v, lo_v, hi_v);
let i_hi = requant_pair_neon(vmovl_s32(vget_high_s32(v)), scale_v, zp_v, lo_v, hi_v);
let i32_all = vcombine_s32(vmovn_s64(i_lo), vmovn_s64(i_hi));
let idx: [u8; 16] = [
0, 4, 8, 12, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff,
];
let gathered = vqtbl1q_u8(vreinterpretq_u8_s32(i32_all), vld1q_u8(idx.as_ptr()));
let packed = vgetq_lane_u32::<0>(vreinterpretq_u32_u8(gathered));
core::ptr::write_unaligned(dst as *mut u32, packed);
}
}
#[cfg(feature = "int8")]
impl KernelSimd<i8, i8, i32, i32> for Neon {
#[inline(always)]
unsafe fn load_lhs(self, p: *const i8) -> int32x4_t {
unsafe {
let a = [
*p as i32,
*p.add(1) as i32,
*p.add(2) as i32,
*p.add(3) as i32,
];
vld1q_s32(a.as_ptr())
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: i8) -> int32x4_t {
unsafe { vdupq_n_s32(v as i32) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const i32) -> int32x4_t {
unsafe { vld1q_s32(p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut i32, v: int32x4_t) {
unsafe { vst1q_s32(p, v) }
}
const REQUANT_VECTOR: bool = true;
#[inline(always)]
unsafe fn requant_store(
self,
dst: *mut i8,
v: int32x4_t,
scale: f64,
zp: i32,
lo: i32,
hi: i32,
) {
unsafe { requant_store_neon(dst, v, scale, zp, lo, hi) }
}
}
#[cfg(feature = "complex")]
impl_complex_simd!(Neon, f32, float32x4_t, 4);
#[cfg(feature = "complex")]
impl_complex_simd!(Neon, f64, float64x2_t, 2);