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

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

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

Go to the source code of this file.

Functions

static int ck_nearest_int_q8_0 (float fval)
 
static int ck_q8_0_debug_ref (void)
 
static int ck_q8_0_fp32_m4n4_enabled (void)
 
static int ck_q8_0_q8_0_debug_ref (void)
 
float dot_q8_0 (const void *w_q8_0, const float *x, int K)
 
void gemm_nt_q8_0 (const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
 
static void gemm_nt_q8_0_rowloop (const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
 Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias.
 
void gemm_q8_0 (float *Y, const void *W, const float *X, int M, int N, int K)
 Matrix-matrix multiply with Q8_0 weights.
 
void gemm_q8_0_backward (float *dX, const void *W, const float *dY, int M, int N, int K)
 Batched backward pass.
 
void gemm_q8_0_q8_0_m2n4 (float *C, const void *W, const void *A_q8, int M, int N, int K)
 
void gemm_q8_0_q8_0_m2n4_strided (float *C, int ldc, const void *W, const void *A_q8, int M, int N, int K)
 
void gemv_q8_0 (float *y, const void *W, const float *x, int M, int K)
 Auto-dispatch GEMV for Q8_0 weights based on CPU features.
 
void gemv_q8_0_backward (float *dX, const void *W, const float *dY, int M, int K)
 Auto-dispatch backward.
 
void gemv_q8_0_backward_ref (float *dX, const void *W, const float *dY, int M, int K)
 Backward pass: compute input gradient (scalar reference)
 
void gemv_q8_0_parallel_simd (float *y, const void *W, const float *x, int M, int K, int ith, int nth)
 Parallel SIMD GEMV for Q8_0 weights x FP32 input with prefetching.
 
void gemv_q8_0_q8_0 (float *y, const void *W, const void *x_q8, int M, int K)
 Matrix-vector multiply with Q8_0 weights and Q8_0 input.
 
void gemv_q8_0_q8_0_parallel (float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
 Parallel reference GEMV for Q8_0 x Q8_0.
 
void gemv_q8_0_q8_0_parallel_simd (float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
 Parallel SIMD GEMV for Q8_0 x Q8_0 with prefetching.
 
void gemv_q8_0_q8_0_x4 (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q8_0_ref (float *y, const void *W, const float *x, int M, int K)
 Matrix-vector multiply with Q8_0 weights (scalar reference)
 
void quantize_batch_q8_0 (const float *x, void *vy, int num_rows, int k)
 Batch quantize FP32 to Q8_0 format (row-major output)
 
void quantize_batch_q8_k (const float *x, void *vy, int num_rows, int k)
 Batch quantize FP32 to Q8_K format (row-major output)
 
void quantize_row_q8_0 (const float *x, void *vy, int k)
 Quantize FP32 to Q8_0 format (scalar reference)
 
void quantize_row_q8_k (const float *x, void *vy, int k)
 
void vec_dot_q8_0_q8_0 (int n, float *s, const void *vx, const void *vy)
 Auto-dispatch quantized dot product Q8_0 x Q8_0.
 
void vec_dot_q8_0_q8_0_ref (int n, float *s, const void *vx, const void *vy)
 Quantized dot product: Q8_0 weights x Q8_0 input (scalar reference)
 

Detailed Description

GEMM/GEMV kernels with Q8_0 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

Q8_0 Format:

  • 32 weights per block
  • 1 FP16 scale per block
  • 34 bytes per 32 weights = 8.5 bits/weight
  • Weights stored as signed 8-bit integers

Operations: Forward: Y = W @ X (W is Q8_0, X and Y are FP32) Backward: dX = W^T @ dY (gradient w.r.t. input)

Note: Q8_0 is often used for activation quantization or as an intermediate format. Higher precision than Q4_0/Q4_K.

Definition in file gemm_kernels_q8_0.c.

Function Documentation

◆ ck_nearest_int_q8_0()

static int ck_nearest_int_q8_0 ( float  fval)
inlinestatic

Definition at line 75 of file gemm_kernels_q8_0.c.

75 {
76 /* Match llama.cpp's deterministic nearest-even helper. */
77 float val = fval + 12582912.f;
78 int i;
79 memcpy(&i, &val, sizeof(int));
80 return (i & 0x007fffff) - 0x00400000;
81}

Referenced by quantize_row_q8_0().

◆ ck_q8_0_debug_ref()

static int ck_q8_0_debug_ref ( void  )
static

Definition at line 46 of file gemm_kernels_q8_0.c.

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

Referenced by gemv_q8_0().

◆ ck_q8_0_fp32_m4n4_enabled()

static int ck_q8_0_fp32_m4n4_enabled ( void  )
static

Definition at line 66 of file gemm_kernels_q8_0.c.

67{
68 static int cached = -1;
69 if (cached < 0) {
70 cached = ck_env_truthy_or_qwen3vl_ocr_profile("CK_ENABLE_Q80_FP32_M4N4");
71 }
72 return cached;
73}
static int ck_env_truthy_or_qwen3vl_ocr_profile(const char *name)

References ck_env_truthy_or_qwen3vl_ocr_profile().

Referenced by gemm_nt_q8_0().

◆ ck_q8_0_q8_0_debug_ref()

static int ck_q8_0_q8_0_debug_ref ( void  )
static

Definition at line 56 of file gemm_kernels_q8_0.c.

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

Referenced by gemm_q8_0_q8_0_m2n4_strided(), gemv_q8_0_q8_0_x4(), and vec_dot_q8_0_q8_0().

◆ dot_q8_0()

float dot_q8_0 ( const void *  w_q8_0,
const float *  x,
int  K 
)

Definition at line 1008 of file gemm_kernels_q8_0.c.

1009{
1010 float result;
1011 gemv_q8_0(&result, w_q8_0, x, 1, K);
1012 return result;
1013}
void gemv_q8_0(float *y, const void *W, const float *x, int M, int K)
Auto-dispatch GEMV for Q8_0 weights based on CPU features.

References gemv_q8_0().

◆ gemm_nt_q8_0()

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

Definition at line 879 of file gemm_kernels_q8_0.c.

884{
885#if defined(__AVX512F__)
886 if (ck_q8_0_fp32_m4n4_enabled() && M >= 4 && N >= 4 && K % QK8_0 == 0) {
887 gemm_nt_q8_0_m4n4_avx512(A, B, bias, C, M, N, K);
888 return;
889 }
890#endif
891 gemm_nt_q8_0_rowloop(A, B, bias, C, M, N, K);
892}
#define QK8_0
static void gemm_nt_q8_0_rowloop(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias.
static int ck_q8_0_fp32_m4n4_enabled(void)
#define C(color)
Definition show_config.c:39

References C, ck_q8_0_fp32_m4n4_enabled(), gemm_nt_q8_0_rowloop(), and QK8_0.

Referenced by ck_gemm_nt_quant(), qwen2_0_5b_decode_decode_token(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_8_decode(), and qwen2_0_5b_decode_layer_9_decode().

◆ gemm_nt_q8_0_rowloop()

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

Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias.

Parameters
AInput matrix [M x K], row-major FP32
BWeight matrix in Q8_0 format, [N x K] stored row-major
biasOptional bias [N], NULL if not used
COutput [M x N], row-major FP32
MBatch size (number of tokens)
NOutput dimension (number of rows in B)
KInput dimension

Definition at line 750 of file gemm_kernels_q8_0.c.

755{
756 for (int m = 0; m < M; m++) {
757 gemv_q8_0(&C[(size_t)m * N], B, &A[(size_t)m * K], N, K);
758 if (bias) {
759 for (int n = 0; n < N; n++) C[(size_t)m * N + n] += bias[n];
760 }
761 }
762}

References C, and gemv_q8_0().

Referenced by gemm_nt_q8_0().

◆ gemm_q8_0()

void gemm_q8_0 ( float *  Y,
const void *  W,
const float *  X,
int  M,
int  N,
int  K 
)

Matrix-matrix multiply with Q8_0 weights.

Definition at line 725 of file gemm_kernels_q8_0.c.

729{
730 for (int n = 0; n < N; n++) {
731 gemv_q8_0(&Y[n * M], W, &X[n * K], M, K);
732 }
733}

References gemv_q8_0().

◆ gemm_q8_0_backward()

void gemm_q8_0_backward ( float *  dX,
const void *  W,
const float *  dY,
int  M,
int  N,
int  K 
)

Batched backward pass.

Definition at line 994 of file gemm_kernels_q8_0.c.

998{
999 for (int n = 0; n < N; n++) {
1000 gemv_q8_0_backward(&dX[n * K], W, &dY[n * M], M, K);
1001 }
1002}
void gemv_q8_0_backward(float *dX, const void *W, const float *dY, int M, int K)
Auto-dispatch backward.

References gemv_q8_0_backward().

◆ gemm_q8_0_q8_0_m2n4()

void gemm_q8_0_q8_0_m2n4 ( float *  C,
const void *  W,
const void *  A_q8,
int  M,
int  N,
int  K 
)

Definition at line 1613 of file gemm_kernels_q8_0.c.

1617{
1618 gemm_q8_0_q8_0_m2n4_strided(C, N, W, A_q8, M, N, K);
1619}
void gemm_q8_0_q8_0_m2n4_strided(float *C, int ldc, const void *W, const void *A_q8, int M, int N, int K)

References C, and gemm_q8_0_q8_0_m2n4_strided().

Referenced by gemm_nt_q8_0_q8_0_m2n4().

◆ gemm_q8_0_q8_0_m2n4_strided()

void gemm_q8_0_q8_0_m2n4_strided ( float *  C,
int  ldc,
const void *  W,
const void *  A_q8,
int  M,
int  N,
int  K 
)

Definition at line 1502 of file gemm_kernels_q8_0.c.

1507{
1508#if defined(__AVX2__) || defined(__AVX512F__)
1509 if (ck_q8_0_q8_0_debug_ref() || (K % QK8_0) != 0) {
1510 const int nb = K / QK8_0;
1511 const block_q8_0 *a = (const block_q8_0 *)A_q8;
1512 for (int m = 0; m < M; ++m) {
1513 gemv_q8_0_q8_0(C + (size_t)m * (size_t)ldc, W,
1514 a + (size_t)m * (size_t)nb, N, K);
1515 }
1516 return;
1517 }
1518
1519 const block_q8_0 *w = (const block_q8_0 *)W;
1520 const block_q8_0 *a = (const block_q8_0 *)A_q8;
1521 const int nb = K / QK8_0;
1522 int m = 0;
1523 for (; m + 1 < M; m += 2) {
1524 const block_q8_0 *a0 = a + (size_t)(m + 0) * (size_t)nb;
1525 const block_q8_0 *a1 = a + (size_t)(m + 1) * (size_t)nb;
1526 int n = 0;
1527 for (; n + 3 < N; n += 4) {
1528 const block_q8_0 *w0 = w + (size_t)(n + 0) * (size_t)nb;
1529 const block_q8_0 *w1 = w + (size_t)(n + 1) * (size_t)nb;
1530 const block_q8_0 *w2 = w + (size_t)(n + 2) * (size_t)nb;
1531 const block_q8_0 *w3 = w + (size_t)(n + 3) * (size_t)nb;
1532 __m256 acc00 = _mm256_setzero_ps();
1533 __m256 acc01 = _mm256_setzero_ps();
1534 __m256 acc02 = _mm256_setzero_ps();
1535 __m256 acc03 = _mm256_setzero_ps();
1536 __m256 acc10 = _mm256_setzero_ps();
1537 __m256 acc11 = _mm256_setzero_ps();
1538 __m256 acc12 = _mm256_setzero_ps();
1539 __m256 acc13 = _mm256_setzero_ps();
1540
1541 for (int ib = 0; ib < nb; ++ib) {
1542 const __m256i qa0 = _mm256_loadu_si256((const __m256i *)a0[ib].qs);
1543 const __m256i qa1 = _mm256_loadu_si256((const __m256i *)a1[ib].qs);
1544 const __m256i qw0 = _mm256_loadu_si256((const __m256i *)w0[ib].qs);
1545 const __m256i qw1 = _mm256_loadu_si256((const __m256i *)w1[ib].qs);
1546 const __m256i qw2 = _mm256_loadu_si256((const __m256i *)w2[ib].qs);
1547 const __m256i qw3 = _mm256_loadu_si256((const __m256i *)w3[ib].qs);
1548 const float da0 = CK_FP16_TO_FP32(a0[ib].d);
1549 const float da1 = CK_FP16_TO_FP32(a1[ib].d);
1550 const float dw0 = CK_FP16_TO_FP32(w0[ib].d);
1551 const float dw1 = CK_FP16_TO_FP32(w1[ib].d);
1552 const float dw2 = CK_FP16_TO_FP32(w2[ib].d);
1553 const float dw3 = CK_FP16_TO_FP32(w3[ib].d);
1554 const __m256 p00 = mul_sum_i8_pairs_float_q8_0_avx2(qw0, qa0);
1555 const __m256 p01 = mul_sum_i8_pairs_float_q8_0_avx2(qw1, qa0);
1556 const __m256 p02 = mul_sum_i8_pairs_float_q8_0_avx2(qw2, qa0);
1557 const __m256 p03 = mul_sum_i8_pairs_float_q8_0_avx2(qw3, qa0);
1558 const __m256 p10 = mul_sum_i8_pairs_float_q8_0_avx2(qw0, qa1);
1559 const __m256 p11 = mul_sum_i8_pairs_float_q8_0_avx2(qw1, qa1);
1560 const __m256 p12 = mul_sum_i8_pairs_float_q8_0_avx2(qw2, qa1);
1561 const __m256 p13 = mul_sum_i8_pairs_float_q8_0_avx2(qw3, qa1);
1562#if defined(__FMA__)
1563 acc00 = _mm256_fmadd_ps(_mm256_set1_ps(dw0 * da0), p00, acc00);
1564 acc01 = _mm256_fmadd_ps(_mm256_set1_ps(dw1 * da0), p01, acc01);
1565 acc02 = _mm256_fmadd_ps(_mm256_set1_ps(dw2 * da0), p02, acc02);
1566 acc03 = _mm256_fmadd_ps(_mm256_set1_ps(dw3 * da0), p03, acc03);
1567 acc10 = _mm256_fmadd_ps(_mm256_set1_ps(dw0 * da1), p10, acc10);
1568 acc11 = _mm256_fmadd_ps(_mm256_set1_ps(dw1 * da1), p11, acc11);
1569 acc12 = _mm256_fmadd_ps(_mm256_set1_ps(dw2 * da1), p12, acc12);
1570 acc13 = _mm256_fmadd_ps(_mm256_set1_ps(dw3 * da1), p13, acc13);
1571#else
1572 acc00 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw0 * da0), p00), acc00);
1573 acc01 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw1 * da0), p01), acc01);
1574 acc02 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw2 * da0), p02), acc02);
1575 acc03 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw3 * da0), p03), acc03);
1576 acc10 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw0 * da1), p10), acc10);
1577 acc11 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw1 * da1), p11), acc11);
1578 acc12 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw2 * da1), p12), acc12);
1579 acc13 = _mm256_add_ps(_mm256_mul_ps(_mm256_set1_ps(dw3 * da1), p13), acc13);
1580#endif
1581 }
1582
1583 C[(size_t)(m + 0) * (size_t)ldc + (n + 0)] = hsum_float_8_q8_0(acc00);
1584 C[(size_t)(m + 0) * (size_t)ldc + (n + 1)] = hsum_float_8_q8_0(acc01);
1585 C[(size_t)(m + 0) * (size_t)ldc + (n + 2)] = hsum_float_8_q8_0(acc02);
1586 C[(size_t)(m + 0) * (size_t)ldc + (n + 3)] = hsum_float_8_q8_0(acc03);
1587 C[(size_t)(m + 1) * (size_t)ldc + (n + 0)] = hsum_float_8_q8_0(acc10);
1588 C[(size_t)(m + 1) * (size_t)ldc + (n + 1)] = hsum_float_8_q8_0(acc11);
1589 C[(size_t)(m + 1) * (size_t)ldc + (n + 2)] = hsum_float_8_q8_0(acc12);
1590 C[(size_t)(m + 1) * (size_t)ldc + (n + 3)] = hsum_float_8_q8_0(acc13);
1591 }
1592 if (n < N) {
1593 gemv_q8_0_q8_0(C + (size_t)(m + 0) * (size_t)ldc + n,
1594 w + (size_t)n * (size_t)nb, a0, N - n, K);
1595 gemv_q8_0_q8_0(C + (size_t)(m + 1) * (size_t)ldc + n,
1596 w + (size_t)n * (size_t)nb, a1, N - n, K);
1597 }
1598 }
1599 if (m < M) {
1600 gemv_q8_0_q8_0_x4(C + (size_t)m * (size_t)ldc, W,
1601 a + (size_t)m * (size_t)nb, N, K);
1602 }
1603#else
1604 const int nb = K / QK8_0;
1605 const block_q8_0 *a = (const block_q8_0 *)A_q8;
1606 for (int m = 0; m < M; ++m) {
1607 gemv_q8_0_q8_0(C + (size_t)m * (size_t)ldc, W,
1608 a + (size_t)m * (size_t)nb, N, K);
1609 }
1610#endif
1611}
#define CK_FP16_TO_FP32(x)
static int ck_q8_0_q8_0_debug_ref(void)
void gemv_q8_0_q8_0_x4(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q8_0_q8_0(float *y, const void *W, const void *x_q8, int M, int K)
Matrix-vector multiply with Q8_0 weights and Q8_0 input.

References C, CK_FP16_TO_FP32, ck_q8_0_q8_0_debug_ref(), gemv_q8_0_q8_0(), gemv_q8_0_q8_0_x4(), and QK8_0.

Referenced by gemm_nt_q8_0_q8_0_m2n4_tile(), and gemm_q8_0_q8_0_m2n4().

◆ gemv_q8_0()

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

Auto-dispatch GEMV for Q8_0 weights based on CPU features.

Dispatch priority (best available):

  1. AVX-512 (512-bit vectors) - Intel Skylake-X+
  2. AVX2+FMA (256-bit vectors) - Intel Haswell+
  3. AVX (256-bit vectors) - Intel Sandy Bridge+
  4. SSE4.1 (128-bit vectors) - Intel Nehalem+
  5. Reference (scalar) - Fallback

Uses ck_features.h for standardized feature detection.

Parameters
yOutput vector [M]
WWeight matrix in Q8_0 format [M x K]
xInput vector [K]
MNumber of output rows
KNumber of input columns (hidden dimension)

Definition at line 694 of file gemm_kernels_q8_0.c.

698{
699 if (ck_q8_0_debug_ref()) {
700 gemv_q8_0_ref(y, W, x, M, K);
701 return;
702 }
703
704// Dispatch order: AVX512 > AVX2 > AVX > SSE > ref
705#if defined(__AVX512F__)
706 gemv_q8_0_avx512(y, W, x, M, K);
707#elif defined(__AVX2__)
708 gemv_q8_0_avx2(y, W, x, M, K);
709#elif defined(__AVX__)
710 gemv_q8_0_avx(y, W, x, M, K);
711#elif defined(__SSE4_1__)
712 gemv_q8_0_sse(y, W, x, M, K);
713#else
714 gemv_q8_0_ref(y, W, x, M, K);
715#endif
716}
static int ck_q8_0_debug_ref(void)
void gemv_q8_0_ref(float *y, const void *W, const float *x, int M, int K)
Matrix-vector multiply with Q8_0 weights (scalar reference)

References ck_q8_0_debug_ref(), and gemv_q8_0_ref().

Referenced by dot_q8_0(), gemm_nt_q8_0_rowloop(), gemm_q8_0(), and gemv_q8_0_q8_0_contract().

◆ gemv_q8_0_backward()

void gemv_q8_0_backward ( float *  dX,
const void *  W,
const float *  dY,
int  M,
int  K 
)

Auto-dispatch backward.

Definition at line 979 of file gemm_kernels_q8_0.c.

983{
984#ifdef __AVX512F__
985 gemv_q8_0_backward_avx512(dX, W, dY, M, K);
986#else
987 gemv_q8_0_backward_ref(dX, W, dY, M, K);
988#endif
989}
void gemv_q8_0_backward_ref(float *dX, const void *W, const float *dY, int M, int K)
Backward pass: compute input gradient (scalar reference)

References gemv_q8_0_backward_ref().

Referenced by gemm_q8_0_backward().

◆ gemv_q8_0_backward_ref()

void gemv_q8_0_backward_ref ( float *  dX,
const void *  W,
const float *  dY,
int  M,
int  K 
)

Backward pass: compute input gradient (scalar reference)

Parameters
dXOutput gradient w.r.t. input [K]
WWeight matrix in Q8_0 format [M x K]
dYGradient w.r.t. output [M]
MNumber of output rows
KNumber of columns (input dimension)

Definition at line 907 of file gemm_kernels_q8_0.c.

911{
912 const block_q8_0 *blocks = (const block_q8_0 *)W;
913 const int blocks_per_row = K / QK8_0;
914
915 /* Zero output gradient */
916 memset(dX, 0, K * sizeof(float));
917
918 /* Accumulate: dX += W^T @ dY */
919 for (int row = 0; row < M; row++) {
920 const float dy = dY[row];
921
922 for (int b = 0; b < blocks_per_row; b++) {
923 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
924 const float d = CK_FP16_TO_FP32(block->d);
925 float *dxp = &dX[b * QK8_0];
926
927 for (int i = 0; i < QK8_0; i++) {
928 dxp[i] += d * (float)block->qs[i] * dy;
929 }
930 }
931 }
932}
int8_t qs[32]

References CK_FP16_TO_FP32, block_q8_0::d, QK8_0, and block_q8_0::qs.

Referenced by gemv_q8_0_backward().

◆ gemv_q8_0_parallel_simd()

void gemv_q8_0_parallel_simd ( float *  y,
const void *  W,
const float *  x,
int  M,
int  K,
int  ith,
int  nth 
)

Parallel SIMD GEMV for Q8_0 weights x FP32 input with prefetching.

Definition at line 1716 of file gemm_kernels_q8_0.c.

1721{
1722 if (!y || !W || !x || M <= 0 || K <= 0) return;
1723 if (ith < 0 || nth <= 0 || ith >= nth) return;
1724
1725 const int dr = (M + nth - 1) / nth;
1726 const int r0 = dr * ith;
1727 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1728
1729 if (r0 >= M) return;
1730
1731 const block_q8_0 *blocks = (const block_q8_0 *)W;
1732 const int blocks_per_row = K / QK8_0;
1733
1734#if defined(__AVX__) || defined(__SSE4_1__)
1735 const int PREFETCH_ROWS = 4;
1736 for (int p = 0; p < PREFETCH_ROWS && r0 + p < r1; ++p) {
1737 const char *row_ptr = (const char *)(blocks + (r0 + p) * blocks_per_row);
1738 _mm_prefetch(row_ptr, _MM_HINT_T0);
1739 _mm_prefetch(row_ptr + 64, _MM_HINT_T0);
1740 }
1741
1742 for (int row = r0; row < r1; ++row) {
1743 if (row + PREFETCH_ROWS < r1) {
1744 const char *pf = (const char *)(blocks + (row + PREFETCH_ROWS) * blocks_per_row);
1745 _mm_prefetch(pf, _MM_HINT_T0);
1746 _mm_prefetch(pf + 64, _MM_HINT_T0);
1747 }
1748
1749 /* Dispatch to best available SIMD for single row */
1750#if defined(__AVX512F__)
1751 gemv_q8_0_avx512(&y[row],
1752 (const char *)blocks + row * blocks_per_row * sizeof(block_q8_0),
1753 x, 1, K);
1754#elif defined(__AVX2__)
1755 gemv_q8_0_avx2(&y[row],
1756 (const char *)blocks + row * blocks_per_row * sizeof(block_q8_0),
1757 x, 1, K);
1758#elif defined(__AVX__)
1759 gemv_q8_0_avx(&y[row],
1760 (const char *)blocks + row * blocks_per_row * sizeof(block_q8_0),
1761 x, 1, K);
1762#elif defined(__SSE4_1__)
1763 gemv_q8_0_sse(&y[row],
1764 (const char *)blocks + row * blocks_per_row * sizeof(block_q8_0),
1765 x, 1, K);
1766#else
1767 gemv_q8_0_ref(&y[row],
1768 (const char *)blocks + row * blocks_per_row * sizeof(block_q8_0),
1769 x, 1, K);
1770#endif
1771 }
1772#else
1773 for (int row = r0; row < r1; row++) {
1774 gemv_q8_0_ref(&y[row],
1775 (const char *)blocks + row * blocks_per_row * sizeof(block_q8_0),
1776 x, 1, K);
1777 }
1778#endif
1779}

References gemv_q8_0_ref(), and QK8_0.

◆ gemv_q8_0_q8_0()

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

Matrix-vector multiply with Q8_0 weights and Q8_0 input.

Parameters
yOutput vector [M]
WWeight matrix in Q8_0 format [M x K]
x_q8Input vector in Q8_0 format [K]
MNumber of output rows
KNumber of columns (must be multiple of 32)

Definition at line 1405 of file gemm_kernels_q8_0.c.

1409{
1410 const block_q8_0 *w_blocks = (const block_q8_0 *)W;
1411 const block_q8_0 *x_blocks = (const block_q8_0 *)x_q8;
1412 const int blocks_per_row = K / QK8_0;
1413
1414 for (int row = 0; row < M; row++) {
1415 vec_dot_q8_0_q8_0(K, &y[row],
1416 &w_blocks[row * blocks_per_row],
1417 x_blocks);
1418 }
1419}
void vec_dot_q8_0_q8_0(int n, float *s, const void *vx, const void *vy)
Auto-dispatch quantized dot product Q8_0 x Q8_0.

References QK8_0, and vec_dot_q8_0_q8_0().

Referenced by ck_moe_q4k_mixed_route_work(), ck_test_gemv_q8_0(), ck_test_gemv_q8_0_q8_0(), gemm_q8_0_q8_0_m2n4_strided(), gemv_q8_0_q8_0_x4(), moe_swiglu_expert_forward_q4k_q8_0_workspace(), moe_swiglu_shared_forward_q4k_q8_0_gated_workspace(), and moe_swiglu_shared_forward_q8_0_gated_workspace().

◆ gemv_q8_0_q8_0_parallel()

void gemv_q8_0_q8_0_parallel ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K,
int  ith,
int  nth 
)

