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

Q4_K (weights) x Q8_K (activations) kernels for inference. More...

#include <assert.h>
#include <immintrin.h>
#include <math.h>
#include <stdlib.h>
#include <string.h>
#include "ck_threadpool.h"
#include "ckernel_quant.h"

Go to the source code of this file.

Functions

static int ck_nearest_int (float fval)
 
static int ck_q4k_q8k_force_ref (void)
 
static float dot_q4_k_q8_k_ref (const block_q4_K *w, const block_q8_K *x, 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 gemm_q4_k_q8_k (float *Y, const void *W, const void *X_q8, int M, int N, 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 void gemm_q4_k_q8_k_thread_fn (int ith, int nth, void *args)
 
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_avx (float *y, const void *W, const void *x_q8, int M, int K)
 
void gemv_q4_k_q8_k_avx2 (float *y, const void *W, const void *x_q8, int M, 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 gemv_q4_k_q8_k_ref (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)
 
void gemv_q4_k_q8_k_vnni (float *y, const void *W, const void *x_q8, int M, int K)
 
void quantize_batch_q8_k_4row_nearest_even (const float *x, void *vy, int num_rows, int k)
 
void quantize_row_q8_k (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_avx2 (const float *x, void *vy, int k)
 
void quantize_row_q8_k_avx512 (const float *x, void *vy, int k)
 
void quantize_row_q8_k_ref (const float *x, void *vy, int k)
 
void quantize_row_q8_k_sse (const float *x, void *vy, int k)
 

Detailed Description

Q4_K (weights) x Q8_K (activations) kernels for inference.

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 decode-style matvec/matmul where weights are Q4_K and the activations are quantized on-the-fly to Q8_K. This is inference-only; no backward pass is provided here.

Definition in file gemm_kernels_q4k_q8k.c.

Function Documentation

◆ ck_nearest_int()

static int ck_nearest_int ( float  fval)
inlinestatic

Definition at line 48 of file gemm_kernels_q4k_q8k.c.

48 {
49 /* Bit-level round-to-nearest from llama.cpp (fast + deterministic). */
50 float val = fval + 12582912.f;
51 int i;
52 memcpy(&i, &val, sizeof(int));
53 return (i & 0x007fffff) - 0x00400000;
54}

Referenced by quantize_row_q8_k_ref().

◆ ck_q4k_q8k_force_ref()

static int ck_q4k_q8k_force_ref ( void  )
static

Definition at line 220 of file gemm_kernels_q4k_q8k.c.

221{
222 static int cached = -1;
223 if (cached < 0) {
224 const char *env = getenv("CK_DEBUG_Q4K_Q8_REF");
225 cached = (env && env[0] && env[0] != '0') ? 1 : 0;
226 }
227 return cached;
228}

Referenced by gemv_q4_k_q8_k().

◆ dot_q4_k_q8_k_ref()

static float dot_q4_k_q8_k_ref ( const block_q4_K w,
const block_q8_K x,
int  k 
)
static

Definition at line 163 of file gemm_kernels_q4k_q8k.c.

166{
167 const int nb = k / QK_K;
168 float sumf = 0.0f;
169
170 for (int i = 0; i < nb; ++i) {
171 uint8_t sc[8], m_val[8];
172 unpack_q4_k_scales(w[i].scales, sc, m_val);
173
174 const float d = CK_FP16_TO_FP32(w[i].d) * x[i].d;
175 const float dmin = CK_FP16_TO_FP32(w[i].dmin) * x[i].d;
176
177 int sumi = 0;
178 for (int j = 0; j < QK_K / 16; ++j) {
179 sumi += (int)x[i].bsums[j] * (int)m_val[j / 2];
180 }
181
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;
187 int32_t lo = 0;
188 int32_t hi = 0;
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];
192 }
193 scaled_sum += (int32_t)sc[2 * group] * lo;
194 scaled_sum += (int32_t)sc[2 * group + 1] * hi;
195 }
196 sumf += d * (float)scaled_sum - dmin * (float)sumi;
197 }
198 return sumf;
199}
#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.
#define QK_K
uint8_t qs[256/2]
int8_t qs[256]

References CK_FP16_TO_FP32, block_q8_K::d, QK_K, block_q4_K::qs, block_q8_K::qs, and unpack_q4_k_scales().

Referenced by gemv_q4_k_q8_k_parallel(), and gemv_q4_k_q8_k_ref().

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

Definition at line 397 of file gemm_kernels_q4k_q8k.c.

402{
403 if (!A_q8 || !B || !C) {
404 return;
405 }
406 if (M <= 0 || N <= 0 || K <= 0) {
407 return;
408 }
409
410 gemm_q4_k_q8_k(C, B, A_q8, /*M_out=*/N, /*N_batch=*/M, K);
411
412 if (!bias) {
413 return;
414 }
415
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) {
419 row[j] += bias[j];
420 }
421 }
422}
void gemm_q4_k_q8_k(float *Y, const void *W, const void *X_q8, int M, int N, int K)
#define C(color)
Definition show_config.c:39

