16#define CK_Q51_STACK_Q8_BLOCKS 256
32static const uint32_t ck_q51_high_nibble_lut[16] = {
33 0x00000000U, 0x00000010U, 0x00001000U, 0x00001010U,
34 0x00100000U, 0x00100010U, 0x00101000U, 0x00101010U,
35 0x10000000U, 0x10000010U, 0x10001000U, 0x10001010U,
36 0x10100000U, 0x10100010U, 0x10101000U, 0x10101010U,
39static inline uint32_t ck_q51_high_nibble_word(uint32_t bits4) {
40 return ck_q51_high_nibble_lut[bits4 & 0x0fU];
43static inline __m128i ck_q51_high_bits_16_avx2(uint32_t bits16) {
44 const uint32_t w0 = ck_q51_high_nibble_word(bits16);
45 const uint32_t w1 = ck_q51_high_nibble_word(bits16 >> 4);
46 const uint32_t w2 = ck_q51_high_nibble_word(bits16 >> 8);
47 const uint32_t w3 = ck_q51_high_nibble_word(bits16 >> 12);
48 return _mm_set_epi32((
int)w3, (
int)w2, (
int)w1, (
int)w0);
51static inline int ck_q51_hsum256_epi32(__m256i v) {
52 const __m128i lo = _mm256_castsi256_si128(v);
53 const __m128i hi = _mm256_extracti128_si256(v, 1);
54 __m128i sum = _mm_add_epi32(lo, hi);
55 sum = _mm_hadd_epi32(sum, sum);
56 sum = _mm_hadd_epi32(sum, sum);
57 return _mm_cvtsi128_si32(sum);
64 const int nb = k /
QK8_1;
65 for (
int b = 0; b < nb; ++b) {
66 const float *xb = x + (size_t)b *
QK8_1;
68 for (
int j = 0; j <
QK8_1; ++j) {
69 float av = xb[j] >= 0.0f ? xb[j] : -xb[j];
70 if (av > amax) amax = av;
73 const float d = amax / 127.0f;
74 const float id = (d != 0.0f) ? (1.0f / d) : 0.0f;
78 for (
int j = 0; j <
QK8_1; ++j) {
79 int q = (int)roundf(xb[j] *
id);
80 y[b].qs[j] = (int8_t)q;
88static inline __m256i dot_q5_1_q8_1_block_sumi_avx2(
const block_q5_1 *w,
89 const block_q8_1 *x) {
91 memcpy(&qh, w->
qh,
sizeof(qh));
93 const __m128i qpacked = _mm_loadu_si128((
const __m128i *)(
const void *)w->
qs);
94 const __m128i low_mask = _mm_set1_epi8(0x0f);
95 const __m128i qlo = _mm_or_si128(_mm_and_si128(qpacked, low_mask),
96 ck_q51_high_bits_16_avx2(qh));
97 const __m128i qhi = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(qpacked, 4), low_mask),
98 ck_q51_high_bits_16_avx2(qh >> 16));
99 const __m256i q5 = _mm256_inserti128_si256(_mm256_castsi128_si256(qlo), qhi, 1);
100 const __m256i q8 = _mm256_loadu_si256((
const __m256i *)(
const void *)x->qs);
101#if defined(__AVXVNNI__)
102 return _mm256_dpbusd_epi32(_mm256_setzero_si256(), q5, q8);
104 const __m256i prod16 = _mm256_maddubs_epi16(q5, q8);
105 return _mm256_madd_epi16(prod16, _mm256_set1_epi16(1));
112static float dot_q5_1_q8_1_block_avx2(
const block_q5_1 *w,
const block_q8_1 *x) {
113 const __m256i sumi = dot_q5_1_q8_1_block_sumi_avx2(w, x);
119 return (wd * xd) * (float)ck_q51_hsum256_epi32(sumi) + wm * xs;
122static inline void dot_q5_1_q8_1_block_m4_avx2(
124 const block_q8_1 *x0,
125 const block_q8_1 *x1,
126 const block_q8_1 *x2,
127 const block_q8_1 *x3,
130 memcpy(&qh, w->
qh,
sizeof(qh));
132 const __m128i qpacked = _mm_loadu_si128((
const __m128i *)(
const void *)w->
qs);
133 const __m128i low_mask = _mm_set1_epi8(0x0f);
134 const __m128i qlo = _mm_or_si128(_mm_and_si128(qpacked, low_mask),
135 ck_q51_high_bits_16_avx2(qh));
136 const __m128i qhi = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(qpacked, 4), low_mask),
137 ck_q51_high_bits_16_avx2(qh >> 16));
138 const __m256i q5 = _mm256_inserti128_si256(_mm256_castsi128_si256(qlo), qhi, 1);
139 const block_q8_1 *rows[4] = {x0, x1, x2, x3};
143 for (
int row = 0; row < 4; ++row) {
144 const __m256i q8 = _mm256_loadu_si256(
145 (
const __m256i *)(
const void *)rows[row]->qs);
146#if defined(__AVXVNNI__)
147 const __m256i sumi = _mm256_dpbusd_epi32(_mm256_setzero_si256(), q5, q8);
149 const __m256i prod16 = _mm256_maddubs_epi16(q5, q8);
150 const __m256i sumi = _mm256_madd_epi16(prod16, _mm256_set1_epi16(1));
154 out[row] = (wd * xd) * (
float)ck_q51_hsum256_epi32(sumi) + wm * xs;
158static inline void dot_q5_1_q8_1_block_m8_avx2(
160 const block_q8_1 *
const rows[8],
163 memcpy(&qh, w->
qh,
sizeof(qh));
165 const __m128i qpacked = _mm_loadu_si128((
const __m128i *)(
const void *)w->
qs);
166 const __m128i low_mask = _mm_set1_epi8(0x0f);
167 const __m128i qlo = _mm_or_si128(_mm_and_si128(qpacked, low_mask),
168 ck_q51_high_bits_16_avx2(qh));
169 const __m128i qhi = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(qpacked, 4), low_mask),
170 ck_q51_high_bits_16_avx2(qh >> 16));
171 const __m256i q5 = _mm256_inserti128_si256(_mm256_castsi128_si256(qlo), qhi, 1);
175 for (
int row = 0; row < 8; ++row) {
176 const __m256i q8 = _mm256_loadu_si256(
177 (
const __m256i *)(
const void *)rows[row]->qs);
178#if defined(__AVXVNNI__)
179 const __m256i sumi = _mm256_dpbusd_epi32(_mm256_setzero_si256(), q5, q8);
181 const __m256i prod16 = _mm256_maddubs_epi16(q5, q8);
182 const __m256i sumi = _mm256_madd_epi16(prod16, _mm256_set1_epi16(1));
186 out[row] = (wd * xd) * (
float)ck_q51_hsum256_epi32(sumi) + wm * xs;
195 return dot_q5_1_q8_1_block_avx2(w, x);
198 memcpy(&qh, w->
qh,
sizeof(qh));
202 for (
int j = 0; j <
QK5_1 / 2; ++j) {
203 const uint8_t xh0 = (uint8_t)(((qh >> (j + 0)) << 4) & 0x10);
204 const uint8_t xh1 = (uint8_t)(((qh >> (j + 12)) ) & 0x10);
205 const int32_t q0 = (int32_t)((w->
qs[j] & 0x0F) | xh0);
206 const int32_t q1 = (int32_t)((w->
qs[j] >> 4) | xh1);
207 sumi0 += q0 * (int32_t)x->qs[j];
208 sumi1 += q1 * (int32_t)x->qs[j +
QK5_1 / 2];
215 return (wd * xd) * (float)(sumi0 + sumi1) + wm * xs;
224 if (!y || !W || !x_q8 || M <= 0 || K <= 0 || (K %
QK5_1) != 0) {
229 const block_q8_1 *x = (
const block_q8_1 *)x_q8;
230 const int blocks_per_row = K /
QK5_1;
232 for (
int row = 0; row < M; ++row) {
233 const block_q5_1 *w_row = &blocks[row * blocks_per_row];
235 for (
int b = 0; b < blocks_per_row; ++b) {
248 if (!A_q8 || !B || !
C || M <= 0 || N <= 0 || K <= 0 || (K %
QK5_1) != 0) {
252 const block_q8_1 *A = (
const block_q8_1 *)A_q8;
254 const int blocks_per_row = K /
QK5_1;
256 for (
int m = 0; m < M; ++m) {
257 const block_q8_1 *a_row = &A[m * blocks_per_row];
258 for (
int n = 0; n < N; ++n) {
259 const block_q5_1 *w_row = &W[n * blocks_per_row];
261 for (
int b = 0; b < blocks_per_row; ++b) {
264 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
275 if (!y || !W || !x || M <= 0 || K <= 0 || (K %
QK5_1) != 0) {
279 const int blocks_per_row = K /
QK5_1;
297 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0 || (K %
QK5_1) != 0) {
301 const int blocks_per_row = K /
QK5_1;
308 for (
int m = 0; m < M; ++m) {
311 float *c_row = &
C[(size_t)m * (
size_t)N];
314 for (; n + 7 < N; n += 8) {
315 const block_q5_1 *w0 = &W[(size_t)(n + 0) * (size_t)blocks_per_row];
316 const block_q5_1 *w1 = &W[(size_t)(n + 1) * (size_t)blocks_per_row];
317 const block_q5_1 *w2 = &W[(size_t)(n + 2) * (size_t)blocks_per_row];
318 const block_q5_1 *w3 = &W[(size_t)(n + 3) * (size_t)blocks_per_row];
319 const block_q5_1 *w4 = &W[(size_t)(n + 4) * (size_t)blocks_per_row];
320 const block_q5_1 *w5 = &W[(size_t)(n + 5) * (size_t)blocks_per_row];
321 const block_q5_1 *w6 = &W[(size_t)(n + 6) * (size_t)blocks_per_row];
322 const block_q5_1 *w7 = &W[(size_t)(n + 7) * (size_t)blocks_per_row];
332 for (
int b = 0; b < blocks_per_row; ++b) {
333 const block_q8_1 *x = &a_q8[b];
344 c_row[n + 0] = s0 + (bias ? bias[n + 0] : 0.0f);
345 c_row[n + 1] = s1 + (bias ? bias[n + 1] : 0.0f);
346 c_row[n + 2] = s2 + (bias ? bias[n + 2] : 0.0f);
347 c_row[n + 3] = s3 + (bias ? bias[n + 3] : 0.0f);
348 c_row[n + 4] = s4 + (bias ? bias[n + 4] : 0.0f);
349 c_row[n + 5] = s5 + (bias ? bias[n + 5] : 0.0f);
350 c_row[n + 6] = s6 + (bias ? bias[n + 6] : 0.0f);
351 c_row[n + 7] = s7 + (bias ? bias[n + 7] : 0.0f);
355 const block_q5_1 *w_row = &W[(size_t)n * (
size_t)blocks_per_row];
357 for (
int b = 0; b < blocks_per_row; ++b) {
360 c_row[n] = sum + (bias ? bias[n] : 0.0f);
374 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0 || (K %
QK5_1) != 0) {
378 const int blocks_per_row = K /
QK5_1;
386 for (; m + 3 < M; m += 4) {
388 for (
int row = 0; row < 4; ++row) {
390 &A[(
size_t)(m + row) * (
size_t)K], activation_q8[row], K);
393 for (
int n = 0; n < N; ++n) {
395 &weights[(size_t)n * (
size_t)blocks_per_row];
396 float sums[4] = {0.0f, 0.0f, 0.0f, 0.0f};
397 for (
int block = 0; block < blocks_per_row; ++block) {
399 dot_q5_1_q8_1_block_m4_avx2(
401 &activation_q8[0][block], &activation_q8[1][block],
402 &activation_q8[2][block], &activation_q8[3][block], partial);
403 sums[0] += partial[0];
404 sums[1] += partial[1];
405 sums[2] += partial[2];
406 sums[3] += partial[3];
408 const float add = bias ? bias[n] : 0.0f;
409 C[(size_t)(m + 0) * (size_t)N + n] = sums[0] + add;
410 C[(size_t)(m + 1) * (size_t)N + n] = sums[1] + add;
411 C[(size_t)(m + 2) * (size_t)N + n] = sums[2] + add;
412 C[(size_t)(m + 3) * (size_t)N + n] = sums[3] + add;
418 A + (
size_t)m * (
size_t)K, B, bias,
419 C + (
size_t)m * (
size_t)N, M - m, N, K);
435 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0 || (K %
QK5_1) != 0) {
439 const int blocks_per_row = K /
QK5_1;
447 for (; m + 7 < M; m += 8) {
449 for (
int row = 0; row < 8; ++row) {
451 &A[(
size_t)(m + row) * (
size_t)K], activation_q8[row], K);
454 for (
int n = 0; n < N; ++n) {
456 &weights[(size_t)n * (
size_t)blocks_per_row];
457 float sums[8] = {0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f};
458 for (
int block = 0; block < blocks_per_row; ++block) {
459 const block_q8_1 *rows[8] = {
460 &activation_q8[0][block], &activation_q8[1][block],
461 &activation_q8[2][block], &activation_q8[3][block],
462 &activation_q8[4][block], &activation_q8[5][block],
463 &activation_q8[6][block], &activation_q8[7][block],
466 dot_q5_1_q8_1_block_m8_avx2(&weight_row[block], rows, partial);
467 for (
int row = 0; row < 8; ++row) {
468 sums[row] += partial[row];
471 const float add = bias ? bias[n] : 0.0f;
472 for (
int row = 0; row < 8; ++row) {
473 C[(size_t)(m + row) * (size_t)N + n] = sums[row] + add;
480 A + (
size_t)m * (
size_t)K, B, bias,
481 C + (
size_t)m * (
size_t)N, M - m, N, K);
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
#define CK_FP32_TO_FP16(x)
void gemm_nt_q5_1_q8_1_m4(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q5_1_q8_1_ref(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_nt_q5_1_q8_1_ref(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
static void quantize_row_q8_1_scalar(const float *x, block_q8_1 *y, int k)
void gemm_nt_q5_1_q8_1(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_1_q8_1_m8(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q5_1_q8_1(float *y, const void *W, const float *x, int M, int K)
#define CK_Q51_STACK_Q8_BLOCKS
static float dot_q5_1_q8_1_block(const block_q5_1 *w, const block_q8_1 *x)