use core::arch::wasm32::*;
#[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 Simd128;
impl Simd for Simd128 {
#[inline(always)]
unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
#[target_feature(enable = "simd128")]
fn inner<R>(f: impl FnOnce() -> R) -> R {
f()
}
inner(f)
}
}
impl SimdOps<f32> for Simd128 {
type Reg = v128;
const LANES: usize = 4;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
f32x4_splat(0.0)
}
#[inline(always)]
unsafe fn splat(self, v: f32) -> Self::Reg {
f32x4_splat(v)
}
#[inline(always)]
unsafe fn loadu(self, p: *const f32) -> Self::Reg {
unsafe { v128_load(p as *const v128) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f32, v: Self::Reg) {
unsafe { v128_store(p as *mut v128, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f32x4_mul(a, b)
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f32x4_add(a, b)
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
f32x4_add(f32x4_mul(a, b), c)
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
f32x4_sub(c, f32x4_mul(a, b))
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f32x4_pmax(b, a)
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f32x4_pmin(b, a)
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f32 {
f32x4_extract_lane::<0>(v)
+ f32x4_extract_lane::<1>(v)
+ f32x4_extract_lane::<2>(v)
+ f32x4_extract_lane::<3>(v)
}
}
impl SimdOps<f64> for Simd128 {
type Reg = v128;
const LANES: usize = 2;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
f64x2_splat(0.0)
}
#[inline(always)]
unsafe fn splat(self, v: f64) -> Self::Reg {
f64x2_splat(v)
}
#[inline(always)]
unsafe fn loadu(self, p: *const f64) -> Self::Reg {
unsafe { v128_load(p as *const v128) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f64, v: Self::Reg) {
unsafe { v128_store(p as *mut v128, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f64x2_mul(a, b)
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f64x2_add(a, b)
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
f64x2_add(f64x2_mul(a, b), c)
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
f64x2_sub(c, f64x2_mul(a, b))
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f64x2_pmax(b, a)
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
f64x2_pmin(b, a)
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f64 {
f64x2_extract_lane::<0>(v) + f64x2_extract_lane::<1>(v)
}
}
#[cfg(feature = "half")]
impl KernelSimd<f16, f16, f32, f16> for Simd128 {
#[inline(always)]
unsafe fn load_lhs(self, p: *const f16) -> v128 {
unsafe {
let a = [
(*p).widen(),
(*p.add(1)).widen(),
(*p.add(2)).widen(),
(*p.add(3)).widen(),
];
v128_load(a.as_ptr() as *const v128)
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: f16) -> v128 {
f32x4_splat(v.widen())
}
#[inline(always)]
unsafe fn load_out(self, p: *const f16) -> v128 {
unsafe { <Self as KernelSimd<f16, f16, f32, f16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut f16, v: v128) {
unsafe {
let mut t = [0.0f32; 4];
v128_store(t.as_mut_ptr() as *mut v128, v);
for (i, &x) in t.iter().enumerate() {
*p.add(i) = f16::narrow(x);
}
}
}
}
#[cfg(feature = "half")]
impl KernelSimd<bf16, bf16, f32, bf16> for Simd128 {
#[inline(always)]
unsafe fn load_lhs(self, p: *const bf16) -> v128 {
unsafe {
let a = [
(*p).widen(),
(*p.add(1)).widen(),
(*p.add(2)).widen(),
(*p.add(3)).widen(),
];
v128_load(a.as_ptr() as *const v128)
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: bf16) -> v128 {
f32x4_splat(v.widen())
}
#[inline(always)]
unsafe fn load_out(self, p: *const bf16) -> v128 {
unsafe { <Self as KernelSimd<bf16, bf16, f32, bf16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut bf16, v: v128) {
unsafe {
let mut t = [0.0f32; 4];
v128_store(t.as_mut_ptr() as *mut v128, v);
for (i, &x) in t.iter().enumerate() {
*p.add(i) = bf16::narrow(x);
}
}
}
}
#[cfg(feature = "int8")]
impl SimdOps<i32> for Simd128 {
type Reg = v128;
const LANES: usize = 4;
#[inline(always)]
unsafe fn zero(self) -> v128 {
i32x4_splat(0)
}
#[inline(always)]
unsafe fn splat(self, v: i32) -> v128 {
i32x4_splat(v)
}
#[inline(always)]
unsafe fn loadu(self, p: *const i32) -> v128 {
unsafe { v128_load(p as *const v128) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut i32, v: v128) {
unsafe { v128_store(p as *mut v128, v) }
}
#[inline(always)]
unsafe fn mul(self, a: v128, b: v128) -> v128 {
i32x4_mul(a, b)
}
#[inline(always)]
unsafe fn add(self, a: v128, b: v128) -> v128 {
i32x4_add(a, b)
}
#[inline(always)]
unsafe fn mul_add(self, a: v128, b: v128, c: v128) -> v128 {
i32x4_add(i32x4_mul(a, b), c)
}
#[inline(always)]
unsafe fn fnma(self, a: v128, b: v128, c: v128) -> v128 {
i32x4_sub(c, i32x4_mul(a, b))
}
#[inline(always)]
unsafe fn reduce_sum(self, v: v128) -> i32 {
i32x4_extract_lane::<0>(v)
.wrapping_add(i32x4_extract_lane::<1>(v))
.wrapping_add(i32x4_extract_lane::<2>(v))
.wrapping_add(i32x4_extract_lane::<3>(v))
}
}
#[cfg(feature = "int8")]
impl KernelSimd<i8, i8, i32, i32> for Simd128 {
#[inline(always)]
unsafe fn load_lhs(self, p: *const i8) -> v128 {
unsafe {
let a = [
*p as i32,
*p.add(1) as i32,
*p.add(2) as i32,
*p.add(3) as i32,
];
v128_load(a.as_ptr() as *const v128)
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: i8) -> v128 {
i32x4_splat(v as i32)
}
#[inline(always)]
unsafe fn load_out(self, p: *const i32) -> v128 {
unsafe { v128_load(p as *const v128) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut i32, v: v128) {
unsafe { v128_store(p as *mut v128, v) }
}
}
#[cfg(feature = "complex")]
impl_complex_simd!(Simd128, f32, v128, 4);
#[cfg(feature = "complex")]
impl_complex_simd!(Simd128, f64, v128, 2);