← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q5_k.c File Reference

GEMM/GEMV kernels with Q5_K quantized weights. More...

#include <stdint.h>
#include <stddef.h>
#include <string.h>
#include <stdlib.h>
#include "ckernel_quant.h"

Go to the source code of this file.

Macros

#define CK_Q5K_STACK_Q8_BLOCKS   128
 
#define QK_K   256
 

Functions

void ck_q5_k_prepare_weight (const void *src, void *dst, int N, int K)
 
size_t ck_q5_k_prepared_block_size (void)
 
static int ck_q5k_debug_fp32_fallback (void)
 
static int ck_q5k_debug_generic_dot (void)
 
static float dot_q5_k_q8_k_row (const block_q5_K *w, const block_q8_K *x, int nb)
 
void gemm_nt_q5_k (const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
 
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_m4 (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_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)
 
void gemm_nt_q5_k_ref (const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
 
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)
 
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)
 
void gemv_q5_k (float *y, const void *W, const float *x, int M, int K)
 
void gemv_q5_k_q8_k (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q5_k_q8_k_ref (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q5_k_ref (float *y, const void *W, const float *x, int M, int K)
 
static void gemv_q5_k_ref_fp32 (float *y, const void *W, const float *x, int M, int K)
 
static uint8_t q5_k_quant_value (const block_q5_K *block, int subblock, int i)
 
void quantize_row_q8_k (const float *x, void *vy, int k)
 
static void unpack_q5_k_scales (const uint8_t *scales, uint8_t *sc, uint8_t *m)
 

Detailed Description

GEMM/GEMV kernels with Q5_K quantized weights.

CK-ENGINE KERNEL RULES:

  1. NO malloc/free - memory via bump allocator, pointers passed in
  2. NO OpenMP - parallelization at orchestrator/codegen layer
  3. API must define: inputs, outputs, workspace, and memory layouts
  4. Pure computation - deterministic, no side effects

After changes: make test && make llamacpp-parity-full

Implements matrix multiplication where:

  • Activations (input): FP32 (quantized internally to Q8_K for dot path)
  • Weights: Q5_K (5-bit super-block quant)
  • Output: FP32

Q5_K Format (256 weights per super-block):

  • d: FP16 super-block scale
  • dmin: FP16 super-block minimum
  • scales[12]: 8 sub-block scales + 8 sub-block mins (6 bits each, packed)
  • qh[32]: high bits for 256 weights (1 bit each)
  • qs[128]: low 4 bits for 256 weights (4 bits each)

Total: 2 + 2 + 12 + 32 + 128 = 176 bytes per 256 weights = 5.5 bits/weight

Dequantization formula (matches llama.cpp): w = d * scale * q - dmin * mins where q = qs_val | (qh_bit << 4) = 5-bit value [0, 31]

Definition in file gemm_kernels_q5_k.c.

Macro Definition Documentation

◆ CK_Q5K_STACK_Q8_BLOCKS

#define CK_Q5K_STACK_Q8_BLOCKS   128

Definition at line 46 of file gemm_kernels_q5_k.c.

◆ QK_K

#define QK_K   256

Definition at line 45 of file gemm_kernels_q5_k.c.

Function Documentation

◆ ck_q5_k_prepare_weight()

void ck_q5_k_prepare_weight ( const void *  src,
void *  dst,
int  N,
int  K 
)

Definition at line 131 of file gemm_kernels_q5_k.c.

132{
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;
140 unpack_q5_k_scales(input[b].scales, output[b].scales, output[b].mins);
141 for (int sb = 0; sb < 8; ++sb) {
142 for (int i = 0; i < 32; ++i) {
143 output[b].qs[sb * 32 + i] = q5_k_quant_value(&input[b], sb, i);
144 }
145 }
146 }
147}
static uint8_t q5_k_quant_value(const block_q5_K *block, int subblock, int i)
static void unpack_q5_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
#define QK_K

References q5_k_quant_value(), QK_K, and unpack_q5_k_scales().

◆ ck_q5_k_prepared_block_size()

size_t ck_q5_k_prepared_block_size ( void  )

Definition at line 126 of file gemm_kernels_q5_k.c.

127{
128 return sizeof(block_q5_K_prepared);
129}

◆ ck_q5k_debug_fp32_fallback()

static int ck_q5k_debug_fp32_fallback ( void  )
static

Definition at line 48 of file gemm_kernels_q5_k.c.

49{
50 static int cached = -1;
51 if (cached < 0) {
52 const char *env = getenv("CK_DEBUG_Q5K_FP32_FALLBACK");
53 cached = (env && env[0] && env[0] != '0') ? 1 : 0;
54 }
55 return cached;
56}

Referenced by gemm_nt_q5_k_ref(), and gemv_q5_k_ref().

◆ ck_q5k_debug_generic_dot()

static int ck_q5k_debug_generic_dot ( void  )
static

Definition at line 58 of file gemm_kernels_q5_k.c.

59{
60 static int cached = -1;
61 if (cached < 0) {
62 const char *env = getenv("CK_DEBUG_Q5K_GENERIC_DOT");
63 cached = (env && env[0] && env[0] != '0') ? 1 : 0;
64 }
65 return cached;
66}

Referenced by dot_q5_k_q8_k_row().

◆ dot_q5_k_q8_k_row()

static float dot_q5_k_q8_k_row ( const block_q5_K *  w,
const block_q8_K x,
int  nb 
)
static

Definition at line 478 of file gemm_kernels_q5_k.c.

478 {
479#if defined(__AVX2__)
481 return dot_q5_k_q8_k_row_avx2(w, x, nb);
482 }
483#endif
484
485 static const uint32_t kmask1 = 0x3f3f3f3fU;
486 static const uint32_t kmask2 = 0x0f0f0f0fU;
487 static const uint32_t kmask3 = 0x03030303U;
488
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];
492
493 int8_t aux8[QK_K];
494 int16_t aux16[8];
495 float sums[8] = {0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f};
496 int32_t aux32[8];
497
498 float sumf = 0.0f;
499 for (int b = 0; b < nb; ++b) {
500 const block_q5_K *wb = &w[b];
501 const block_q8_K *xb = &x[b];
502 const uint8_t *q4 = wb->qs;
503 const uint8_t *hm = wb->qh;
504 int8_t *a = aux8;
505 uint8_t m = 1;
506 memset(aux32, 0, sizeof(aux32));
507
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);
511 a += 32;
512 m <<= 1;
513
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);
516 a += 32;
517 m <<= 1;
518
519 q4 += 32;
520 }
521
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);
526 utmp[2] = uaux;
527 utmp[0] &= kmask1;
528
529 int sumi = 0;
530 for (int j = 0; j < QK_K / 16; ++j) {
531 sumi += (int)xb->bsums[j] * (int)mins[j / 2];
532 }
533
534 a = aux8;
535 const int8_t *q8 = xb->qs;
536 int is = 0;
537 for (int j = 0; j < QK_K / 32; ++j) {
538 const int32_t scale = (int32_t)scales[is++];
539
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];
542 q8 += 8; a += 8;
543
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];
546 q8 += 8; a += 8;
547
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];
550 q8 += 8; a += 8;
551
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];
554 q8 += 8; a += 8;
555 }
556
557 const float d = CK_FP16_TO_FP32(wb->d) * xb->d;
558 for (int l = 0; l < 8; ++l) {
559 sums[l] += d * (float)aux32[l];
560 }
561 const float dmin = CK_FP16_TO_FP32(wb->dmin) * xb->d;
562 sumf -= dmin * (float)sumi;
563 }
564
565 for (int l = 0; l < 8; ++l) {
566 sumf += sums[l];
567 }
568 return sumf;
569}
#define CK_FP16_TO_FP32(x)
static int ck_q5k_debug_generic_dot(void)
int8_t qs[256]
int16_t bsums[256/16]