Parallel reference GEMV for Q8_0 x Q8_0.

Definition at line 1631 of file gemm_kernels_q8_0.c.

1636{
1637 if (!y || !W || !x_q8 || M <= 0 || K <= 0) return;
1638 if (ith < 0 || nth <= 0 || ith >= nth) return;
1639
1640 const int dr = (M + nth - 1) / nth;
1641 const int r0 = dr * ith;
1642 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1643
1644 if (r0 >= M) return;
1645
1646 const block_q8_0 *w_blocks = (const block_q8_0 *)W;
1647 const block_q8_0 *x_blocks = (const block_q8_0 *)x_q8;
1648 const int blocks_per_row = K / QK8_0;
1649
1650 for (int row = r0; row < r1; row++) {
1651 vec_dot_q8_0_q8_0(K, &y[row],
1652 &w_blocks[row * blocks_per_row],
1653 x_blocks);
1654 }
1655}

References QK8_0, and vec_dot_q8_0_q8_0().

◆ gemv_q8_0_q8_0_parallel_simd()

void gemv_q8_0_q8_0_parallel_simd ( float *  y,
const void *  W,
const void *  x_q8,
int  M,
int  K,
int  ith,
int  nth 
)

Parallel SIMD GEMV for Q8_0 x Q8_0 with prefetching.

