Enable AVX-VNNI 256-bit path for IQ4_NL R4 matmul (#1467)
Co-authored-by: Adam Caldwell <accaldwell@users.noreply.github.com>
This commit is contained in:
parent
77f8060ca6
commit
a56a7863c7
|
|
@ -892,7 +892,9 @@ static void mul_mat_iq4_nl_r4_q8_2(int n, const void * vx, size_t bx, const Data
|
|||
GGML_ASSERT(nrc_x%4 == 0);
|
||||
Q8<nrc_y, block_q8_2_x4> q8(info);
|
||||
auto m4 = _mm256_set1_epi8(0xf);
|
||||
#ifndef HAVE_VNNI256
|
||||
auto m1 = _mm256_set1_epi16(1);
|
||||
#endif
|
||||
auto values128 = _mm_loadu_si128((const __m128i *)iq4k_values);
|
||||
auto values = MM256_SET_M128I(values128, values128);
|
||||
int nb = n / QK4_NL;
|
||||
|
|
@ -910,6 +912,16 @@ static void mul_mat_iq4_nl_r4_q8_2(int n, const void * vx, size_t bx, const Data
|
|||
qs[3] = _mm256_shuffle_epi8(values, _mm256_and_si256(_mm256_srli_epi16(bits2, 4), m4));
|
||||
return scales;
|
||||
};
|
||||
#ifdef HAVE_VNNI256
|
||||
auto dot = [&qs] (__m256i y) {
|
||||
auto sumi = _mm256_setzero_si256();
|
||||
sumi = _mm256_dpbusd_epi32(sumi, _mm256_sign_epi8(qs[0], qs[0]), _mm256_sign_epi8(_mm256_shuffle_epi32(y, 0x00), qs[0]));
|
||||
sumi = _mm256_dpbusd_epi32(sumi, _mm256_sign_epi8(qs[1], qs[1]), _mm256_sign_epi8(_mm256_shuffle_epi32(y, 0x55), qs[1]));
|
||||
sumi = _mm256_dpbusd_epi32(sumi, _mm256_sign_epi8(qs[2], qs[2]), _mm256_sign_epi8(_mm256_shuffle_epi32(y, 0xaa), qs[2]));
|
||||
sumi = _mm256_dpbusd_epi32(sumi, _mm256_sign_epi8(qs[3], qs[3]), _mm256_sign_epi8(_mm256_shuffle_epi32(y, 0xff), qs[3]));
|
||||
return sumi;
|
||||
};
|
||||
#else
|
||||
auto dot = [&qs, &m1] (__m256i y) {
|
||||
auto u1 = _mm256_sign_epi8(qs[0], qs[0]);
|
||||
auto u2 = _mm256_sign_epi8(qs[1], qs[1]);
|
||||
|
|
@ -923,6 +935,7 @@ static void mul_mat_iq4_nl_r4_q8_2(int n, const void * vx, size_t bx, const Data
|
|||
_mm256_madd_epi16(m1, _mm256_maddubs_epi16(u2, _mm256_sign_epi8(_mm256_shuffle_epi32(y, 0xff), qs[3]))));
|
||||
return _mm256_add_epi32(sumi1, sumi2);
|
||||
};
|
||||
#endif
|
||||
for (int ix = 0; ix < nrc_x; ix += 4) {
|
||||
const block_iq4_nl_r4 * iq4 = (const block_iq4_nl_r4 *)((const char *)vx + ix*bx);
|
||||
for (int ib4 = 0; ib4 < nb/4; ++ib4) {
|
||||
|
|
|
|||
Loading…
Reference in New Issue