References block_q8_K::bsums, CK_FP16_TO_FP32, ck_q5k_debug_generic_dot(), block_q8_K::d, QK_K, and block_q8_K::qs.

Referenced by gemm_nt_q5_k_q8_k_ref(), and gemv_q5_k_q8_k_ref().

◆ gemm_nt_q5_k()

void gemm_nt_q5_k ( const float *  A,
const void *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1001 of file gemm_kernels_q5_k.c.

1006{
1007#if defined(__AVX512F__)
1008 /* TODO: AVX-512 implementation */
1009 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1010#elif defined(__AVX2__)
1011 /* TODO: AVX-2 implementation */
1012 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1013#elif defined(__AVX__)
1014 /* TODO: AVX implementation */
1015 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1016#elif defined(__SSE4_1__)
1017 /* TODO: SSE4.1 implementation */
1018 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1019#else
1020 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1021#endif
1022}
void gemm_nt_q5_k_ref(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
#define C(color)
Definition show_config.c:39

References C, and gemm_nt_q5_k_ref().

◆ gemm_nt_q5_k_prepared()

void gemm_nt_q5_k_prepared ( const float *  A,
const void *  B_prepared,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 705 of file gemm_kernels_q5_k.c.

710{
711#if !defined(__AVX2__)
712 (void)A; (void)B_prepared; (void)bias; (void)C;
713 (void)M; (void)N; (void)K;
714#else
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;
718 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) return;
719 const block_q5_K_prepared *W = (const block_q5_K_prepared *)B_prepared;
720 for (int m = 0; m < M; ++m) {
721 block_q8_K a_q8[blocks_per_row];
722 quantize_row_q8_k(A + (size_t)m * K, a_q8, K);
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);
727 }
728 }
729#endif
730}
void quantize_row_q8_k(const float *x, void *vy, int k)
#define CK_Q5K_STACK_Q8_BLOCKS