References C, and gemm_q4_k_q8_k().

Referenced by ck_attention_project_head_major_q4_k_q8_k(), ck_layer_forward_rmsnorm_swiglu_decode_q4_k(), ck_mlp_swiglu_forward_q4_k_q8_k(), ck_mlp_swiglu_forward_q4_k_q8_k_prefill(), ck_qkv_project_head_major_token_q4_k_q8_k(), ck_test_gemm_q4_k(), gemm_nt_q8_k_mlp_dispatch(), gemm_nt_q8_k_qkv_dispatch(), model_forward_prefill_impl(), model_forward_prefill_impl(), model_forward_prefill_impl(), model_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), and qwen2_0_5b_decode_forward_prefill_impl().

◆ gemm_q4_k_q8_k()

void gemm_q4_k_q8_k ( float *  Y,
const void *  W,
const void *  X_q8,
int  M,
int  N,
int  K 
)

Definition at line 354 of file gemm_kernels_q4k_q8k.c.

358{
359 if (!Y || !W || !X_q8 || M <= 0 || N <= 0 || K <= 0) {
360 return;
361 }
362
363 const block_q8_K *X = (const block_q8_K *)X_q8;
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;
367
368 if (work_items >= 4096u && M >= 512 && N > 1) {
369 ck_threadpool_t *pool = ck_threadpool_global();
370 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
371 int active_threads = pool_threads;
372 if (active_threads > M) {
373 active_threads = M;
374 }
375 if (active_threads > 1) {
376 gemm_q4_k_q8_k_work_t work = {
377 .Y = Y,
378 .W = W,
379 .X = X,
380 .M_out = M,
381 .N_batch = N,
382 .K = K,
383 .blocks_per_vec = blocks_per_vec,
384 .blocks_per_row = blocks_per_row,
385 };
386 ck_threadpool_dispatch_n(pool, active_threads, gemm_q4_k_q8_k_thread_fn, &work);
387 return;
388 }
389 }
390
391 for (int n = 0; n < N; ++n) {
392 const block_q8_K *x_row = X + (size_t)n * (size_t)blocks_per_vec;
393 gemv_q4_k_q8_k(&Y[n * M], W, x_row, M, K);
394 }
395}
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)
void gemv_q4_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
static void gemm_q4_k_q8_k_thread_fn(int ith, int nth, void *args)

References ck_threadpool_dispatch_n(), ck_threadpool_global(), ck_threadpool_n_threads(), gemm_q4_k_q8_k_thread_fn(), gemv_q4_k_q8_k(), and QK_K.

Referenced by gemm_nt_q4_k_q8_k().

◆ gemm_q4_k_q8_k_ref()

void gemm_q4_k_q8_k_ref ( float *  Y,
const void *  W,
const void *  X_q8,
int  M,
int  N,
int  K 
)

Definition at line 297 of file gemm_kernels_q4k_q8k.c.

301{
302 if (!Y || !W || !X_q8 || M <= 0 || N <= 0 || K <= 0) {
303 return;
304 }
305
306 const block_q8_K *X = (const block_q8_K *)X_q8;
307 const int blocks_per_vec = K / QK_K;
308
309 for (int n = 0; n < N; ++n) {
310 const block_q8_K *x_row = X + (size_t)n * (size_t)blocks_per_vec;
311 gemv_q4_k_q8_k_ref(&Y[n * M], W, x_row, M, K);
312 }
313}
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

References gemv_q4_k_q8_k_ref(), and QK_K.

◆ gemm_q4_k_q8_k_thread_fn()

static void gemm_q4_k_q8_k_thread_fn ( int  ith,
int  nth,
void *  args 
)
static

Definition at line 326 of file gemm_kernels_q4k_q8k.c.