Each thread processes rows [r0, r1) where r0 = ith * ceil(M/nth). Prefetches upcoming weight rows to hide memory latency.

Definition at line 1663 of file gemm_kernels_q8_0.c.

1668{
1669 if (!y || !W || !x_q8 || M <= 0 || K <= 0) return;
1670 if (ith < 0 || nth <= 0 || ith >= nth) return;
1671
1672 const int dr = (M + nth - 1) / nth;
1673 const int r0 = dr * ith;
1674 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
1675
1676 if (r0 >= M) return;
1677
1678 const block_q8_0 *w_blocks = (const block_q8_0 *)W;
1679 const block_q8_0 *x_blocks = (const block_q8_0 *)x_q8;
1680 const int blocks_per_row = K / QK8_0;
1681
1682#if defined(__AVX__) || defined(__SSE4_1__)
1683 /* Prefetch first few rows */
1684 const int PREFETCH_ROWS = 4;
1685 for (int p = 0; p < PREFETCH_ROWS && r0 + p < r1; ++p) {
1686 const char *row_ptr = (const char *)(w_blocks + (r0 + p) * blocks_per_row);
1687 _mm_prefetch(row_ptr, _MM_HINT_T0);
1688 _mm_prefetch(row_ptr + 64, _MM_HINT_T0);
1689 }
1690
1691 for (int row = r0; row < r1; ++row) {
1692 /* Prefetch upcoming rows */
1693 if (row + PREFETCH_ROWS < r1) {
1694 const char *pf = (const char *)(w_blocks + (row + PREFETCH_ROWS) * blocks_per_row);
1695 _mm_prefetch(pf, _MM_HINT_T0);
1696 _mm_prefetch(pf + 64, _MM_HINT_T0);
1697 }
1698
1699 vec_dot_q8_0_q8_0(K, &y[row],
1700 &w_blocks[row * blocks_per_row],
1701 x_blocks);
1702 }
1703#else
1704 /* Fallback: no prefetching */
1705 for (int row = r0; row < r1; row++) {
1706 vec_dot_q8_0_q8_0(K, &y[row],
1707 &w_blocks[row * blocks_per_row],
1708 x_blocks);
1709 }
1710#endif
1711}

