use crate::repack::common::*;
use crate::repack::q8_0x4::*;
use crate::Q8_0_BLOCK_ELEMS;
use std::arch::x86_64::*;
use super::{bcast8, load_f16x4, rows4_from_pairs};
#[target_feature(enable = "avx2,fma")]
pub unsafe fn gemm_q8_0x4_q8_0_avx2(
packed: &[u8],
tile: &Q8ActsX4,
n_cols: usize,
out: &mut [f32],
) {
let nb = n_cols / Q8_0_BLOCK_ELEMS;
let na = tile.na;
let ones = _mm256_set1_epi16(1);
let mut acc = [_mm_setzero_ps(); Q8K_ACTS_X4_NC];
for l in 0..nb {
let blk = packed.as_ptr().add(l * Q8_0X4_BLOCK_BYTES);
let d4 = load_f16x4(blk);
let qs = blk.add(8);
let acts = tile.qs.as_ptr().add(l * Q8_0_BLOCK_ELEMS * 4);
for a in 0..na {
let mut i32acc = _mm256_setzero_si256();
for k in 0..(Q8_0_BLOCK_ELEMS / 8) {
let w = _mm256_loadu_si256(qs.add(k * 32) as *const __m256i);
let av = bcast8(acts.add(k * 32 + a * 8));
let ax = _mm256_sign_epi8(w, w);
let sy = _mm256_sign_epi8(av, w);
i32acc = _mm256_add_epi32(
i32acc,
_mm256_madd_epi16(_mm256_maddubs_epi16(ax, sy), ones),
);
}
let da = _mm_set1_ps(tile.d[l * 4 + a]);
acc[a] = _mm_fmadd_ps(
_mm_cvtepi32_ps(rows4_from_pairs(i32acc)),
_mm_mul_ps(d4, da),
acc[a],
);
}
}
for a in 0..na {
let mut v = [0f32; 4];
_mm_storeu_ps(v.as_mut_ptr(), acc[a]);
for (j, got) in v.iter().enumerate() {
out[j * na + a] = *got;
}
}
}