327{
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) {
330 return;
331 }
332
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) {
337 return;
338 }
339
340 const block_q4_K *blocks = (const block_q4_K *)a->W;
341 const block_q4_K *w_start = blocks + (size_t)r0 * (size_t)a->blocks_per_row;
342 const int rows = r1 - r0;
343
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;
346 gemv_q4_k_q8_k(a->Y + (size_t)n * (size_t)a->M_out + (size_t)r0,
347 w_start,
348 x_row,
349 rows,
350 a->K);
351 }
352}

References gemv_q4_k_q8_k().

Referenced by gemm_q4_k_q8_k().

◆ gemv_q4_k_q8_k()

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

Definition at line 273 of file gemm_kernels_q4k_q8k.c.

277{
278 if (ck_q4k_q8k_force_ref()) {
279 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
280 return;
281 }
282#if defined(__AVX512VNNI__) && defined(__AVX512VL__) && !defined(CK_NO_AVX512_VNNI)
283 /* VNNI: Best for decode (single token) - INT8 dot product acceleration */
284 gemv_q4_k_q8_k_vnni(y, W, x_q8, M, K);
285#elif defined(__AVX2__)
286 gemv_q4_k_q8_k_avx2(y, W, x_q8, M, K);
287#elif defined(__AVX__)
288 /* AVX version uses maddubs_epi16 (more efficient than SSE) */
289 gemv_q4_k_q8_k_avx(y, W, x_q8, M, K);
290#elif defined(__SSE4_1__)
291 gemv_q4_k_q8_k_sse(y, W, x_q8, M, K);
292#else
293 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
294#endif
295}
void gemv_q4_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
void gemv_q4_k_q8_k_vnni(float *y, const void *W, const void *x_q8, int M, int K)
static int ck_q4k_q8k_force_ref(void)
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)

References ck_q4k_q8k_force_ref(), gemv_q4_k_q8_k_avx(), gemv_q4_k_q8_k_avx2(), gemv_q4_k_q8_k_ref(), gemv_q4_k_q8_k_sse(), and gemv_q4_k_q8_k_vnni().

Referenced by ck_moe_q4k_llama_projection(), ck_moe_q4k_mixed_route_work(), ck_moe_q4k_q5k_route_work(), ck_test_gemv_q4_k(), ck_test_vec_dot_q4_k_q8_k(), fused_rmsnorm_linear_q4k(), gemm_q4_k_q8_k(), gemm_q4_k_q8_k_compact_rows4(), gemm_q4_k_q8_k_thread_fn(), gemv_q4_k(), model_decode_token(), model_decode_token(), model_layer_0_decode(), model_layer_0_decode(), model_layer_10_decode(), model_layer_10_decode(), model_layer_11_decode(), model_layer_11_decode(), model_layer_12_decode(), model_layer_12_decode(), model_layer_13_decode(), model_layer_13_decode(), model_layer_14_decode(), model_layer_14_decode(), model_layer_15_decode(), model_layer_15_decode(), model_layer_16_decode(), model_layer_16_decode(), model_layer_17_decode(), model_layer_17_decode(), model_layer_18_decode(), model_layer_18_decode(), model_layer_19_decode(), model_layer_19_decode(), model_layer_1_decode(), model_layer_1_decode(), model_layer_20_decode(), model_layer_20_decode(), model_layer_21_decode(), model_layer_21_decode(), model_layer_22_decode(), model_layer_22_decode(), model_layer_23_decode(), model_layer_23_decode(), model_layer_2_decode(), model_layer_2_decode(), model_layer_3_decode(), model_layer_3_decode(), model_layer_4_decode(), model_layer_4_decode(), model_layer_5_decode(), model_layer_5_decode(), model_layer_6_decode(), model_layer_6_decode(), model_layer_7_decode(), model_layer_7_decode(), model_layer_8_decode(), model_layer_8_decode(), model_layer_9_decode(), model_layer_9_decode(), moe_swiglu_expert_forward_q4k_q4k_workspace(), moe_swiglu_expert_forward_q4k_q5k_workspace(), moe_swiglu_expert_forward_q4k_q6k_workspace(), moe_swiglu_shared_forward_q4k_q4k_workspace(), moe_swiglu_shared_forward_q4k_q6k_workspace(), qwen2_0_5b_decode_decode_token(), qwen2_0_5b_decode_decode_token(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), and unfused_rmsnorm_linear_q4k_ref().

◆ gemv_q4_k_q8_k_avx()

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

Definition at line 251 of file gemm_kernels_q4k_avx.c.