References QK8_0, and vec_dot_q8_0_q8_0().

◆ gemv_q8_0_q8_0_x4()

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

Definition at line 1428 of file gemm_kernels_q8_0.c.

1432{
1433#if defined(__AVX2__) || defined(__AVX512F__)
1434 if (ck_q8_0_q8_0_debug_ref() || (K % QK8_0) != 0) {
1435 gemv_q8_0_q8_0(y, W, x_q8, M, K);
1436 return;
1437 }
1438
1439 const block_q8_0 *w = (const block_q8_0 *)W;
1440 const block_q8_0 *x = (const block_q8_0 *)x_q8;
1441 const int nb = K / QK8_0;
1442 int row = 0;
1443 for (; row + 3 < M; row += 4) {
1444 __m256 acc0 = _mm256_setzero_ps();
1445 __m256 acc1 = _mm256_setzero_ps();
1446 __m256 acc2 = _mm256_setzero_ps();
1447 __m256 acc3 = _mm256_setzero_ps();
1448 const block_q8_0 *w0 = w + (size_t)(row + 0) * (size_t)nb;
1449 const block_q8_0 *w1 = w + (size_t)(row + 1) * (size_t)nb;
1450 const block_q8_0 *w2 = w + (size_t)(row + 2) * (size_t)nb;
1451 const block_q8_0 *w3 = w + (size_t)(row + 3) * (size_t)nb;
1452
1453 for (int ib = 0; ib < nb; ++ib) {
1454 const __m256i qx = _mm256_loadu_si256((const __m256i *)x[ib].qs);
1455 const float dx = CK_FP16_TO_FP32(x[ib].d);
1456 const __m256i qw0 = _mm256_loadu_si256((const __m256i *)w0[ib].qs);
1457 const __m256i qw1 = _mm256_loadu_si256((const __m256i *)w1[ib].qs);
1458 const __m256i qw2 = _mm256_loadu_si256((const __m256i *)w2[ib].qs);
1459 const __m256i qw3 = _mm256_loadu_si256((const __m256i *)w3[ib].qs);
1460 const __m256 p0 = mul_sum_i8_pairs_float_q8_0_avx2(qw0, qx);
1461 const __m256 p1 = mul_sum_i8_pairs_float_q8_0_avx2(qw1, qx);
1462 const __m256 p2 = mul_sum_i8_pairs_float_q8_0_avx2(qw2, qx);
1463 const __m256 p3 = mul_sum_i8_pairs_float_q8_0_avx2(qw3, qx);
1464 const __m256 d0 = _mm256_set1_ps(CK_FP16_TO_FP32(w0[ib].d) * dx);
1465 const __m256 d1 = _mm256_set1_ps(CK_FP16_TO_FP32(w1[ib].d) * dx);
1466 const __m256 d2 = _mm256_set1_ps(CK_FP16_TO_FP32(w2[ib].d) * dx);
1467 const __m256 d3 = _mm256_set1_ps(CK_FP16_TO_FP32(w3[ib].d) * dx);
1468#if defined(__FMA__)
1469 acc0 = _mm256_fmadd_ps(d0, p0, acc0);
1470 acc1 = _mm256_fmadd_ps(d1, p1, acc1);
1471 acc2 = _mm256_fmadd_ps(d2, p2, acc2);
1472 acc3 = _mm256_fmadd_ps(d3, p3, acc3);
1473#else
1474 acc0 = _mm256_add_ps(_mm256_mul_ps(d0, p0), acc0);
1475 acc1 = _mm256_add_ps(_mm256_mul_ps(d1, p1), acc1);
1476 acc2 = _mm256_add_ps(_mm256_mul_ps(d2, p2), acc2);
1477 acc3 = _mm256_add_ps(_mm256_mul_ps(d3, p3), acc3);
1478#endif
1479 }
1480 y[row + 0] = hsum_float_8_q8_0(acc0);
1481 y[row + 1] = hsum_float_8_q8_0(acc1);
1482 y[row + 2] = hsum_float_8_q8_0(acc2);
1483 y[row + 3] = hsum_float_8_q8_0(acc3);
1484 }
1485 if (row < M) {
1486 gemv_q8_0_q8_0(y + row, w + (size_t)row * (size_t)nb,
1487 x, M - row, K);
1488 }
1489#else
1490 gemv_q8_0_q8_0(y, W, x_q8, M, K);
1491#endif
1492}

