Skip to content
Merged
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
148 changes: 101 additions & 47 deletions ggml/src/ggml-cpu/arch/x86/repack.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
#include <cassert>
#include <cstdlib> // for qsort
#include <cstdio> // for GGML_ASSERT
#include <vector>

#define GGML_CPU_CLANG_WORKAROUND
#include "../../repack.h"
Expand Down Expand Up @@ -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 <int T>
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<int8_t> qy;
static thread_local std::vector<float> dy;
static thread_local std::vector<float> 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<TR>(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;
Expand Down