255{
256 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
257}
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

Referenced by gemv_q4_k_q8_k().

◆ gemv_q4_k_q8_k_avx2()

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

Definition at line 118 of file gemm_kernels_q4k_q8k_avx2.c.

122{
123#if defined(__AVX2__)
124 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
125 return;
126 }
127
128 const block_q4_K *blocks = (const block_q4_K *)W;
129 const block_q8_K *x = (const block_q8_K *)x_q8;
130 const int blocks_per_row = K / QK_K;
131
132 for (int row = 0; row < M; ++row) {
133 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
134 y[row] = dot_q4_k_q8_k_avx2(w_row, x, K);
135 }
136#else
137 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
138#endif
139}
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

Referenced by gemv_q4_k_q8_k().

◆ gemv_q4_k_q8_k_parallel()

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

Definition at line 240 of file gemm_kernels_q4k_q8k.c.

245{
246 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
247 return;
248 }
249 if (ith < 0 || nth <= 0 || ith >= nth) {
250 return;
251 }
252
253 /* Compute row range for this thread */
254 const int dr = (M + nth - 1) / nth;
255 const int r0 = dr * ith;
256 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
257
258 if (r0 >= M) {
259 return; /* This thread has no work */
260 }
261
262 const block_q4_K *blocks = (const block_q4_K *)W;
263 const block_q8_K *x = (const block_q8_K *)x_q8;
264 const int blocks_per_row = K / QK_K;
265
266 /* Only process rows [r0, r1) */
267 for (int row = r0; row < r1; ++row) {
268 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
269 y[row] = dot_q4_k_q8_k_ref(w_row, x, K);
270 }
271}
static float dot_q4_k_q8_k_ref(const block_q4_K *w, const block_q8_K *x, int k)

References dot_q4_k_q8_k_ref(), and QK_K.

Referenced by gemv_q4_k_q8_k_parallel_simd().

◆ gemv_q4_k_q8_k_ref()

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

Definition at line 201 of file gemm_kernels_q4k_q8k.c.

205{
206 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
207 return;
208 }
209
210 const block_q4_K *blocks = (const block_q4_K *)W;
211 const block_q8_K *x = (const block_q8_K *)x_q8;
212 const int blocks_per_row = K / QK_K;
213
214 for (int row = 0; row < M; ++row) {
215 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
216 y[row] = dot_q4_k_q8_k_ref(w_row, x, K);
217 }
218}

References dot_q4_k_q8_k_ref(), and QK_K.

Referenced by gemm_q4_k_q8_k_ref(), gemv_q4_k_q8_k(), gemv_q4_k_q8_k_amx(), gemv_q4_k_q8_k_avx(), gemv_q4_k_q8_k_avx2(), and gemv_q4_k_q8_k_vnni().

◆ gemv_q4_k_q8_k_sse()

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

Definition at line 33 of file gemm_kernels_q4k_sse.c.