References C, CK_Q5K_STACK_Q8_BLOCKS, QK_K, and quantize_row_q8_k().

◆ gemm_nt_q5_k_prepared_m4()

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 
)

Definition at line 732 of file gemm_kernels_q5_k.c.

737{
738#if !defined(__AVX2__)
739 (void)A; (void)B_prepared; (void)bias; (void)C;
740 (void)M; (void)N; (void)K;
741#else
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;
745 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) return;
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;
749 block_q8_K a_q8[4][blocks_per_row];
750 const block_q8_K *row_ptrs[4] = {
751 a_q8[0], a_q8[1], a_q8[2], a_q8[3],
752 };
753 for (int r = 0; r < rows; ++r) {
754 quantize_row_q8_k(A + (size_t)(m + r) * K, a_q8[r], K);
755 }
756 for (int n = 0; n < N; ++n) {
757 float sums[4];
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);
763 }
764 }
765 }
766#endif
767}

References C, CK_Q5K_STACK_Q8_BLOCKS, QK_K, and quantize_row_q8_k().

◆ gemm_nt_q5_k_prepared_q8_m4_nrange()

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 
)

Definition at line 769 of file gemm_kernels_q5_k.c.

775{
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;
779#else
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;
784 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) return;
785 const block_q8_K *A = (const block_q8_K *)A_q8;
786 const block_q5_K_prepared *W =
787 (const block_q5_K_prepared *)B_prepared;
788
789 for (int m = 0; m < M; m += 4) {
790 const int rows = M - m < 4 ? M - m : 4;
791 const block_q8_K *row_ptrs[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,
796 };
797 for (int n = n_begin; n < n_end; ++n) {
798 float sums[4];
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);
805 }
806 }
807 }
808#endif
809}

References C, CK_Q5K_STACK_Q8_BLOCKS, and QK_K.

◆ gemm_nt_q5_k_q8_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 
)

Definition at line 914 of file gemm_kernels_q5_k.c.

919{
920#if defined(__AVX512F__)
921 /* TODO: AVX-512 implementation */
922 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
923#elif defined(__AVX2__)
924 /* TODO: AVX-2 implementation */
925 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
926#elif defined(__AVX__)
927 /* TODO: AVX implementation */
928 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
929#elif defined(__SSE4_1__)
930 /* TODO: SSE4.1 implementation */
931 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
932#else
933 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
934#endif
935}
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)

References C, and gemm_nt_q5_k_q8_k_ref().

◆ gemm_nt_q5_k_q8_k_ref()

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 
)

Definition at line 678 of file gemm_kernels_q5_k.c.

