34#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__)
38#define CK_Q4K_STACK_Q8_BLOCKS 128
41void gemv_q4_k_q8_k(
float *y,
const void *W,
const void *x_q8,
int M,
int K);
45 static int cached = -1;
47 const char *env = getenv(
"CK_DEBUG_Q4K_Q8_CONTRACT");
48 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
75 const int blocks_per_row = K /
QK_K;
77 for (
int row = 0; row < M; row++) {
80 for (
int b = 0; b < blocks_per_row; b++) {
81 const block_q4_K *block = &blocks[row * blocks_per_row + b];
94 for (
int iter = 0; iter < 4; iter++) {
95 const float d1 = d * (float)sc[2*iter];
96 const float m1 = dmin * (float)m[2*iter];
97 const float d2 = d * (float)sc[2*iter + 1];
98 const float m2 = dmin * (float)m[2*iter + 1];
99 const uint8_t *qs = &block->
qs[iter * 32];
100 const float *xp = &x[b *
QK_K + iter * 64];
103 for (
int l = 0; l < 32; l++) {
104 const int8_t q = (qs[l] & 0x0F);
105 sum += (d1 * (float)q - m1) * xp[l];
108 for (
int l = 0; l < 32; l++) {
109 const int8_t q = (qs[l] >> 4);
110 sum += (d2 * (float)q - m2) * xp[l + 32];
125void gemv_q4_k_avx512(
float *y,
131 const int blocks_per_row = K /
QK_K;
133 for (
int row = 0; row < M; row++) {
134 __m512 acc = _mm512_setzero_ps();
136 for (
int b = 0; b < blocks_per_row; b++) {
137 const block_q4_K *block = &blocks[row * blocks_per_row + b];
141 uint8_t sc[8], m_arr[8];
144 const __m512i mask_lo = _mm512_set1_epi32(0x0F);
149 for (
int iter = 0; iter < 4; iter++) {
150 const float d1 = d * (float)sc[2*iter];
151 const float m1 = dmin * (float)m_arr[2*iter];
152 const float d2 = d * (float)sc[2*iter + 1];
153 const float m2 = dmin * (float)m_arr[2*iter + 1];
155 const __m512 vscale1 = _mm512_set1_ps(d1);
156 const __m512 vmin1 = _mm512_set1_ps(m1);
157 const __m512 vscale2 = _mm512_set1_ps(d2);
158 const __m512 vmin2 = _mm512_set1_ps(m2);
160 const uint8_t *qs = &block->
qs[iter * 32];
161 const float *xp = &x[b *
QK_K + iter * 64];
165 for (
int chunk = 0; chunk < 2; chunk++) {
166 __m128i packed = _mm_loadu_si128((
const __m128i *)&qs[chunk * 16]);
167 __m512i bytes = _mm512_cvtepu8_epi32(packed);
168 __m512i lo = _mm512_and_epi32(bytes, mask_lo);
170 __m512 w = _mm512_fnmadd_ps(_mm512_set1_ps(1.0f), vmin1,
171 _mm512_mul_ps(_mm512_cvtepi32_ps(lo), vscale1));
172 __m512 x_vec = _mm512_loadu_ps(&xp[chunk * 16]);
173 acc = _mm512_fmadd_ps(w, x_vec, acc);
177 for (
int chunk = 0; chunk < 2; chunk++) {
178 __m128i packed = _mm_loadu_si128((
const __m128i *)&qs[chunk * 16]);
179 __m512i bytes = _mm512_cvtepu8_epi32(packed);
180 __m512i hi = _mm512_srli_epi32(bytes, 4);
182 __m512 w = _mm512_fnmadd_ps(_mm512_set1_ps(1.0f), vmin2,
183 _mm512_mul_ps(_mm512_cvtepi32_ps(hi), vscale2));
184 __m512 x_vec = _mm512_loadu_ps(&xp[32 + chunk * 16]);
185 acc = _mm512_fmadd_ps(w, x_vec, acc);
191 y[row] = _mm512_reduce_add_ps(acc);
204#if defined(__AVX__) && !defined(__AVX512F__)
211void gemv_q4_k_avx(
float *y,
217 const int blocks_per_row = K /
QK_K;
219 for (
int row = 0; row < M; row++) {
221 __m256 acc0 = _mm256_setzero_ps();
222 __m256 acc1 = _mm256_setzero_ps();
223 __m256 acc2 = _mm256_setzero_ps();
224 __m256 acc3 = _mm256_setzero_ps();
226 for (
int b = 0; b < blocks_per_row; b++) {
227 const block_q4_K *block = &blocks[row * blocks_per_row + b];
232 uint8_t sc[8], m_arr[8];
236 for (
int iter = 0; iter < 4; iter++) {
237 const float d1 = d * (float)sc[2*iter];
238 const float m1 = dmin * (float)m_arr[2*iter];
239 const float d2 = d * (float)sc[2*iter + 1];
240 const float m2 = dmin * (float)m_arr[2*iter + 1];
241 const uint8_t *qs = &block->
qs[iter * 32];
242 const float *xp = &x[b *
QK_K + iter * 64];
245 __m256 vd1 = _mm256_set1_ps(d1);
246 __m256 vm1 = _mm256_set1_ps(m1);
247 __m256 vd2 = _mm256_set1_ps(d2);
248 __m256 vm2 = _mm256_set1_ps(m2);
251 for (
int g = 0; g < 4; g++) {
254 for (
int i = 0; i < 8; i++) {
255 dq[i] = d1 * (float)(qs[g*8 + i] & 0x0F) - m1;
257 __m256 vw = _mm256_loadu_ps(dq);
258 __m256 vx = _mm256_loadu_ps(&xp[g*8]);
261 __m256 prod = _mm256_mul_ps(vw, vx);
262 acc0 = _mm256_add_ps(acc0, prod);
266 for (
int g = 0; g < 4; g++) {
269 for (
int i = 0; i < 8; i++) {
270 dq[i] = d2 * (float)(qs[g*8 + i] >> 4) - m2;
272 __m256 vw = _mm256_loadu_ps(dq);
273 __m256 vx = _mm256_loadu_ps(&xp[32 + g*8]);
275 __m256 prod = _mm256_mul_ps(vw, vx);
276 acc1 = _mm256_add_ps(acc1, prod);
282 __m256 sum01 = _mm256_add_ps(acc0, acc1);
283 __m256 sum23 = _mm256_add_ps(acc2, acc3);
284 __m256 sum = _mm256_add_ps(sum01, sum23);
287 __m128 hi = _mm256_extractf128_ps(sum, 1);
288 __m128 lo = _mm256_castps256_ps128(sum);
289 __m128 sum128 = _mm_add_ps(hi, lo);
290 sum128 = _mm_hadd_ps(sum128, sum128);
291 sum128 = _mm_hadd_ps(sum128, sum128);
293 y[row] = _mm_cvtss_f32(sum128);
307 const int nb = K /
QK_K;
316 gemv_q4_k_avx512(y, W, x, M, K);
317#elif defined(__AVX__)
318 gemv_q4_k_avx(y, W, x, M, K);
348 for (
int n = 0; n < N; n++) {
349 gemv_q4_k(&Y[n * M], W, &X[n * K], M, K);
359void gemm_q4_k_avx512(
float *Y,
365 const int blocks_per_row = K /
QK_K;
368 const int N4 = N / 4 * 4;
370 for (
int row = 0; row < M; row++) {
372 for (
int n = 0; n < N4; n += 4) {
373 __m512 acc0 = _mm512_setzero_ps();
374 __m512 acc1 = _mm512_setzero_ps();
375 __m512 acc2 = _mm512_setzero_ps();
376 __m512 acc3 = _mm512_setzero_ps();
378 for (
int b = 0; b < blocks_per_row; b++) {
379 const block_q4_K *block = &blocks[row * blocks_per_row + b];
383 uint8_t sc[8], m_arr[8];
386 for (
int sub = 0; sub < 8; sub++) {
387 const float scale = d * (float)sc[sub];
388 const float min_val = dmin * (float)m_arr[sub];
389 const __m512 vscale = _mm512_set1_ps(scale);
390 const __m512 vmin = _mm512_set1_ps(min_val);
391 const __m512i offset = _mm512_set1_epi32(8);
392 const __m512i mask_lo = _mm512_set1_epi32(0x0F);
394 const uint8_t *qs = &block->
qs[sub * 16];
395 const int x_offset = b *
QK_K + sub * 32;
398 __m128i packed = _mm_loadu_si128((
const __m128i *)qs);
399 __m512i bytes = _mm512_cvtepu8_epi32(packed);
401 __m512i lo = _mm512_sub_epi32(_mm512_and_epi32(bytes, mask_lo), offset);
402 __m512i hi = _mm512_sub_epi32(_mm512_srli_epi32(bytes, 4), offset);
404 __m512 w_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(lo), vscale, vmin);
405 __m512 w_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(hi), vscale, vmin);
409 for (
int bn = 0; bn < 4; bn++) {
410 const float *xp = &X[(n + bn) * K + x_offset];
412 __m512 x_even = _mm512_set_ps(
413 xp[30], xp[28], xp[26], xp[24], xp[22], xp[20], xp[18], xp[16],
414 xp[14], xp[12], xp[10], xp[8], xp[6], xp[4], xp[2], xp[0]);
415 __m512 x_odd = _mm512_set_ps(
416 xp[31], xp[29], xp[27], xp[25], xp[23], xp[21], xp[19], xp[17],
417 xp[15], xp[13], xp[11], xp[9], xp[7], xp[5], xp[3], xp[1]);
419 __m512 *acc = (bn == 0) ? &acc0 : (bn == 1) ? &acc1 :
420 (bn == 2) ? &acc2 : &acc3;
421 *acc = _mm512_fmadd_ps(w_lo, x_even, *acc);
422 *acc = _mm512_fmadd_ps(w_hi, x_odd, *acc);
427 Y[(n + 0) * M + row] = _mm512_reduce_add_ps(acc0);
428 Y[(n + 1) * M + row] = _mm512_reduce_add_ps(acc1);
429 Y[(n + 2) * M + row] = _mm512_reduce_add_ps(acc2);
430 Y[(n + 3) * M + row] = _mm512_reduce_add_ps(acc3);
434 for (
int n = N4; n < N; n++) {
435 __m512 acc = _mm512_setzero_ps();
437 for (
int b = 0; b < blocks_per_row; b++) {
438 const block_q4_K *block = &blocks[row * blocks_per_row + b];
442 uint8_t sc[8], m_arr[8];
445 for (
int sub = 0; sub < 8; sub++) {
446 const float scale = d * (float)sc[sub];
447 const float min_val = dmin * (float)m_arr[sub];
448 const __m512 vscale = _mm512_set1_ps(scale);
449 const __m512 vmin = _mm512_set1_ps(min_val);
450 const __m512i offset = _mm512_set1_epi32(8);
451 const __m512i mask_lo = _mm512_set1_epi32(0x0F);
453 const uint8_t *qs = &block->
qs[sub * 16];
454 const float *xp = &X[n * K + b *
QK_K + sub * 32];
456 __m128i packed = _mm_loadu_si128((
const __m128i *)qs);
457 __m512i bytes = _mm512_cvtepu8_epi32(packed);
459 __m512i lo = _mm512_sub_epi32(_mm512_and_epi32(bytes, mask_lo), offset);
460 __m512i hi = _mm512_sub_epi32(_mm512_srli_epi32(bytes, 4), offset);
462 __m512 w_lo = _mm512_fmadd_ps(_mm512_cvtepi32_ps(lo), vscale, vmin);
463 __m512 w_hi = _mm512_fmadd_ps(_mm512_cvtepi32_ps(hi), vscale, vmin);
465 __m512 x_even = _mm512_set_ps(
466 xp[30], xp[28], xp[26], xp[24], xp[22], xp[20], xp[18], xp[16],
467 xp[14], xp[12], xp[10], xp[8], xp[6], xp[4], xp[2], xp[0]);
468 __m512 x_odd = _mm512_set_ps(
469 xp[31], xp[29], xp[27], xp[25], xp[23], xp[21], xp[19], xp[17],
470 xp[15], xp[13], xp[11], xp[9], xp[7], xp[5], xp[3], xp[1]);
472 acc = _mm512_fmadd_ps(w_lo, x_even, acc);
473 acc = _mm512_fmadd_ps(w_hi, x_odd, acc);
477 Y[n * M + row] = _mm512_reduce_add_ps(acc);
509float dot_q4_k(
const void *w_q4k,
const float *x,
int K)
542 const int blocks_per_row = K /
QK_K;
545 memset(dX, 0, K *
sizeof(
float));
549 for (
int row = 0; row < M; row++) {
550 const float dy = dY[row];
552 for (
int b = 0; b < blocks_per_row; b++) {
553 const block_q4_K *block = &blocks[row * blocks_per_row + b];
561 for (
int iter = 0; iter < 4; iter++) {
562 const float d1 = d * (float)sc[2 * iter];
563 const float m1 = dmin * (float)m[2 * iter];
564 const float d2 = d * (float)sc[2 * iter + 1];
565 const float m2 = dmin * (float)m[2 * iter + 1];
567 const uint8_t *qs = &block->
qs[iter * 32];
568 float *dxp = &dX[b *
QK_K + iter * 64];
571 for (
int l = 0; l < 32; l++) {
572 const int q = (qs[l] & 0x0F);
573 const float w = d1 * (float)q - m1;
578 for (
int l = 0; l < 32; l++) {
579 const int q = (qs[l] >> 4);
580 const float w = d2 * (float)q - m2;
581 dxp[32 + l] += w * dy;
594void gemv_q4_k_backward_avx512(
float *dX,
600 const int blocks_per_row = K /
QK_K;
601 const __m512i mask_lo = _mm512_set1_epi32(0x0F);
604 memset(dX, 0, K *
sizeof(
float));
606 for (
int row = 0; row < M; row++) {
607 const __m512 vdy = _mm512_set1_ps(dY[row]);
609 for (
int b = 0; b < blocks_per_row; b++) {
610 const block_q4_K *block = &blocks[row * blocks_per_row + b];
614 uint8_t sc[8], m_arr[8];
618 for (
int iter = 0; iter < 4; iter++) {
619 const float d1 = d * (float)sc[2 * iter];
620 const float m1 = dmin * (float)m_arr[2 * iter];
621 const float d2 = d * (float)sc[2 * iter + 1];
622 const float m2 = dmin * (float)m_arr[2 * iter + 1];
624 const __m512 vd1 = _mm512_set1_ps(d1);
625 const __m512 vm1 = _mm512_set1_ps(m1);
626 const __m512 vd2 = _mm512_set1_ps(d2);
627 const __m512 vm2 = _mm512_set1_ps(m2);
629 const uint8_t *qs = &block->
qs[iter * 32];
630 float *dxp = &dX[b *
QK_K + iter * 64];
633 for (
int chunk = 0; chunk < 2; chunk++) {
634 __m128i packed = _mm_loadu_si128((
const __m128i *)&qs[chunk * 16]);
635 __m512i bytes = _mm512_cvtepu8_epi32(packed);
636 __m512i lo = _mm512_and_epi32(bytes, mask_lo);
638 __m512 w = _mm512_fnmadd_ps(_mm512_set1_ps(1.0f), vm1,
639 _mm512_mul_ps(_mm512_cvtepi32_ps(lo), vd1));
640 __m512 grad = _mm512_mul_ps(w, vdy);
641 __m512 existing = _mm512_loadu_ps(&dxp[chunk * 16]);
642 _mm512_storeu_ps(&dxp[chunk * 16], _mm512_add_ps(existing, grad));
646 for (
int chunk = 0; chunk < 2; chunk++) {
647 __m128i packed = _mm_loadu_si128((
const __m128i *)&qs[chunk * 16]);
648 __m512i bytes = _mm512_cvtepu8_epi32(packed);
649 __m512i hi = _mm512_srli_epi32(bytes, 4);
651 __m512 w = _mm512_fnmadd_ps(_mm512_set1_ps(1.0f), vm2,
652 _mm512_mul_ps(_mm512_cvtepi32_ps(hi), vd2));
653 __m512 grad = _mm512_mul_ps(w, vdy);
654 __m512 existing = _mm512_loadu_ps(&dxp[32 + chunk * 16]);
655 _mm512_storeu_ps(&dxp[32 + chunk * 16], _mm512_add_ps(existing, grad));
672 gemv_q4_k_backward_avx512(dX, W, dY, M, K);
686 for (
int n = 0; n < N; n++) {
714 if (!A || !B || !
C) {
717 if (M <= 0 || N <= 0 || K <= 0) {
730 for (
int i = 0; i < M; ++i) {
731 float *row =
C + (size_t)i * (
size_t)N;
732 for (
int j = 0; j < N; ++j) {
Quantization block structures for weight-only quantization.
#define GGML_FP16_TO_FP32
#define CK_FP16_TO_FP32(x)
static void unpack_q4_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
Unpack Q4_K sub-block scales and mins.
void gemm_q4_k_ref(float *Y, const void *W, const float *X, int M, int N, int K)
Matrix-matrix multiply with Q4_K weights (scalar reference)
float dot_q4_k(const void *w_q4k, const float *x, int K)
Compute dot product of Q4_K row with FP32 vector.
void gemm_q4_k_backward(float *dX, const void *W, const float *dY, int M, int N, int K)
Batched backward pass.
void gemm_nt_q4_k(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q4_k_backward(float *dX, const void *W, const float *dY, int M, int K)
Auto-dispatch backward.
void gemv_q4_k_ref(float *y, const void *W, const float *x, int M, int K)
Matrix-vector multiply with Q4_K weights (scalar reference)
#define CK_Q4K_STACK_Q8_BLOCKS
void gemv_q4_k_backward_ref(float *dX, const void *W, const float *dY, int M, int K)
Backward pass: compute input gradient (scalar reference)
void quantize_row_q8_k(const float *x, void *vy, int k)
void gemv_q4_k(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV based on available SIMD.
void gemv_q4_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_q4_k(float *Y, const void *W, const float *X, int M, int N, int K)
Auto-dispatch GEMM based on available SIMD.
static int ck_q4k_debug_q8_contract(void)