37{
38 const block_q4_K *blocks = (const block_q4_K *)W;
39 const block_q8_K *x = (const block_q8_K *)x_q8;
40 const int blocks_per_row = K / QK_K;
41
42 const __m128i mask_low = _mm_set1_epi8(0x0F);
43
44 for (int row = 0; row < M; ++row) {
45 float sumf = 0.0f;
46 const block_q4_K *w_row = blocks + row * blocks_per_row;
47
48 for (int i = 0; i < blocks_per_row; ++i) {
49 const block_q4_K *b4 = &w_row[i];
50 const block_q8_K *b8 = &x[i];
51
52 // Unpack scales (same as ref)
53 uint8_t sc[8], m_val[8];
54 unpack_q4_k_scales(b4->scales, sc, m_val);
55
56 float d = CK_FP16_TO_FP32(b4->d) * b8->d;
57 float dmin = CK_FP16_TO_FP32(b4->dmin) * b8->d;
58
59 int is = 0;
60 int q_offset = 0;
61
62 // Process 4 chunks of 64 elements (256 total)
63 for (int j = 0; j < QK_K; j += 64) {
64 // We process 32 bytes of qs (covering 64 elements via low/high nibbles)
65 // We access qs[0..31] relative to q_offset
66
67 // Accumulators for this 64-element chunk
68 __m128i acc_lo = _mm_setzero_si128();
69 __m128i acc_hi = _mm_setzero_si128();
70
71 // Inner loop: 2 iters of 16 bytes (32 elements)
72 for (int l = 0; l < 32; l += 16) {
73 // Load 16 bytes of Q4
74 __m128i q4_vec = _mm_loadu_si128((const __m128i *)(b4->qs + q_offset + l));
75
76 // Low nibbles -> correspond to q8_lo (elements j+l .. j+l+15)
77 __m128i q4_lo = _mm_and_si128(q4_vec, mask_low);
78
79 // High nibbles -> correspond to q8_hi (elements j+32+l .. j+32+l+15)
80 __m128i q4_hi = _mm_and_si128(_mm_srli_epi16(q4_vec, 4), mask_low);
81
82 // Load Q8
83 __m128i q8_lo_vec = _mm_loadu_si128((const __m128i *)(b8->qs + j + l));
84 __m128i q8_hi_vec = _mm_loadu_si128((const __m128i *)(b8->qs + j + 32 + l));
85
86 // Expand and Multiply-Add: Q4(u8) * Q8(s8) -> i32
87 // Since Q4 is u8 and Q8 is s8, we use intermediate i16
88
89 // LO PART
90 __m128i q4_lo_16_L = _mm_cvtepu8_epi16(q4_lo); // lower 8 -> 16
91 __m128i q8_lo_16_L = _mm_cvtepi8_epi16(q8_lo_vec);
92 __m128i prod_lo_L = _mm_madd_epi16(q4_lo_16_L, q8_lo_16_L); // i32
93 acc_lo = _mm_add_epi32(acc_lo, prod_lo_L);
94
95 __m128i q4_lo_16_H = _mm_cvtepu8_epi16(_mm_srli_si128(q4_lo, 8)); // upper 8 -> 16
96 __m128i q8_lo_16_H = _mm_cvtepi8_epi16(_mm_srli_si128(q8_lo_vec, 8));
97 __m128i prod_lo_H = _mm_madd_epi16(q4_lo_16_H, q8_lo_16_H); // i32
98 acc_lo = _mm_add_epi32(acc_lo, prod_lo_H);
99
100 // HI PART
101 __m128i q4_hi_16_L = _mm_cvtepu8_epi16(q4_hi);
102 __m128i q8_hi_16_L = _mm_cvtepi8_epi16(q8_hi_vec);
103 __m128i prod_hi_L = _mm_madd_epi16(q4_hi_16_L, q8_hi_16_L);
104 acc_hi = _mm_add_epi32(acc_hi, prod_hi_L);
105
106 __m128i q4_hi_16_H = _mm_cvtepu8_epi16(_mm_srli_si128(q4_hi, 8));
107 __m128i q8_hi_16_H = _mm_cvtepi8_epi16(_mm_srli_si128(q8_hi_vec, 8));
108 __m128i prod_hi_H = _mm_madd_epi16(q4_hi_16_H, q8_hi_16_H);
109 acc_hi = _mm_add_epi32(acc_hi, prod_hi_H);
110 }
111
112 int32_t sum_q4q8_lo = hsum_epi32_sse(acc_lo);
113 int32_t sum_q4q8_hi = hsum_epi32_sse(acc_hi);
114
115 /* bsums: each bsum is 16 elements */
116 int32_t bsum_lo = (int32_t)b8->bsums[j / 16] +
117 (int32_t)b8->bsums[j / 16 + 1];
118 int32_t bsum_hi = (int32_t)b8->bsums[(j + 32) / 16] +
119 (int32_t)b8->bsums[(j + 32) / 16 + 1];
120
121 sumf += d * (float)sc[is] * (float)sum_q4q8_lo;
122 sumf -= dmin * (float)m_val[is] * (float)bsum_lo;
123 sumf += d * (float)sc[is + 1] * (float)sum_q4q8_hi;
124 sumf -= dmin * (float)m_val[is + 1] * (float)bsum_hi;
125
126 q_offset += 32;
127 is += 2;
128 }
129 }
130 y[row] = sumf;
131 }
132}
static int32_t hsum_epi32_sse(__m128i v)
uint8_t scales[12]
int16_t bsums[256/16]

Referenced by gemv_q4_k_q8_k().

◆ gemv_q4_k_q8_k_vnni()

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

Definition at line 3534 of file gemm_kernels_q4k_q8k_vnni.c.