References CK_FP16_TO_FP32, ck_q8_0_q8_0_debug_ref(), gemv_q8_0_q8_0(), and QK8_0.

Referenced by gemm_q8_0_q8_0_m2n4_strided(), and gemv_q8_0_q8_0_contract().

◆ gemv_q8_0_ref()

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

Matrix-vector multiply with Q8_0 weights (scalar reference)

Parameters
yOutput vector [M]
WWeight matrix in Q8_0 format [M x K]
xInput vector [K]
MNumber of output rows
KNumber of columns (must be multiple of 32)

Definition at line 316 of file gemm_kernels_q8_0.c.

320{
321 const block_q8_0 *blocks = (const block_q8_0 *)W;
322 const int blocks_per_row = K / QK8_0;
323
324 for (int row = 0; row < M; row++) {
325 float sum = 0.0f;
326
327 for (int b = 0; b < blocks_per_row; b++) {
328 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
329 const float d = CK_FP16_TO_FP32(block->d);
330 const float *xp = &x[b * QK8_0];
331
332 for (int i = 0; i < QK8_0; i++) {
333 sum += d * (float)block->qs[i] * xp[i];
334 }
335 }
336
337 y[row] = sum;
338 }
339}

References CK_FP16_TO_FP32, block_q8_0::d, QK8_0, and block_q8_0::qs.

