diff --git a/ggml/src/ggml-cpu/arch/x86/quants.c b/ggml/src/ggml-cpu/arch/x86/quants.c index ea54cfe44ce..bb9996f7f78 100644 --- a/ggml/src/ggml-cpu/arch/x86/quants.c +++ b/ggml/src/ggml-cpu/arch/x86/quants.c @@ -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)); @@ -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); @@ -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); diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 491316f7491..d79763fabb6 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -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,