#![allow(unsafe_op_in_unsafe_fn, clippy::missing_safety_doc)]
use core::arch::x86_64::*;
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn dot_product_16x_f32_avx512(weights: &[[f32; 16]], state: &[f32]) -> [f32; 16] {
let len = state.len();
let mut acc0 = _mm512_setzero_ps();
let mut acc1 = _mm512_setzero_ps();
let mut i = 0;
unsafe {
while i + 2 <= len {
let w0 = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
let s0 = _mm512_set1_ps(*state.get_unchecked(i));
acc0 = _mm512_fmadd_ps(w0, s0, acc0);
let w1 = _mm512_loadu_ps(weights.as_ptr().add(i + 1) as *const f32);
let s1 = _mm512_set1_ps(*state.get_unchecked(i + 1));
acc1 = _mm512_fmadd_ps(w1, s1, acc1);
i += 2;
}
while i < len {
let s = _mm512_set1_ps(*state.get_unchecked(i));
let w = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
acc0 = _mm512_fmadd_ps(w, s, acc0);
i += 1;
}
acc0 = _mm512_add_ps(acc0, acc1);
let mut out = [0.0f32; 16];
_mm512_storeu_ps(out.as_mut_ptr(), acc0);
out
}
}
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn dot_product_16x_f32_dual_avx512(
weights: &[[f32; 16]],
state_f0: &[f32],
state_f1: &[f32],
) -> ([f32; 16], [f32; 16]) {
let len = core::cmp::min(
weights.len(),
core::cmp::min(state_f0.len(), state_f1.len()),
);
let mut acc_f0_0 = _mm512_setzero_ps();
let mut acc_f0_1 = _mm512_setzero_ps();
let mut acc_f1_0 = _mm512_setzero_ps();
let mut acc_f1_1 = _mm512_setzero_ps();
let mut i = 0;
unsafe {
while i + 2 <= len {
let w0 = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
let s_f0_0 = _mm512_set1_ps(*state_f0.get_unchecked(i));
let s_f1_0 = _mm512_set1_ps(*state_f1.get_unchecked(i));
acc_f0_0 = _mm512_fmadd_ps(w0, s_f0_0, acc_f0_0);
acc_f1_0 = _mm512_fmadd_ps(w0, s_f1_0, acc_f1_0);
let w1 = _mm512_loadu_ps(weights.as_ptr().add(i + 1) as *const f32);
let s_f0_1 = _mm512_set1_ps(*state_f0.get_unchecked(i + 1));
let s_f1_1 = _mm512_set1_ps(*state_f1.get_unchecked(i + 1));
acc_f0_1 = _mm512_fmadd_ps(w1, s_f0_1, acc_f0_1);
acc_f1_1 = _mm512_fmadd_ps(w1, s_f1_1, acc_f1_1);
i += 2;
}
while i < len {
let w = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
let s_f0 = _mm512_set1_ps(*state_f0.get_unchecked(i));
let s_f1 = _mm512_set1_ps(*state_f1.get_unchecked(i));
acc_f0_0 = _mm512_fmadd_ps(w, s_f0, acc_f0_0);
acc_f1_0 = _mm512_fmadd_ps(w, s_f1, acc_f1_0);
i += 1;
}
acc_f0_0 = _mm512_add_ps(acc_f0_0, acc_f0_1);
acc_f1_0 = _mm512_add_ps(acc_f1_0, acc_f1_1);
let mut out_f0 = [0.0f32; 16];
let mut out_f1 = [0.0f32; 16];
_mm512_storeu_ps(out_f0.as_mut_ptr(), acc_f0_0);
_mm512_storeu_ps(out_f1.as_mut_ptr(), acc_f1_0);
(out_f0, out_f1)
}
}
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn dot_product_16x_f32_accumulate_avx512(
weights: &[[f32; 16]],
state: &[f32],
init: &[f32; 16],
) -> [f32; 16] {
let len = state.len();
let mut acc0 = _mm512_loadu_ps(init.as_ptr());
let mut acc1 = _mm512_setzero_ps();
let mut i = 0;
unsafe {
while i + 2 <= len {
let w0 = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
let s0 = _mm512_set1_ps(*state.get_unchecked(i));
acc0 = _mm512_fmadd_ps(w0, s0, acc0);
let w1 = _mm512_loadu_ps(weights.as_ptr().add(i + 1) as *const f32);
let s1 = _mm512_set1_ps(*state.get_unchecked(i + 1));
acc1 = _mm512_fmadd_ps(w1, s1, acc1);
i += 2;
}
while i < len {
let s = _mm512_set1_ps(*state.get_unchecked(i));
let w = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
acc0 = _mm512_fmadd_ps(w, s, acc0);
i += 1;
}
acc0 = _mm512_add_ps(acc0, acc1);
let mut out = [0.0f32; 16];
_mm512_storeu_ps(out.as_mut_ptr(), acc0);
out
}
}
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn dot_product_16x_f32_dual_accumulate_avx512(
weights: &[[f32; 16]],
state_f0: &[f32],
state_f1: &[f32],
init_f0: &[f32; 16],
init_f1: &[f32; 16],
) -> ([f32; 16], [f32; 16]) {
let len = core::cmp::min(
weights.len(),
core::cmp::min(state_f0.len(), state_f1.len()),
);
let mut acc_f0_0 = _mm512_loadu_ps(init_f0.as_ptr());
let mut acc_f0_1 = _mm512_setzero_ps();
let mut acc_f1_0 = _mm512_loadu_ps(init_f1.as_ptr());
let mut acc_f1_1 = _mm512_setzero_ps();
let mut i = 0;
unsafe {
while i + 2 <= len {
let w0 = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
let s_f0_0 = _mm512_set1_ps(*state_f0.get_unchecked(i));
let s_f1_0 = _mm512_set1_ps(*state_f1.get_unchecked(i));
acc_f0_0 = _mm512_fmadd_ps(w0, s_f0_0, acc_f0_0);
acc_f1_0 = _mm512_fmadd_ps(w0, s_f1_0, acc_f1_0);
let w1 = _mm512_loadu_ps(weights.as_ptr().add(i + 1) as *const f32);
let s_f0_1 = _mm512_set1_ps(*state_f0.get_unchecked(i + 1));
let s_f1_1 = _mm512_set1_ps(*state_f1.get_unchecked(i + 1));
acc_f0_1 = _mm512_fmadd_ps(w1, s_f0_1, acc_f0_1);
acc_f1_1 = _mm512_fmadd_ps(w1, s_f1_1, acc_f1_1);
i += 2;
}
while i < len {
let w = _mm512_loadu_ps(weights.as_ptr().add(i) as *const f32);
let s_f0 = _mm512_set1_ps(*state_f0.get_unchecked(i));
let s_f1 = _mm512_set1_ps(*state_f1.get_unchecked(i));
acc_f0_0 = _mm512_fmadd_ps(w, s_f0, acc_f0_0);
acc_f1_0 = _mm512_fmadd_ps(w, s_f1, acc_f1_0);
i += 1;
}
acc_f0_0 = _mm512_add_ps(acc_f0_0, acc_f0_1);
acc_f1_0 = _mm512_add_ps(acc_f1_0, acc_f1_1);
let mut out_f0 = [0.0f32; 16];
let mut out_f1 = [0.0f32; 16];
_mm512_storeu_ps(out_f0.as_mut_ptr(), acc_f0_0);
_mm512_storeu_ps(out_f1.as_mut_ptr(), acc_f1_0);
(out_f0, out_f1)
}
}