#[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;
#[cfg(feature = "int8")]
use super::VNNI_A_BIAS;
use super::{Simd, SimdOps};
#[derive(Copy, Clone, Default)]
pub struct Avx512F;
impl Simd for Avx512F {
#[inline(always)]
unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
#[target_feature(enable = "avx512f")]
unsafe fn inner<R>(f: impl FnOnce() -> R) -> R {
f()
}
unsafe { inner(f) }
}
}
impl SimdOps<f32> for Avx512F {
type Reg = __m512;
const LANES: usize = 16;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { _mm512_setzero_ps() }
}
#[inline(always)]
unsafe fn splat(self, v: f32) -> Self::Reg {
unsafe { _mm512_set1_ps(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const f32) -> Self::Reg {
unsafe { _mm512_loadu_ps(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f32, v: Self::Reg) {
unsafe { _mm512_storeu_ps(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_mul_ps(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_add_ps(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm512_fmadd_ps(a, b, c) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm512_fnmadd_ps(a, b, c) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_max_ps(a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_min_ps(a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f32 {
unsafe { _mm512_reduce_add_ps(v) }
}
}
impl SimdOps<f64> for Avx512F {
type Reg = __m512d;
const LANES: usize = 8;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { _mm512_setzero_pd() }
}
#[inline(always)]
unsafe fn splat(self, v: f64) -> Self::Reg {
unsafe { _mm512_set1_pd(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const f64) -> Self::Reg {
unsafe { _mm512_loadu_pd(p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut f64, v: Self::Reg) {
unsafe { _mm512_storeu_pd(p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_mul_pd(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_add_pd(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm512_fmadd_pd(a, b, c) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { _mm512_fnmadd_pd(a, b, c) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_max_pd(a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { _mm512_min_pd(a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> f64 {
unsafe { _mm512_reduce_add_pd(v) }
}
}
#[cfg(feature = "half")]
impl KernelSimd<f16, f16, f32, f16> for Avx512F {
#[inline(always)]
unsafe fn load_lhs(self, p: *const f16) -> __m512 {
unsafe { _mm512_cvtph_ps(_mm256_loadu_si256(p as *const __m256i)) }
}
#[inline(always)]
unsafe fn splat_rhs(self, v: f16) -> __m512 {
unsafe { _mm512_cvtph_ps(_mm256_set1_epi16(v.to_bits() as i16)) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const f16) -> __m512 {
unsafe { <Self as KernelSimd<f16, f16, f32, f16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut f16, v: __m512) {
unsafe {
let h = _mm512_cvtps_ph::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(v);
_mm256_storeu_si256(p as *mut __m256i, h);
}
}
}
#[cfg(feature = "half")]
impl KernelSimd<bf16, bf16, f32, bf16> for Avx512F {
#[inline(always)]
unsafe fn load_lhs(self, p: *const bf16) -> __m512 {
unsafe {
let w = _mm256_loadu_si256(p as *const __m256i); _mm512_castsi512_ps(_mm512_slli_epi32::<16>(_mm512_cvtepu16_epi32(w)))
}
}
#[inline(always)]
unsafe fn splat_rhs(self, v: bf16) -> __m512 {
unsafe { _mm512_set1_ps(f32::from_bits((v.to_bits() as u32) << 16)) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const bf16) -> __m512 {
unsafe { <Self as KernelSimd<bf16, bf16, f32, bf16>>::load_lhs(self, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut bf16, v: __m512) {
unsafe {
let bits = _mm512_castps_si512(v);
let lsb = _mm512_and_si512(_mm512_srli_epi32::<16>(bits), _mm512_set1_epi32(1));
let bias = _mm512_add_epi32(lsb, _mm512_set1_epi32(0x7FFF));
let rounded = _mm512_srli_epi32::<16>(_mm512_add_epi32(bits, bias));
let abs = _mm512_and_si512(bits, _mm512_set1_epi32(0x7FFF_FFFFu32 as i32));
let nan = _mm512_cmpgt_epi32_mask(abs, _mm512_set1_epi32(0x7F80_0000));
let nan_out = _mm512_or_si512(_mm512_srli_epi32::<16>(bits), _mm512_set1_epi32(0x0040));
let out = _mm512_mask_blend_epi32(nan, rounded, nan_out);
_mm256_storeu_si256(p as *mut __m256i, _mm512_cvtepi32_epi16(out));
}
}
}
#[cfg(feature = "int8")]
impl SimdOps<i32> for Avx512F {
type Reg = __m512i;
const LANES: usize = 16;
#[inline(always)]
unsafe fn zero(self) -> __m512i {
unsafe { _mm512_setzero_si512() }
}
#[inline(always)]
unsafe fn splat(self, v: i32) -> __m512i {
unsafe { _mm512_set1_epi32(v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const i32) -> __m512i {
unsafe { _mm512_loadu_si512(p as *const __m512i) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut i32, v: __m512i) {
unsafe { _mm512_storeu_si512(p as *mut __m512i, v) }
}
#[inline(always)]
unsafe fn mul(self, a: __m512i, b: __m512i) -> __m512i {
unsafe { _mm512_mullo_epi32(a, b) }
}
#[inline(always)]
unsafe fn add(self, a: __m512i, b: __m512i) -> __m512i {
unsafe { _mm512_add_epi32(a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: __m512i, b: __m512i, c: __m512i) -> __m512i {
unsafe { _mm512_add_epi32(_mm512_mullo_epi32(a, b), c) }
}
#[inline(always)]
unsafe fn fnma(self, a: __m512i, b: __m512i, c: __m512i) -> __m512i {
unsafe { _mm512_sub_epi32(c, _mm512_mullo_epi32(a, b)) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: __m512i) -> i32 {
unsafe { _mm512_reduce_add_epi32(v) }
}
}
#[cfg(feature = "int8")]
#[inline(always)]
unsafe fn requant_half_avx512f(
x: __m256i,
scale_v: __m512d,
zp_v: __m512d,
lo_v: __m512d,
hi_v: __m512d,
) -> __m256i {
unsafe {
let t = _mm512_cvtepi32_pd(x);
let t = _mm512_mul_pd(t, scale_v);
let t = _mm512_roundscale_pd::<{ _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC }>(t);
let u = _mm512_add_pd(t, zp_v);
let u = _mm512_max_pd(u, lo_v);
let u = _mm512_min_pd(u, hi_v);
_mm512_cvtpd_epi32(u)
}
}
#[cfg(feature = "int8")]
#[inline(always)]
unsafe fn requant_store_avx512f(dst: *mut i8, v: __m512i, scale: f64, zp: i32, lo: i32, hi: i32) {
unsafe {
let scale_v = _mm512_set1_pd(scale);
let zp_v = _mm512_set1_pd(zp as f64);
let lo_v = _mm512_set1_pd(lo as f64);
let hi_v = _mm512_set1_pd(hi as f64);
let lo8 = requant_half_avx512f(_mm512_castsi512_si256(v), scale_v, zp_v, lo_v, hi_v);
let hi8 =
requant_half_avx512f(_mm512_extracti64x4_epi64::<1>(v), scale_v, zp_v, lo_v, hi_v);
let combined = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(lo8), hi8);
_mm_storeu_si128(dst as *mut __m128i, _mm512_cvtepi32_epi8(combined));
}
}
#[cfg(feature = "int8")]
impl KernelSimd<i8, i8, i32, i32> for Avx512F {
#[inline(always)]
unsafe fn load_lhs(self, p: *const i8) -> __m512i {
unsafe { _mm512_cvtepi8_epi32(_mm_loadu_si128(p as *const __m128i)) }
}
#[inline(always)]
unsafe fn splat_rhs(self, v: i8) -> __m512i {
unsafe { _mm512_set1_epi32(v as i32) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const i32) -> __m512i {
unsafe { _mm512_loadu_si512(p as *const __m512i) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut i32, v: __m512i) {
unsafe { _mm512_storeu_si512(p as *mut __m512i, v) }
}
const REQUANT_VECTOR: bool = true;
#[inline(always)]
unsafe fn requant_store(self, dst: *mut i8, v: __m512i, scale: f64, zp: i32, lo: i32, hi: i32) {
unsafe { requant_store_avx512f(dst, v, scale, zp, lo, hi) }
}
}
#[cfg(any(feature = "int8", feature = "half"))]
macro_rules! delegate_simdops {
($tok:ty => $src:ty, $t:ty) => {
impl SimdOps<$t> for $tok {
type Reg = <$src as SimdOps<$t>>::Reg;
const LANES: usize = <$src as SimdOps<$t>>::LANES;
#[inline(always)]
unsafe fn zero(self) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::zero(<$src as Default>::default()) }
}
#[inline(always)]
unsafe fn splat(self, v: $t) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::splat(<$src as Default>::default(), v) }
}
#[inline(always)]
unsafe fn loadu(self, p: *const $t) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::loadu(<$src as Default>::default(), p) }
}
#[inline(always)]
unsafe fn storeu(self, p: *mut $t, v: Self::Reg) {
unsafe { <$src as SimdOps<$t>>::storeu(<$src as Default>::default(), p, v) }
}
#[inline(always)]
unsafe fn mul(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::mul(<$src as Default>::default(), a, b) }
}
#[inline(always)]
unsafe fn add(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::add(<$src as Default>::default(), a, b) }
}
#[inline(always)]
unsafe fn mul_add(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::mul_add(<$src as Default>::default(), a, b, c) }
}
#[inline(always)]
unsafe fn fnma(self, a: Self::Reg, b: Self::Reg, c: Self::Reg) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::fnma(<$src as Default>::default(), a, b, c) }
}
#[inline(always)]
unsafe fn max(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::max(<$src as Default>::default(), a, b) }
}
#[inline(always)]
unsafe fn min(self, a: Self::Reg, b: Self::Reg) -> Self::Reg {
unsafe { <$src as SimdOps<$t>>::min(<$src as Default>::default(), a, b) }
}
#[inline(always)]
unsafe fn reduce_sum(self, v: Self::Reg) -> $t {
unsafe { <$src as SimdOps<$t>>::reduce_sum(<$src as Default>::default(), v) }
}
}
};
}
#[cfg(feature = "int8")]
#[derive(Copy, Clone, Default)]
pub struct Avx512Vnni;
#[cfg(feature = "int8")]
impl Simd for Avx512Vnni {
#[inline(always)]
unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
unsafe fn inner<R>(f: impl FnOnce() -> R) -> R {
f()
}
unsafe { inner(f) }
}
}
#[cfg(feature = "int8")]
delegate_simdops!(Avx512Vnni => Avx512F, i32);
#[cfg(feature = "int8")]
impl KernelSimd<i8, i8, i32, i32> for Avx512Vnni {
#[inline(always)]
unsafe fn load_lhs(self, p: *const i8) -> __m512i {
unsafe { <Avx512F as KernelSimd<i8, i8, i32, i32>>::load_lhs(Avx512F, p) }
}
#[inline(always)]
unsafe fn splat_rhs(self, v: i8) -> __m512i {
unsafe { <Avx512F as KernelSimd<i8, i8, i32, i32>>::splat_rhs(Avx512F, v) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const i32) -> __m512i {
unsafe { <Avx512F as KernelSimd<i8, i8, i32, i32>>::load_out(Avx512F, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut i32, v: __m512i) {
unsafe { <Avx512F as KernelSimd<i8, i8, i32, i32>>::store_out(Avx512F, p, v) }
}
const REQUANT_VECTOR: bool = true;
#[inline(always)]
unsafe fn requant_store(self, dst: *mut i8, v: __m512i, scale: f64, zp: i32, lo: i32, hi: i32) {
unsafe { requant_store_avx512f(dst, v, scale, zp, lo, hi) }
}
#[allow(clippy::needless_range_loop)]
#[inline(always)]
unsafe fn dot_accumulate<const MR_REG: usize, const NR: usize>(
self,
kc: usize,
a: *const i8,
b: *const i8,
acc: &mut [[__m512i; MR_REG]; NR],
) {
unsafe {
let mr = MR_REG * 16;
let nquads = kc.div_ceil(4);
let mut colsum = [0i32; NR];
for q in 0..nquads {
for j in 0..NR {
let base = q * NR * 4 + j * 4;
let mut s = 0i32;
for t in 0..4 {
s += *b.add(base + t) as i32;
}
colsum[j] = colsum[j].wrapping_add(s);
}
}
for q in 0..nquads {
let a_regs: [__m512i; MR_REG] = core::array::from_fn(|i| {
_mm512_loadu_si512(a.add(q * mr * 4 + i * 64) as *const __m512i)
});
for j in 0..NR {
let bj = _mm512_set1_epi32(
(b.add(q * NR * 4 + j * 4) as *const i32).read_unaligned(),
);
for i in 0..MR_REG {
acc[j][i] = _mm512_dpbusd_epi32(acc[j][i], a_regs[i], bj);
}
}
}
for j in 0..NR {
let corr = _mm512_set1_epi32(VNNI_A_BIAS.wrapping_mul(colsum[j]));
for i in 0..MR_REG {
acc[j][i] = _mm512_sub_epi32(acc[j][i], corr);
}
}
}
}
}
#[cfg(feature = "half")]
#[derive(Copy, Clone, Default)]
pub struct Avx512Bf16;
#[cfg(feature = "half")]
impl Simd for Avx512Bf16 {
#[inline(always)]
unsafe fn vectorize<R>(self, f: impl FnOnce() -> R) -> R {
#[target_feature(enable = "avx512f,avx512bf16")]
unsafe fn inner<R>(f: impl FnOnce() -> R) -> R {
f()
}
unsafe { inner(f) }
}
}
#[cfg(feature = "half")]
delegate_simdops!(Avx512Bf16 => Avx512F, f32);
#[cfg(feature = "half")]
impl KernelSimd<bf16, bf16, f32, bf16> for Avx512Bf16 {
#[inline(always)]
unsafe fn load_lhs(self, p: *const bf16) -> __m512 {
unsafe { <Avx512F as KernelSimd<bf16, bf16, f32, bf16>>::load_lhs(Avx512F, p) }
}
#[inline(always)]
unsafe fn splat_rhs(self, v: bf16) -> __m512 {
unsafe { <Avx512F as KernelSimd<bf16, bf16, f32, bf16>>::splat_rhs(Avx512F, v) }
}
#[inline(always)]
unsafe fn load_out(self, p: *const bf16) -> __m512 {
unsafe { <Avx512F as KernelSimd<bf16, bf16, f32, bf16>>::load_out(Avx512F, p) }
}
#[inline(always)]
unsafe fn store_out(self, p: *mut bf16, v: __m512) {
unsafe { <Avx512F as KernelSimd<bf16, bf16, f32, bf16>>::store_out(Avx512F, p, v) }
}
#[allow(clippy::needless_range_loop)]
#[inline(always)]
unsafe fn dot_accumulate<const MR_REG: usize, const NR: usize>(
self,
kc: usize,
a: *const bf16,
b: *const bf16,
acc: &mut [[__m512; MR_REG]; NR],
) {
unsafe {
let mr = MR_REG * 16;
let npairs = kc.div_ceil(2);
for p2 in 0..npairs {
let a_regs: [__m512bh; MR_REG] = core::array::from_fn(|i| {
core::mem::transmute::<__m512i, __m512bh>(_mm512_loadu_si512(
a.add(p2 * mr * 2 + i * 32) as *const __m512i,
))
});
for j in 0..NR {
let bj = core::mem::transmute::<__m512i, __m512bh>(_mm512_set1_epi32(
(b.add(p2 * NR * 2 + j * 2) as *const i32).read_unaligned(),
));
for i in 0..MR_REG {
acc[j][i] = _mm512_dpbf16_ps(acc[j][i], a_regs[i], bj);
}
}
}
}
}
}
#[cfg(feature = "complex")]
impl_complex_simd!(Avx512F, f32, __m512, 16);
#[cfg(feature = "complex")]
impl_complex_simd!(Avx512F, f64, __m512d, 8);