3538{
3539#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3540 const char *fast_env = getenv("CK_ENABLE_Q4K_Q8K_VNNI_FAST");
3541 const int fast_disabled = fast_env && fast_env[0] && fast_env[0] == '0';
3542 if (!fast_disabled && !ck_strict_parity_enabled()) {
3543 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
3544 return;
3545 }
3546
3547 const block_q4_K *blocks = (const block_q4_K *)W;
3548 const block_q8_K *x = (const block_q8_K *)x_q8;
3549 const int blocks_per_row = K / QK_K;
3550
3551 for (int row = 0; row < M; ++row) {
3552 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
3553 float sum = 0.0f;
3554 for (int b = 0; b < blocks_per_row; ++b) {
3555 sum += dot_q4_k_q8_k_vnni_block(&w_row[b], &x[b]);
3556 }
3557 y[row] = sum;
3558 }
3559 return;
3560 }
3561#endif
3562
3563 /* Strict/debug parity keeps the llama-style scalar accumulation path.
3564 * Production AVX-512 hosts use VNNI by default; set
3565 * CK_ENABLE_Q4K_Q8K_VNNI_FAST=0 or CK_DEBUG_Q4K_Q8_REF=1 when attributing
3566 * borderline logit movement against scalar/reference behavior.
3567 */
3568 gemv_q4_k_q8_k_ref(y, W, x_q8, M, K);
3569}
int ck_strict_parity_enabled(void)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)

Referenced by gemv_q4_k_q8_k().

◆ quantize_batch_q8_k_4row_nearest_even()

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

Definition at line 140 of file gemm_kernels_q4k_q8k.c.

141 {
142 if (!x || !vy || num_rows <= 0 || k <= 0) {
143 return;
144 }
145 assert(k % QK_K == 0);
146
147 block_q8_K *y = (block_q8_K *)vy;
148 const int blocks_per_row = k / QK_K;
149
150 /* Q8_K bytes are a numerical ABI. The previous four-row AVX2 path used
151 * a different max/scale evaluation order and changed real Qwen3-VL
152 * visual-prefix scales by one FP32 ULP. Keep this public grouped ABI on
153 * the canonical row provider until an optimized implementation is
154 * byte-exact against llama.cpp on both synthetic and production inputs. */
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,
159 k);
160 }
161}
void quantize_row_q8_k(const float *x, void *vy, int k)

References QK_K, and quantize_row_q8_k().

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

References quantize_row_q8_k_avx(), quantize_row_q8_k_avx2(), quantize_row_q8_k_avx512(), quantize_row_q8_k_ref(), and quantize_row_q8_k_sse().

