use crate::repack::common::*;
use crate::repack::q4_0x4::*;
use crate::Q4_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_q4_0x4_q8_0_avx2(
packed: &[u8],
tile: &Q8ActsX4,
n_cols: usize,
out: &mut [f32],
) {
let nb = n_cols / Q4_0_BLOCK_ELEMS;
let na = tile.na;
let m4b = _mm256_set1_epi8(0x0F);
let unxor = _mm256_set1_epi8(0x88u8 as i8);
let eight = _mm256_set1_epi8(8);
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 * Q4_0X4_BLOCK_BYTES);
let d4 = load_f16x4(blk);
let qs = blk.add(8);
let acts = tile.qs.as_ptr().add(l * Q4_0_BLOCK_ELEMS * 4);
for a in 0..na {
let mut i16acc = _mm256_setzero_si256();
for k in 0..(Q4_0_BLOCK_ELEMS / 16) {
let x = _mm256_loadu_si256(qs.add(k * 32) as *const __m256i);
let orig = _mm256_xor_si256(x, unxor);
let n0 = _mm256_and_si256(orig, m4b);
let n1 = _mm256_and_si256(_mm256_srli_epi16(orig, 4), m4b);
let a0 = bcast8(acts.add(k * 32 + a * 8));
let a1 = bcast8(acts.add((k + 2) * 32 + a * 8));
let p =
_mm256_add_epi16(_mm256_maddubs_epi16(n0, a0), _mm256_maddubs_epi16(n1, a1));
let bias = _mm256_add_epi16(
_mm256_maddubs_epi16(eight, a0),
_mm256_maddubs_epi16(eight, a1),
);
i16acc = _mm256_add_epi16(i16acc, _mm256_sub_epi16(p, bias));
}
let da = _mm_set1_ps(tile.d[l * 4 + a]);
acc[a] = _mm_fmadd_ps(
_mm_cvtepi32_ps(rows4_from_pairs(_mm256_madd_epi16(i16acc, ones))),
_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;
}
}
}