40#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__)
46#define CK_Q5K_STACK_Q8_BLOCKS 128
50 static int cached = -1;
52 const char *env = getenv(
"CK_DEBUG_Q5K_FP32_FALLBACK");
53 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
60 static int cached = -1;
62 const char *env = getenv(
"CK_DEBUG_Q5K_GENERIC_DOT");
63 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
90_Static_assert(
sizeof(block_q5_K_prepared) == 276,
91 "Q5_K prepared-size contract changed");
98 sc[0] = scales[0] & 0x3F;
99 sc[1] = scales[1] & 0x3F;
100 sc[2] = scales[2] & 0x3F;
101 sc[3] = scales[3] & 0x3F;
103 m[0] = scales[4] & 0x3F;
104 m[1] = scales[5] & 0x3F;
105 m[2] = scales[6] & 0x3F;
106 m[3] = scales[7] & 0x3F;
108 sc[4] = (scales[8] & 0x0F) | ((scales[0] >> 6) << 4);
109 sc[5] = (scales[9] & 0x0F) | ((scales[1] >> 6) << 4);
110 sc[6] = (scales[10] & 0x0F) | ((scales[2] >> 6) << 4);
111 sc[7] = (scales[11] & 0x0F) | ((scales[3] >> 6) << 4);
113 m[4] = (scales[8] >> 4) | ((scales[4] >> 6) << 4);
114 m[5] = (scales[9] >> 4) | ((scales[5] >> 6) << 4);
115 m[6] = (scales[10] >> 4) | ((scales[6] >> 6) << 4);
116 m[7] = (scales[11] >> 4) | ((scales[7] >> 6) << 4);
120 const uint8_t *ql = block->qs + (subblock / 2) * 32;
121 const uint8_t low = (subblock & 1) ? (uint8_t)(ql[i] >> 4) : (uint8_t)(ql[i] & 0x0F);
122 const uint8_t high = (block->qh[i] & (uint8_t)(1u << subblock)) ? 16u : 0u;
123 return (uint8_t)(low | high);
128 return sizeof(block_q5_K_prepared);
133 if (!src || !dst || N <= 0 || K <= 0 || (K %
QK_K) != 0)
return;
134 const block_q5_K *input = (
const block_q5_K *)src;
135 block_q5_K_prepared *output = (block_q5_K_prepared *)dst;
136 const size_t blocks = (size_t)N * (
size_t)(K /
QK_K);
137 for (
size_t b = 0; b < blocks; ++b) {
138 output[b].d = input[b].d;
139 output[b].dmin = input[b].dmin;
141 for (
int sb = 0; sb < 8; ++sb) {
142 for (
int i = 0; i < 32; ++i) {
153static inline __m256i ck_q5k_scale_shuffle_avx2(
int i)
155 static const uint8_t k_shuffle[256] = {
156 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1,
157 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3,
158 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5,
159 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7,
160 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9,
161 10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,
162 12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,
163 14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15
165 return _mm256_loadu_si256((
const __m256i *)(
const void *)(k_shuffle + 32 * i));
168static inline __m256i ck_mm256_set_m128i(__m128i hi, __m128i lo)
170 return _mm256_inserti128_si256(_mm256_castsi128_si256(lo), hi, 1);
173static inline float ck_q5k_hsum256_ps(__m256 v)
175 __m128 sum = _mm256_extractf128_ps(v, 1);
176 sum = _mm_add_ps(sum, _mm256_castps256_ps128(v));
177 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
178 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
179 return _mm_cvtss_f32(sum);
182static float dot_q5_k_q8_k_row_avx2(
const block_q5_K *w,
const block_q8_K *x,
int nb) {
183 static const uint32_t kmask1 = 0x3f3f3f3fU;
184 static const uint32_t kmask2 = 0x0f0f0f0fU;
185 static const uint32_t kmask3 = 0x03030303U;
187 const __m256i m4 = _mm256_set1_epi8(0x0f);
188 const __m128i mzero = _mm_setzero_si128();
189 const __m256i mone = _mm256_set1_epi8(1);
191 uint32_t utmp[4] = {0, 0, 0, 0};
192 __m256 acc = _mm256_setzero_ps();
195 for (
int b = 0; b < nb; ++b) {
196 const block_q5_K *wb = &w[b];
198 const uint8_t *q5 = wb->
qs;
199 const int8_t *q8 = xb->
qs;
204 memcpy(utmp, wb->scales, 12);
205 utmp[3] = ((utmp[2] >> 4) & kmask2) | (((utmp[1] >> 6) & kmask3) << 4);
206 const uint32_t uaux = utmp[1] & kmask1;
207 utmp[1] = (utmp[2] & kmask2) | (((utmp[0] >> 6) & kmask3) << 4);
211 const __m256i mins_and_scales =
212 _mm256_cvtepu8_epi16(_mm_set_epi32((
int)utmp[3], (
int)utmp[2], (
int)utmp[1], (
int)utmp[0]));
214 const __m256i q8sums = _mm256_loadu_si256((
const __m256i *)(
const void *)xb->
bsums);
215 const __m128i q8s = _mm_hadd_epi16(_mm256_extracti128_si256(q8sums, 0),
216 _mm256_extracti128_si256(q8sums, 1));
217 const __m128i prod = _mm_madd_epi16(_mm256_extracti128_si256(mins_and_scales, 1), q8s);
218 const __m128i hsum = _mm_hadd_epi32(_mm_hadd_epi32(prod, mzero), mzero);
219 summs += dmin * (float)_mm_extract_epi32(hsum, 0);
221 const __m128i sc128 = _mm256_extracti128_si256(mins_and_scales, 0);
222 const __m256i scales = ck_mm256_set_m128i(sc128, sc128);
223 const __m256i hbits = _mm256_loadu_si256((
const __m256i *)(
const void *)wb->qh);
224 __m256i hmask = mone;
225 __m256i sumi = _mm256_setzero_si256();
228 for (
int j = 0; j <
QK_K / 64; ++j) {
229 const __m256i scale_0 = _mm256_shuffle_epi8(scales, ck_q5k_scale_shuffle_avx2(2 * j + 0));
230 const __m256i scale_1 = _mm256_shuffle_epi8(scales, ck_q5k_scale_shuffle_avx2(2 * j + 1));
232 const __m256i q5bits = _mm256_loadu_si256((
const __m256i *)(
const void *)q5);
235 const __m256i q5l_0 = _mm256_and_si256(q5bits, m4);
236 const __m256i q5h_0 = _mm256_slli_epi16(_mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
237 const __m256i q5_0 = _mm256_add_epi8(q5l_0, q5h_0);
238 hmask = _mm256_slli_epi16(hmask, 1);
240 const __m256i q5l_1 = _mm256_and_si256(_mm256_srli_epi16(q5bits, 4), m4);
241 const __m256i q5h_1 = _mm256_slli_epi16(_mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
242 const __m256i q5_1 = _mm256_add_epi8(q5l_1, q5h_1);
243 hmask = _mm256_slli_epi16(hmask, 1);
245 const __m256i q8_0 = _mm256_loadu_si256((
const __m256i *)(
const void *)q8);
247 const __m256i q8_1 = _mm256_loadu_si256((
const __m256i *)(
const void *)q8);
250 __m256i p16_0 = _mm256_maddubs_epi16(q5_0, q8_0);
251 __m256i p16_1 = _mm256_maddubs_epi16(q5_1, q8_1);
252 p16_0 = _mm256_madd_epi16(scale_0, p16_0);
253 p16_1 = _mm256_madd_epi16(scale_1, p16_1);
254 sumi = _mm256_add_epi32(sumi, _mm256_add_epi32(p16_0, p16_1));
257 acc = _mm256_fmadd_ps(
258 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi), acc);
261 return ck_q5k_hsum256_ps(acc) + summs;
264static void dot_q5_k_q8_k_rows4_avx2(
271 static const uint32_t kmask1 = 0x3f3f3f3fU;
272 static const uint32_t kmask2 = 0x0f0f0f0fU;
273 static const uint32_t kmask3 = 0x03030303U;
274 const __m256i m4 = _mm256_set1_epi8(0x0f);
275 const __m128i mzero = _mm_setzero_si128();
276 const __m256i mone = _mm256_set1_epi8(1);
278 _mm256_setzero_ps(), _mm256_setzero_ps(),
279 _mm256_setzero_ps(), _mm256_setzero_ps(),
281 float summs[4] = {0.0f, 0.0f, 0.0f, 0.0f};
283 for (
int b = 0; b < nb; ++b) {
284 const block_q5_K *wb = &w[b];
285 uint32_t utmp[4] = {0, 0, 0, 0};
286 memcpy(utmp, wb->scales, 12);
287 utmp[3] = ((utmp[2] >> 4) & kmask2) |
288 (((utmp[1] >> 6) & kmask3) << 4);
289 const uint32_t uaux = utmp[1] & kmask1;
290 utmp[1] = (utmp[2] & kmask2) |
291 (((utmp[0] >> 6) & kmask3) << 4);
295 const __m256i mins_and_scales = _mm256_cvtepu8_epi16(
296 _mm_set_epi32((
int)utmp[3], (
int)utmp[2],
297 (
int)utmp[1], (
int)utmp[0]));
298 const __m128i sc128 =
299 _mm256_extracti128_si256(mins_and_scales, 0);
300 const __m256i scales = ck_mm256_set_m128i(sc128, sc128);
301 const __m256i hbits = _mm256_loadu_si256(
302 (
const __m256i *)(
const void *)wb->qh);
304 _mm256_setzero_si256(), _mm256_setzero_si256(),
305 _mm256_setzero_si256(), _mm256_setzero_si256(),
308 for (
int row = 0; row < rows; ++row) {
310 const __m256i q8sums = _mm256_loadu_si256(
311 (
const __m256i *)(
const void *)xb->
bsums);
312 const __m128i q8s = _mm_hadd_epi16(
313 _mm256_extracti128_si256(q8sums, 0),
314 _mm256_extracti128_si256(q8sums, 1));
315 const __m128i prod = _mm_madd_epi16(
316 _mm256_extracti128_si256(mins_and_scales, 1), q8s);
317 const __m128i hsum = _mm_hadd_epi32(
318 _mm_hadd_epi32(prod, mzero), mzero);
320 summs[row] += dmin * (float)_mm_extract_epi32(hsum, 0);
323 const uint8_t *q5 = wb->qs;
324 __m256i hmask = mone;
326 for (
int j = 0; j <
QK_K / 64; ++j) {
327 const __m256i scale_0 = _mm256_shuffle_epi8(
328 scales, ck_q5k_scale_shuffle_avx2(2 * j));
329 const __m256i scale_1 = _mm256_shuffle_epi8(
330 scales, ck_q5k_scale_shuffle_avx2(2 * j + 1));
331 const __m256i q5bits = _mm256_loadu_si256(
332 (
const __m256i *)(
const void *)q5);
334 const __m256i q5l_0 = _mm256_and_si256(q5bits, m4);
335 const __m256i q5h_0 = _mm256_slli_epi16(
336 _mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
337 const __m256i q5_0 = _mm256_add_epi8(q5l_0, q5h_0);
338 hmask = _mm256_slli_epi16(hmask, 1);
339 const __m256i q5l_1 = _mm256_and_si256(
340 _mm256_srli_epi16(q5bits, 4), m4);
341 const __m256i q5h_1 = _mm256_slli_epi16(
342 _mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
343 const __m256i q5_1 = _mm256_add_epi8(q5l_1, q5h_1);
344 hmask = _mm256_slli_epi16(hmask, 1);
346 for (
int row = 0; row < rows; ++row) {
348 const __m256i q8_0 = _mm256_loadu_si256(
349 (
const __m256i *)(
const void *)&xb->
qs[j * 64]);
350 const __m256i q8_1 = _mm256_loadu_si256(
351 (
const __m256i *)(
const void *)&xb->
qs[j * 64 + 32]);
352 __m256i p16_0 = _mm256_maddubs_epi16(q5_0, q8_0);
353 __m256i p16_1 = _mm256_maddubs_epi16(q5_1, q8_1);
354 p16_0 = _mm256_madd_epi16(scale_0, p16_0);
355 p16_1 = _mm256_madd_epi16(scale_1, p16_1);
356 sumi[row] = _mm256_add_epi32(
357 sumi[row], _mm256_add_epi32(p16_0, p16_1));
361 for (
int row = 0; row < rows; ++row) {
363 acc[row] = _mm256_fmadd_ps(
364 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi[row]), acc[row]);
368 for (
int row = 0; row < rows; ++row) {
369 out[row] = ck_q5k_hsum256_ps(acc[row]) + summs[row];
373static float dot_q5_k_prepared_q8_k_row_avx2(
374 const block_q5_K_prepared *w,
const block_q8_K *x,
int nb)
376 const __m128i mzero = _mm_setzero_si128();
377 __m256 acc = _mm256_setzero_ps();
380 for (
int b = 0; b < nb; ++b) {
381 const block_q5_K_prepared *wb = &w[b];
386 const __m128i mins8 = _mm_loadl_epi64((
const __m128i *)(
const void *)wb->mins);
387 const __m256i mins16 = _mm256_cvtepu8_epi16(mins8);
388 const __m256i q8sums = _mm256_loadu_si256((
const __m256i *)(
const void *)xb->
bsums);
389 const __m128i q8s = _mm_hadd_epi16(
390 _mm256_extracti128_si256(q8sums, 0),
391 _mm256_extracti128_si256(q8sums, 1));
392 const __m128i prod = _mm_madd_epi16(
393 _mm256_castsi256_si128(mins16), q8s);
394 const __m128i hsum = _mm_hadd_epi32(_mm_hadd_epi32(prod, mzero), mzero);
395 summs += dmin * (float)_mm_extract_epi32(hsum, 0);
397 __m256i sumi = _mm256_setzero_si256();
398 for (
int sb = 0; sb < 8; ++sb) {
399 const __m256i q5 = _mm256_loadu_si256(
400 (
const __m256i *)(
const void *)(wb->qs + sb * 32));
401 const __m256i q8 = _mm256_loadu_si256(
402 (
const __m256i *)(
const void *)(xb->
qs + sb * 32));
403 __m256i p16 = _mm256_maddubs_epi16(q5, q8);
404 const __m256i scale = _mm256_set1_epi16((int16_t)wb->scales[sb]);
405 p16 = _mm256_madd_epi16(scale, p16);
406 sumi = _mm256_add_epi32(sumi, p16);
408 acc = _mm256_fmadd_ps(
409 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi), acc);
411 return ck_q5k_hsum256_ps(acc) + summs;
414static void dot_q5_k_prepared_q8_k_m4_avx2(
415 const block_q5_K_prepared *w,
417 int rows,
int nb,
float out[4])
419 const __m128i mzero = _mm_setzero_si128();
421 _mm256_setzero_ps(), _mm256_setzero_ps(),
422 _mm256_setzero_ps(), _mm256_setzero_ps(),
424 float summs[4] = {0.0f, 0.0f, 0.0f, 0.0f};
426 for (
int b = 0; b < nb; ++b) {
427 const block_q5_K_prepared *wb = &w[b];
428 const __m128i mins8 = _mm_loadl_epi64((
const __m128i *)(
const void *)wb->mins);
429 const __m256i mins16 = _mm256_cvtepu8_epi16(mins8);
431 _mm256_setzero_si256(), _mm256_setzero_si256(),
432 _mm256_setzero_si256(), _mm256_setzero_si256(),
435 for (
int r = 0; r < rows; ++r) {
437 const __m256i q8sums = _mm256_loadu_si256(
438 (
const __m256i *)(
const void *)xb->
bsums);
439 const __m128i q8s = _mm_hadd_epi16(
440 _mm256_extracti128_si256(q8sums, 0),
441 _mm256_extracti128_si256(q8sums, 1));
442 const __m128i prod = _mm_madd_epi16(
443 _mm256_castsi256_si128(mins16), q8s);
444 const __m128i hsum = _mm_hadd_epi32(
445 _mm_hadd_epi32(prod, mzero), mzero);
447 summs[r] += dmin * (float)_mm_extract_epi32(hsum, 0);
450 for (
int sb = 0; sb < 8; ++sb) {
451 const __m256i q5 = _mm256_loadu_si256(
452 (
const __m256i *)(
const void *)(wb->qs + sb * 32));
453 const __m256i scale = _mm256_set1_epi16((int16_t)wb->scales[sb]);
454 for (
int r = 0; r < rows; ++r) {
455 const __m256i q8 = _mm256_loadu_si256(
456 (
const __m256i *)(
const void *)(x[r][b].qs + sb * 32));
457 __m256i p16 = _mm256_maddubs_epi16(q5, q8);
458 p16 = _mm256_madd_epi16(scale, p16);
459 sumi[r] = _mm256_add_epi32(sumi[r], p16);
462 for (
int r = 0; r < rows; ++r) {
464 acc[r] = _mm256_fmadd_ps(
465 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi[r]), acc[r]);
468 for (
int r = 0; r < rows; ++r) {
469 out[r] = ck_q5k_hsum256_ps(acc[r]) + summs[r];
481 return dot_q5_k_q8_k_row_avx2(w, x, nb);
485 static const uint32_t kmask1 = 0x3f3f3f3fU;
486 static const uint32_t kmask2 = 0x0f0f0f0fU;
487 static const uint32_t kmask3 = 0x03030303U;
489 uint32_t utmp[4] = {0, 0, 0, 0};
490 const uint8_t *scales = (
const uint8_t *)&utmp[0];
491 const uint8_t *mins = (
const uint8_t *)&utmp[2];
495 float sums[8] = {0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f};
499 for (
int b = 0; b < nb; ++b) {
500 const block_q5_K *wb = &w[b];
502 const uint8_t *q4 = wb->
qs;
503 const uint8_t *hm = wb->qh;
506 memset(aux32, 0,
sizeof(aux32));
508 for (
int j = 0; j <
QK_K / 64; ++j) {
509 for (
int l = 0; l < 32; ++l) a[l] = (int8_t)(q4[l] & 0xF);
510 for (
int l = 0; l < 32; ++l) a[l] += (hm[l] & m ? 16 : 0);
514 for (
int l = 0; l < 32; ++l) a[l] = (int8_t)(q4[l] >> 4);
515 for (
int l = 0; l < 32; ++l) a[l] += (hm[l] & m ? 16 : 0);
522 memcpy(utmp, wb->scales, 12);
523 utmp[3] = ((utmp[2] >> 4) & kmask2) | (((utmp[1] >> 6) & kmask3) << 4);
524 const uint32_t uaux = utmp[1] & kmask1;
525 utmp[1] = (utmp[2] & kmask2) | (((utmp[0] >> 6) & kmask3) << 4);
530 for (
int j = 0; j <
QK_K / 16; ++j) {
531 sumi += (int)xb->
bsums[j] * (
int)mins[j / 2];
535 const int8_t *q8 = xb->
qs;
537 for (
int j = 0; j <
QK_K / 32; ++j) {
538 const int32_t scale = (int32_t)scales[is++];
540 for (
int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
541 for (
int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
544 for (
int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
545 for (
int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
548 for (
int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
549 for (
int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
552 for (
int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
553 for (
int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
558 for (
int l = 0; l < 8; ++l) {
559 sums[l] += d * (float)aux32[l];
562 sumf -= dmin * (float)sumi;
565 for (
int l = 0; l < 8; ++l) {
574 const block_q5_K *blocks = (
const block_q5_K *)W;
575 const int blocks_per_row = K /
QK_K;
577 for (
int m = 0; m < M; m++) {
578 const float *x_row = x;
581 for (
int b = 0; b < blocks_per_row; b++) {
582 const block_q5_K *block = &blocks[m * blocks_per_row + b];
585 uint8_t sc_arr[8], m_arr[8];
588 for (
int sb = 0; sb < 8; sb++) {
589 const float d_sub = d * (float)sc_arr[sb];
590 const float m_sub = dmin * (float)m_arr[sb];
592 for (
int i = 0; i < 32; i++) {
594 sum += (d_sub * (float)q - m_sub) * x_row[b *
QK_K + sb * 32 + i];
609 const block_q5_K *blocks = (
const block_q5_K *)B;
610 const int blocks_per_col = K /
QK_K;
612 for (
int m = 0; m < M; m++) {
613 const float *a_row = &A[m * K];
615 for (
int n = 0; n < N; n++) {
617 const block_q5_K *w_row = &blocks[n * blocks_per_col];
618 for (
int b = 0; b < blocks_per_col; b++) {
619 const block_q5_K *block = &w_row[b];
622 uint8_t sc_arr[8], m_arr[8];
625 for (
int sb = 0; sb < 8; sb++) {
626 const float d_sub = d * (float)sc_arr[sb];
627 const float m_sub = dmin * (float)m_arr[sb];
629 for (
int i = 0; i < 32; i++) {
631 sum += (d_sub * (float)q - m_sub) * a_row[b *
QK_K + sb * 32 + i];
636 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
661 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
668 const block_q5_K *blocks = (
const block_q5_K *)W;
670 const int blocks_per_row = K /
QK_K;
672 for (
int m = 0; m < M; ++m) {
673 const block_q5_K *w_row = &blocks[m * blocks_per_row];
684 if (!A_q8 || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
692 const block_q5_K *W = (
const block_q5_K *)B;
693 const int blocks_per_row = K /
QK_K;
695 for (
int m = 0; m < M; ++m) {
696 const block_q8_K *a_row = &A[m * blocks_per_row];
697 for (
int n = 0; n < N; ++n) {
698 const block_q5_K *w_row = &W[n * blocks_per_row];
700 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
706 const void *B_prepared,
711#if !defined(__AVX2__)
712 (void)A; (void)B_prepared; (void)bias; (void)
C;
713 (void)M; (void)N; (void)K;
715 if (!A || !B_prepared || !
C || M <= 0 || N <= 0 || K <= 0 ||
716 (K %
QK_K) != 0)
return;
717 const int blocks_per_row = K /
QK_K;
719 const block_q5_K_prepared *W = (
const block_q5_K_prepared *)B_prepared;
720 for (
int m = 0; m < M; ++m) {
723 for (
int n = 0; n < N; ++n) {
724 const float sum = dot_q5_k_prepared_q8_k_row_avx2(
725 W + (
size_t)n * blocks_per_row, a_q8, blocks_per_row);
726 C[(size_t)m * N + n] = sum + (bias ? bias[n] : 0.0f);
733 const void *B_prepared,
738#if !defined(__AVX2__)
739 (void)A; (void)B_prepared; (void)bias; (void)
C;
740 (void)M; (void)N; (void)K;
742 if (!A || !B_prepared || !
C || M <= 0 || N <= 0 || K <= 0 ||
743 (K %
QK_K) != 0)
return;
744 const int blocks_per_row = K /
QK_K;
746 const block_q5_K_prepared *W = (
const block_q5_K_prepared *)B_prepared;
747 for (
int m = 0; m < M; m += 4) {
748 const int rows = M - m < 4 ? M - m : 4;
751 a_q8[0], a_q8[1], a_q8[2], a_q8[3],
753 for (
int r = 0; r < rows; ++r) {
756 for (
int n = 0; n < N; ++n) {
758 dot_q5_k_prepared_q8_k_m4_avx2(
759 W + (
size_t)n * blocks_per_row,
760 row_ptrs, rows, blocks_per_row, sums);
761 for (
int r = 0; r < rows; ++r) {
762 C[(size_t)(m + r) * N + n] = sums[r] + (bias ? bias[n] : 0.0f);
770 const void *B_prepared,
774 int n_begin,
int n_end)
776#if !defined(__AVX2__)
777 (void)A_q8; (void)B_prepared; (void)bias; (void)
C;
778 (void)M; (void)N; (void)K; (void)n_begin; (void)n_end;
780 if (!A_q8 || !B_prepared || !
C || M <= 0 || N <= 0 || K <= 0 ||
781 (K %
QK_K) != 0 || n_begin < 0 || n_end > N ||
782 n_begin >= n_end)
return;
783 const int blocks_per_row = K /
QK_K;
786 const block_q5_K_prepared *W =
787 (
const block_q5_K_prepared *)B_prepared;
789 for (
int m = 0; m < M; m += 4) {
790 const int rows = M - m < 4 ? M - m : 4;
792 A + (size_t)(m + 0) * blocks_per_row,
793 A + (size_t)(m + (rows > 1 ? 1 : 0)) * blocks_per_row,
794 A + (size_t)(m + (rows > 2 ? 2 : 0)) * blocks_per_row,
795 A + (size_t)(m + (rows > 3 ? 3 : 0)) * blocks_per_row,
797 for (
int n = n_begin; n < n_end; ++n) {
799 dot_q5_k_prepared_q8_k_m4_avx2(
800 W + (
size_t)n * blocks_per_row,
801 row_ptrs, rows, blocks_per_row, sums);
802 for (
int r = 0; r < rows; ++r) {
803 C[(size_t)(m + r) * N + n] =
804 sums[r] + (bias ? bias[n] : 0.0f);
821 if (!y || !W || !x || M <= 0 || K <= 0) {
833 const block_q5_K *blocks = (
const block_q5_K *)W;
834 const int blocks_per_row = K /
QK_K;
861 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
873 const block_q5_K *blocks = (
const block_q5_K *)B;
874 const int blocks_per_col = K /
QK_K;
880 for (
int m = 0; m < M; ++m) {
881 const float *a_row = &A[m * K];
897#if defined(__AVX512F__)
900#elif defined(__AVX2__)
903#elif defined(__AVX__)
906#elif defined(__SSE4_1__)
920#if defined(__AVX512F__)
923#elif defined(__AVX2__)
926#elif defined(__AVX__)
929#elif defined(__SSE4_1__)
940 const void *
const input_rows[4],
945 if (!output || !weights || !input_rows || rows <= 0 || rows > 4 ||
946 output_stride < output_dim || output_dim <= 0 || input_dim <= 0 ||
947 (input_dim %
QK_K) != 0) {
950 for (
int row = 0; row < rows; ++row) {
951 if (!input_rows[row])
return;
955 const block_q5_K *blocks = (
const block_q5_K *)weights;
956 const int blocks_per_row = input_dim /
QK_K;
959 (
const block_q8_K *)input_rows[rows > 1 ? 1 : 0],
960 (
const block_q8_K *)input_rows[rows > 2 ? 2 : 0],
961 (
const block_q8_K *)input_rows[rows > 3 ? 3 : 0],
963 for (
int n = 0; n < output_dim; ++n) {
965 dot_q5_k_q8_k_rows4_avx2(
966 blocks + (
size_t)n * (
size_t)blocks_per_row,
967 inputs, rows, blocks_per_row, values);
968 for (
int row = 0; row < rows; ++row) {
969 output[(size_t)row * (
size_t)output_stride + (size_t)n] =
974 for (
int row = 0; row < rows; ++row) {
976 output + (
size_t)row * (
size_t)output_stride,
977 weights, input_rows[row], output_dim, input_dim);
982void gemv_q5_k(
float *y,
const void *W,
const float *x,
int M,
int K)
984#if defined(__AVX512F__)
987#elif defined(__AVX2__)
990#elif defined(__AVX__)
993#elif defined(__SSE4_1__)
1005 int M,
int N,
int K)
1007#if defined(__AVX512F__)
1010#elif defined(__AVX2__)
1013#elif defined(__AVX__)
1016#elif defined(__SSE4_1__)
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
static void gemm_nt_q5_k_ref_fp32(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
static int ck_q5k_debug_generic_dot(void)
void gemv_q5_k_ref(float *y, const void *W, const float *x, int M, int K)
void ck_q5_k_prepare_weight(const void *src, void *dst, int N, int K)
void gemm_q5_k_q8_k_compact_rows4(float *output, int output_stride, const void *weights, const void *const input_rows[4], int rows, int output_dim, int input_dim)
static uint8_t q5_k_quant_value(const block_q5_K *block, int subblock, int i)
void gemv_q5_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_nt_q5_k_prepared_m4(const float *A, const void *B_prepared, const float *bias, float *C, int M, int N, int K)
void quantize_row_q8_k(const float *x, void *vy, int k)
void gemv_q5_k(float *y, const void *W, const float *x, int M, int K)
static int ck_q5k_debug_fp32_fallback(void)
#define CK_Q5K_STACK_Q8_BLOCKS
void gemm_nt_q5_k_ref(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_k_q8_k(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_k_q8_k_ref(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
static float dot_q5_k_q8_k_row(const block_q5_K *w, const block_q8_K *x, int nb)
static void gemv_q5_k_ref_fp32(float *y, const void *W, const float *x, int M, int K)
size_t ck_q5_k_prepared_block_size(void)
static void unpack_q5_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
void gemm_nt_q5_k_prepared(const float *A, const void *B_prepared, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_k_prepared_q8_m4_nrange(const void *A_q8, const void *B_prepared, const float *bias, float *C, int M, int N, int K, int n_begin, int n_end)
void gemm_nt_q5_k(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q5_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)