41#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__) || defined(__SSSE3__)
44#if defined(__ARM_NEON) || defined(__aarch64__)
54_Static_assert(
sizeof(block_q6_K_prepared) == 274,
55 "Q6_K prepared-size contract changed");
59 return sizeof(block_q6_K_prepared);
64#if defined(__AVX512F__) && defined(__AVX512BW__) && \
65 defined(__AVX512VNNI__)
66 return "q6_k_prepared_avx512_vnni_exact";
67#elif defined(__AVX2__)
68 return "q6_k_prepared_avx2_exact";
70 return "q6_k_prepared_unavailable";
76 if (!src || !dst || N <= 0 || K <= 0 || (K %
QK_K) != 0)
return;
78 block_q6_K_prepared *output = (block_q6_K_prepared *)dst;
79 const size_t blocks = (size_t)N * (
size_t)(K /
QK_K);
81 for (
size_t b = 0; b < blocks; ++b) {
82 output[b].
d = input[b].
d;
83 memcpy(output[b].scales, input[b].scales,
sizeof(output[b].scales));
84 for (
int n = 0; n <
QK_K; n += 128) {
85 const uint8_t *ql = input[b].
ql + n / 2;
86 const uint8_t *qh = input[b].
qh + n / 4;
87 for (
int l = 0; l < 32; ++l) {
88 output[b].qs[n + l + 0] =
89 (uint8_t)((ql[l] & 0x0f) | (((qh[l] >> 0) & 3) << 4));
90 output[b].qs[n + l + 32] =
91 (uint8_t)((ql[l + 32] & 0x0f) | (((qh[l] >> 2) & 3) << 4));
92 output[b].qs[n + l + 64] =
93 (uint8_t)((ql[l] >> 4) | (((qh[l] >> 4) & 3) << 4));
94 output[b].qs[n + l + 96] =
95 (uint8_t)((ql[l + 32] >> 4) | (((qh[l] >> 6) & 3) << 4));
110 static int cached = -1;
112 const char *env = getenv(
"CK_DEBUG_Q6K_Q8K_REF");
113 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
139 const int nb = K /
QK_K;
141 float sums[8] = {0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f};
143 for (
int i = 0; i < nb; ++i) {
146 const uint8_t *ql = w[i].
ql;
147 const uint8_t *qh = w[i].
qh;
148 const int8_t *sc = w[i].
scales;
149 const int8_t *q8 = x[i].
qs;
150 int32_t aux32[8] = {0, 0, 0, 0, 0, 0, 0, 0};
153 for (
int n = 0; n <
QK_K; n += 128) {
159 for (
int l = 0; l < 32; ++l) {
161 const int is = l / 16;
165 const int8_t q1 = (int8_t)((ql[l + 0] & 0xF) | (((qh[l] >> 0) & 3) << 4)) - 32;
167 const int8_t q2 = (int8_t)((ql[l + 32] & 0xF) | (((qh[l] >> 2) & 3) << 4)) - 32;
169 const int8_t q3 = (int8_t)((ql[l + 0] >> 4) | (((qh[l] >> 4) & 3) << 4)) - 32;
171 const int8_t q4 = (int8_t)((ql[l + 32] >> 4) | (((qh[l] >> 6) & 3) << 4)) - 32;
173 aux32[l & 7] += (int)sc[is + 0] * (
int)q1 * (int)q8[l + 0];
174 aux32[l & 7] += (int)sc[is + 2] * (
int)q2 * (int)q8[l + 32];
175 aux32[l & 7] += (int)sc[is + 4] * (
int)q3 * (int)q8[l + 64];
176 aux32[l & 7] += (int)sc[is + 6] * (
int)q4 * (int)q8[l + 96];
184 for (
int l = 0; l < 8; ++l) {
185 sums[l] += d * (float)aux32[l];
189 for (
int l = 0; l < 8; ++l) {
195#if defined(__ARM_NEON) || defined(__aarch64__)
196static float dot_q6_k_q8_k_neon(
const block_q6_K *w,
200 const int nb = K /
QK_K;
203 for (
int i = 0; i < nb; ++i) {
206 const uint8_t *ql = w[i].
ql;
207 const uint8_t *qh = w[i].
qh;
208 const int8_t *sc = w[i].
scales;
209 const int8_t *q8 = x[i].
qs;
214 for (
int n = 0; n <
QK_K; n += 128) {
215 for (
int l = 0; l < 32; ++l) {
216 const int is = l / 16;
218 const int8_t q1 = (int8_t)((ql[l + 0] & 0xF) | (((qh[l] >> 0) & 3) << 4)) - 32;
219 const int8_t q2 = (int8_t)((ql[l + 32] & 0xF) | (((qh[l] >> 2) & 3) << 4)) - 32;
220 const int8_t q3 = (int8_t)((ql[l + 0] >> 4) | (((qh[l] >> 4) & 3) << 4)) - 32;
221 const int8_t q4 = (int8_t)((ql[l + 32] >> 4) | (((qh[l] >> 6) & 3) << 4)) - 32;
224 wvals[base + l + 0] = q1;
225 wvals[base + l + 32] = q2;
226 wvals[base + l + 64] = q3;
227 wvals[base + l + 96] = q4;
229 svals[base + l + 0] = sc[is + 0];
230 svals[base + l + 32] = sc[is + 2];
231 svals[base + l + 64] = sc[is + 4];
232 svals[base + l + 96] = sc[is + 6];
240 int32x4_t acc = vdupq_n_s32(0);
241 for (
int j = 0; j <
QK_K; j += 16) {
242 const int8x16_t wv = vld1q_s8(&wvals[j]);
243 const int8x16_t sv = vld1q_s8(&svals[j]);
244 const int8x16_t xv = vld1q_s8(&q8[j]);
246 const int16x8_t ws0 = vmull_s8(vget_low_s8(wv), vget_low_s8(sv));
247 const int16x8_t ws1 = vmull_s8(vget_high_s8(wv), vget_high_s8(sv));
248 const int16x8_t x0 = vmovl_s8(vget_low_s8(xv));
249 const int16x8_t x1 = vmovl_s8(vget_high_s8(xv));
251 int32x4_t p0 = vmull_s16(vget_low_s16(ws0), vget_low_s16(x0));
252 p0 = vmlal_s16(p0, vget_high_s16(ws0), vget_high_s16(x0));
254 int32x4_t p1 = vmull_s16(vget_low_s16(ws1), vget_low_s16(x1));
255 p1 = vmlal_s16(p1, vget_high_s16(ws1), vget_high_s16(x1));
257 acc = vaddq_s32(acc, p0);
258 acc = vaddq_s32(acc, p1);
262 vst1q_s32(lanes, acc);
263 sumf += d * (float)(lanes[0] + lanes[1] + lanes[2] + lanes[3]);
275 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
281 const int blocks_per_row = K /
QK_K;
283 for (
int row = 0; row < M; ++row) {
284 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
296#if defined(__SSSE3__)
299static const int8_t q6k_scale_shuffle[8][16] = {
300 { 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1 },
301 { 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3 },
302 { 4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5 },
303 { 6, 6, 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7 },
304 { 8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9 },
305 {10,10,10,10,10,10,10,10,11,11,11,11,11,11,11,11 },
306 {12,12,12,12,12,12,12,12,13,13,13,13,13,13,13,13 },
307 {14,14,14,14,14,14,14,14,15,15,15,15,15,15,15,15 },
310static float dot_q6_k_q8_k_sse(
const block_q6_K *w,
314 const int nb = K /
QK_K;
315 const __m128i m3 = _mm_set1_epi8(3);
316 const __m128i m15 = _mm_set1_epi8(15);
318 __m128 acc = _mm_setzero_ps();
320 for (
int i = 0; i < nb; ++i) {
323 const uint8_t *ql = w[i].
ql;
324 const uint8_t *qh = w[i].
qh;
325 const int8_t *q8 = x[i].
qs;
328 const __m128i scales = _mm_loadu_si128((
const __m128i *)w[i].scales);
329 const __m128i q8sums_0 = _mm_loadu_si128((
const __m128i *)x[i].bsums);
330 const __m128i q8sums_1 = _mm_loadu_si128((
const __m128i *)x[i].bsums + 1);
333 const __m128i scales_16_0 = _mm_cvtepi8_epi16(scales);
334 const __m128i scales_16_1 = _mm_cvtepi8_epi16(_mm_bsrli_si128(scales, 8));
335 const __m128i q8sclsub_0 = _mm_slli_epi32(_mm_madd_epi16(q8sums_0, scales_16_0), 5);
336 const __m128i q8sclsub_1 = _mm_slli_epi32(_mm_madd_epi16(q8sums_1, scales_16_1), 5);
338 __m128i sumi_0 = _mm_setzero_si128();
339 __m128i sumi_1 = _mm_setzero_si128();
344 for (
int j = 0; j <
QK_K / 128; ++j) {
346 const __m128i q4bitsH_0 = _mm_loadu_si128((
const __m128i *)qh);
348 const __m128i q4bitsH_1 = _mm_loadu_si128((
const __m128i *)qh);
352 const __m128i q4h_0 = _mm_slli_epi16(_mm_and_si128(q4bitsH_0, m3), 4);
353 const __m128i q4h_1 = _mm_slli_epi16(_mm_and_si128(q4bitsH_1, m3), 4);
354 const __m128i q4h_2 = _mm_slli_epi16(_mm_and_si128(q4bitsH_0, _mm_set1_epi8(12)), 2);
355 const __m128i q4h_3 = _mm_slli_epi16(_mm_and_si128(q4bitsH_1, _mm_set1_epi8(12)), 2);
356 const __m128i q4h_4 = _mm_and_si128(q4bitsH_0, _mm_set1_epi8(48));
357 const __m128i q4h_5 = _mm_and_si128(q4bitsH_1, _mm_set1_epi8(48));
358 const __m128i q4h_6 = _mm_srli_epi16(_mm_and_si128(q4bitsH_0, _mm_set1_epi8(-64)), 2);
359 const __m128i q4h_7 = _mm_srli_epi16(_mm_and_si128(q4bitsH_1, _mm_set1_epi8(-64)), 2);
362 const __m128i q4bits1_0 = _mm_loadu_si128((
const __m128i *)ql);
364 const __m128i q4bits1_1 = _mm_loadu_si128((
const __m128i *)ql);
366 const __m128i q4bits2_0 = _mm_loadu_si128((
const __m128i *)ql);
368 const __m128i q4bits2_1 = _mm_loadu_si128((
const __m128i *)ql);
372 const __m128i q4_0 = _mm_or_si128(_mm_and_si128(q4bits1_0, m15), q4h_0);
373 const __m128i q4_1 = _mm_or_si128(_mm_and_si128(q4bits1_1, m15), q4h_1);
374 const __m128i q4_2 = _mm_or_si128(_mm_and_si128(q4bits2_0, m15), q4h_2);
375 const __m128i q4_3 = _mm_or_si128(_mm_and_si128(q4bits2_1, m15), q4h_3);
376 const __m128i q4_4 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits1_0, 4), m15), q4h_4);
377 const __m128i q4_5 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits1_1, 4), m15), q4h_5);
378 const __m128i q4_6 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits2_0, 4), m15), q4h_6);
379 const __m128i q4_7 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits2_1, 4), m15), q4h_7);
382 const __m128i q8_0 = _mm_loadu_si128((
const __m128i *)q8);
384 const __m128i q8_1 = _mm_loadu_si128((
const __m128i *)q8);
386 const __m128i q8_2 = _mm_loadu_si128((
const __m128i *)q8);
388 const __m128i q8_3 = _mm_loadu_si128((
const __m128i *)q8);
390 const __m128i q8_4 = _mm_loadu_si128((
const __m128i *)q8);
392 const __m128i q8_5 = _mm_loadu_si128((
const __m128i *)q8);
394 const __m128i q8_6 = _mm_loadu_si128((
const __m128i *)q8);
396 const __m128i q8_7 = _mm_loadu_si128((
const __m128i *)q8);
400 __m128i p16_0 = _mm_maddubs_epi16(q4_0, q8_0);
401 __m128i p16_1 = _mm_maddubs_epi16(q4_1, q8_1);
402 __m128i p16_2 = _mm_maddubs_epi16(q4_2, q8_2);
403 __m128i p16_3 = _mm_maddubs_epi16(q4_3, q8_3);
404 __m128i p16_4 = _mm_maddubs_epi16(q4_4, q8_4);
405 __m128i p16_5 = _mm_maddubs_epi16(q4_5, q8_5);
406 __m128i p16_6 = _mm_maddubs_epi16(q4_6, q8_6);
407 __m128i p16_7 = _mm_maddubs_epi16(q4_7, q8_7);
410 const __m128i scale_0 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 0]));
411 const __m128i scale_1 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 1]));
412 const __m128i scale_2 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 2]));
413 const __m128i scale_3 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 3]));
417 p16_0 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_0), p16_0);
418 p16_1 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_0, 8)), p16_1);
419 p16_2 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_1), p16_2);
420 p16_3 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_1, 8)), p16_3);
421 p16_4 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_2), p16_4);
422 p16_5 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_2, 8)), p16_5);
423 p16_6 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_3), p16_6);
424 p16_7 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_3, 8)), p16_7);
427 sumi_0 = _mm_add_epi32(sumi_0, _mm_add_epi32(p16_0, p16_2));
428 sumi_1 = _mm_add_epi32(sumi_1, _mm_add_epi32(p16_1, p16_3));
429 sumi_0 = _mm_add_epi32(sumi_0, _mm_add_epi32(p16_4, p16_6));
430 sumi_1 = _mm_add_epi32(sumi_1, _mm_add_epi32(p16_5, p16_7));
434 sumi_0 = _mm_sub_epi32(sumi_0, q8sclsub_0);
435 sumi_1 = _mm_sub_epi32(sumi_1, q8sclsub_1);
438 __m128i sumi = _mm_add_epi32(sumi_0, sumi_1);
439 __m128 sumf_vec = _mm_mul_ps(_mm_set1_ps(d), _mm_cvtepi32_ps(sumi));
442 sumf_vec = _mm_hadd_ps(sumf_vec, sumf_vec);
443 sumf_vec = _mm_hadd_ps(sumf_vec, sumf_vec);
444 acc = _mm_add_ss(acc, sumf_vec);
447 return _mm_cvtss_f32(acc);
455 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
461 const int blocks_per_row = K /
QK_K;
463 for (
int row = 0; row < M; ++row) {
464 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
465 y[row] = dot_q6_k_q8_k_sse(w_row, x, K);
477#if defined(__AVX__) && !defined(__AVX2__)
479static float dot_q6_k_q8_k_avx(
const block_q6_K *w,
483 const int nb = K /
QK_K;
484 const __m128i m3 = _mm_set1_epi8(3);
485 const __m128i m15 = _mm_set1_epi8(15);
487 __m128 acc = _mm_setzero_ps();
489 for (
int i = 0; i < nb; ++i) {
492 const uint8_t *ql = w[i].
ql;
493 const uint8_t *qh = w[i].
qh;
494 const int8_t *q8 = x[i].
qs;
498 _mm_prefetch((
const char *)&w[i + 1], _MM_HINT_T0);
499 _mm_prefetch((
const char *)&x[i + 1], _MM_HINT_T0);
503 const __m128i scales = _mm_loadu_si128((
const __m128i *)w[i].scales);
504 const __m128i q8sums_0 = _mm_loadu_si128((
const __m128i *)x[i].bsums);
505 const __m128i q8sums_1 = _mm_loadu_si128((
const __m128i *)x[i].bsums + 1);
508 const __m128i scales_16_0 = _mm_cvtepi8_epi16(scales);
509 const __m128i scales_16_1 = _mm_cvtepi8_epi16(_mm_bsrli_si128(scales, 8));
510 const __m128i q8sclsub_0 = _mm_slli_epi32(_mm_madd_epi16(q8sums_0, scales_16_0), 5);
511 const __m128i q8sclsub_1 = _mm_slli_epi32(_mm_madd_epi16(q8sums_1, scales_16_1), 5);
513 __m128i sumi_0 = _mm_setzero_si128();
514 __m128i sumi_1 = _mm_setzero_si128();
519 for (
int j = 0; j <
QK_K / 128; ++j) {
521 const __m128i q4bitsH_0 = _mm_loadu_si128((
const __m128i *)qh);
523 const __m128i q4bitsH_1 = _mm_loadu_si128((
const __m128i *)qh);
527 const __m128i q4h_0 = _mm_slli_epi16(_mm_and_si128(q4bitsH_0, m3), 4);
528 const __m128i q4h_1 = _mm_slli_epi16(_mm_and_si128(q4bitsH_1, m3), 4);
529 const __m128i q4h_2 = _mm_slli_epi16(_mm_and_si128(q4bitsH_0, _mm_set1_epi8(12)), 2);
530 const __m128i q4h_3 = _mm_slli_epi16(_mm_and_si128(q4bitsH_1, _mm_set1_epi8(12)), 2);
531 const __m128i q4h_4 = _mm_and_si128(q4bitsH_0, _mm_set1_epi8(48));
532 const __m128i q4h_5 = _mm_and_si128(q4bitsH_1, _mm_set1_epi8(48));
533 const __m128i q4h_6 = _mm_srli_epi16(_mm_and_si128(q4bitsH_0, _mm_set1_epi8(-64)), 2);
534 const __m128i q4h_7 = _mm_srli_epi16(_mm_and_si128(q4bitsH_1, _mm_set1_epi8(-64)), 2);
537 const __m128i q4bits1_0 = _mm_loadu_si128((
const __m128i *)ql);
539 const __m128i q4bits1_1 = _mm_loadu_si128((
const __m128i *)ql);
541 const __m128i q4bits2_0 = _mm_loadu_si128((
const __m128i *)ql);
543 const __m128i q4bits2_1 = _mm_loadu_si128((
const __m128i *)ql);
547 const __m128i q4_0 = _mm_or_si128(_mm_and_si128(q4bits1_0, m15), q4h_0);
548 const __m128i q4_1 = _mm_or_si128(_mm_and_si128(q4bits1_1, m15), q4h_1);
549 const __m128i q4_2 = _mm_or_si128(_mm_and_si128(q4bits2_0, m15), q4h_2);
550 const __m128i q4_3 = _mm_or_si128(_mm_and_si128(q4bits2_1, m15), q4h_3);
551 const __m128i q4_4 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits1_0, 4), m15), q4h_4);
552 const __m128i q4_5 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits1_1, 4), m15), q4h_5);
553 const __m128i q4_6 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits2_0, 4), m15), q4h_6);
554 const __m128i q4_7 = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(q4bits2_1, 4), m15), q4h_7);
557 const __m128i q8_0 = _mm_loadu_si128((
const __m128i *)q8);
559 const __m128i q8_1 = _mm_loadu_si128((
const __m128i *)q8);
561 const __m128i q8_2 = _mm_loadu_si128((
const __m128i *)q8);
563 const __m128i q8_3 = _mm_loadu_si128((
const __m128i *)q8);
565 const __m128i q8_4 = _mm_loadu_si128((
const __m128i *)q8);
567 const __m128i q8_5 = _mm_loadu_si128((
const __m128i *)q8);
569 const __m128i q8_6 = _mm_loadu_si128((
const __m128i *)q8);
571 const __m128i q8_7 = _mm_loadu_si128((
const __m128i *)q8);
575 __m128i p16_0 = _mm_maddubs_epi16(q4_0, q8_0);
576 __m128i p16_1 = _mm_maddubs_epi16(q4_1, q8_1);
577 __m128i p16_2 = _mm_maddubs_epi16(q4_2, q8_2);
578 __m128i p16_3 = _mm_maddubs_epi16(q4_3, q8_3);
579 __m128i p16_4 = _mm_maddubs_epi16(q4_4, q8_4);
580 __m128i p16_5 = _mm_maddubs_epi16(q4_5, q8_5);
581 __m128i p16_6 = _mm_maddubs_epi16(q4_6, q8_6);
582 __m128i p16_7 = _mm_maddubs_epi16(q4_7, q8_7);
585 const __m128i scale_0 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 0]));
586 const __m128i scale_1 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 1]));
587 const __m128i scale_2 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 2]));
588 const __m128i scale_3 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)q6k_scale_shuffle[is + 3]));
592 p16_0 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_0), p16_0);
593 p16_1 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_0, 8)), p16_1);
594 p16_2 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_1), p16_2);
595 p16_3 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_1, 8)), p16_3);
596 p16_4 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_2), p16_4);
597 p16_5 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_2, 8)), p16_5);
598 p16_6 = _mm_madd_epi16(_mm_cvtepi8_epi16(scale_3), p16_6);
599 p16_7 = _mm_madd_epi16(_mm_cvtepi8_epi16(_mm_bsrli_si128(scale_3, 8)), p16_7);
602 sumi_0 = _mm_add_epi32(sumi_0, _mm_add_epi32(p16_0, p16_2));
603 sumi_1 = _mm_add_epi32(sumi_1, _mm_add_epi32(p16_1, p16_3));
604 sumi_0 = _mm_add_epi32(sumi_0, _mm_add_epi32(p16_4, p16_6));
605 sumi_1 = _mm_add_epi32(sumi_1, _mm_add_epi32(p16_5, p16_7));
609 sumi_0 = _mm_sub_epi32(sumi_0, q8sclsub_0);
610 sumi_1 = _mm_sub_epi32(sumi_1, q8sclsub_1);
613 __m128i sumi = _mm_add_epi32(sumi_0, sumi_1);
614 __m128 sumf_vec = _mm_mul_ps(_mm_set1_ps(d), _mm_cvtepi32_ps(sumi));
617 sumf_vec = _mm_hadd_ps(sumf_vec, sumf_vec);
618 sumf_vec = _mm_hadd_ps(sumf_vec, sumf_vec);
619 acc = _mm_add_ss(acc, sumf_vec);
622 return _mm_cvtss_f32(acc);
630 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
636 const int blocks_per_row = K /
QK_K;
638 for (
int row = 0; row < M; ++row) {
639 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
640 y[row] = dot_q6_k_q8_k_avx(w_row, x, K);
655static inline float ck_q6k_hsum_float_8(
const __m256 x)
657 __m128 res = _mm256_extractf128_ps(x, 1);
658 res = _mm_add_ps(res, _mm256_castps256_ps128(x));
659 res = _mm_add_ps(res, _mm_movehl_ps(res, res));
660 res = _mm_add_ss(res, _mm_movehdup_ps(res));
661 return _mm_cvtss_f32(res);
677static inline __m128i get_scale_shuffle_avx2(
int i) {
678 static const uint8_t patterns[8][16] = {
679 { 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1 },
680 { 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3 },
681 { 4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5 },
682 { 6, 6, 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7 },
683 { 8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9 },
684 {10,10,10,10,10,10,10,10,11,11,11,11,11,11,11,11 },
685 {12,12,12,12,12,12,12,12,13,13,13,13,13,13,13,13 },
686 {14,14,14,14,14,14,14,14,15,15,15,15,15,15,15,15 },
688 return _mm_loadu_si128((
const __m128i *)patterns[i]);
691static float dot_q6_k_q8_k_avx2(
const block_q6_K *w,
695 const int nb = K /
QK_K;
696 const __m256i m3 = _mm256_set1_epi8(3);
697 const __m256i m15 = _mm256_set1_epi8(15);
699 __m256 acc = _mm256_setzero_ps();
701 for (
int i = 0; i < nb; ++i) {
704 const uint8_t *q4 = w[i].
ql;
705 const uint8_t *qh = w[i].
qh;
706 const int8_t *q8 = x[i].
qs;
708 const __m256i q8sums = _mm256_loadu_si256((
const __m256i *)x[i].bsums);
709 const __m128i scales = _mm_loadu_si128((
const __m128i *)w[i].scales);
710 const __m256i scales_16 = _mm256_cvtepi8_epi16(scales);
711 const __m256i q8sclsub = _mm256_slli_epi32(
712 _mm256_madd_epi16(q8sums, scales_16), 5);
714 __m256i sumi = _mm256_setzero_si256();
717 for (
int j = 0; j <
QK_K / 128; ++j) {
718 const __m256i q4bits1 = _mm256_loadu_si256((
const __m256i *)q4);
720 const __m256i q4bits2 = _mm256_loadu_si256((
const __m256i *)q4);
722 const __m256i q4bitsH = _mm256_loadu_si256((
const __m256i *)qh);
725 const __m256i q4h_0 = _mm256_slli_epi16(_mm256_and_si256(q4bitsH, m3), 4);
726 const __m256i q4h_1 = _mm256_slli_epi16(
727 _mm256_and_si256(q4bitsH, _mm256_set1_epi8(12)), 2);
728 const __m256i q4h_2 = _mm256_and_si256(q4bitsH, _mm256_set1_epi8(48));
729 const __m256i q4h_3 = _mm256_srli_epi16(
730 _mm256_and_si256(q4bitsH, _mm256_set1_epi8(-64)), 2);
732 const __m256i q4_0 = _mm256_or_si256(_mm256_and_si256(q4bits1, m15), q4h_0);
733 const __m256i q4_1 = _mm256_or_si256(_mm256_and_si256(q4bits2, m15), q4h_1);
734 const __m256i q4_2 = _mm256_or_si256(
735 _mm256_and_si256(_mm256_srli_epi16(q4bits1, 4), m15), q4h_2);
736 const __m256i q4_3 = _mm256_or_si256(
737 _mm256_and_si256(_mm256_srli_epi16(q4bits2, 4), m15), q4h_3);
739 const __m256i q8_0 = _mm256_loadu_si256((
const __m256i *)q8);
741 const __m256i q8_1 = _mm256_loadu_si256((
const __m256i *)q8);
743 const __m256i q8_2 = _mm256_loadu_si256((
const __m256i *)q8);
745 const __m256i q8_3 = _mm256_loadu_si256((
const __m256i *)q8);
748 __m256i p16_0 = _mm256_maddubs_epi16(q4_0, q8_0);
749 __m256i p16_1 = _mm256_maddubs_epi16(q4_1, q8_1);
750 __m256i p16_2 = _mm256_maddubs_epi16(q4_2, q8_2);
751 __m256i p16_3 = _mm256_maddubs_epi16(q4_3, q8_3);
753 const __m128i scale_0 = _mm_shuffle_epi8(
754 scales, get_scale_shuffle_avx2(is + 0));
755 const __m128i scale_1 = _mm_shuffle_epi8(
756 scales, get_scale_shuffle_avx2(is + 1));
757 const __m128i scale_2 = _mm_shuffle_epi8(
758 scales, get_scale_shuffle_avx2(is + 2));
759 const __m128i scale_3 = _mm_shuffle_epi8(
760 scales, get_scale_shuffle_avx2(is + 3));
763 p16_0 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_0), p16_0);
764 p16_1 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_1), p16_1);
765 p16_2 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_2), p16_2);
766 p16_3 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_3), p16_3);
768 sumi = _mm256_add_epi32(sumi, _mm256_add_epi32(p16_0, p16_1));
769 sumi = _mm256_add_epi32(sumi, _mm256_add_epi32(p16_2, p16_3));
772 sumi = _mm256_sub_epi32(sumi, q8sclsub);
773 acc = _mm256_fmadd_ps(_mm256_broadcast_ss(&d), _mm256_cvtepi32_ps(sumi), acc);
776 return ck_q6k_hsum_float_8(acc);
788static void dot_q6_k_q8_k_avx2_m4(
const block_q6_K *w,
795 const int nb = K /
QK_K;
796 const __m256i m3 = _mm256_set1_epi8(3);
797 const __m256i m15 = _mm256_set1_epi8(15);
799 _mm256_setzero_ps(), _mm256_setzero_ps(),
800 _mm256_setzero_ps(), _mm256_setzero_ps()
803 for (
int i = 0; i < nb; ++i) {
804 const uint8_t *q4 = w[i].
ql;
805 const uint8_t *qh = w[i].
qh;
806 const __m128i scales = _mm_loadu_si128((
const __m128i *)w[i].scales);
807 const __m256i scales_16 = _mm256_cvtepi8_epi16(scales);
810 _mm256_setzero_si256(), _mm256_setzero_si256(),
811 _mm256_setzero_si256(), _mm256_setzero_si256()
815 for (
int j = 0; j <
QK_K / 128; ++j) {
816 const __m256i q4bits1 = _mm256_loadu_si256((
const __m256i *)q4);
818 const __m256i q4bits2 = _mm256_loadu_si256((
const __m256i *)q4);
820 const __m256i q4bitsH = _mm256_loadu_si256((
const __m256i *)qh);
823 const __m256i q4h_0 = _mm256_slli_epi16(_mm256_and_si256(q4bitsH, m3), 4);
824 const __m256i q4h_1 = _mm256_slli_epi16(
825 _mm256_and_si256(q4bitsH, _mm256_set1_epi8(12)), 2);
826 const __m256i q4h_2 = _mm256_and_si256(q4bitsH, _mm256_set1_epi8(48));
827 const __m256i q4h_3 = _mm256_srli_epi16(
828 _mm256_and_si256(q4bitsH, _mm256_set1_epi8(-64)), 2);
830 const __m256i q6_0 = _mm256_or_si256(_mm256_and_si256(q4bits1, m15), q4h_0);
831 const __m256i q6_1 = _mm256_or_si256(_mm256_and_si256(q4bits2, m15), q4h_1);
832 const __m256i q6_2 = _mm256_or_si256(
833 _mm256_and_si256(_mm256_srli_epi16(q4bits1, 4), m15), q4h_2);
834 const __m256i q6_3 = _mm256_or_si256(
835 _mm256_and_si256(_mm256_srli_epi16(q4bits2, 4), m15), q4h_3);
837 const __m128i scale_0 = _mm_shuffle_epi8(
838 scales, get_scale_shuffle_avx2(is + 0));
839 const __m128i scale_1 = _mm_shuffle_epi8(
840 scales, get_scale_shuffle_avx2(is + 1));
841 const __m128i scale_2 = _mm_shuffle_epi8(
842 scales, get_scale_shuffle_avx2(is + 2));
843 const __m128i scale_3 = _mm_shuffle_epi8(
844 scales, get_scale_shuffle_avx2(is + 3));
846 const __m256i scale16_0 = _mm256_cvtepi8_epi16(scale_0);
847 const __m256i scale16_1 = _mm256_cvtepi8_epi16(scale_1);
848 const __m256i scale16_2 = _mm256_cvtepi8_epi16(scale_2);
849 const __m256i scale16_3 = _mm256_cvtepi8_epi16(scale_3);
851 for (
int r = 0; r < rows; ++r) {
852 const block_q8_K *xb = x + (size_t)r * (
size_t)x_row_blocks + i;
853 const int8_t *q8 = xb->
qs + j * 128;
854 const __m256i q8_0 = _mm256_loadu_si256((
const __m256i *)(q8 + 0));
855 const __m256i q8_1 = _mm256_loadu_si256((
const __m256i *)(q8 + 32));
856 const __m256i q8_2 = _mm256_loadu_si256((
const __m256i *)(q8 + 64));
857 const __m256i q8_3 = _mm256_loadu_si256((
const __m256i *)(q8 + 96));
859 __m256i p16_0 = _mm256_maddubs_epi16(q6_0, q8_0);
860 __m256i p16_1 = _mm256_maddubs_epi16(q6_1, q8_1);
861 __m256i p16_2 = _mm256_maddubs_epi16(q6_2, q8_2);
862 __m256i p16_3 = _mm256_maddubs_epi16(q6_3, q8_3);
863 p16_0 = _mm256_madd_epi16(scale16_0, p16_0);
864 p16_1 = _mm256_madd_epi16(scale16_1, p16_1);
865 p16_2 = _mm256_madd_epi16(scale16_2, p16_2);
866 p16_3 = _mm256_madd_epi16(scale16_3, p16_3);
868 sumi[r] = _mm256_add_epi32(
869 sumi[r], _mm256_add_epi32(p16_0, p16_1));
870 sumi[r] = _mm256_add_epi32(
871 sumi[r], _mm256_add_epi32(p16_2, p16_3));
875 for (
int r = 0; r < rows; ++r) {
876 const block_q8_K *xb = x + (size_t)r * (
size_t)x_row_blocks + i;
877 const __m256i q8sums =
878 _mm256_loadu_si256((
const __m256i *)xb->
bsums);
879 const __m256i q8sclsub = _mm256_slli_epi32(
880 _mm256_madd_epi16(q8sums, scales_16), 5);
881 sumi[r] = _mm256_sub_epi32(sumi[r], q8sclsub);
882 const float d = wd * xb->
d;
883 acc[r] = _mm256_fmadd_ps(
884 _mm256_broadcast_ss(&d), _mm256_cvtepi32_ps(sumi[r]), acc[r]);
888 for (
int r = 0; r < rows; ++r) {
889 out[r] = ck_q6k_hsum_float_8(acc[r]);
893static float dot_q6_k_prepared_q8_k_avx2(
894 const block_q6_K_prepared *w,
const block_q8_K *x,
int K)
896 const int nb = K /
QK_K;
897 __m256 acc = _mm256_setzero_ps();
899 for (
int i = 0; i < nb; ++i) {
900 const __m128i scales =
901 _mm_loadu_si128((
const __m128i *)(
const void *)w[i].scales);
902 const __m256i scales_16 = _mm256_cvtepi8_epi16(scales);
903 const __m256i q8sums =
904 _mm256_loadu_si256((
const __m256i *)(
const void *)x[i].bsums);
905 const __m256i q8sclsub = _mm256_slli_epi32(
906 _mm256_madd_epi16(q8sums, scales_16), 5);
907 __m256i sumi = _mm256_setzero_si256();
909 for (
int j = 0; j < 2; ++j) {
911 for (
int lane = 0; lane < 4; ++lane) {
912 const int group = j * 4 + lane;
913 const __m256i q6 = _mm256_loadu_si256(
914 (
const __m256i *)(
const void *)(w[i].qs + group * 32));
915 const __m256i q8 = _mm256_loadu_si256(
916 (
const __m256i *)(
const void *)(x[i].qs + group * 32));
917 p16[lane] = _mm256_maddubs_epi16(q6, q8);
918 const __m128i scale = _mm_shuffle_epi8(
919 scales, get_scale_shuffle_avx2(group));
920 p16[lane] = _mm256_madd_epi16(
921 _mm256_cvtepi8_epi16(scale), p16[lane]);
923 sumi = _mm256_add_epi32(
924 sumi, _mm256_add_epi32(p16[0], p16[1]));
925 sumi = _mm256_add_epi32(
926 sumi, _mm256_add_epi32(p16[2], p16[3]));
929 sumi = _mm256_sub_epi32(sumi, q8sclsub);
931 acc = _mm256_fmadd_ps(
932 _mm256_broadcast_ss(&d), _mm256_cvtepi32_ps(sumi), acc);
935 return ck_q6k_hsum_float_8(acc);
938#if defined(__AVX512F__) && defined(__AVX512BW__) && \
939 defined(__AVX512VNNI__)
946static float dot_q6_k_prepared_q8_k_avx512_vnni(
947 const block_q6_K_prepared *w,
const block_q8_K *x,
int K)
949 const int nb = K /
QK_K;
950 const __m512i scale_indices = _mm512_setr_epi32(
951 0, 0, 0, 0, 1, 1, 1, 1,
952 2, 2, 2, 2, 3, 3, 3, 3);
953 __m256 acc = _mm256_setzero_ps();
955 for (
int i = 0; i < nb; ++i) {
956 const __m128i scales =
957 _mm_loadu_si128((
const __m128i *)(
const void *)w[i].scales);
958 const __m512i scales_32 = _mm512_cvtepi8_epi32(scales);
959 const __m256i scales_16 = _mm256_cvtepi8_epi16(scales);
960 const __m256i q8sums =
961 _mm256_loadu_si256((
const __m256i *)(
const void *)x[i].bsums);
962 const __m256i q8sclsub = _mm256_slli_epi32(
963 _mm256_madd_epi16(q8sums, scales_16), 5);
964 __m256i sumi = _mm256_setzero_si256();
966 for (
int chunk = 0; chunk < 4; ++chunk) {
967 const __m512i q6 = _mm512_loadu_si512(
968 (
const void *)(w[i].qs + chunk * 64));
969 const __m512i q8 = _mm512_loadu_si512(
970 (
const void *)(x[i].qs + chunk * 64));
971 const __m512i products = _mm512_dpbusd_epi32(
972 _mm512_setzero_si512(), q6, q8);
973 const __m512i indices = _mm512_add_epi32(
974 scale_indices, _mm512_set1_epi32(chunk * 4));
975 const __m512i scaled = _mm512_mullo_epi32(
976 products, _mm512_permutexvar_epi32(indices, scales_32));
977 const __m512i upper_half = _mm512_shuffle_i32x4(
978 scaled, scaled, _MM_SHUFFLE(3, 2, 3, 2));
979 const __m256i pair_sum = _mm256_add_epi32(
980 _mm512_castsi512_si256(scaled),
981 _mm512_castsi512_si256(upper_half));
982 sumi = _mm256_add_epi32(sumi, pair_sum);
985 sumi = _mm256_sub_epi32(sumi, q8sclsub);
987 acc = _mm256_fmadd_ps(
988 _mm256_broadcast_ss(&d), _mm256_cvtepi32_ps(sumi), acc);
991 return ck_q6k_hsum_float_8(acc);
1000 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
1006 const int blocks_per_row = K /
QK_K;
1008 for (
int row = 0; row < M; ++row) {
1009 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
1010 y[row] = dot_q6_k_q8_k_avx2(w_row, x, K);
1022#if defined(__AVX512F__) && defined(__AVX512BW__) && defined(__AVX512VBMI__)
1029static float dot_q6_k_q8_k_avx512_vbmi(
const block_q6_K *w,
1033 const int nb = K /
QK_K;
1034 const __m512i m4 = _mm512_set1_epi8(0xF);
1035 const __m512i m2 = _mm512_set1_epi8(3);
1036 const __m512i m32s = _mm512_set1_epi8(32);
1038 __m512 acc = _mm512_setzero_ps();
1040 for (
int i = 0; i < nb; ++i) {
1043 const uint8_t *ql = w[i].
ql;
1044 const uint8_t *qh = w[i].
qh;
1045 const int8_t *q8 = x[i].
qs;
1046 const int8_t *sc = w[i].
scales;
1048 __m512i sumi = _mm512_setzero_si512();
1052 const __m512i q4bits1 = _mm512_loadu_si512((
const __m512i *)ql);
1053 const __m512i q4bits2 = _mm512_loadu_si512((
const __m512i *)(ql + 64));
1056 const __m512i q4bitsH = _mm512_loadu_si512((
const __m512i *)qh);
1060 const __m512i q4h_0 = _mm512_slli_epi16(_mm512_and_si512(q4bitsH, m2), 4);
1062 const __m512i q4h_1 = _mm512_slli_epi16(_mm512_and_si512(_mm512_srli_epi16(q4bitsH, 2), m2), 4);
1064 const __m512i q4h_2 = _mm512_slli_epi16(_mm512_and_si512(_mm512_srli_epi16(q4bitsH, 4), m2), 4);
1066 const __m512i q4h_3 = _mm512_slli_epi16(_mm512_and_si512(_mm512_srli_epi16(q4bitsH, 6), m2), 4);
1070 const __m512i q6_0 = _mm512_or_si512(_mm512_and_si512(q4bits1, m4), q4h_0);
1071 const __m512i q6_1 = _mm512_or_si512(_mm512_and_si512(q4bits2, m4), q4h_1);
1073 const __m512i q6_2 = _mm512_or_si512(_mm512_and_si512(_mm512_srli_epi16(q4bits1, 4), m4), q4h_2);
1074 const __m512i q6_3 = _mm512_or_si512(_mm512_and_si512(_mm512_srli_epi16(q4bits2, 4), m4), q4h_3);
1077 const __m512i q8_0 = _mm512_loadu_si512((
const __m512i *)q8);
1078 const __m512i q8_1 = _mm512_loadu_si512((
const __m512i *)(q8 + 64));
1079 const __m512i q8_2 = _mm512_loadu_si512((
const __m512i *)(q8 + 128));
1080 const __m512i q8_3 = _mm512_loadu_si512((
const __m512i *)(q8 + 192));
1083 __m512i q8s_0 = _mm512_maddubs_epi16(m32s, q8_0);
1084 __m512i q8s_1 = _mm512_maddubs_epi16(m32s, q8_1);
1085 __m512i q8s_2 = _mm512_maddubs_epi16(m32s, q8_2);
1086 __m512i q8s_3 = _mm512_maddubs_epi16(m32s, q8_3);
1089 __m512i p16_0 = _mm512_maddubs_epi16(q6_0, q8_0);
1090 __m512i p16_1 = _mm512_maddubs_epi16(q6_1, q8_1);
1091 __m512i p16_2 = _mm512_maddubs_epi16(q6_2, q8_2);
1092 __m512i p16_3 = _mm512_maddubs_epi16(q6_3, q8_3);
1095 p16_0 = _mm512_sub_epi16(p16_0, q8s_0);
1096 p16_1 = _mm512_sub_epi16(p16_1, q8s_1);
1097 p16_2 = _mm512_sub_epi16(p16_2, q8s_2);
1098 p16_3 = _mm512_sub_epi16(p16_3, q8s_3);
1103 const __m128i scales_128 = _mm_loadu_si128((
const __m128i *)sc);
1107 const __m512i scale_idx_0 = _mm512_set_epi8(
1108 3,3,3,3,3,3,3,3,3,3,3,3,3,3,3,3,
1109 2,2,2,2,2,2,2,2,2,2,2,2,2,2,2,2,
1110 1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,1,
1111 0,0,0,0,0,0,0,0,0,0,0,0,0,0,0,0);
1112 const __m512i scale_idx_1 = _mm512_set_epi8(
1113 7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,7,
1114 6,6,6,6,6,6,6,6,6,6,6,6,6,6,6,6,
1115 5,5,5,5,5,5,5,5,5,5,5,5,5,5,5,5,
1116 4,4,4,4,4,4,4,4,4,4,4,4,4,4,4,4);
1117 const __m512i scale_idx_2 = _mm512_set_epi8(
1118 11,11,11,11,11,11,11,11,11,11,11,11,11,11,11,11,
1119 10,10,10,10,10,10,10,10,10,10,10,10,10,10,10,10,
1120 9,9,9,9,9,9,9,9,9,9,9,9,9,9,9,9,
1121 8,8,8,8,8,8,8,8,8,8,8,8,8,8,8,8);
1122 const __m512i scale_idx_3 = _mm512_set_epi8(
1123 15,15,15,15,15,15,15,15,15,15,15,15,15,15,15,15,
1124 14,14,14,14,14,14,14,14,14,14,14,14,14,14,14,14,
1125 13,13,13,13,13,13,13,13,13,13,13,13,13,13,13,13,
1126 12,12,12,12,12,12,12,12,12,12,12,12,12,12,12,12);
1129 const __m512i scales_512 = _mm512_broadcast_i32x4(scales_128);
1130 const __m512i sc_0 = _mm512_permutexvar_epi8(scale_idx_0, scales_512);
1131 const __m512i sc_1 = _mm512_permutexvar_epi8(scale_idx_1, scales_512);
1132 const __m512i sc_2 = _mm512_permutexvar_epi8(scale_idx_2, scales_512);
1133 const __m512i sc_3 = _mm512_permutexvar_epi8(scale_idx_3, scales_512);
1137 __m512i p32_0 = _mm512_madd_epi16(_mm512_cvtepi8_epi16(_mm512_castsi512_si256(sc_0)), p16_0);
1138 __m512i p32_1 = _mm512_madd_epi16(_mm512_cvtepi8_epi16(_mm512_castsi512_si256(sc_1)), p16_1);
1139 __m512i p32_2 = _mm512_madd_epi16(_mm512_cvtepi8_epi16(_mm512_castsi512_si256(sc_2)), p16_2);
1140 __m512i p32_3 = _mm512_madd_epi16(_mm512_cvtepi8_epi16(_mm512_castsi512_si256(sc_3)), p16_3);
1143 sumi = _mm512_add_epi32(sumi, p32_0);
1144 sumi = _mm512_add_epi32(sumi, p32_1);
1145 sumi = _mm512_add_epi32(sumi, p32_2);
1146 sumi = _mm512_add_epi32(sumi, p32_3);
1149 acc = _mm512_fmadd_ps(_mm512_set1_ps(d), _mm512_cvtepi32_ps(sumi), acc);
1152 return _mm512_reduce_add_ps(acc);
1160 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
1166 const int blocks_per_row = K /
QK_K;
1168 for (
int row = 0; row < M; ++row) {
1169 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
1170 y[row] = dot_q6_k_q8_k_avx512_vbmi(w_row, x, K);
1176#if defined(__AVX512F__) && defined(__AVX512BW__)
1185static float dot_q6_k_q8_k_avx512(
const block_q6_K *w,
1189 const int nb = K /
QK_K;
1190 const __m256i m4 = _mm256_set1_epi8(0xF);
1191 const __m256i m2 = _mm256_set1_epi8(3);
1192 const __m256i m32s = _mm256_set1_epi8(32);
1195 __m256 acc = _mm256_setzero_ps();
1197 for (
int i = 0; i < nb; ++i) {
1200 const uint8_t *q4 = w[i].
ql;
1201 const uint8_t *qh = w[i].
qh;
1202 const int8_t *q8 = x[i].
qs;
1204 const __m128i scales = _mm_loadu_si128((
const __m128i *)w[i].scales);
1207 __m256i sumi = _mm256_setzero_si256();
1211 for (
int j = 0; j <
QK_K / 128; ++j) {
1213 static const uint8_t patterns[8][16] = {
1214 { 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1 },
1215 { 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3 },
1216 { 4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5 },
1217 { 6, 6, 6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7 },
1218 { 8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9 },
1219 {10,10,10,10,10,10,10,10,11,11,11,11,11,11,11,11 },
1220 {12,12,12,12,12,12,12,12,13,13,13,13,13,13,13,13 },
1221 {14,14,14,14,14,14,14,14,15,15,15,15,15,15,15,15 },
1224 const __m128i scale_0 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)patterns[is + 0]));
1225 const __m128i scale_1 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)patterns[is + 1]));
1226 const __m128i scale_2 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)patterns[is + 2]));
1227 const __m128i scale_3 = _mm_shuffle_epi8(scales, _mm_loadu_si128((
const __m128i *)patterns[is + 3]));
1231 const __m256i q4bits1 = _mm256_loadu_si256((
const __m256i *)q4);
1233 const __m256i q4bits2 = _mm256_loadu_si256((
const __m256i *)q4);
1235 const __m256i q4bitsH = _mm256_loadu_si256((
const __m256i *)qh);
1239 const __m256i q4h_0 = _mm256_slli_epi16(_mm256_and_si256(q4bitsH, m2), 4);
1240 const __m256i q4h_1 = _mm256_slli_epi16(_mm256_and_si256(_mm256_srli_epi16(q4bitsH, 2), m2), 4);
1241 const __m256i q4h_2 = _mm256_slli_epi16(_mm256_and_si256(_mm256_srli_epi16(q4bitsH, 4), m2), 4);
1242 const __m256i q4h_3 = _mm256_slli_epi16(_mm256_and_si256(_mm256_srli_epi16(q4bitsH, 6), m2), 4);
1245 const __m256i q4_0 = _mm256_or_si256(_mm256_and_si256(q4bits1, m4), q4h_0);
1246 const __m256i q4_1 = _mm256_or_si256(_mm256_and_si256(q4bits2, m4), q4h_1);
1247 const __m256i q4_2 = _mm256_or_si256(_mm256_and_si256(_mm256_srli_epi16(q4bits1, 4), m4), q4h_2);
1248 const __m256i q4_3 = _mm256_or_si256(_mm256_and_si256(_mm256_srli_epi16(q4bits2, 4), m4), q4h_3);
1251 const __m256i q8_0 = _mm256_loadu_si256((
const __m256i *)q8);
1253 const __m256i q8_1 = _mm256_loadu_si256((
const __m256i *)q8);
1255 const __m256i q8_2 = _mm256_loadu_si256((
const __m256i *)q8);
1257 const __m256i q8_3 = _mm256_loadu_si256((
const __m256i *)q8);
1261 __m256i q8s_0 = _mm256_maddubs_epi16(m32s, q8_0);
1262 __m256i q8s_1 = _mm256_maddubs_epi16(m32s, q8_1);
1263 __m256i q8s_2 = _mm256_maddubs_epi16(m32s, q8_2);
1264 __m256i q8s_3 = _mm256_maddubs_epi16(m32s, q8_3);
1267 __m256i p16_0 = _mm256_maddubs_epi16(q4_0, q8_0);
1268 __m256i p16_1 = _mm256_maddubs_epi16(q4_1, q8_1);
1269 __m256i p16_2 = _mm256_maddubs_epi16(q4_2, q8_2);
1270 __m256i p16_3 = _mm256_maddubs_epi16(q4_3, q8_3);
1273 p16_0 = _mm256_sub_epi16(p16_0, q8s_0);
1274 p16_1 = _mm256_sub_epi16(p16_1, q8s_1);
1275 p16_2 = _mm256_sub_epi16(p16_2, q8s_2);
1276 p16_3 = _mm256_sub_epi16(p16_3, q8s_3);
1279 p16_0 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_0), p16_0);
1280 p16_1 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_1), p16_1);
1281 p16_2 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_2), p16_2);
1282 p16_3 = _mm256_madd_epi16(_mm256_cvtepi8_epi16(scale_3), p16_3);
1285 sumi = _mm256_add_epi32(sumi, _mm256_add_epi32(p16_0, p16_1));
1286 sumi = _mm256_add_epi32(sumi, _mm256_add_epi32(p16_2, p16_3));
1290 acc = _mm256_fmadd_ps(_mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi), acc);
1293 return ck_q6k_hsum_float_8(acc);
1301 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
1307 const int blocks_per_row = K /
QK_K;
1309 for (
int row = 0; row < M; ++row) {
1310 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
1311 y[row] = dot_q6_k_q8_k_avx512(w_row, x, K);
1326 if (!s || !vx || !vy || n <= 0) {
1351#if defined(__AVX2__)
1357#elif defined(__AVX__)
1360#elif defined(__SSE4_1__)
1370 return "q6_k_q8_k_ref";
1372#if defined(__AVX2__)
1373 return "q6_k_q8_k_avx2";
1374#elif defined(__AVX__)
1375 return "q6_k_q8_k_avx";
1376#elif defined(__SSE4_1__)
1377 return "q6_k_q8_k_sse";
1379 return "q6_k_q8_k_ref";
1404 if (!y || !W || !x_q8 || M <= 0 || K <= 0)
return;
1405 if (ith < 0 || nth <= 0 || ith >= nth)
return;
1408 const int dr = (M + nth - 1) / nth;
1409 const int r0 = dr * ith;
1410 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1412 if (r0 >= M)
return;
1416 const int blocks_per_row = K /
QK_K;
1418 for (
int row = r0; row < r1; ++row) {
1419 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
1436 if (!y || !W || !x_q8 || M <= 0 || K <= 0)
return;
1437 if (ith < 0 || nth <= 0 || ith >= nth)
return;
1439 const int dr = (M + nth - 1) / nth;
1440 const int r0 = dr * ith;
1441 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1443 if (r0 >= M)
return;
1447 const int blocks_per_row = K /
QK_K;
1450 for (
int row = r0; row < r1; ++row) {
1451 const block_q6_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
1452#if defined(__AVX2__)
1454 : dot_q6_k_q8_k_avx2(w_row, x, K);
1455#elif defined(__AVX__)
1457 : dot_q6_k_q8_k_avx(w_row, x, K);
1458#elif defined(__SSE4_1__)
1460 : dot_q6_k_q8_k_sse(w_row, x, K);
1484 int M,
int N,
int K)
1486 if (!Y || !W || !X_q8 || M <= 0 || N <= 0 || K <= 0) {
1491 const int blocks_per_vec = K /
QK_K;
1493 for (
int n = 0; n < N; ++n) {
1494 const block_q8_K *x_row = X + (size_t)n * (
size_t)blocks_per_vec;
1519 int M,
int N,
int K)
1521 if (!A_q8 || !B || !
C) {
1524 if (M <= 0 || N <= 0 || K <= 0) {
1534 const int blocks_per_vec = K /
QK_K;
1535 const int blocks_per_row = K /
QK_K;
1537 for (
int m = 0; m < M; ++m) {
1538 const block_q8_K *a_row = A + (size_t)m * (
size_t)blocks_per_vec;
1539 float *c_row =
C + (size_t)m * (
size_t)N;
1540 for (
int n = 0; n < N; ++n) {
1541 const block_q6_K *w_row = W + (size_t)n * (
size_t)blocks_per_row;
1542 const float b = bias ? bias[n] : 0.0f;
1550 const void *B_prepared,
1553 int M,
int N,
int K,
1556 int use_avx512_vnni)
1558#if !defined(__AVX2__)
1559 (void)A_q8; (void)B_prepared; (void)bias; (void)
C;
1560 (void)M; (void)N; (void)K; (void)m0; (void)m1; (void)n0; (void)n1;
1561 (void)use_avx512_vnni;
1563#if !defined(__AVX512F__) || !defined(__AVX512BW__) || \
1564 !defined(__AVX512VNNI__)
1565 (void)use_avx512_vnni;
1567 if (!A_q8 || !B_prepared || !
C || M <= 0 || N <= 0 || K <= 0 ||
1568 (K %
QK_K) != 0)
return;
1573 if (m0 >= m1 || n0 >= n1)
return;
1576 const block_q6_K_prepared *W =
1577 (
const block_q6_K_prepared *)B_prepared;
1578 const int blocks_per_row = K /
QK_K;
1579 for (
int n = n0; n < n1; ++n) {
1580 const block_q6_K_prepared *w_row =
1581 W + (size_t)n * (
size_t)blocks_per_row;
1582 const float b = bias ? bias[n] : 0.0f;
1583 for (
int m = m0; m < m1; ++m) {
1585 A + (size_t)m * (
size_t)blocks_per_row;
1586#if defined(__AVX512F__) && defined(__AVX512BW__) && \
1587 defined(__AVX512VNNI__)
1588 const float dot = use_avx512_vnni
1589 ? dot_q6_k_prepared_q8_k_avx512_vnni(w_row, a_row, K)
1590 : dot_q6_k_prepared_q8_k_avx2(w_row, a_row, K);
1592 const float dot = dot_q6_k_prepared_q8_k_avx2(w_row, a_row, K);
1594 C[(size_t)m * (
size_t)N + (size_t)n] =
1602 const void *B_prepared,
1605 int M,
int N,
int K,
1609 int use_avx512_vnni = 0;
1610#if defined(__AVX512F__) && defined(__AVX512BW__) && \
1611 defined(__AVX512VNNI__)
1612 use_avx512_vnni = 1;
1615 A_q8, B_prepared, bias,
C, M, N, K,
1616 m0, m1, n0, n1, use_avx512_vnni);
1621 const void *B_prepared,
1624 int M,
int N,
int K)
1627 A_q8, B_prepared, bias,
C, M, N, K,
1632 const void *B_prepared,
1635 int M,
int N,
int K)
1638 A_q8, B_prepared, bias,
C, M, N, K, 0, M, 0, N);
1648#if defined(__AVX2__)
1649 return dot_q6_k_q8_k_avx2(w, x, K);
1650#elif defined(__AVX__)
1651 return dot_q6_k_q8_k_avx(w, x, K);
1652#elif defined(__SSE4_1__)
1653 return dot_q6_k_q8_k_sse(w, x, K);
1669 int M,
int N,
int K,
1673 if (!A_q8 || !B || !
C) {
1676 if (M <= 0 || N <= 0 || K <= 0 || K %
QK_K != 0) {
1683 if (m0 >= m1 || n0 >= n1) {
1689 const int blocks_per_vec = K /
QK_K;
1690 const int blocks_per_row = K /
QK_K;
1692 for (
int n = n0; n < n1; ++n) {
1693 const block_q6_K *w_row = W + (size_t)n * (
size_t)blocks_per_row;
1694 const float b = bias ? bias[n] : 0.0f;
1695 for (
int m = m0; m < m1; ++m) {
1696 const block_q8_K *a_row = A + (size_t)m * (
size_t)blocks_per_vec;
1706 int M,
int N,
int K,
1710 if (!A_q8 || !B || !
C || M <= 0 || N <= 0 || K <= 0 ||
1718 if (m0 >= m1 || n0 >= n1)
return;
1720#if defined(__AVX2__)
1724 const int blocks_per_vec = K /
QK_K;
1725 for (
int n = n0; n < n1; ++n) {
1727 W + (size_t)n * (
size_t)blocks_per_vec;
1728 const float b = bias ? bias[n] : 0.0f;
1730 for (; m + 4 <= m1; m += 4) {
1732 dot_q6_k_q8_k_avx2_m4(
1733 w_row, A + (
size_t)m * (
size_t)blocks_per_vec,
1734 blocks_per_vec, 4, K, values);
1735 for (
int r = 0; r < 4; ++r) {
1736 C[(size_t)(m + r) * (size_t)N + (
size_t)n] = values[r] + b;
1741 const int rows = m1 - m;
1742 dot_q6_k_q8_k_avx2_m4(
1743 w_row, A + (
size_t)m * (
size_t)blocks_per_vec,
1744 blocks_per_vec, rows, K, values);
1745 for (
int r = 0; r < rows; ++r) {
1746 C[(size_t)(m + r) * (size_t)N + (
size_t)n] = values[r] + b;
1754 A_q8, B, bias,
C, M, N, K, m0, m1, n0, n1);
1767 int M,
int N,
int K)
1769 enum { TILE_M = 8, TILE_N = 16 };
1770 for (
int n0 = 0; n0 < N; n0 += TILE_N) {
1771 const int n1 = (n0 + TILE_N < N) ? (n0 + TILE_N) : N;
1772 for (
int m0 = 0; m0 < M; m0 += TILE_M) {
1773 const int m1 = (m0 + TILE_M < M) ? (m0 + TILE_M) : M;
1774 gemm_nt_q6_k_q8_k_tile(A_q8, B, bias,
C, M, N, K, m0, m1, n0, n1);
int ck_strict_parity_enabled(void)
Quantization block structures for weight-only quantization.
#define GGML_FP16_TO_FP32
void gemm_nt_q6_k_q8_k_tile(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K, int m0, int m1, int n0, int n1)
Compute one C[m0:m1, n0:n1] tile for Q8_K activations x Q6_K weights.
void gemv_q6_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q6_k_q8_k_sse(float *y, const void *W, const void *x_q8, int M, int K)
size_t ck_q6_k_prepared_block_size(void)
void gemv_q6_k_q8_k_avx(float *y, const void *W, const void *x_q8, int M, int K)
const char * ck_q6_k_q8_k_provider_name(void)
void gemm_nt_q6_k_q8_k_m4_tile(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K, int m0, int m1, int n0, int n1)
static float ck_dot_q6_k_q8_k_fast_or_ref(const block_q6_K *w, const block_q8_K *x, int K)
void gemm_q6_k_q8_k(float *Y, const void *W, const void *X_q8, int M, int N, int K)
GEMM: Y = W @ X^T where W is Q6_K and X is Q8_K.
static void gemm_nt_q6_k_q8_k_prepared_tile_impl(const void *A_q8, const void *B_prepared, const float *bias, float *C, int M, int N, int K, int m0, int m1, int n0, int n1, int use_avx512_vnni)
void gemv_q6_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
GEMV: y = W @ x where W is Q6_K and x is Q8_K.
void gemv_q6_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
void vec_dot_q6_k_q8_k(int n, float *s, const void *vx, const void *vy)
Q6_K x Q8_K dot product (single row)
void gemm_nt_q6_k_q8_k_prepared_avx512_vnni(const void *A_q8, const void *B_prepared, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q6_k_q8_k_prepared(const void *A_q8, const void *B_prepared, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q6_k_q8_k_tiled(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
Experimental single-thread tiled NT GEMM wrapper.
void gemv_q6_k_q8_k_avx512_vbmi(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q6_k_q8_k_parallel(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
Parallel reference GEMV for Q6_K × Q8_K.
void gemm_nt_q6_k_q8_k_prepared_tile(const void *A_q8, const void *B_prepared, const float *bias, float *C, int M, int N, int K, int m0, int m1, int n0, int n1)
static int ck_q6k_q8k_force_ref(void)
void ck_q6_k_prepare_weight(const void *src, void *dst, int N, int K)
const char * ck_q6_k_prepared_provider_name(void)
void gemm_nt_q6_k_q8_k(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
NT GEMM: C = A @ B^T where A is Q8_K and B is Q6_K.
void gemv_q6_k_q8_k_avx512(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q6_k_q8_k_parallel_simd(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
Parallel SIMD GEMV for Q6_K × Q8_K.
static float dot_q6_k_q8_k_ref(const block_q6_K *w, const block_q8_K *x, int K)
Scalar dot product for Q6_K x Q8_K.