Referenced by gemv_q8_0(), and gemv_q8_0_parallel_simd().

◆ quantize_batch_q8_0()

void quantize_batch_q8_0 ( const float *  x,
void *  vy,
int  num_rows,
int  k 
)

Batch quantize FP32 to Q8_0 format (row-major output)

Quantizes multiple rows of FP32 data to Q8_0 format, placing each row's Q8_0 output at the correct byte offset for GEMM compatibility.

Memory layout: Input: [num_rows, k] FP32, row-major (stride = k * sizeof(float)) Output: [num_rows, q8_row_bytes] Q8_0, row-major (stride = q8_row_bytes)

where q8_row_bytes = (k / 32) * sizeof(block_q8_0) = (k / 32) * 34

Parameters
xInput FP32 values [num_rows * k]
vyOutput Q8_0 blocks [num_rows * (k/32) blocks]
num_rowsNumber of rows (batch size / tokens)
kElements per row (must be multiple of 32)

Definition at line 256 of file gemm_kernels_q8_0.c.

257{
258 const size_t row_bytes_in = (size_t)k * sizeof(float);
259 const size_t row_bytes_out = (size_t)(k / QK8_0) * sizeof(block_q8_0);
260
261 uint8_t *out = (uint8_t *)vy;
262 const uint8_t *in = (const uint8_t *)x;
263
264 for (int row = 0; row < num_rows; ++row) {
266 (const float *)(in + row * row_bytes_in),
267 (void *)(out + row * row_bytes_out),
268 k
269 );
270 }
271}
void quantize_row_q8_0(const float *x, void *vy, int k)
Quantize FP32 to Q8_0 format (scalar reference)

