From 941e5c5dae5ec5176aa978ff8ab1dad089ad3988 Mon Sep 17 00:00:00 2001 From: lenny76 Date: Thu, 24 Sep 2026 00:09:36 +0200 Subject: [PATCH] ggml-cpu : PQ2_0 AVX-512 VNNI gemm on tiles of 8 activation rows Prompt processing spent most of its time re-expanding the 2-bit weights every 4 rows, gathering each q8_0 row out of the x4 interleave, and recomputing dpbusd(ones, q) for the codes-1 offset per 32-block, row and column group. Rows are now processed 8 at a time, de-interleaved once per call with their scales, and the offset is subtracted once per 128-block from a precomputed per-row sum. Co-Authored-By: Claude Opus 5.5 --- ggml/src/ggml-cpu/arch/x86/repack.cpp | 148 ++++++++++++++++++-------- 1 file changed, 101 insertions(+), 47 deletions(-) diff --git a/ggml/src/ggml-cpu/arch/x86/repack.cpp b/ggml/src/ggml-cpu/arch/x86/repack.cpp index 42aa92a5af16..582d689c51e2 100644 --- a/ggml/src/ggml-cpu/arch/x86/repack.cpp +++ b/ggml/src/ggml-cpu/arch/x86/repack.cpp @@ -14,6 +14,7 @@ #include #include // for qsort #include // for GGML_ASSERT +#include #define GGML_CPU_CLANG_WORKAROUND #include "../../repack.h" @@ -6905,71 +6906,124 @@ void ggml_gemv_pq2_0_4x8_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const v ggml_gemv_pq2_0_4x8_q8_0_generic(n, s, bs, vx, vy, nr, nc); } +#if defined(__AVX512F__) && defined(__AVX512BW__) && defined(__AVX512DQ__) && defined(__AVX512VNNI__) +// T activation rows against all columns of vx. qy: T de-interleaved q8_0 rows of n bytes, dy: their q8_0 scales +// (T x nb*4), cy: per row and 128-block sum over the 4 sub-blocks of scale * sum(q) / 8, the correction for +// codes-1, subtracted from each of the 8 dword lanes of a column. +template +static void ggml_gemm_pq2_0_4x8_q8_0_rows(int nb, float * GGML_RESTRICT s, size_t bs, const block_pq2_0x4 * GGML_RESTRICT vx, int nc, + const int8_t * GGML_RESTRICT qy, const float * GGML_RESTRICT dy, const float * GGML_RESTRICT cy) { + const int n = nb * QK_PQ2_0; + const int nk = nb * (QK_PQ2_0 / QK8_0); + + const __m512i idx01 = _mm512_set_epi32(1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0); + const __m512i idx23 = _mm512_set_epi32(3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2, 2); + + for (int x = 0; x < nc / 4; x++) { + const block_pq2_0x4 * GGML_RESTRICT b_ptr = vx + x * nb; + + __m512 acc01[T]; + __m512 acc23[T]; + for (int t = 0; t < T; t++) { + acc01[t] = _mm512_setzero_ps(); + acc23[t] = _mm512_setzero_ps(); + } + + for (int l = 0; l < nb; l++) { + const __m512 d0 = _mm512_castps128_ps512(_mm_cvtph_ps(_mm_loadl_epi64((const __m128i *) b_ptr[l].d))); + const __m512 d0v01 = _mm512_permutexvar_ps(idx01, d0); + const __m512 d0v23 = _mm512_permutexvar_ps(idx23, d0); + + for (int k = 0; k < QK_PQ2_0 / QK8_0; ++k) { + const int kb = l * (QK_PQ2_0 / QK8_0) + k; + + __m512i w01, w23; + __q2_0_expand_x4((const uint8_t *) b_ptr[l].qs + 32 * k, &w01, &w23); + + for (int t = 0; t < T; t++) { + const __m512i qa = _mm512_broadcast_i64x4(_mm256_loadu_si256((const __m256i *) (qy + (size_t) t * n + kb * QK8_0))); + const __m512i i01 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), w01, qa); + const __m512i i23 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), w23, qa); + const __m512 dd = _mm512_set1_ps(dy[t * nk + kb]); + acc01[t] = _mm512_fmadd_ps(_mm512_cvtepi32_ps(i01), _mm512_mul_ps(d0v01, dd), acc01[t]); + acc23[t] = _mm512_fmadd_ps(_mm512_cvtepi32_ps(i23), _mm512_mul_ps(d0v23, dd), acc23[t]); + } + } + + // codes are stored as w+1: subtract d0 * sum_k(d1_k * sum(q_k)) + for (int t = 0; t < T; t++) { + const __m512 c = _mm512_set1_ps(cy[t * nb + l]); + acc01[t] = _mm512_fnmadd_ps(d0v01, c, acc01[t]); + acc23[t] = _mm512_fnmadd_ps(d0v23, c, acc23[t]); + } + } + + for (int t = 0; t < T; t++) { + _mm_storeu_ps(s + t * bs + x * 4, __acc_pair_reduce_ps(acc01[t], acc23[t])); + } + } +} +#endif // AVX512 VNNI + void ggml_gemm_pq2_0_4x8_q8_0(int n, float * GGML_RESTRICT s, size_t bs, const void * GGML_RESTRICT vx, const void * GGML_RESTRICT vy, int nr, int nc) { #if defined(__AVX512F__) && defined(__AVX512BW__) && defined(__AVX512DQ__) && defined(__AVX512VNNI__) { const int qk = QK_PQ2_0; const int nb = n / qk; + const int nk = n / QK8_0; assert(n % qk == 0); assert(nr % 4 == 0); assert(nc % 4 == 0); - const __m512i ones = _mm512_set1_epi8(1); - const __m512i idx01 = _mm512_set_epi32(1, 1, 1, 1, 1, 1, 1, 1, 0, 0, 0, 0, 0, 0, 0, 0); - const __m512i idx23 = _mm512_set_epi32(3, 3, 3, 3, 3, 3, 3, 3, 2, 2, 2, 2, 2, 2, 2, 2); - - // qword gathers of row m from a block_q8_0x4 (row-interleaved in 8-byte chunks) - __m512i rowidx[4]; - for (int m = 0; m < 4; ++m) { - rowidx[m] = _mm512_set_epi64(12 + m, 8 + m, 4 + m, m, 12 + m, 8 + m, 4 + m, m); - } - - for (int y = 0; y < nr / 4; y++) { - const block_q8_0x4 * a_ptr = (const block_q8_0x4 *) vy + (4 * y * nb); - for (int x = 0; x < nc / 4; x++) { - const block_pq2_0x4 * b_ptr = (const block_pq2_0x4 *) vx + (x * nb); - - __m512 accf01[4]; - __m512 accf23[4]; - for (int m = 0; m < 4; m++) { - accf01[m] = _mm512_setzero_ps(); - accf23[m] = _mm512_setzero_ps(); - } - - for (int l = 0; l < nb; l++) { - const __m512 d0 = _mm512_castps128_ps512(_mm_cvtph_ps(_mm_loadl_epi64((const __m128i *) b_ptr[l].d))); - const __m512 d0v01 = _mm512_permutexvar_ps(idx01, d0); - const __m512 d0v23 = _mm512_permutexvar_ps(idx23, d0); - - for (int k = 0; k < QK_PQ2_0 / QK8_0; ++k) { - const block_q8_0x4 * GGML_RESTRICT a_blk = a_ptr + 4 * l + k; - - __m512i w01, w23; - __q2_0_expand_x4((const uint8_t *) b_ptr[l].qs + 32 * k, &w01, &w23); + constexpr int TR = 8; // activation rows per tile - const __m512i qa_lo = _mm512_loadu_si512((const void *) a_blk->qs); - const __m512i qa_hi = _mm512_loadu_si512((const void *) (a_blk->qs + 64)); + static thread_local std::vector qy; + static thread_local std::vector dy; + static thread_local std::vector cy; + qy.resize((size_t) TR * n); + dy.resize((size_t) TR * nk); + cy.resize((size_t) TR * nb); - for (int m = 0; m < 4; ++m) { - const __m512i qa = _mm512_permutex2var_epi64(qa_lo, rowidx[m], qa_hi); - const __m512i sq = _mm512_dpbusd_epi32(_mm512_setzero_si512(), ones, qa); - const __m512i i01 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), w01, qa); - const __m512i i23 = _mm512_dpbusd_epi32(_mm512_setzero_si512(), w23, qa); + const block_q8_0x4 * a_all = (const block_q8_0x4 *) vy; + const block_pq2_0x4 * b_all = (const block_pq2_0x4 *) vx; - const __m512 f01 = _mm512_cvtepi32_ps(_mm512_sub_epi32(i01, sq)); - const __m512 f23 = _mm512_cvtepi32_ps(_mm512_sub_epi32(i23, sq)); + for (int y0 = 0; y0 < nr; y0 += TR) { + const int rows = nr - y0 < TR ? nr - y0 : TR; // 8 or 4 - const __m512 dd = _mm512_set1_ps(GGML_CPU_FP16_TO_FP32(a_blk->d[m])); - accf01[m] = _mm512_fmadd_ps(f01, _mm512_mul_ps(d0v01, dd), accf01[m]); - accf23[m] = _mm512_fmadd_ps(f23, _mm512_mul_ps(d0v23, dd), accf23[m]); + // de-interleave the rows of this tile: row m of a block_q8_0x4 is 4 chunks of 8 bytes at qs[c*32 + m*8] + for (int g = 0; g < rows / 4; g++) { + const block_q8_0x4 * a_ptr = a_all + (size_t) (y0 / 4 + g) * nk; + for (int m = 0; m < 4; m++) { + const int t = 4 * g + m; + int8_t * q = qy.data() + (size_t) t * n; + for (int l = 0; l < nb; l++) { + float c = 0.0f; + for (int k = 0; k < qk / QK8_0; k++) { + const int kb = l * (qk / QK8_0) + k; + const block_q8_0x4 * a_blk = a_ptr + kb; + const __m256i r = _mm256_setr_epi64x( + *(const int64_t *) (a_blk->qs + 0 * 32 + m * 8), *(const int64_t *) (a_blk->qs + 1 * 32 + m * 8), + *(const int64_t *) (a_blk->qs + 2 * 32 + m * 8), *(const int64_t *) (a_blk->qs + 3 * 32 + m * 8)); + _mm256_storeu_si256((__m256i *) (q + kb * QK8_0), r); + // sum of the 32 int8: dpbusd with ones as the unsigned operand + const __m256i sq = _mm256_madd_epi16(_mm256_maddubs_epi16(_mm256_set1_epi8(1), r), _mm256_set1_epi16(1)); + const __m128i s4 = _mm_add_epi32(_mm256_castsi256_si128(sq), _mm256_extracti128_si256(sq, 1)); + const __m128i s2 = _mm_add_epi32(s4, _mm_unpackhi_epi64(s4, s4)); + const int sum = _mm_cvtsi128_si32(_mm_add_epi32(s2, _mm_shuffle_epi32(s2, 1))); + const float d1 = GGML_CPU_FP16_TO_FP32(a_blk->d[m]); + dy[t * nk + kb] = d1; + c += d1 * (float) sum; } + cy[t * nb + l] = c * 0.125f; // spread over the 8 dword lanes of a column } } + } - for (int m = 0; m < 4; m++) { - _mm_storeu_ps(s + (y * 4 + m) * bs + x * 4, __acc_pair_reduce_ps(accf01[m], accf23[m])); - } + if (rows == TR) { + ggml_gemm_pq2_0_4x8_q8_0_rows(nb, s + (size_t) y0 * bs, bs, b_all, nc, qy.data(), dy.data(), cy.data()); + } else { + ggml_gemm_pq2_0_4x8_q8_0_rows<4>(nb, s + (size_t) y0 * bs, bs, b_all, nc, qy.data(), dy.data(), cy.data()); } } return;