683{
684 if (!A_q8 || !B || !C || M <= 0 || N <= 0 || K <= 0) {
685 return;
686 }
687 if (K % QK_K != 0) {
688 return;
689 }
690
691 const block_q8_K *A = (const block_q8_K *)A_q8;
692 const block_q5_K *W = (const block_q5_K *)B;
693 const int blocks_per_row = K / QK_K;
694
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];
699 const float sum = dot_q5_k_q8_k_row(w_row, a_row, blocks_per_row);
700 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
701 }
702 }
703}
static float dot_q5_k_q8_k_row(const block_q5_K *w, const block_q8_K *x, int nb)

References C, dot_q5_k_q8_k_row(), and QK_K.

Referenced by gemm_nt_q5_k_q8_k(), and gemm_nt_q5_k_ref().

◆ gemm_nt_q5_k_ref()

void gemm_nt_q5_k_ref ( const float *  A,
const void *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 855 of file gemm_kernels_q5_k.c.

860{
861 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0) {
862 return;
863 }
865 gemm_nt_q5_k_ref_fp32(A, B, bias, C, M, N, K);
866 return;
867 }
868 if (K % QK_K != 0) {
869 gemm_nt_q5_k_ref_fp32(A, B, bias, C, M, N, K);
870 return;
871 }
872
873 const block_q5_K *blocks = (const block_q5_K *)B;
874 const int blocks_per_col = K / QK_K;
875 if (blocks_per_col > CK_Q5K_STACK_Q8_BLOCKS) {
876 gemm_nt_q5_k_ref_fp32(A, B, bias, C, M, N, K);
877 return;
878 }
879
880 for (int m = 0; m < M; ++m) {
881 const float *a_row = &A[m * K];
883 quantize_row_q8_k(a_row, a_q8, K);
884 gemm_nt_q5_k_q8_k_ref(a_q8, blocks, bias, &C[m * N], 1, N, K);
885 }
886}
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_fp32_fallback(void)

References C, ck_q5k_debug_fp32_fallback(), CK_Q5K_STACK_Q8_BLOCKS, gemm_nt_q5_k_q8_k_ref(), gemm_nt_q5_k_ref_fp32(), QK_K, and quantize_row_q8_k().

Referenced by gemm_nt_q5_k().

◆ gemm_nt_q5_k_ref_fp32()

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

Definition at line 603 of file gemm_kernels_q5_k.c.

608{
609 const block_q5_K *blocks = (const block_q5_K *)B;
610 const int blocks_per_col = K / QK_K;
611
612 for (int m = 0; m < M; m++) {
613 const float *a_row = &A[m * K];
614
615 for (int n = 0; n < N; n++) {
616 float sum = 0.0f;
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];
620 const float d = CK_FP16_TO_FP32(block->d);
621 const float dmin = CK_FP16_TO_FP32(block->dmin);
622 uint8_t sc_arr[8], m_arr[8];
623 unpack_q5_k_scales(block->scales, sc_arr, m_arr);
624
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];
628
629 for (int i = 0; i < 32; i++) {
630 const uint8_t q = q5_k_quant_value(block, sb, i);
631 sum += (d_sub * (float)q - m_sub) * a_row[b * QK_K + sb * 32 + i];
632 }
633 }
634 }
635
636 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
637 }
638 }
639}

References C, CK_FP16_TO_FP32, q5_k_quant_value(), QK_K, and unpack_q5_k_scales().

Referenced by gemm_nt_q5_k_ref().

◆ gemm_q5_k_q8_k_compact_rows4()

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 
)

Definition at line 937 of file gemm_kernels_q5_k.c.

944{
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) {
948 return;
949 }
950 for (int row = 0; row < rows; ++row) {
951 if (!input_rows[row]) return;
952 }
953
954#if defined(__AVX2__)
955 const block_q5_K *blocks = (const block_q5_K *)weights;
956 const int blocks_per_row = input_dim / QK_K;
957 const block_q8_K *inputs[4] = {
958 (const block_q8_K *)input_rows[0],
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],
962 };
963 for (int n = 0; n < output_dim; ++n) {
964 float values[4];
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] =
970 values[row];
971 }
972 }
973#else
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);
978 }
979#endif
980}
void gemv_q5_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q5_k_q8_k(), and QK_K.

Referenced by ck_moe_q4k_q5k_bucket_work().

