Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
124 changes: 124 additions & 0 deletions ggml/src/ggml-cpu/arch/x86/quants.c
Original file line number Diff line number Diff line change
Expand Up @@ -48,6 +48,13 @@ static inline float hsum_float_8(const __m256 x) {
return _mm_cvtss_f32(res);
}

// horizontally add 4 floats
static inline float hsum_float_4(__m128 x) {
x = _mm_add_ps(x, _mm_movehl_ps(x, x));
x = _mm_add_ss(x, _mm_movehdup_ps(x));
return _mm_cvtss_f32(x);
}

// horizontally add 8 int32_t
static inline int hsum_i32_8(const __m256i a) {
const __m128i sum128 = _mm_add_epi32(_mm256_castsi256_si128(a), _mm256_extractf128_si256(a, 1));
Expand Down Expand Up @@ -2037,7 +2044,11 @@ void ggml_vec_dot_q3_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi

void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, size_t bx, const void * GGML_RESTRICT vy, size_t by, int nrc) {
assert(n % QK_K == 0);
#if defined __AVX2__
assert((nrc == 2) || (nrc == 1));
#else
assert(nrc == 1);
#endif
UNUSED(nrc);
UNUSED(bx);
UNUSED(by);
Expand All @@ -2055,6 +2066,119 @@ void ggml_vec_dot_q4_K_q8_K(int n, float * GGML_RESTRICT s, size_t bs, const voi
uint32_t utmp[4];

#if defined __AVX2__
if (nrc == 2) {
// 2x2 register-blocked kernel: two src0 rows against two src1 columns.
// The per-block scales/mins and the 4-bit weights of a row are unpacked
// once and reused by both columns, so a 2x2 output tile is much cheaper
// than four independent dot products. The per-block accumulation order
// is the same as the 1x1 path below, so results are bit-identical.
const block_q4_K * GGML_RESTRICT x0 = x;
const block_q4_K * GGML_RESTRICT x1 = (const block_q4_K *)((const uint8_t *) vx + bx);
const block_q8_K * GGML_RESTRICT y0 = y;
const block_q8_K * GGML_RESTRICT y1 = (const block_q8_K *)((const uint8_t *) vy + by);

const __m256i m4 = _mm256_set1_epi8(0xF);

__m256 acc00 = _mm256_setzero_ps(), acc10 = _mm256_setzero_ps();
__m256 acc01 = _mm256_setzero_ps(), acc11 = _mm256_setzero_ps();
__m128 accm00 = _mm_setzero_ps(), accm10 = _mm_setzero_ps();
__m128 accm01 = _mm_setzero_ps(), accm11 = _mm_setzero_ps();

for (int i = 0; i < nb; ++i) {
// unpack the 6-bit scales/mins of both src0 rows
memcpy(utmp, x0[i].scales, 12);
utmp[3] = ((utmp[2] >> 4) & kmask2) | (((utmp[1] >> 6) & kmask3) << 4);
uint32_t uaux = utmp[1] & kmask1;
utmp[1] = (utmp[2] & kmask2) | (((utmp[0] >> 6) & kmask3) << 4);
utmp[2] = uaux;
utmp[0] &= kmask1;
const __m256i ms0 = _mm256_cvtepu8_epi16(_mm_set_epi32(utmp[3], utmp[2], utmp[1], utmp[0]));

memcpy(utmp, x1[i].scales, 12);
utmp[3] = ((utmp[2] >> 4) & kmask2) | (((utmp[1] >> 6) & kmask3) << 4);
uaux = utmp[1] & kmask1;
utmp[1] = (utmp[2] & kmask2) | (((utmp[0] >> 6) & kmask3) << 4);
utmp[2] = uaux;
utmp[0] &= kmask1;
const __m256i ms1 = _mm256_cvtepu8_epi16(_mm_set_epi32(utmp[3], utmp[2], utmp[1], utmp[0]));

// mins contribution, shared by both columns
const __m256i q8sums0 = _mm256_loadu_si256((const __m256i*)y0[i].bsums);
const __m128i q8s0 = _mm_hadd_epi16(_mm256_extracti128_si256(q8sums0, 0), _mm256_extracti128_si256(q8sums0, 1));
const __m256i q8sums1 = _mm256_loadu_si256((const __m256i*)y1[i].bsums);
const __m128i q8s1 = _mm_hadd_epi16(_mm256_extracti128_si256(q8sums1, 0), _mm256_extracti128_si256(q8sums1, 1));

const __m128i mins0 = _mm256_extracti128_si256(ms0, 1);
const __m128i mins1 = _mm256_extracti128_si256(ms1, 1);

accm00 = _mm_fmadd_ps(_mm_set1_ps(-y0[i].d * GGML_CPU_FP16_TO_FP32(x0[i].dmin)), _mm_cvtepi32_ps(_mm_madd_epi16(mins0, q8s0)), accm00);
accm10 = _mm_fmadd_ps(_mm_set1_ps(-y0[i].d * GGML_CPU_FP16_TO_FP32(x1[i].dmin)), _mm_cvtepi32_ps(_mm_madd_epi16(mins1, q8s0)), accm10);
accm01 = _mm_fmadd_ps(_mm_set1_ps(-y1[i].d * GGML_CPU_FP16_TO_FP32(x0[i].dmin)), _mm_cvtepi32_ps(_mm_madd_epi16(mins0, q8s1)), accm01);
accm11 = _mm_fmadd_ps(_mm_set1_ps(-y1[i].d * GGML_CPU_FP16_TO_FP32(x1[i].dmin)), _mm_cvtepi32_ps(_mm_madd_epi16(mins1, q8s1)), accm11);

const __m128i sc0 = _mm256_extracti128_si256(ms0, 0);
const __m256i scales0 = MM256_SET_M128I(sc0, sc0);
const __m128i sc1 = _mm256_extracti128_si256(ms1, 0);
const __m256i scales1 = MM256_SET_M128I(sc1, sc1);

const uint8_t * GGML_RESTRICT q4r0 = x0[i].qs;
const uint8_t * GGML_RESTRICT q4r1 = x1[i].qs;
const int8_t * GGML_RESTRICT q8c0 = y0[i].qs;
const int8_t * GGML_RESTRICT q8c1 = y1[i].qs;

__m256i sumi00 = _mm256_setzero_si256(), sumi10 = _mm256_setzero_si256();
__m256i sumi01 = _mm256_setzero_si256(), sumi11 = _mm256_setzero_si256();

for (int j = 0; j < QK_K/64; ++j) {
const __m256i shuf_l = get_scale_shuffle_k4(2*j+0);
const __m256i shuf_h = get_scale_shuffle_k4(2*j+1);

{ // src0 row 0 against both columns
const __m256i sl = _mm256_shuffle_epi8(scales0, shuf_l);
const __m256i sh = _mm256_shuffle_epi8(scales0, shuf_h);

const __m256i bits = _mm256_loadu_si256((const __m256i*)(q4r0 + 32*j));
const __m256i ql = _mm256_and_si256(bits, m4);
const __m256i qh = _mm256_and_si256(_mm256_srli_epi16(bits, 4), m4);

sumi00 = _mm256_add_epi32(sumi00, _mm256_add_epi32(
_mm256_madd_epi16(sl, _mm256_maddubs_epi16(ql, _mm256_loadu_si256((const __m256i*)(q8c0 + 64*j)))),
_mm256_madd_epi16(sh, _mm256_maddubs_epi16(qh, _mm256_loadu_si256((const __m256i*)(q8c0 + 64*j + 32))))));
sumi01 = _mm256_add_epi32(sumi01, _mm256_add_epi32(
_mm256_madd_epi16(sl, _mm256_maddubs_epi16(ql, _mm256_loadu_si256((const __m256i*)(q8c1 + 64*j)))),
_mm256_madd_epi16(sh, _mm256_maddubs_epi16(qh, _mm256_loadu_si256((const __m256i*)(q8c1 + 64*j + 32))))));
}

{ // src0 row 1 against both columns
const __m256i sl = _mm256_shuffle_epi8(scales1, shuf_l);
const __m256i sh = _mm256_shuffle_epi8(scales1, shuf_h);

const __m256i bits = _mm256_loadu_si256((const __m256i*)(q4r1 + 32*j));
const __m256i ql = _mm256_and_si256(bits, m4);
const __m256i qh = _mm256_and_si256(_mm256_srli_epi16(bits, 4), m4);

sumi10 = _mm256_add_epi32(sumi10, _mm256_add_epi32(
_mm256_madd_epi16(sl, _mm256_maddubs_epi16(ql, _mm256_loadu_si256((const __m256i*)(q8c0 + 64*j)))),
_mm256_madd_epi16(sh, _mm256_maddubs_epi16(qh, _mm256_loadu_si256((const __m256i*)(q8c0 + 64*j + 32))))));
sumi11 = _mm256_add_epi32(sumi11, _mm256_add_epi32(
_mm256_madd_epi16(sl, _mm256_maddubs_epi16(ql, _mm256_loadu_si256((const __m256i*)(q8c1 + 64*j)))),
_mm256_madd_epi16(sh, _mm256_maddubs_epi16(qh, _mm256_loadu_si256((const __m256i*)(q8c1 + 64*j + 32))))));
}
}

acc00 = _mm256_fmadd_ps(_mm256_set1_ps(y0[i].d * GGML_CPU_FP16_TO_FP32(x0[i].d)), _mm256_cvtepi32_ps(sumi00), acc00);
acc10 = _mm256_fmadd_ps(_mm256_set1_ps(y0[i].d * GGML_CPU_FP16_TO_FP32(x1[i].d)), _mm256_cvtepi32_ps(sumi10), acc10);
acc01 = _mm256_fmadd_ps(_mm256_set1_ps(y1[i].d * GGML_CPU_FP16_TO_FP32(x0[i].d)), _mm256_cvtepi32_ps(sumi01), acc01);
acc11 = _mm256_fmadd_ps(_mm256_set1_ps(y1[i].d * GGML_CPU_FP16_TO_FP32(x1[i].d)), _mm256_cvtepi32_ps(sumi11), acc11);
}

s[0] = hsum_float_8(acc00) + hsum_float_4(accm00);
s[1] = hsum_float_8(acc10) + hsum_float_4(accm10);
s[bs + 0] = hsum_float_8(acc01) + hsum_float_4(accm01);
s[bs + 1] = hsum_float_8(acc11) + hsum_float_4(accm11);

return;
}

const __m256i m4 = _mm256_set1_epi8(0xF);

Expand Down
2 changes: 1 addition & 1 deletion ggml/src/ggml-cpu/ggml-cpu.c
Original file line number Diff line number Diff line change
Expand Up @@ -311,7 +311,7 @@ static const struct ggml_type_traits_cpu type_traits_cpu[GGML_TYPE_COUNT] = {
.from_float = quantize_row_q4_K,
.vec_dot = ggml_vec_dot_q4_K_q8_K,
.vec_dot_type = GGML_TYPE_Q8_K,
#if defined (__ARM_FEATURE_MATMUL_INT8)
#if defined (__ARM_FEATURE_MATMUL_INT8) || defined (__AVX2__)
.nrows = 2,
#else
.nrows = 1,
Expand Down