50 float val = fval + 12582912.f;
52 memcpy(&i, &val,
sizeof(
int));
53 return (i & 0x007fffff) - 0x00400000;
57#pragma clang fp contract(off)
58#elif defined(__GNUC__)
62 if (!x || !vy || k <= 0) {
65 assert(k %
QK_K == 0);
66 const int nb = k /
QK_K;
69 for (
int i = 0; i < nb; ++i) {
72 for (
int j = 0; j <
QK_K; ++j) {
73 float ax = fabsf(x[j]);
81 memset(y[i].qs, 0,
sizeof(y[i].qs));
82 memset(y[i].bsums, 0,
sizeof(y[i].bsums));
87 const float iscale = -127.0f / max;
88 for (
int j = 0; j <
QK_K; ++j) {
91 float scaled = iscale * x[j];
99 y[i].
qs[j] = (int8_t)v;
102 for (
int j = 0; j <
QK_K / 16; ++j) {
104 const int8_t *qs = &y[i].
qs[j * 16];
105 for (
int ii = 0; ii < 16; ++ii) {
108 y[i].
bsums[j] = (int16_t)sum;
111 y[i].
d = 1.0f / iscale;
122 const char *ref_env = getenv(
"CK_DEBUG_Q8K_REF");
123 if (ref_env && atoi(ref_env) != 0) {
127#if defined(__AVX512F__) && defined(__AVX512BW__)
129#elif defined(__AVX2__)
131#elif defined(__AVX__)
133#elif defined(__SSE4_1__)
141 int num_rows,
int k) {
142 if (!x || !vy || num_rows <= 0 || k <= 0) {
145 assert(k %
QK_K == 0);
148 const int blocks_per_row = k /
QK_K;
155 for (
int row = 0; row < num_rows; ++row) {
157 x + (
size_t)row * (
size_t)k,
158 y + (
size_t)row * (
size_t)blocks_per_row,
167 const int nb = k /
QK_K;
170 for (
int i = 0; i < nb; ++i) {
171 uint8_t sc[8], m_val[8];
178 for (
int j = 0; j <
QK_K / 16; ++j) {
179 sumi += (int)x[i].bsums[j] * (
int)m_val[j / 2];
182 int32_t scaled_sum = 0;
183 for (
int group = 0; group < 4; ++group) {
184 const uint8_t *qs = &w[i].
qs[group * 32];
185 const int8_t *q8_lo = &x[i].
qs[group * 64];
186 const int8_t *q8_hi = q8_lo + 32;
189 for (
int l = 0; l < 32; ++l) {
190 lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo[l];
191 hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi[l];
193 scaled_sum += (int32_t)sc[2 * group] * lo;
194 scaled_sum += (int32_t)sc[2 * group + 1] * hi;
196 sumf += d * (float)scaled_sum - dmin * (
float)sumi;
206 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
212 const int blocks_per_row = K /
QK_K;
214 for (
int row = 0; row < M; ++row) {
215 const block_q4_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
222 static int cached = -1;
224 const char *env = getenv(
"CK_DEBUG_Q4K_Q8_REF");
225 cached = (env && env[0] && env[0] !=
'0') ? 1 : 0;
246 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
249 if (ith < 0 || nth <= 0 || ith >= nth) {
254 const int dr = (M + nth - 1) / nth;
255 const int r0 = dr * ith;
256 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
264 const int blocks_per_row = K /
QK_K;
267 for (
int row = r0; row < r1; ++row) {
268 const block_q4_K *w_row = blocks + (size_t)row * (
size_t)blocks_per_row;
282#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && !defined(CK_NO_AVX512_VNNI)
285#elif defined(__AVX2__)
287#elif defined(__AVX__)
290#elif defined(__SSE4_1__)
302 if (!Y || !W || !X_q8 || M <= 0 || N <= 0 || K <= 0) {
307 const int blocks_per_vec = K /
QK_K;
309 for (
int n = 0; n < N; ++n) {
310 const block_q8_K *x_row = X + (size_t)n * (
size_t)blocks_per_vec;
324} gemm_q4_k_q8_k_work_t;
328 gemm_q4_k_q8_k_work_t *a = (gemm_q4_k_q8_k_work_t *)args;
329 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
333 const int dr = (a->M_out + nth - 1) / nth;
334 const int r0 = dr * ith;
335 const int r1 = (r0 + dr < a->M_out) ? (r0 + dr) : a->M_out;
336 if (r0 >= a->M_out) {
341 const block_q4_K *w_start = blocks + (size_t)r0 * (
size_t)a->blocks_per_row;
342 const int rows = r1 - r0;
344 for (
int n = 0; n < a->N_batch; ++n) {
345 const block_q8_K *x_row = a->X + (size_t)n * (
size_t)a->blocks_per_vec;
359 if (!Y || !W || !X_q8 || M <= 0 || N <= 0 || K <= 0) {
364 const int blocks_per_vec = K /
QK_K;
365 const int blocks_per_row = K /
QK_K;
366 const size_t work_items = (size_t)M * (
size_t)N;
368 if (work_items >= 4096u && M >= 512 && N > 1) {
371 int active_threads = pool_threads;
372 if (active_threads > M) {
375 if (active_threads > 1) {
376 gemm_q4_k_q8_k_work_t work = {
383 .blocks_per_vec = blocks_per_vec,
384 .blocks_per_row = blocks_per_row,
391 for (
int n = 0; n < N; ++n) {
392 const block_q8_K *x_row = X + (size_t)n * (
size_t)blocks_per_vec;
403 if (!A_q8 || !B || !
C) {
406 if (M <= 0 || N <= 0 || K <= 0) {
416 for (
int i = 0; i < M; ++i) {
417 float *row =
C + (size_t)i * (
size_t)N;
418 for (
int j = 0; j < N; ++j) {
Persistent pthread thread pool for CK-Engine inference.
void ck_threadpool_dispatch_n(ck_threadpool_t *pool, int active_threads, ck_work_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
Quantization block structures for weight-only quantization.
#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 gemv_q4_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
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 gemv_q4_k_q8_k_vnni(float *y, const void *W, const void *x_q8, int M, int K)
void quantize_row_q8_k(const float *x, void *vy, int k)
void gemm_nt_q4_k_q8_k(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q4_k_q8_k_parallel(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
void quantize_batch_q8_k_4row_nearest_even(const float *x, void *vy, int num_rows, int k)
void gemm_q4_k_q8_k_ref(float *Y, const void *W, const void *X_q8, int M, int N, int K)
static int ck_nearest_int(float fval)
void quantize_row_q8_k_avx(const float *x, void *vy, int k)
void gemv_q4_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_q4_k_q8_k(float *Y, const void *W, const void *X_q8, int M, int N, int K)
static int ck_q4k_q8k_force_ref(void)
void quantize_row_q8_k_sse(const float *x, void *vy, int k)
static float dot_q4_k_q8_k_ref(const block_q4_K *w, const block_q8_K *x, int k)
static void gemm_q4_k_q8_k_thread_fn(int ith, int nth, void *args)
void quantize_row_q8_k_ref(const float *x, void *vy, int k)
void gemv_q4_k_q8_k_avx(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q4_k_q8_k_sse(float *y, const void *W, const void *x_q8, int M, int K)
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)