Referenced by ck_attention_project_head_major_q4_k_q8_k(), ck_layer_forward_rmsnorm_swiglu_decode_q4_k(), ck_mlp_swiglu_forward_q4_k_q8_k(), ck_mlp_swiglu_forward_q4_k_q8_k_prefill(), ck_moe_q4k_mixed_route_parallel(), ck_moe_q4k_mixed_route_work(), ck_moe_q4k_q5k_bucket_work(), ck_moe_q4k_q5k_quantize_work(), ck_moe_q4k_q5k_route_parallel(), ck_moe_q4k_q5k_route_work(), ck_moe_shared_q4k_gated_workspace(), ck_qkv_project_head_major_q4_k_q8_k(), ck_test_gemm_q4_k(), ck_test_gemm_q6_k(), ck_test_gemv_q4_k(), ck_test_gemv_q6_k(), ck_test_quantize_q8_k(), decode_layer_parallel(), fused_mlp_swiglu_prefill_w1w2_quant(), fused_rmsnorm_qkv_prefill_head_major_quant(), gemm_nt_q5_0_sse_v2(), gemm_nt_q5_k_prepared(), gemm_nt_q5_k_prepared_m4(), gemm_nt_q5_k_ref(), gemm_nt_q6_k_sse(), gemv_q4_k(), gemv_q5_k_ref(), hyper_connection_mix_quantized(), mlp_parallel(), model_decode_token(), model_decode_token(), model_forward_prefill_impl(), model_forward_prefill_impl(), model_forward_prefill_impl(), model_forward_prefill_impl(), model_layer_0_decode(), model_layer_0_decode(), model_layer_10_decode(), model_layer_10_decode(), model_layer_11_decode(), model_layer_11_decode(), model_layer_12_decode(), model_layer_12_decode(), model_layer_13_decode(), model_layer_13_decode(), model_layer_14_decode(), model_layer_14_decode(), model_layer_15_decode(), model_layer_15_decode(), model_layer_16_decode(), model_layer_16_decode(), model_layer_17_decode(), model_layer_17_decode(), model_layer_18_decode(), model_layer_18_decode(), model_layer_19_decode(), model_layer_19_decode(), model_layer_1_decode(), model_layer_1_decode(), model_layer_20_decode(), model_layer_20_decode(), model_layer_21_decode(), model_layer_21_decode(), model_layer_22_decode(), model_layer_22_decode(), model_layer_23_decode(), model_layer_23_decode(), model_layer_2_decode(), model_layer_2_decode(), model_layer_3_decode(), model_layer_3_decode(), model_layer_4_decode(), model_layer_4_decode(), model_layer_5_decode(), model_layer_5_decode(), model_layer_6_decode(), model_layer_6_decode(), model_layer_7_decode(), model_layer_7_decode(), model_layer_8_decode(), model_layer_8_decode(), model_layer_9_decode(), model_layer_9_decode(), moe_swiglu_expert_forward_q4k_q4k_workspace(), moe_swiglu_expert_forward_q4k_q5_0_workspace(), moe_swiglu_expert_forward_q4k_q5k_workspace(), moe_swiglu_expert_forward_q4k_q6k_workspace(), moe_swiglu_expert_forward_q4k_q8_0_workspace(), moe_swiglu_shared_forward_q4k_q4k_workspace(), moe_swiglu_shared_forward_q4k_q6k_workspace(), quantize_batch_q8_k(), quantize_batch_q8_k_4row_nearest_even(), qwen2_0_5b_decode_decode_token(), qwen2_0_5b_decode_decode_token(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_forward_prefill_impl(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), swiglu_forward_q8_k(), and unfused_rmsnorm_linear_q4k_ref().

◆ quantize_row_q8_k_avx()

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

Definition at line 20 of file quantize_row_q8_k_avx.c.

20 {
21 /* Reuse the parity-clean SIMD implementation until a wider AVX variant is
22 * worth maintaining separately. */
23#if defined(__SSE4_1__)
24 quantize_row_q8_k_sse(x, vy, k);
25#else
26 quantize_row_q8_k_ref(x, vy, k);
27#endif
28}
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)

References quantize_row_q8_k_ref(), and quantize_row_q8_k_sse().

Referenced by quantize_row_q8_k().

◆ quantize_row_q8_k_avx2()

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

Definition at line 14 of file quantize_row_q8_k_avx2.c.

14 {
15 quantize_row_q8_k_ref(x, vy, k);
16}
void quantize_row_q8_k_ref(const float *x, void *vy, int k)

References quantize_row_q8_k_ref().

Referenced by quantize_row_q8_k().

◆ quantize_row_q8_k_avx512()

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

Definition at line 8 of file quantize_row_q8_k_avx512.c.

8 {
9 /*
10 * Q8_K bytes are part of the numerical ABI consumed by Q4_K/Q6_K dots.
11 * Keep AVX-512 on the reference contract until a vector implementation is
12 * byte-exact for every block, including multi-block activation rows.
13 */
14 quantize_row_q8_k_ref(x, vy, k);
15}
void quantize_row_q8_k_ref(const float *x, void *vy, int k)

References quantize_row_q8_k_ref().

Referenced by quantize_row_q8_k().

◆ quantize_row_q8_k_ref()

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

Definition at line 61 of file gemm_kernels_q4k_q8k.c.

61 {
62 if (!x || !vy || k <= 0) {
63 return;
64 }
65 assert(k % QK_K == 0);
66 const int nb = k / QK_K;
67 block_q8_K *y = (block_q8_K *)vy;
68
69 for (int i = 0; i < nb; ++i) {
70 float max = 0.0f;
71 float amax = 0.0f;
72 for (int j = 0; j < QK_K; ++j) {
73 float ax = fabsf(x[j]);
74 if (ax > amax) {
75 amax = ax;
76 max = x[j];
77 }
78 }
79 if (!amax) {
80 y[i].d = 0.0f;
81 memset(y[i].qs, 0, sizeof(y[i].qs));
82 memset(y[i].bsums, 0, sizeof(y[i].bsums));
83 x += QK_K;
84 continue;
85 }
86
87 const float iscale = -127.0f / max;
88 for (int j = 0; j < QK_K; ++j) {
89 /* llama.cpp rounds the multiply before adding nearest_int's magic
90 * constant. Contracting both operations changes Q8_K tie cases. */
91 float scaled = iscale * x[j];
92 int v = ck_nearest_int(scaled);
93 if (v > 127) {
94 v = 127;
95 }
96 if (v < -128) {
97 v = -128;
98 }
99 y[i].qs[j] = (int8_t)v;
100 }
101
102 for (int j = 0; j < QK_K / 16; ++j) {
103 int sum = 0;
104 const int8_t *qs = &y[i].qs[j * 16];
105 for (int ii = 0; ii < 16; ++ii) {
106 sum += qs[ii];
107 }
108 y[i].bsums[j] = (int16_t)sum;
109 }
110
111 y[i].d = 1.0f / iscale;
112 x += QK_K;
113 }
114}
static int ck_nearest_int(float fval)

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