◆ gemv_q5_k()

void gemv_q5_k ( float *  y,
const void *  W,
const float *  x,
int  M,
int  K 
)

Definition at line 982 of file gemm_kernels_q5_k.c.

983{
984#if defined(__AVX512F__)
985 /* TODO: AVX-512 implementation */
986 gemv_q5_k_ref(y, W, x, M, K);
987#elif defined(__AVX2__)
988 /* TODO: AVX-2 implementation */
989 gemv_q5_k_ref(y, W, x, M, K);
990#elif defined(__AVX__)
991 /* TODO: AVX implementation */
992 gemv_q5_k_ref(y, W, x, M, K);
993#elif defined(__SSE4_1__)
994 /* TODO: SSE4.1 implementation */
995 gemv_q5_k_ref(y, W, x, M, K);
996#else
997 gemv_q5_k_ref(y, W, x, M, K);
998#endif
999}
void gemv_q5_k_ref(float *y, const void *W, const float *x, int M, int K)

References gemv_q5_k_ref().

◆ gemv_q5_k_q8_k()

void gemv_q5_k_q8_k ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K 
)

Definition at line 892 of file gemm_kernels_q5_k.c.

896{
897#if defined(__AVX512F__)
898 /* TODO: AVX-512 implementation */
899 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
900#elif defined(__AVX2__)
901 /* TODO: AVX-2 implementation */
902 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
903#elif defined(__AVX__)
904 /* TODO: AVX implementation */
905 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
906#elif defined(__SSE4_1__)
907 /* TODO: SSE4.1 implementation */
908 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
909#else
910 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
911#endif
912}
void gemv_q5_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q5_k_q8_k_ref().

Referenced by ck_moe_q4k_q5k_route_work(), gemm_q5_k_q8_k_compact_rows4(), and moe_swiglu_expert_forward_q4k_q5k_workspace().

◆ gemv_q5_k_q8_k_ref()

void gemv_q5_k_q8_k_ref ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K 
)

Definition at line 656 of file gemm_kernels_q5_k.c.

660{
661 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
662 return;
663 }
664 if (K % QK_K != 0) {
665 return;
666 }
667
668 const block_q5_K *blocks = (const block_q5_K *)W;
669 const block_q8_K *x = (const block_q8_K *)x_q8;
670 const int blocks_per_row = K / QK_K;
671
672 for (int m = 0; m < M; ++m) {
673 const block_q5_K *w_row = &blocks[m * blocks_per_row];
674 y[m] = dot_q5_k_q8_k_row(w_row, x, blocks_per_row);
675 }
676}

References dot_q5_k_q8_k_row(), and QK_K.

Referenced by gemv_q5_k_q8_k(), and gemv_q5_k_ref().

◆ gemv_q5_k_ref()

void gemv_q5_k_ref ( float *  y,
const void *  W,
const float *  x,
int  M,
int  K 
)

Definition at line 819 of file gemm_kernels_q5_k.c.

820{
821 if (!y || !W || !x || M <= 0 || K <= 0) {
822 return;
823 }
825 gemv_q5_k_ref_fp32(y, W, x, M, K);
826 return;
827 }
828 if (K % QK_K != 0) {
829 gemv_q5_k_ref_fp32(y, W, x, M, K);
830 return;
831 }
832
833 const block_q5_K *blocks = (const block_q5_K *)W;
834 const int blocks_per_row = K / QK_K;
835 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) {
836 gemv_q5_k_ref_fp32(y, W, x, M, K);
837 return;
838 }
839
841 /* Q8_K bytes are part of the numerical ABI. Use the shared provider,
842 * whose FP-contraction policy is validated against llama.cpp. */
843 quantize_row_q8_k(x, x_q8, K);
844 gemv_q5_k_q8_k_ref(y, blocks, x_q8, M, K);
845}
static void gemv_q5_k_ref_fp32(float *y, const void *W, const float *x, int M, int K)

References ck_q5k_debug_fp32_fallback(), CK_Q5K_STACK_Q8_BLOCKS, gemv_q5_k_q8_k_ref(), gemv_q5_k_ref_fp32(), QK_K, and quantize_row_q8_k().