References QK8_0, and quantize_row_q8_0().

◆ quantize_batch_q8_k()

void quantize_batch_q8_k ( const float *  x,
void *  vy,
int  num_rows,
int  k 
)

Batch quantize FP32 to Q8_K format (row-major output)

Same as quantize_batch_q8_0 but for Q8_K format (super-blocks).

Parameters
xInput FP32 values [num_rows * k]
vyOutput Q8_K blocks
num_rowsNumber of rows (batch size / tokens)
kElements per row (must be multiple of 256)

Definition at line 283 of file gemm_kernels_q8_0.c.

284{
285 /* Q8_K: 256 elements per super-block, each block is larger */
286 const size_t row_bytes_in = (size_t)k * sizeof(float);
287 /* Q8_K block size = 2 (d) + 256 (qs) + 32 (bsums/2) = ~274 bytes for 256 elements */
288 /* Actual: sizeof(block_q8_K) from ckernel_quant.h */
289 const size_t row_bytes_out = (size_t)(k / 256) * sizeof(block_q8_K);
290
291 uint8_t *out = (uint8_t *)vy;
292 const uint8_t *in = (const uint8_t *)x;
293
294 for (int row = 0; row < num_rows; ++row) {
296 (const float *)(in + row * row_bytes_in),
297 (void *)(out + row * row_bytes_out),
298 k
299 );
300 }
301}
void quantize_row_q8_k(const float *x, void *vy, int k)

References quantize_row_q8_k().

◆ quantize_row_q8_0()

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

Quantize FP32 to Q8_0 format (scalar reference)

Parameters
xInput FP32 values
vyOutput Q8_0 blocks
kNumber of elements (must be multiple of 32)

Definition at line 125 of file gemm_kernels_q8_0.c.

126{
127 block_q8_0 *y = (block_q8_0 *)vy;
128 const int nb = k / QK8_0; /* QK8_0 = 32 */
129
130#if defined(__AVX__)
131 const __m256 sign_bit = _mm256_set1_ps(-0.0f);
132
133 for (int i = 0; i < nb; i++) {
134 __m256 v0 = _mm256_loadu_ps(x + 0);
135 __m256 v1 = _mm256_loadu_ps(x + 8);
136 __m256 v2 = _mm256_loadu_ps(x + 16);
137 __m256 v3 = _mm256_loadu_ps(x + 24);
138 x += QK8_0;
139
140 __m256 max_abs = _mm256_andnot_ps(sign_bit, v0);
141 max_abs = _mm256_max_ps(max_abs, _mm256_andnot_ps(sign_bit, v1));
142 max_abs = _mm256_max_ps(max_abs, _mm256_andnot_ps(sign_bit, v2));
143 max_abs = _mm256_max_ps(max_abs, _mm256_andnot_ps(sign_bit, v3));
144
145 __m128 max4 = _mm_max_ps(_mm256_extractf128_ps(max_abs, 1),
146 _mm256_castps256_ps128(max_abs));
147 max4 = _mm_max_ps(max4, _mm_movehl_ps(max4, max4));
148 max4 = _mm_max_ss(max4, _mm_movehdup_ps(max4));
149 const float max_scalar = _mm_cvtss_f32(max4);
150
151#if defined(__INTEL_LLVM_COMPILER)
152 const float d = ck_q8_0_div_rounded_f32(max_scalar, 127.0f);
153 const float id = max_scalar != 0.0f
154 ? ck_q8_0_div_rounded_f32(127.0f, max_scalar)
155 : 0.0f;
156#else
157 const float d = max_scalar / 127.0f;
158 const float id = max_scalar != 0.0f ? 127.0f / max_scalar : 0.0f;
159#endif
160 y[i].d = CK_FP32_TO_FP16(d);
161
162 const __m256 mul = _mm256_set1_ps(id);
163 v0 = _mm256_mul_ps(v0, mul);
164 v1 = _mm256_mul_ps(v1, mul);
165 v2 = _mm256_mul_ps(v2, mul);
166 v3 = _mm256_mul_ps(v3, mul);
167
168 /* Match llama.cpp x86 Q8 quantization: nearest-even rounding. */
169 v0 = _mm256_round_ps(v0, _MM_ROUND_NEAREST);
170 v1 = _mm256_round_ps(v1, _MM_ROUND_NEAREST);
171 v2 = _mm256_round_ps(v2, _MM_ROUND_NEAREST);
172 v3 = _mm256_round_ps(v3, _MM_ROUND_NEAREST);
173
174 __m256i i0 = _mm256_cvtps_epi32(v0);
175 __m256i i1 = _mm256_cvtps_epi32(v1);
176 __m256i i2 = _mm256_cvtps_epi32(v2);
177 __m256i i3 = _mm256_cvtps_epi32(v3);
178
179#if defined(__AVX2__)
180 i0 = _mm256_packs_epi32(i0, i1);
181 i2 = _mm256_packs_epi32(i2, i3);
182 i0 = _mm256_packs_epi16(i0, i2);
183
184 const __m256i perm = _mm256_setr_epi32(0, 4, 1, 5, 2, 6, 3, 7);
185 i0 = _mm256_permutevar8x32_epi32(i0, perm);
186 _mm256_storeu_si256((__m256i *)y[i].qs, i0);
187#else
188 __m128i ni0 = _mm256_castsi256_si128(i0);
189 __m128i ni1 = _mm256_extractf128_si256(i0, 1);
190 __m128i ni2 = _mm256_castsi256_si128(i1);
191 __m128i ni3 = _mm256_extractf128_si256(i1, 1);
192 __m128i ni4 = _mm256_castsi256_si128(i2);
193 __m128i ni5 = _mm256_extractf128_si256(i2, 1);
194 __m128i ni6 = _mm256_castsi256_si128(i3);
195 __m128i ni7 = _mm256_extractf128_si256(i3, 1);
196
197 ni0 = _mm_packs_epi32(ni0, ni1);
198 ni2 = _mm_packs_epi32(ni2, ni3);
199 ni4 = _mm_packs_epi32(ni4, ni5);
200 ni6 = _mm_packs_epi32(ni6, ni7);
201
202 ni0 = _mm_packs_epi16(ni0, ni2);
203 ni4 = _mm_packs_epi16(ni4, ni6);
204
205 _mm_storeu_si128((__m128i *)(y[i].qs + 0), ni0);
206 _mm_storeu_si128((__m128i *)(y[i].qs + 16), ni4);
207#endif
208 }
209#else
210 for (int i = 0; i < nb; i++) {
211 const float *xb = x + i * QK8_0;
212
213 /* Find max absolute value in block */
214 float amax = 0.0f;
215 for (int j = 0; j < QK8_0; j++) {
216 float av = xb[j] >= 0 ? xb[j] : -xb[j];
217 if (av > amax) amax = av;
218 }
219
220 /* Compute scale: d = max / 127 */
221 float d = amax / 127.0f;
222 float id = d != 0.0f ? 127.0f / amax : 0.0f;
223
224 /* Store scale as FP16 */
225 y[i].d = CK_FP32_TO_FP16(d);
226
227 /* Quantize values */
228 for (int j = 0; j < QK8_0; j++) {
229 float v = xb[j] * id;
230 int q = ck_nearest_int_q8_0(v);
231 if (q > 127) q = 127;
232 if (q < -127) q = -127;
233 y[i].qs[j] = (int8_t)q;
234 }
235 }
236#endif
237}
#define CK_FP32_TO_FP16(x)
static int ck_nearest_int_q8_0(float fval)
int32_t id
Definition tokenizer.h:316