Referenced by quantize_row_q8_k(), quantize_row_q8_k_avx(), quantize_row_q8_k_avx2(), and quantize_row_q8_k_avx512().

◆ quantize_row_q8_k_sse()

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

Definition at line 22 of file quantize_row_q8_k_sse.c.

22 {
23 if (!x || !vy || k <= 0) {
24 return;
25 }
26 assert(k % QK_K == 0);
27
28 const int nb = k / QK_K;
29 block_q8_K *y = (block_q8_K *)vy;
30
31 for (int i = 0; i < nb; ++i) {
32 /* Keep the exact signed-max selection contract from llama.cpp/ref. */
33 float max = 0.0f;
34 float amax = 0.0f;
35 for (int j = 0; j < QK_K; ++j) {
36 const float xv = x[j];
37 const float ax = fabsf(xv);
38 if (ax > amax) {
39 amax = ax;
40 max = xv;
41 }
42 }
43
44 if (amax == 0.0f) {
45 y[i].d = 0.0f;
46 memset(y[i].qs, 0, sizeof(y[i].qs));
47 memset(y[i].bsums, 0, sizeof(y[i].bsums));
48 x += QK_K;
49 continue;
50 }
51
52 const float iscale = -127.0f / max;
53 const __m128 v_iscale = _mm_set1_ps(iscale);
54 const __m128 v_magic = _mm_set1_ps(12582912.0f);
55 const __m128i v_mantissa = _mm_set1_epi32(0x007fffff);
56 const __m128i v_bias = _mm_set1_epi32(0x00400000);
57 const __m128i v_min = _mm_set1_epi32(-128);
58 const __m128i v_max = _mm_set1_epi32(127);
59
60 for (int j = 0; j < QK_K; j += 16) {
61 const __m128 x0 = _mm_loadu_ps(x + j + 0);
62 const __m128 x1 = _mm_loadu_ps(x + j + 4);
63 const __m128 x2 = _mm_loadu_ps(x + j + 8);
64 const __m128 x3 = _mm_loadu_ps(x + j + 12);
65
66 __m128i q0 = _mm_sub_epi32(
67 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x0, v_iscale), v_magic)), v_mantissa),
68 v_bias);
69 __m128i q1 = _mm_sub_epi32(
70 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x1, v_iscale), v_magic)), v_mantissa),
71 v_bias);
72 __m128i q2 = _mm_sub_epi32(
73 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x2, v_iscale), v_magic)), v_mantissa),
74 v_bias);
75 __m128i q3 = _mm_sub_epi32(
76 _mm_and_si128(_mm_castps_si128(_mm_add_ps(_mm_mul_ps(x3, v_iscale), v_magic)), v_mantissa),
77 v_bias);
78
79 q0 = _mm_min_epi32(_mm_max_epi32(q0, v_min), v_max);
80 q1 = _mm_min_epi32(_mm_max_epi32(q1, v_min), v_max);
81 q2 = _mm_min_epi32(_mm_max_epi32(q2, v_min), v_max);
82 q3 = _mm_min_epi32(_mm_max_epi32(q3, v_min), v_max);
83
84 const __m128i q01 = _mm_packs_epi32(q0, q1);
85 const __m128i q23 = _mm_packs_epi32(q2, q3);
86 const __m128i q0123 = _mm_packs_epi16(q01, q23);
87
88 _mm_storeu_si128((__m128i *)(y[i].qs + j), q0123);
89
90 int sum = 0;
91 for (int ii = 0; ii < 16; ++ii) {
92 sum += y[i].qs[j + ii];
93 }
94 y[i].bsums[j / 16] = (int16_t)sum;
95 }
96
97 y[i].d = 1.0f / iscale;
98 x += QK_K;
99 }
100}

Referenced by quantize_row_q8_k().