Referenced by gemv_q5_k().

◆ gemv_q5_k_ref_fp32()

static void gemv_q5_k_ref_fp32 ( float *  y,
const void *  W,
const float *  x,
int  M,
int  K 
)
static

Definition at line 572 of file gemm_kernels_q5_k.c.

573{
574 const block_q5_K *blocks = (const block_q5_K *)W;
575 const int blocks_per_row = K / QK_K;
576
577 for (int m = 0; m < M; m++) {
578 const float *x_row = x;
579 float sum = 0.0f;
580
581 for (int b = 0; b < blocks_per_row; b++) {
582 const block_q5_K *block = &blocks[m * blocks_per_row + b];
583 const float d = CK_FP16_TO_FP32(block->d);
584 const float dmin = CK_FP16_TO_FP32(block->dmin);
585 uint8_t sc_arr[8], m_arr[8];
586 unpack_q5_k_scales(block->scales, sc_arr, m_arr);
587
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];
591
592 for (int i = 0; i < 32; i++) {
593 const uint8_t q = q5_k_quant_value(block, sb, i);
594 sum += (d_sub * (float)q - m_sub) * x_row[b * QK_K + sb * 32 + i];
595 }
596 }
597 }
598
599 y[m] = sum;
600 }
601}

References CK_FP16_TO_FP32, q5_k_quant_value(), QK_K, and unpack_q5_k_scales().

Referenced by gemv_q5_k_ref().

◆ q5_k_quant_value()

static uint8_t q5_k_quant_value ( const block_q5_K *  block,
int  subblock,
int  i 
)
inlinestatic

Definition at line 119 of file gemm_kernels_q5_k.c.

119 {
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);
124}

Referenced by ck_q5_k_prepare_weight(), gemm_nt_q5_k_ref_fp32(), and gemv_q5_k_ref_fp32().

◆ quantize_row_q8_k()

void quantize_row_q8_k ( const float *  x,
void *  vy,
int  k 
)

Definition at line 121 of file gemm_kernels_q4k_q8k.c.

121 {
122 const char *ref_env = getenv("CK_DEBUG_Q8K_REF");
123 if (ref_env && atoi(ref_env) != 0) {
124 quantize_row_q8_k_ref(x, vy, k);
125 return;
126 }
127#if defined(__AVX512F__) && defined(__AVX512BW__)
128 quantize_row_q8_k_avx512(x, vy, k);
129#elif defined(__AVX2__)
130 quantize_row_q8_k_avx2(x, vy, k);
131#elif defined(__AVX__)
132 quantize_row_q8_k_avx(x, vy, k);
133#elif defined(__SSE4_1__)
134 quantize_row_q8_k_sse(x, vy, k);
135#else
136 quantize_row_q8_k_ref(x, vy, k);
137#endif
138}
void quantize_row_q8_k_avx512(const float *x, void *vy, int k)
void quantize_row_q8_k_avx2(const float *x, void *vy, int k)
void quantize_row_q8_k_avx(const float *x, void *vy, int k)
void quantize_row_q8_k_sse(const float *x, void *vy, int k)
void quantize_row_q8_k_ref(const float *x, void *vy, int k)

Referenced by gemm_nt_q5_k_prepared(), gemm_nt_q5_k_prepared_m4(), gemm_nt_q5_k_ref(), and gemv_q5_k_ref().

◆ unpack_q5_k_scales()

static void unpack_q5_k_scales ( const uint8_t *  scales,
uint8_t *  sc,
uint8_t *  m 
)
inlinestatic

Definition at line 95 of file gemm_kernels_q5_k.c.

97 {
98 sc[0] = scales[0] & 0x3F;
99 sc[1] = scales[1] & 0x3F;
100 sc[2] = scales[2] & 0x3F;
101 sc[3] = scales[3] & 0x3F;
102
103 m[0] = scales[4] & 0x3F;
104 m[1] = scales[5] & 0x3F;
105 m[2] = scales[6] & 0x3F;
106 m[3] = scales[7] & 0x3F;
107
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);
112
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);
117}

Referenced by ck_q5_k_prepare_weight(), gemm_nt_q5_k_ref_fp32(), and gemv_q5_k_ref_fp32().