#[cfg(target_arch = "x86")]
use core::arch::x86::*;
#[cfg(target_arch = "x86_64")]
use core::arch::x86_64::*;
#[cfg(feature = "half")]
use half::{bf16, f16};
#[cfg(any(feature = "half", feature = "int8"))]
use super::KernelSimd;
use super::{Simd, SimdOps};
#[derive(Copy, Clone, Default)]
pub struct Fma;
impl Simd for Fma {
#[inline(always)]
unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
#[target_feature(enable = "avx2,fma,f16c")]
unsafe fn inner<R>(f: impl FnOnce() -> R) -> R {
f()
}
unsafe { inner(f) }
}
}
impl SimdOps<f32> for Fma {
type Reg = __m256;
const LANES: usize = 8;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { _mm256_setzero_ps() }
}
#[inline(always)]
unsafe fn splat(self, v: f32) -> Self::Reg {
unsafe { _mm256_set1_ps(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const f32) -> Self::Reg {
unsafe { _mm256_loadu_ps(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f32, v: Self::Reg) {
unsafe { _mm256_storeu_ps(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_mul_ps(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_add_ps(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm256_fmadd_ps(a, b, c) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm256_fnmadd_ps(a, b, c) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_max_ps(a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_min_ps(a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f32 {
unsafe {
let hi = _mm256_extractf128_ps(v, 1);
let lo = _mm256_castps256_ps128(v);
let s = _mm_add_ps(lo, hi); let shuf = _mm_movehdup_ps(s); let sums = _mm_add_ps(s, shuf); let hi2 = _mm_movehl_ps(shuf, sums); let r = _mm_add_ss(sums, hi2);
_mm_cvtss_f32(r)
}
}
}
impl SimdOps<f64> for Fma {
type Reg = __m256d;
const LANES: usize = 4;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { _mm256_setzero_pd() }
}
#[inline(always)]
unsafe fn splat(self, v: f64) -> Self::Reg {
unsafe { _mm256_set1_pd(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const f64) -> Self::Reg {
unsafe { _mm256_loadu_pd(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f64, v: Self::Reg) {
unsafe { _mm256_storeu_pd(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_mul_pd(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_add_pd(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm256_fmadd_pd(a, b, c) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm256_fnmadd_pd(a, b, c) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_max_pd(a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm256_min_pd(a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f64 {
unsafe {
let hi = _mm256_extractf128_pd(v, 1);
let lo = _mm256_castpd256_pd128(v);
let s = _mm_add_pd(lo, hi); let sh = _mm_unpackhi_pd(s, s); let r = _mm_add_sd(s, sh);
_mm_cvtsd_f64(r)
}
}
}
#[cfg(feature = "half")]
impl KernelSimd<f16, f16, f32, f16> for Fma {
#[inline(always)]
unsafe fn load_lhs(self, p: *const f16) -> __m256 {
unsafe { _mm256_cvtph_ps(_mm_loadu_si128(p as *const __m128i)) }
}
#[inline(always)]
unsafe fn splat_rhs(self, v: f16) -> __m256 {
unsafe {
let lo = _mm_cvtph_ps(_mm_cvtsi32_si128(v.to_bits() as i32)); _mm256_broadcastss_ps(lo)
}
}
#[inline(always)]
unsafe fn load_out(self, p: *const f16) -> __m256 {
unsafe { <Self as KernelSimd<f16, f16, f32, f16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut f16, v: __m256) {
unsafe {
let h = _mm256_cvtps_ph::<_MM_FROUND_TO_NEAREST_INT>(v);
_mm_storeu_si128(p as *mut __m128i, h);
}
}
}
#[cfg(feature = "half")]
impl KernelSimd<bf16, bf16, f32, bf16> for Fma {
#[inline(always)]
unsafe fn load_lhs(self, p: *const bf16) -> __m256 {
unsafe {
let w = _mm_loadu_si128(p as *const __m128i); _mm256_castsi256_ps(_mm256_slli_epi32::<16>(_mm256_cvtepu16_epi32(w)))
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: bf16) -> __m256 {
unsafe { _mm256_set1_ps(f32::from_bits((v.to_bits() as u32) << 16)) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const bf16) -> __m256 {
unsafe { <Self as KernelSimd<bf16, bf16, f32, bf16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut bf16, v: __m256) {
unsafe {
let bits = _mm256_castps_si256(v);
let lsb = _mm256_and_si256(_mm256_srli_epi32::<16>(bits), _mm256_set1_epi32(1));
let bias = _mm256_add_epi32(lsb, _mm256_set1_epi32(0x7FFF));
let rounded = _mm256_srli_epi32::<16>(_mm256_add_epi32(bits, bias));
let abs = _mm256_and_si256(bits, _mm256_set1_epi32(0x7FFF_FFFFu32 as i32));
let is_nan = _mm256_cmpgt_epi32(abs, _mm256_set1_epi32(0x7F80_0000));
let nan_out = _mm256_or_si256(_mm256_srli_epi32::<16>(bits), _mm256_set1_epi32(0x0040));
let out = _mm256_blendv_epi8(rounded, nan_out, is_nan);
let lo = _mm256_castsi256_si128(out);
let hi = _mm256_extracti128_si256::<1>(out);
_mm_storeu_si128(p as *mut __m128i, _mm_packus_epi32(lo, hi));
}
}
}
#[cfg(feature = "int8")]
impl SimdOps<i32> for Fma {
type Reg = __m256i;
const LANES: usize = 8;
#[inline(always)]
unsafe fn zero(self) -> __m256i {
unsafe { _mm256_setzero_si256() }
}
#[inline(always)]
unsafe fn splat(self, v: i32) -> __m256i {
unsafe { _mm256_set1_epi32(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const i32) -> __m256i {
unsafe { _mm256_loadu_si256(p as *const __m256i) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut i32, v: __m256i) {
unsafe { _mm256_storeu_si256(p as *mut __m256i, v) }
}
#[inline(always)]
unsafe fn mul(self, a: __m256i, b: __m256i) -> __m256i {
unsafe { _mm256_mullo_epi32(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: __m256i, b: __m256i) -> __m256i {
unsafe { _mm256_add_epi32(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: __m256i, b: __m256i, c: __m256i) -> __m256i {
unsafe { _mm256_add_epi32(_mm256_mullo_epi32(a, b), c) }
}
#[inline(always)]
unsafe fn fnma(self, a: __m256i, b: __m256i, c: __m256i) -> __m256i {
unsafe { _mm256_sub_epi32(c, _mm256_mullo_epi32(a, b)) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: __m256i) -> i32 {
unsafe {
let hi = _mm256_extracti128_si256::<1>(v);
let lo = _mm256_castsi256_si128(v);
let s = _mm_add_epi32(lo, hi); let sh = _mm_shuffle_epi32::<0b01_00_11_10>(s); let s = _mm_add_epi32(s, sh); let sh = _mm_shuffle_epi32::<0b00_00_00_01>(s); _mm_cvtsi128_si32(_mm_add_epi32(s, sh))
}
}
}
#[cfg(feature = "int8")]
#[inline(always)]
unsafe fn requant_quad_fma(
x: __m128i,
scale_v: __m256d,
zp_v: __m256d,
lo_v: __m256d,
hi_v: __m256d,
) -> __m128i {
unsafe {
let t = _mm256_cvtepi32_pd(x);
let t = _mm256_mul_pd(t, scale_v);
let t = _mm256_round_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(t);
let u = _mm256_add_pd(t, zp_v);
let u = _mm256_max_pd(u, lo_v);
let u = _mm256_min_pd(u, hi_v);
_mm256_cvtpd_epi32(u)
}
}
#[cfg(feature = "int8")]
#[inline(always)]
unsafe fn requant_store_fma(dst: *mut i8, v: __m256i, scale: f64, zp: i32, lo: i32, hi: i32) {
unsafe {
let scale_v = _mm256_set1_pd(scale);
let zp_v = _mm256_set1_pd(zp as f64);
let lo_v = _mm256_set1_pd(lo as f64);
let hi_v = _mm256_set1_pd(hi as f64);
let i_lo = requant_quad_fma(_mm256_castsi256_si128(v), scale_v, zp_v, lo_v, hi_v);
let i_hi = requant_quad_fma(_mm256_extracti128_si256::<1>(v), scale_v, zp_v, lo_v, hi_v);
let mask = _mm_set_epi8(-1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 12, 8, 4, 0);
let lo_u32 = _mm_cvtsi128_si32(_mm_shuffle_epi8(i_lo, mask)) as u32;
let hi_u32 = _mm_cvtsi128_si32(_mm_shuffle_epi8(i_hi, mask)) as u32;
let packed = (lo_u32 as u64) | ((hi_u32 as u64) << 32);
core::ptr::write_unaligned(dst as *mut u64, packed);
}
}
#[cfg(feature = "int8")]
impl KernelSimd<i8, i8, i32, i32> for Fma {
#[inline(always)]
unsafe fn load_lhs(self, p: *const i8) -> __m256i {
unsafe { _mm256_cvtepi8_epi32(_mm_loadl_epi64(p as *const __m128i)) }
}
#[inline(always)]
unsafe fn splat_rhs(self, v: i8) -> __m256i {
unsafe { _mm256_set1_epi32(v as i32) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const i32) -> __m256i {
unsafe { _mm256_loadu_si256(p as *const __m256i) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut i32, v: __m256i) {
unsafe { _mm256_storeu_si256(p as *mut __m256i, v) }
}
const REQUANT_VECTOR: bool = true;
#[inline(always)]
unsafe fn requant_store(self, dst: *mut i8, v: __m256i, scale: f64, zp: i32, lo: i32, hi: i32) {
unsafe { requant_store_fma(dst, v, scale, zp, lo, hi) }
}
}
#[cfg(feature = "complex")]
impl_complex_simd!(Fma, f32, __m256, 8);
#[cfg(feature = "complex")]
impl_complex_simd!(Fma, f64, __m256d, 4);