References CK_FP32_TO_FP16, ck_nearest_int_q8_0(), block_q8_0::d, id, QK8_0, and block_q8_0::qs.

Referenced by ck_moe_q4k_mixed_route_work(), ck_moe_shared_q4k_gated_workspace(), ck_moe_swiglu_nvfp4_projection(), ck_test_gemm_q5_0(), ck_test_gemm_q8_0(), ck_test_gemv_q5_0(), ck_test_gemv_q5_0_q8_0(), ck_test_gemv_q8_0(), ck_test_gemv_q8_0_q8_0(), fused_mlp_swiglu_prefill_w1w2_quant(), fused_rmsnorm_qkv_prefill_head_major_quant(), gemv_fused_q5_0_bias_parallel_omp(), gemv_q5_0_from_fp32(), gemv_q8_0_from_fp32(), gemv_q8_0_q8_0_contract(), hyper_connection_mix_quantized(), mega_fused_attention_decode_q5_0(), mega_fused_attention_decode_q5_0_parallel_simd(), moe_swiglu_expert_forward_q4k_q5_0_workspace(), moe_swiglu_expert_forward_q4k_q8_0_workspace(), moe_swiglu_shared_forward_q8_0_gated_workspace(), quantize_attn_out_head_major_q8_0(), quantize_attn_out_head_major_q8_0(), quantize_attn_out_head_major_q8_0(), and quantize_batch_q8_0().

◆ 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 quantize_batch_q8_k(), and quantize_batch_q8_k_4row_nearest_even().

◆ vec_dot_q8_0_q8_0()

void vec_dot_q8_0_q8_0 ( int  n,
float *  s,
const void *  vx,
const void *  vy 
)

Auto-dispatch quantized dot product Q8_0 x Q8_0.

Definition at line 1368 of file gemm_kernels_q8_0.c.

1369{
1370 if (ck_q8_0_q8_0_debug_ref()) {
1371 vec_dot_q8_0_q8_0_ref(n, s, vx, vy);
1372 return;
1373 }
1374#ifdef __AVX512F__
1375 vec_dot_q8_0_q8_0_avx512(n, s, vx, vy);
1376#elif defined(__AVX2__)
1377 vec_dot_q8_0_q8_0_avx2(n, s, vx, vy);
1378#elif defined(__ARM_NEON) || defined(__aarch64__)
1379 vec_dot_q8_0_q8_0_neon(n, s, vx, vy);
1380#elif defined(__AVX__)
1381 vec_dot_q8_0_q8_0_avx(n, s, vx, vy);
1382#elif defined(__SSE4_1__)
1383 vec_dot_q8_0_q8_0_sse(n, s, vx, vy);
1384#else
1385 vec_dot_q8_0_q8_0_ref(n, s, vx, vy);
1386#endif
1387}
void vec_dot_q8_0_q8_0_ref(int n, float *s, const void *vx, const void *vy)
Quantized dot product: Q8_0 weights x Q8_0 input (scalar reference)

References ck_q8_0_q8_0_debug_ref(), and vec_dot_q8_0_q8_0_ref().

Referenced by ck_test_vec_dot_q8_0_q8_0(), gemv_q8_0_from_fp32(), gemv_q8_0_q8_0(), gemv_q8_0_q8_0_parallel(), gemv_q8_0_q8_0_parallel_omp(), gemv_q8_0_q8_0_parallel_simd(), out_proj_head_major_q8_0_q8_0(), and out_proj_head_major_q8_0_q8_0().

◆ vec_dot_q8_0_q8_0_ref()

void vec_dot_q8_0_q8_0_ref ( int  n,
float *  s,
const void *  vx,
const void *  vy 
)

Quantized dot product: Q8_0 weights x Q8_0 input (scalar reference)

Parameters
nNumber of elements (must be multiple of 32)
sOutput: scalar dot product result
vxQ8_0 quantized weights
vyQ8_0 quantized input

Definition at line 1140 of file gemm_kernels_q8_0.c.

1141{
1142 const int qk = QK8_0; /* 32 */
1143 const int nb = n / qk;
1144
1145 const block_q8_0 *x = (const block_q8_0 *)vx;
1146 const block_q8_0 *y = (const block_q8_0 *)vy;
1147
1148 float sumf = 0.0f;
1149
1150 for (int ib = 0; ib < nb; ib++) {
1151 int sumi = 0;
1152
1153 for (int j = 0; j < qk; j++) {
1154 sumi += x[ib].qs[j] * y[ib].qs[j];
1155 }
1156
1157 sumf += sumi * (CK_FP16_TO_FP32(x[ib].d) * CK_FP16_TO_FP32(y[ib].d));
1158 }
1159
1160 *s = sumf;
1161}

References CK_FP16_TO_FP32, QK8_0, and block_q8_0::qs.

Referenced by gemv_q8_0_q8_0_ref_rows(), and vec_dot_q8_0_q8_0().