← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q8_0.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels_q8_0.c
3 * @brief GEMM/GEMV kernels with Q8_0 quantized weights
4 *
5 * CK-ENGINE KERNEL RULES:
6 * =======================
7 * 1. NO malloc/free - memory via bump allocator, pointers passed in
8 * 2. NO OpenMP - parallelization at orchestrator/codegen layer
9 * 3. API must define: inputs, outputs, workspace, and memory layouts
10 * 4. Pure computation - deterministic, no side effects
11 *
12 * After changes: make test && make llamacpp-parity-full
13 *
14 * Q8_0 Format:
15 * - 32 weights per block
16 * - 1 FP16 scale per block
17 * - 34 bytes per 32 weights = 8.5 bits/weight
18 * - Weights stored as signed 8-bit integers
19 *
20 * Operations:
21 * Forward: Y = W @ X (W is Q8_0, X and Y are FP32)
22 * Backward: dX = W^T @ dY (gradient w.r.t. input)
23 *
24 * Note: Q8_0 is often used for activation quantization or as an
25 * intermediate format. Higher precision than Q4_0/Q4_K.
26 */
27
28#include <stdint.h>
29#include <stddef.h>
30#include <stdlib.h>
31#include <string.h>
32#include "ckernel_quant.h"
33#include "ck_features.h"
34#include "ck_speed_profiles.h"
35
36/* Include SIMD headers based on available extensions */
37#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__)
38#include <immintrin.h>
39#endif
40#if defined(__ARM_NEON) || defined(__aarch64__)
41#include <arm_neon.h>
42#endif
43
44void quantize_row_q8_k(const float *x, void *vy, int k);
45
46static int ck_q8_0_debug_ref(void)
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}
55
56static int ck_q8_0_q8_0_debug_ref(void)
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}
65
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}
74
75static inline int ck_nearest_int_q8_0(float fval) {
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}
82
83#if defined(__INTEL_LLVM_COMPILER)
84#if defined(__clang__)
85#define CK_Q80_NOINLINE_OPTNONE __attribute__((noinline, optnone))
86#elif defined(__GNUC__)
87#define CK_Q80_NOINLINE_OPTNONE __attribute__((noinline, optimize("O0")))
88#else
89#define CK_Q80_NOINLINE_OPTNONE
90#endif
91
92static CK_Q80_NOINLINE_OPTNONE float
93ck_q8_0_div_rounded_f32(float numerator, float denominator)
94{
95 /*
96 * Preserve llama.cpp's scalar IEEE-754 division before the FP16 scale
97 * conversion. ICX can otherwise strength-reduce division by 127 at -O3;
98 * inputs near an FP16 midpoint then select the adjacent Q8_0 scale.
99 */
100 volatile float n = numerator;
101 volatile float d = denominator;
102 volatile float result = n / d;
103 return result;
104}
105#endif
106
107/* ============================================================================
108 * Q8_0 Quantization
109 *
110 * Quantizes FP32 values to Q8_0 format (32 elements per block).
111 * Each block has:
112 * - 1 FP16 scale (computed as max abs value / 127)
113 * - 32 int8 quantized values
114 *
115 * This matches llama.cpp's quantize_row_q8_0.
116 * ============================================================================ */
117
118/**
119 * @brief Quantize FP32 to Q8_0 format (scalar reference)
120 *
121 * @param x Input FP32 values
122 * @param vy Output Q8_0 blocks
123 * @param k Number of elements (must be multiple of 32)
124 */
125void quantize_row_q8_0(const float *x, void *vy, int k)
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}
238
239/**
240 * @brief Batch quantize FP32 to Q8_0 format (row-major output)
241 *
242 * Quantizes multiple rows of FP32 data to Q8_0 format, placing each row's
243 * Q8_0 output at the correct byte offset for GEMM compatibility.
244 *
245 * Memory layout:
246 * Input: [num_rows, k] FP32, row-major (stride = k * sizeof(float))
247 * Output: [num_rows, q8_row_bytes] Q8_0, row-major (stride = q8_row_bytes)
248 *
249 * where q8_row_bytes = (k / 32) * sizeof(block_q8_0) = (k / 32) * 34
250 *
251 * @param x Input FP32 values [num_rows * k]
252 * @param vy Output Q8_0 blocks [num_rows * (k/32) blocks]
253 * @param num_rows Number of rows (batch size / tokens)
254 * @param k Elements per row (must be multiple of 32)
255 */
256void quantize_batch_q8_0(const float *x, void *vy, int num_rows, int k)
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}
272
273/**
274 * @brief Batch quantize FP32 to Q8_K format (row-major output)
275 *
276 * Same as quantize_batch_q8_0 but for Q8_K format (super-blocks).
277 *
278 * @param x Input FP32 values [num_rows * k]
279 * @param vy Output Q8_K blocks
280 * @param num_rows Number of rows (batch size / tokens)
281 * @param k Elements per row (must be multiple of 256)
282 */
283void quantize_batch_q8_k(const float *x, void *vy, int num_rows, int k)
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}
302
303/* ============================================================================
304 * Forward Pass: GEMV y = W @ x
305 * ============================================================================ */
306
307/**
308 * @brief Matrix-vector multiply with Q8_0 weights (scalar reference)
309 *
310 * @param y Output vector [M]
311 * @param W Weight matrix in Q8_0 format [M x K]
312 * @param x Input vector [K]
313 * @param M Number of output rows
314 * @param K Number of columns (must be multiple of 32)
315 */
316void gemv_q8_0_ref(float *y,
317 const void *W,
318 const float *x,
319 int M, int K)
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}
340
341#ifdef __AVX512F__
342/**
343 * @brief Matrix-vector multiply with Q8_0 weights (AVX-512)
344 */
345void gemv_q8_0_avx512(float *y,
346 const void *W,
347 const float *x,
348 int M, int K)
349{
350 const block_q8_0 *blocks = (const block_q8_0 *)W;
351 const int blocks_per_row = K / QK8_0;
352
353 for (int row = 0; row < M; row++) {
354 __m512 acc = _mm512_setzero_ps();
355
356 for (int b = 0; b < blocks_per_row; b++) {
357 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
358 const __m512 vscale = _mm512_set1_ps(CK_FP16_TO_FP32(block->d));
359 const float *xp = &x[b * QK8_0];
360
361 /* Process 32 weights in two batches of 16 */
362 for (int chunk = 0; chunk < 2; chunk++) {
363 /* Load 16 x int8 weights */
364 __m128i q8 = _mm_loadu_si128((const __m128i *)&block->qs[chunk * 16]);
365
366 /* Sign-extend to 32-bit */
367 __m512i q32 = _mm512_cvtepi8_epi32(q8);
368
369 /* Convert to float and scale */
370 __m512 w = _mm512_mul_ps(_mm512_cvtepi32_ps(q32), vscale);
371
372 /* Load input */
373 __m512 x_vec = _mm512_loadu_ps(&xp[chunk * 16]);
374
375 /* FMA */
376 acc = _mm512_fmadd_ps(w, x_vec, acc);
377 }
378 }
379
380 y[row] = _mm512_reduce_add_ps(acc);
381 }
382}
383#endif
384
385/* ============================================================================
386 * AVX2 Implementation (Haswell+, 256-bit integer operations)
387 *
388 * Q8_0 format: 32 signed int8 weights per block
389 * - d: FP16 scale
390 * - qs: 32 int8 weights
391 * - Dequant: w = d * q
392 *
393 * AVX2 provides _mm256_cvtepi8_epi32 for efficient 8-to-32 sign extension.
394 * Processes 8 weights at a time with full 256-bit FMA.
395 * ============================================================================ */
396
397#if defined(__AVX2__) && !defined(__AVX512F__)
398
399/* Helper: AVX2 horizontal sum of 8 floats */
400static inline float hsum_avx2_q8(__m256 v) {
401 __m128 lo = _mm256_castps256_ps128(v);
402 __m128 hi = _mm256_extractf128_ps(v, 1);
403 lo = _mm_add_ps(lo, hi); /* 4 floats */
404 __m128 shuf = _mm_shuffle_ps(lo, lo, _MM_SHUFFLE(2, 3, 0, 1));
405 __m128 sums = _mm_add_ps(lo, shuf);
406 shuf = _mm_movehl_ps(shuf, sums);
407 sums = _mm_add_ss(sums, shuf);
408 return _mm_cvtss_f32(sums);
409}
410
411/**
412 * @brief Matrix-vector multiply with Q8_0 weights (AVX2 optimized)
413 *
414 * Uses AVX2's _mm256_cvtepi8_epi32 for efficient sign extension.
415 * Processes 8 weights at a time with FMA.
416 */
417void gemv_q8_0_avx2(float *y,
418 const void *W,
419 const float *x,
420 int M, int K)
421{
422 const block_q8_0 *blocks = (const block_q8_0 *)W;
423 const int blocks_per_row = K / QK8_0; /* QK8_0 = 32 */
424
425 for (int row = 0; row < M; row++) {
426 __m256 acc = _mm256_setzero_ps();
427
428 for (int b = 0; b < blocks_per_row; b++) {
429 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
430 const float d = CK_FP16_TO_FP32(block->d);
431 const __m256 vscale = _mm256_set1_ps(d);
432 const float *xp = &x[b * QK8_0];
433
434 /* Process 32 weights in 4 groups of 8 using AVX2 */
435
436 /* Group 0: weights 0-7 */
437 {
438 __m128i q8 = _mm_loadl_epi64((const __m128i *)&block->qs[0]);
439 __m256i q32 = _mm256_cvtepi8_epi32(q8);
440 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
441 __m256 xv = _mm256_loadu_ps(&xp[0]);
442 acc = _mm256_fmadd_ps(wf, xv, acc);
443 }
444
445 /* Group 1: weights 8-15 */
446 {
447 __m128i q8 = _mm_loadl_epi64((const __m128i *)&block->qs[8]);
448 __m256i q32 = _mm256_cvtepi8_epi32(q8);
449 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
450 __m256 xv = _mm256_loadu_ps(&xp[8]);
451 acc = _mm256_fmadd_ps(wf, xv, acc);
452 }
453
454 /* Group 2: weights 16-23 */
455 {
456 __m128i q8 = _mm_loadl_epi64((const __m128i *)&block->qs[16]);
457 __m256i q32 = _mm256_cvtepi8_epi32(q8);
458 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
459 __m256 xv = _mm256_loadu_ps(&xp[16]);
460 acc = _mm256_fmadd_ps(wf, xv, acc);
461 }
462
463 /* Group 3: weights 24-31 */
464 {
465 __m128i q8 = _mm_loadl_epi64((const __m128i *)&block->qs[24]);
466 __m256i q32 = _mm256_cvtepi8_epi32(q8);
467 __m256 wf = _mm256_mul_ps(_mm256_cvtepi32_ps(q32), vscale);
468 __m256 xv = _mm256_loadu_ps(&xp[24]);
469 acc = _mm256_fmadd_ps(wf, xv, acc);
470 }
471 }
472
473 y[row] = hsum_avx2_q8(acc);
474 }
475}
476#endif /* __AVX2__ && !__AVX512F__ */
477
478/* ============================================================================
479 * AVX Implementation with True SIMD (256-bit float + 128-bit integer)
480 *
481 * Q8_0 format: 32 signed int8 weights per block
482 * - d: FP16 scale
483 * - qs: 32 int8 weights
484 * - Dequant: w = d * q
485 *
486 * This is much simpler than Q5_0 since weights are already in int8 format.
487 * We use SSE for integer-to-float conversion and AVX for accumulation.
488 * ============================================================================ */
489
490#if defined(__AVX__) && !defined(__AVX2__) && !defined(__AVX512F__)
491
492/* Helper: SSE horizontal sum of 4 floats */
493static inline float hsum_sse_q8(__m128 v) {
494 __m128 shuf = _mm_shuffle_ps(v, v, _MM_SHUFFLE(2, 3, 0, 1));
495 __m128 sums = _mm_add_ps(v, shuf);
496 shuf = _mm_movehl_ps(shuf, sums);
497 sums = _mm_add_ss(sums, shuf);
498 return _mm_cvtss_f32(sums);
499}
500
501/**
502 * @brief Matrix-vector multiply with Q8_0 weights (AVX + SSE optimized)
503 *
504 * Uses full SIMD: SSE for int8->float conversion, SSE/AVX for dot product.
505 * ~4-6x faster than scalar reference on Ivy Bridge.
506 */
507void gemv_q8_0_avx(float *y,
508 const void *W,
509 const float *x,
510 int M, int K)
511{
512 const block_q8_0 *blocks = (const block_q8_0 *)W;
513 const int blocks_per_row = K / QK8_0; /* QK8_0 = 32 */
514
515 for (int row = 0; row < M; row++) {
516 /* Use 4 SSE accumulators for ILP */
517 __m128 acc0 = _mm_setzero_ps();
518 __m128 acc1 = _mm_setzero_ps();
519 __m128 acc2 = _mm_setzero_ps();
520 __m128 acc3 = _mm_setzero_ps();
521
522 for (int b = 0; b < blocks_per_row; b++) {
523 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
524 const float d = CK_FP16_TO_FP32(block->d);
525 const float *xp = &x[b * QK8_0];
526 const __m128 vscale = _mm_set1_ps(d);
527
528 /* Load 32 int8 weights in 2 SSE loads of 16 bytes each */
529 __m128i q8_0 = _mm_loadu_si128((const __m128i *)&block->qs[0]);
530 __m128i q8_1 = _mm_loadu_si128((const __m128i *)&block->qs[16]);
531
532 /* Process first 16 weights: convert int8 -> int16 -> int32 -> float */
533 /* Chunk 0: weights 0-3 */
534 {
535 __m128i q16 = _mm_cvtepi8_epi16(q8_0); /* 8 int8 -> 8 int16 */
536 __m128i q32 = _mm_cvtepi16_epi32(q16); /* 4 int16 -> 4 int32 */
537 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
538 __m128 vx = _mm_loadu_ps(&xp[0]);
539 acc0 = _mm_add_ps(acc0, _mm_mul_ps(w, vx));
540 }
541
542 /* Chunk 1: weights 4-7 */
543 {
544 __m128i q16 = _mm_cvtepi8_epi16(q8_0);
545 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
546 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
547 __m128 vx = _mm_loadu_ps(&xp[4]);
548 acc1 = _mm_add_ps(acc1, _mm_mul_ps(w, vx));
549 }
550
551 /* Chunk 2: weights 8-11 */
552 {
553 __m128i q8_shifted = _mm_srli_si128(q8_0, 8); /* shift right 8 bytes */
554 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
555 __m128i q32 = _mm_cvtepi16_epi32(q16);
556 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
557 __m128 vx = _mm_loadu_ps(&xp[8]);
558 acc2 = _mm_add_ps(acc2, _mm_mul_ps(w, vx));
559 }
560
561 /* Chunk 3: weights 12-15 */
562 {
563 __m128i q8_shifted = _mm_srli_si128(q8_0, 8);
564 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
565 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
566 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
567 __m128 vx = _mm_loadu_ps(&xp[12]);
568 acc3 = _mm_add_ps(acc3, _mm_mul_ps(w, vx));
569 }
570
571 /* Process second 16 weights (16-31) */
572 /* Chunk 4: weights 16-19 */
573 {
574 __m128i q16 = _mm_cvtepi8_epi16(q8_1);
575 __m128i q32 = _mm_cvtepi16_epi32(q16);
576 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
577 __m128 vx = _mm_loadu_ps(&xp[16]);
578 acc0 = _mm_add_ps(acc0, _mm_mul_ps(w, vx));
579 }
580
581 /* Chunk 5: weights 20-23 */
582 {
583 __m128i q16 = _mm_cvtepi8_epi16(q8_1);
584 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
585 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
586 __m128 vx = _mm_loadu_ps(&xp[20]);
587 acc1 = _mm_add_ps(acc1, _mm_mul_ps(w, vx));
588 }
589
590 /* Chunk 6: weights 24-27 */
591 {
592 __m128i q8_shifted = _mm_srli_si128(q8_1, 8);
593 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
594 __m128i q32 = _mm_cvtepi16_epi32(q16);
595 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
596 __m128 vx = _mm_loadu_ps(&xp[24]);
597 acc2 = _mm_add_ps(acc2, _mm_mul_ps(w, vx));
598 }
599
600 /* Chunk 7: weights 28-31 */
601 {
602 __m128i q8_shifted = _mm_srli_si128(q8_1, 8);
603 __m128i q16 = _mm_cvtepi8_epi16(q8_shifted);
604 __m128i q32 = _mm_cvtepi16_epi32(_mm_srli_si128(q16, 8));
605 __m128 w = _mm_mul_ps(_mm_cvtepi32_ps(q32), vscale);
606 __m128 vx = _mm_loadu_ps(&xp[28]);
607 acc3 = _mm_add_ps(acc3, _mm_mul_ps(w, vx));
608 }
609 }
610
611 /* Combine accumulators and reduce */
612 __m128 sum01 = _mm_add_ps(acc0, acc1);
613 __m128 sum23 = _mm_add_ps(acc2, acc3);
614 __m128 sum = _mm_add_ps(sum01, sum23);
615
616 y[row] = hsum_sse_q8(sum);
617 }
618}
619#endif /* __AVX__ && !__AVX512F__ */
620
621#if defined(__SSE4_1__)
622#include <immintrin.h>
623
624/* Helper macro: extract 4 int8 weights at byte offset, convert to float, multiply with x */
625#define SSE_Q8_BLOCK(q8_reg, offset, xp, d_val, acc) do { \
626 __m128 vx = _mm_loadu_ps(&(xp)[offset]); \
627 __m128i qw = _mm_cvtepi8_epi32(_mm_srli_si128(q8_reg, offset)); \
628 __m128 vw = _mm_cvtepi32_ps(qw); \
629 acc = _mm_add_ps(acc, _mm_mul_ps(_mm_mul_ps(vw, vx), _mm_set1_ps(d_val))); \
630} while(0)
631
632void gemv_q8_0_sse(float *y,
633 const void *W,
634 const float *x,
635 int M, int K)
636{
637 const block_q8_0 *blocks = (const block_q8_0 *)W;
638 const int blocks_per_row = K / QK8_0;
639
640 for (int row = 0; row < M; row++) {
641 __m128 acc = _mm_setzero_ps();
642
643 for (int b = 0; b < blocks_per_row; b++) {
644 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
645 const float d_val = CK_FP16_TO_FP32(block->d);
646 const float *xp = &x[b * QK8_0];
647
648 /* Load 32 weights (signed 8-bit) in two 16-byte chunks */
649 __m128i q8_0 = _mm_loadu_si128((const __m128i *)&block->qs[0]);
650 __m128i q8_1 = _mm_loadu_si128((const __m128i *)&block->qs[16]);
651
652 /* Process first 16 weights (q8_0) - unrolled with compile-time constants */
653 SSE_Q8_BLOCK(q8_0, 0, xp, d_val, acc);
654 SSE_Q8_BLOCK(q8_0, 4, xp, d_val, acc);
655 SSE_Q8_BLOCK(q8_0, 8, xp, d_val, acc);
656 SSE_Q8_BLOCK(q8_0, 12, xp, d_val, acc);
657
658 /* Process second 16 weights (q8_1) - offset xp by 16 */
659 const float *xp1 = xp + 16;
660 SSE_Q8_BLOCK(q8_1, 0, xp1, d_val, acc);
661 SSE_Q8_BLOCK(q8_1, 4, xp1, d_val, acc);
662 SSE_Q8_BLOCK(q8_1, 8, xp1, d_val, acc);
663 SSE_Q8_BLOCK(q8_1, 12, xp1, d_val, acc);
664 }
665
666 /* Horizontal sum */
667 acc = _mm_add_ps(acc, _mm_shuffle_ps(acc, acc, _MM_SHUFFLE(1, 0, 3, 2)));
668 acc = _mm_add_ps(acc, _mm_shuffle_ps(acc, acc, _MM_SHUFFLE(0, 1, 0, 1)));
669 _mm_store_ss(&y[row], acc);
670 }
671}
672
673#undef SSE_Q8_BLOCK
674#endif
675
676/**
677 * @brief Auto-dispatch GEMV for Q8_0 weights based on CPU features
678 *
679 * Dispatch priority (best available):
680 * 1. AVX-512 (512-bit vectors) - Intel Skylake-X+
681 * 2. AVX2+FMA (256-bit vectors) - Intel Haswell+
682 * 3. AVX (256-bit vectors) - Intel Sandy Bridge+
683 * 4. SSE4.1 (128-bit vectors) - Intel Nehalem+
684 * 5. Reference (scalar) - Fallback
685 *
686 * Uses ck_features.h for standardized feature detection.
687 *
688 * @param y Output vector [M]
689 * @param W Weight matrix in Q8_0 format [M x K]
690 * @param x Input vector [K]
691 * @param M Number of output rows
692 * @param K Number of input columns (hidden dimension)
693 */
694void gemv_q8_0(float *y,
695 const void *W,
696 const float *x,
697 int M, int K)
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}
717
718/* ============================================================================
719 * Forward Pass: GEMM Y = W @ X
720 * ============================================================================ */
721
722/**
723 * @brief Matrix-matrix multiply with Q8_0 weights
724 */
725void gemm_q8_0(float *Y,
726 const void *W,
727 const float *X,
728 int M, int N, int K)
729{
730 for (int n = 0; n < N; n++) {
731 gemv_q8_0(&Y[n * M], W, &X[n * K], M, K);
732 }
733}
734
735/* ============================================================================
736 * GEMM NT: C = A @ B^T + bias (B stored as N rows of K elements)
737 * ============================================================================ */
738
739/**
740 * @brief Matrix-matrix multiply: C[M,N] = A[M,K] @ B[N,K]^T + bias
741 *
742 * @param A Input matrix [M x K], row-major FP32
743 * @param B Weight matrix in Q8_0 format, [N x K] stored row-major
744 * @param bias Optional bias [N], NULL if not used
745 * @param C Output [M x N], row-major FP32
746 * @param M Batch size (number of tokens)
747 * @param N Output dimension (number of rows in B)
748 * @param K Input dimension
749 */
750static void gemm_nt_q8_0_rowloop(const float *A,
751 const void *B,
752 const float *bias,
753 float *C,
754 int M, int N, int K)
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}
763
764#if defined(__AVX512F__)
765static inline float hsum512_ps_q80(__m512 v)
766{
767 return _mm512_reduce_add_ps(v);
768}
769
770static void gemm_nt_q8_0_m4n4_avx512(const float *A,
771 const void *B,
772 const float *bias,
773 float *C,
774 int M, int N, int K)
775{
776 const block_q8_0 *blocks = (const block_q8_0 *)B;
777 const int blocks_per_row = K / QK8_0;
778 const int M4 = M & ~3;
779 const int N4 = N & ~3;
780
781 for (int m = 0; m < M4; m += 4) {
782 for (int n = 0; n < N4; n += 4) {
783 __m512 acc00 = _mm512_setzero_ps(), acc01 = _mm512_setzero_ps();
784 __m512 acc02 = _mm512_setzero_ps(), acc03 = _mm512_setzero_ps();
785 __m512 acc10 = _mm512_setzero_ps(), acc11 = _mm512_setzero_ps();
786 __m512 acc12 = _mm512_setzero_ps(), acc13 = _mm512_setzero_ps();
787 __m512 acc20 = _mm512_setzero_ps(), acc21 = _mm512_setzero_ps();
788 __m512 acc22 = _mm512_setzero_ps(), acc23 = _mm512_setzero_ps();
789 __m512 acc30 = _mm512_setzero_ps(), acc31 = _mm512_setzero_ps();
790 __m512 acc32 = _mm512_setzero_ps(), acc33 = _mm512_setzero_ps();
791
792 const block_q8_0 *b0 = blocks + (size_t)(n + 0) * blocks_per_row;
793 const block_q8_0 *b1 = blocks + (size_t)(n + 1) * blocks_per_row;
794 const block_q8_0 *b2 = blocks + (size_t)(n + 2) * blocks_per_row;
795 const block_q8_0 *b3 = blocks + (size_t)(n + 3) * blocks_per_row;
796 const float *a0 = A + (size_t)(m + 0) * K;
797 const float *a1 = A + (size_t)(m + 1) * K;
798 const float *a2 = A + (size_t)(m + 2) * K;
799 const float *a3 = A + (size_t)(m + 3) * K;
800
801 for (int ib = 0; ib < blocks_per_row; ++ib) {
802 const int k0 = ib * QK8_0;
803 for (int chunk = 0; chunk < 2; ++chunk) {
804 const int off = chunk * 16;
805 const __m512 x0 = _mm512_loadu_ps(a0 + k0 + off);
806 const __m512 x1 = _mm512_loadu_ps(a1 + k0 + off);
807 const __m512 x2 = _mm512_loadu_ps(a2 + k0 + off);
808 const __m512 x3 = _mm512_loadu_ps(a3 + k0 + off);
809
810 __m128i q0 = _mm_loadu_si128((const __m128i *)&b0[ib].qs[off]);
811 __m128i q1 = _mm_loadu_si128((const __m128i *)&b1[ib].qs[off]);
812 __m128i q2 = _mm_loadu_si128((const __m128i *)&b2[ib].qs[off]);
813 __m128i q3 = _mm_loadu_si128((const __m128i *)&b3[ib].qs[off]);
814
815 __m512 w0 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q0)),
816 _mm512_set1_ps(CK_FP16_TO_FP32(b0[ib].d)));
817 __m512 w1 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q1)),
818 _mm512_set1_ps(CK_FP16_TO_FP32(b1[ib].d)));
819 __m512 w2 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q2)),
820 _mm512_set1_ps(CK_FP16_TO_FP32(b2[ib].d)));
821 __m512 w3 = _mm512_mul_ps(_mm512_cvtepi32_ps(_mm512_cvtepi8_epi32(q3)),
822 _mm512_set1_ps(CK_FP16_TO_FP32(b3[ib].d)));
823
824 acc00 = _mm512_fmadd_ps(w0, x0, acc00);
825 acc01 = _mm512_fmadd_ps(w1, x0, acc01);
826 acc02 = _mm512_fmadd_ps(w2, x0, acc02);
827 acc03 = _mm512_fmadd_ps(w3, x0, acc03);
828 acc10 = _mm512_fmadd_ps(w0, x1, acc10);
829 acc11 = _mm512_fmadd_ps(w1, x1, acc11);
830 acc12 = _mm512_fmadd_ps(w2, x1, acc12);
831 acc13 = _mm512_fmadd_ps(w3, x1, acc13);
832 acc20 = _mm512_fmadd_ps(w0, x2, acc20);
833 acc21 = _mm512_fmadd_ps(w1, x2, acc21);
834 acc22 = _mm512_fmadd_ps(w2, x2, acc22);
835 acc23 = _mm512_fmadd_ps(w3, x2, acc23);
836 acc30 = _mm512_fmadd_ps(w0, x3, acc30);
837 acc31 = _mm512_fmadd_ps(w1, x3, acc31);
838 acc32 = _mm512_fmadd_ps(w2, x3, acc32);
839 acc33 = _mm512_fmadd_ps(w3, x3, acc33);
840 }
841 }
842
843 const float b00 = bias ? bias[n + 0] : 0.0f;
844 const float b01 = bias ? bias[n + 1] : 0.0f;
845 const float b02 = bias ? bias[n + 2] : 0.0f;
846 const float b03 = bias ? bias[n + 3] : 0.0f;
847 C[(size_t)(m + 0) * N + n + 0] = hsum512_ps_q80(acc00) + b00;
848 C[(size_t)(m + 0) * N + n + 1] = hsum512_ps_q80(acc01) + b01;
849 C[(size_t)(m + 0) * N + n + 2] = hsum512_ps_q80(acc02) + b02;
850 C[(size_t)(m + 0) * N + n + 3] = hsum512_ps_q80(acc03) + b03;
851 C[(size_t)(m + 1) * N + n + 0] = hsum512_ps_q80(acc10) + b00;
852 C[(size_t)(m + 1) * N + n + 1] = hsum512_ps_q80(acc11) + b01;
853 C[(size_t)(m + 1) * N + n + 2] = hsum512_ps_q80(acc12) + b02;
854 C[(size_t)(m + 1) * N + n + 3] = hsum512_ps_q80(acc13) + b03;
855 C[(size_t)(m + 2) * N + n + 0] = hsum512_ps_q80(acc20) + b00;
856 C[(size_t)(m + 2) * N + n + 1] = hsum512_ps_q80(acc21) + b01;
857 C[(size_t)(m + 2) * N + n + 2] = hsum512_ps_q80(acc22) + b02;
858 C[(size_t)(m + 2) * N + n + 3] = hsum512_ps_q80(acc23) + b03;
859 C[(size_t)(m + 3) * N + n + 0] = hsum512_ps_q80(acc30) + b00;
860 C[(size_t)(m + 3) * N + n + 1] = hsum512_ps_q80(acc31) + b01;
861 C[(size_t)(m + 3) * N + n + 2] = hsum512_ps_q80(acc32) + b02;
862 C[(size_t)(m + 3) * N + n + 3] = hsum512_ps_q80(acc33) + b03;
863 }
864 }
865
866 if (N4 < N) {
867 gemm_nt_q8_0_rowloop(A, B, bias, C, M, N, K);
868 return;
869 }
870 for (int m = M4; m < M; ++m) {
871 gemv_q8_0(&C[(size_t)m * N], B, &A[(size_t)m * K], N, K);
872 if (bias) {
873 for (int n = 0; n < N; ++n) C[(size_t)m * N + n] += bias[n];
874 }
875 }
876}
877#endif
878
879void gemm_nt_q8_0(const float *A,
880 const void *B,
881 const float *bias,
882 float *C,
883 int M, int N, int K)
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}
893
894/* ============================================================================
895 * Backward Pass: Gradient w.r.t. Input
896 * ============================================================================ */
897
898/**
899 * @brief Backward pass: compute input gradient (scalar reference)
900 *
901 * @param dX Output gradient w.r.t. input [K]
902 * @param W Weight matrix in Q8_0 format [M x K]
903 * @param dY Gradient w.r.t. output [M]
904 * @param M Number of output rows
905 * @param K Number of columns (input dimension)
906 */
908 const void *W,
909 const float *dY,
910 int M, int K)
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}
933
934#ifdef __AVX512F__
935/**
936 * @brief Backward pass with AVX-512
937 */
938void gemv_q8_0_backward_avx512(float *dX,
939 const void *W,
940 const float *dY,
941 int M, int K)
942{
943 const block_q8_0 *blocks = (const block_q8_0 *)W;
944 const int blocks_per_row = K / QK8_0;
945
946 /* Zero output */
947 memset(dX, 0, K * sizeof(float));
948
949 for (int row = 0; row < M; row++) {
950 const __m512 vdy = _mm512_set1_ps(dY[row]);
951
952 for (int b = 0; b < blocks_per_row; b++) {
953 const block_q8_0 *block = &blocks[row * blocks_per_row + b];
954 const __m512 vscale = _mm512_set1_ps(CK_FP16_TO_FP32(block->d));
955 float *dxp = &dX[b * QK8_0];
956
957 /* Process 32 weights in two batches of 16 */
958 for (int chunk = 0; chunk < 2; chunk++) {
959 /* Load and dequantize weights */
960 __m128i q8 = _mm_loadu_si128((const __m128i *)&block->qs[chunk * 16]);
961 __m512i q32 = _mm512_cvtepi8_epi32(q8);
962 __m512 w = _mm512_mul_ps(_mm512_cvtepi32_ps(q32), vscale);
963
964 /* Compute gradient */
965 __m512 grad = _mm512_mul_ps(w, vdy);
966
967 /* Accumulate */
968 __m512 dx_cur = _mm512_loadu_ps(&dxp[chunk * 16]);
969 _mm512_storeu_ps(&dxp[chunk * 16], _mm512_add_ps(dx_cur, grad));
970 }
971 }
972 }
973}
974#endif
975
976/**
977 * @brief Auto-dispatch backward
978 */
979void gemv_q8_0_backward(float *dX,
980 const void *W,
981 const float *dY,
982 int M, int K)
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}
990
991/**
992 * @brief Batched backward pass
993 */
994void gemm_q8_0_backward(float *dX,
995 const void *W,
996 const float *dY,
997 int M, int N, int K)
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}
1003
1004/* ============================================================================
1005 * Dot Product Utility
1006 * ============================================================================ */
1007
1008float dot_q8_0(const void *w_q8_0, const float *x, int K)
1009{
1010 float result;
1011 gemv_q8_0(&result, w_q8_0, x, 1, K);
1012 return result;
1013}
1014
1015#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
1016/* Match ggml x86 horizontal reduction order for Q8_0 x Q8_0 parity. */
1017static inline float hsum_float_8_q8_0(const __m256 x)
1018{
1019 __m128 res = _mm256_extractf128_ps(x, 1);
1020 res = _mm_add_ps(res, _mm256_castps256_ps128(x));
1021 res = _mm_add_ps(res, _mm_movehl_ps(res, res));
1022 res = _mm_add_ss(res, _mm_movehdup_ps(res));
1023 return _mm_cvtss_f32(res);
1024}
1025
1026#if defined(__AVX2__) || defined(__AVX512F__)
1027static inline __m256 sum_i16_pairs_float_q8_0_avx2(const __m256i x)
1028{
1029 const __m256i ones = _mm256_set1_epi16(1);
1030 const __m256i summed_pairs = _mm256_madd_epi16(ones, x);
1031 return _mm256_cvtepi32_ps(summed_pairs);
1032}
1033
1034static inline __m256 mul_sum_us8_pairs_float_q8_0_avx2(const __m256i ax, const __m256i sy)
1035{
1036#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1037 const __m256i zero = _mm256_setzero_si256();
1038 const __m256i summed_pairs = _mm256_dpbusd_epi32(zero, ax, sy);
1039 return _mm256_cvtepi32_ps(summed_pairs);
1040#elif defined(__AVXVNNI__)
1041 const __m256i zero = _mm256_setzero_si256();
1042 const __m256i summed_pairs = _mm256_dpbusd_avx_epi32(zero, ax, sy);
1043 return _mm256_cvtepi32_ps(summed_pairs);
1044#else
1045 const __m256i dot = _mm256_maddubs_epi16(ax, sy);
1046 return sum_i16_pairs_float_q8_0_avx2(dot);
1047#endif
1048}
1049
1050static inline __m256 mul_sum_i8_pairs_float_q8_0_avx2(const __m256i x, const __m256i y)
1051{
1052#if __AVXVNNIINT8__
1053 const __m256i zero = _mm256_setzero_si256();
1054 const __m256i summed_pairs = _mm256_dpbssd_epi32(zero, x, y);
1055 return _mm256_cvtepi32_ps(summed_pairs);
1056#else
1057 const __m256i ax = _mm256_sign_epi8(x, x);
1058 const __m256i sy = _mm256_sign_epi8(y, x);
1059 return mul_sum_us8_pairs_float_q8_0_avx2(ax, sy);
1060#endif
1061}
1062#elif defined(__AVX__)
1063static inline __m128i mul_add_epi8_sse_q8_0(const __m128i x, const __m128i y)
1064{
1065 const __m128i ax = _mm_sign_epi8(x, x);
1066 const __m128i sy = _mm_sign_epi8(y, x);
1067 return _mm_maddubs_epi16(ax, sy);
1068}
1069
1070static inline __m256 sum_i16_pairs_float_q8_0_avx(const __m128i xh, const __m128i xl)
1071{
1072 const __m128i ones = _mm_set1_epi16(1);
1073 const __m128i summed_pairsl = _mm_madd_epi16(ones, xl);
1074 const __m128i summed_pairsh = _mm_madd_epi16(ones, xh);
1075 const __m256i summed_pairs = _mm256_insertf128_si256(
1076 _mm256_castsi128_si256(summed_pairsl),
1077 summed_pairsh,
1078 1
1079 );
1080 return _mm256_cvtepi32_ps(summed_pairs);
1081}
1082
1083static inline __m256 mul_sum_i8_quad_float_q8_0_avx(const __m128i x_1_0,
1084 const __m128i x_1_1,
1085 const __m128i x_2_0,
1086 const __m128i x_2_1,
1087 const __m128i y_1_0,
1088 const __m128i y_1_1,
1089 const __m128i y_2_0,
1090 const __m128i y_2_1)
1091{
1092 const __m128i mone = _mm_set1_epi16(1);
1093
1094 const __m128i p16_1_0 = mul_add_epi8_sse_q8_0(x_1_0, y_1_0);
1095 const __m128i p16_1_1 = mul_add_epi8_sse_q8_0(x_1_1, y_1_1);
1096 const __m128i p16_2_0 = mul_add_epi8_sse_q8_0(x_2_0, y_2_0);
1097 const __m128i p16_2_1 = mul_add_epi8_sse_q8_0(x_2_1, y_2_1);
1098 const __m128i p_1_0 = _mm_madd_epi16(p16_1_0, mone);
1099 const __m128i p_1_1 = _mm_madd_epi16(p16_1_1, mone);
1100 const __m128i p_2_0 = _mm_madd_epi16(p16_2_0, mone);
1101 const __m128i p_2_1 = _mm_madd_epi16(p16_2_1, mone);
1102 const __m128i p_1 = _mm_add_epi32(p_1_0, p_1_1);
1103 const __m128i p_2 = _mm_add_epi32(p_2_0, p_2_1);
1104 const __m256i packed = _mm256_insertf128_si256(_mm256_castsi128_si256(p_1), p_2, 1);
1105 return _mm256_cvtepi32_ps(packed);
1106}
1107
1108static inline __m256 quad_fp16_delta_float_q8_0_avx(uint16_t x0, uint16_t y0, uint16_t x1, uint16_t y1)
1109{
1110 return _mm256_set_m128(
1111 _mm_set1_ps(CK_FP16_TO_FP32(x1) * CK_FP16_TO_FP32(y1)),
1112 _mm_set1_ps(CK_FP16_TO_FP32(x0) * CK_FP16_TO_FP32(y0))
1113 );
1114}
1115#endif
1116#endif
1117
1118/* ============================================================================
1119 * Quantized Dot Product: Q8_0 x Q8_0
1120 *
1121 * This matches llama.cpp's ggml_vec_dot_q8_0_q8_0 exactly.
1122 * Both weights and input are in Q8_0 format, enabling pure integer dot products.
1123 * Result: sum_blocks( (d_w * d_x) * sum_weights( w8 * x8 ) )
1124 *
1125 * Key difference from gemv_q8_0:
1126 * - gemv_q8_0: Takes FP32 input, dequantizes weights to FP32, FP32 dot
1127 * - vec_dot_q8_0_q8_0: Takes Q8_0 input, does integer dot, scales at end
1128 *
1129 * The quantized path is faster and matches llama.cpp for parity testing.
1130 * ============================================================================ */
1131
1132/**
1133 * @brief Quantized dot product: Q8_0 weights x Q8_0 input (scalar reference)
1134 *
1135 * @param n Number of elements (must be multiple of 32)
1136 * @param s Output: scalar dot product result
1137 * @param vx Q8_0 quantized weights
1138 * @param vy Q8_0 quantized input
1139 */
1140void vec_dot_q8_0_q8_0_ref(int n, float *s, const void *vx, const void *vy)
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}
1162
1163#if defined(__ARM_NEON) || defined(__aarch64__)
1164void vec_dot_q8_0_q8_0_neon(int n, float *s, const void *vx, const void *vy)
1165{
1166 const int qk = QK8_0;
1167 const int nb = n / qk;
1168
1169 const block_q8_0 *x = (const block_q8_0 *)vx;
1170 const block_q8_0 *y = (const block_q8_0 *)vy;
1171
1172 float sumf = 0.0f;
1173
1174 for (int ib = 0; ib < nb; ib++) {
1175 const int8x16_t x0 = vld1q_s8(&x[ib].qs[0]);
1176 const int8x16_t x1 = vld1q_s8(&x[ib].qs[16]);
1177 const int8x16_t y0 = vld1q_s8(&y[ib].qs[0]);
1178 const int8x16_t y1 = vld1q_s8(&y[ib].qs[16]);
1179
1180 int32x4_t acc = vdupq_n_s32(0);
1181
1182 const int16x8_t p0 = vmull_s8(vget_low_s8(x0), vget_low_s8(y0));
1183 const int16x8_t p1 = vmull_s8(vget_high_s8(x0), vget_high_s8(y0));
1184 const int16x8_t p2 = vmull_s8(vget_low_s8(x1), vget_low_s8(y1));
1185 const int16x8_t p3 = vmull_s8(vget_high_s8(x1), vget_high_s8(y1));
1186
1187 acc = vaddq_s32(acc, vpaddlq_s16(p0));
1188 acc = vaddq_s32(acc, vpaddlq_s16(p1));
1189 acc = vaddq_s32(acc, vpaddlq_s16(p2));
1190 acc = vaddq_s32(acc, vpaddlq_s16(p3));
1191
1192 int32_t lanes[4];
1193 vst1q_s32(lanes, acc);
1194 const int sumi = lanes[0] + lanes[1] + lanes[2] + lanes[3];
1195
1196 sumf += (float)sumi * (CK_FP16_TO_FP32(x[ib].d) * CK_FP16_TO_FP32(y[ib].d));
1197 }
1198
1199 *s = sumf;
1200}
1201#endif
1202
1203#if defined(__AVX2__) && !defined(__AVX512F__)
1204void vec_dot_q8_0_q8_0_avx2(int n, float *s, const void *vx, const void *vy)
1205{
1206 const int qk = QK8_0;
1207 const int nb = n / qk;
1208
1209 const block_q8_0 *x = (const block_q8_0 *)vx;
1210 const block_q8_0 *y = (const block_q8_0 *)vy;
1211
1212 int ib = 0;
1213 float sumf = 0.0f;
1214 __m256 acc = _mm256_setzero_ps();
1215
1216 for (; ib < nb; ++ib) {
1217 const __m256 d = _mm256_set1_ps(CK_FP16_TO_FP32(x[ib].d) * CK_FP16_TO_FP32(y[ib].d));
1218 const __m256i qx = _mm256_loadu_si256((const __m256i *)x[ib].qs);
1219 const __m256i qy = _mm256_loadu_si256((const __m256i *)y[ib].qs);
1220 const __m256 q = mul_sum_i8_pairs_float_q8_0_avx2(qx, qy);
1221#if defined(__FMA__)
1222 acc = _mm256_fmadd_ps(d, q, acc);
1223#else
1224 acc = _mm256_add_ps(_mm256_mul_ps(d, q), acc);
1225#endif
1226 }
1227
1228 sumf = hsum_float_8_q8_0(acc);
1229 *s = sumf;
1230}
1231#endif
1232
1233#ifdef __AVX512F__
1234/**
1235 * @brief Quantized dot product Q8_0 x Q8_0 (AVX-512)
1236 */
1237void vec_dot_q8_0_q8_0_avx512(int n, float *s, const void *vx, const void *vy)
1238{
1239 const int qk = QK8_0;
1240 const int nb = n / qk;
1241
1242 const block_q8_0 *x = (const block_q8_0 *)vx;
1243 const block_q8_0 *y = (const block_q8_0 *)vy;
1244
1245 /*
1246 * ggml keeps the Q8_0 dot on its AVX2 eight-lane accumulation tree even
1247 * in an AVX-512 build. Preserve that numerical contract here: reducing
1248 * every block to a scalar first changes FP32 addition order and is not
1249 * bit-exact at production vision widths.
1250 */
1251 __m256 acc = _mm256_setzero_ps();
1252 for (int ib = 0; ib < nb; ++ib) {
1253 const __m256 d = _mm256_set1_ps(
1254 CK_FP16_TO_FP32(x[ib].d) * CK_FP16_TO_FP32(y[ib].d));
1255 const __m256i qx = _mm256_loadu_si256((const __m256i *)x[ib].qs);
1256 const __m256i qy = _mm256_loadu_si256((const __m256i *)y[ib].qs);
1257 const __m256 q = mul_sum_i8_pairs_float_q8_0_avx2(qx, qy);
1258#if defined(__FMA__)
1259 acc = _mm256_fmadd_ps(d, q, acc);
1260#else
1261 acc = _mm256_add_ps(_mm256_mul_ps(d, q), acc);
1262#endif
1263 }
1264
1265 *s = hsum_float_8_q8_0(acc);
1266}
1267#endif
1268
1269#if defined(__AVX__) && !defined(__AVX2__) && !defined(__AVX512F__)
1270/**
1271 * @brief Quantized dot product Q8_0 x Q8_0 (AVX + SSE)
1272 */
1273void vec_dot_q8_0_q8_0_avx(int n, float *s, const void *vx, const void *vy)
1274{
1275 const int qk = QK8_0;
1276 const int nb = n / qk;
1277
1278 const block_q8_0 *x = (const block_q8_0 *)vx;
1279 const block_q8_0 *y = (const block_q8_0 *)vy;
1280
1281 int ib = 0;
1282 __m256 accum = _mm256_setzero_ps();
1283
1284 for (; ib + 1 < nb; ib += 2) {
1285 const __m128i qx_1_0 = _mm_loadu_si128((const __m128i *)x[ib].qs);
1286 const __m128i qx_1_1 = _mm_loadu_si128((const __m128i *)x[ib].qs + 1);
1287 const __m128i qx_2_0 = _mm_loadu_si128((const __m128i *)x[ib + 1].qs);
1288 const __m128i qx_2_1 = _mm_loadu_si128((const __m128i *)x[ib + 1].qs + 1);
1289 const __m128i qy_1_0 = _mm_loadu_si128((const __m128i *)y[ib].qs);
1290 const __m128i qy_1_1 = _mm_loadu_si128((const __m128i *)y[ib].qs + 1);
1291 const __m128i qy_2_0 = _mm_loadu_si128((const __m128i *)y[ib + 1].qs);
1292 const __m128i qy_2_1 = _mm_loadu_si128((const __m128i *)y[ib + 1].qs + 1);
1293
1294 const __m256 p = mul_sum_i8_quad_float_q8_0_avx(
1295 qx_1_0, qx_1_1, qx_2_0, qx_2_1,
1296 qy_1_0, qy_1_1, qy_2_0, qy_2_1
1297 );
1298 const __m256 deltas = quad_fp16_delta_float_q8_0_avx(
1299 x[ib].d, y[ib].d, x[ib + 1].d, y[ib + 1].d
1300 );
1301 accum = _mm256_add_ps(_mm256_mul_ps(deltas, p), accum);
1302 }
1303
1304 float sumf = hsum_float_8_q8_0(accum);
1305 for (; ib < nb; ++ib) {
1306 int sumi = 0;
1307 for (int j = 0; j < qk; ++j) {
1308 sumi += x[ib].qs[j] * y[ib].qs[j];
1309 }
1310 sumf += (float)sumi * (CK_FP16_TO_FP32(x[ib].d) * CK_FP16_TO_FP32(y[ib].d));
1311 }
1312 *s = sumf;
1313}
1314#endif
1315
1316#if defined(__SSE4_1__) && !defined(__AVX__)
1317/**
1318 * @brief Quantized dot product Q8_0 x Q8_0 (SSE4.1)
1319 */
1320void vec_dot_q8_0_q8_0_sse(int n, float *s, const void *vx, const void *vy)
1321{
1322 const int qk = QK8_0;
1323 const int nb = n / qk;
1324
1325 const block_q8_0 *x = (const block_q8_0 *)vx;
1326 const block_q8_0 *y = (const block_q8_0 *)vy;
1327
1328 float sumf = 0.0f;
1329
1330 for (int ib = 0; ib < nb; ib++) {
1331 const float d = CK_FP16_TO_FP32(x[ib].d) * CK_FP16_TO_FP32(y[ib].d);
1332
1333 __m128i acc_lo = _mm_setzero_si128();
1334 __m128i acc_hi = _mm_setzero_si128();
1335
1336 /* Process 32 elements in 4 groups of 8 */
1337 for (int j = 0; j < 32; j += 8) {
1338 /* Load 8 int8 values from each */
1339 __m128i x8 = _mm_loadl_epi64((const __m128i *)&x[ib].qs[j]);
1340 __m128i y8 = _mm_loadl_epi64((const __m128i *)&y[ib].qs[j]);
1341
1342 /* Sign-extend to 16-bit */
1343 __m128i x16 = _mm_cvtepi8_epi16(x8);
1344 __m128i y16 = _mm_cvtepi8_epi16(y8);
1345
1346 /* Multiply and add horizontally: (a0*b0 + a1*b1, a2*b2 + a3*b3, ...) */
1347 __m128i prod = _mm_madd_epi16(x16, y16);
1348
1349 /* Accumulate */
1350 acc_lo = _mm_add_epi32(acc_lo, prod);
1351 }
1352
1353 /* Horizontal sum */
1354 acc_lo = _mm_add_epi32(acc_lo, _mm_shuffle_epi32(acc_lo, _MM_SHUFFLE(1, 0, 3, 2)));
1355 acc_lo = _mm_add_epi32(acc_lo, _mm_shuffle_epi32(acc_lo, _MM_SHUFFLE(0, 1, 0, 1)));
1356 int sumi = _mm_extract_epi32(acc_lo, 0);
1357
1358 sumf += d * (float)sumi;
1359 }
1360
1361 *s = sumf;
1362}
1363#endif
1364
1365/**
1366 * @brief Auto-dispatch quantized dot product Q8_0 x Q8_0
1367 */
1368void vec_dot_q8_0_q8_0(int n, float *s, const void *vx, const void *vy)
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}
1388
1389/* ============================================================================
1390 * Quantized GEMV: y = W @ x where W is Q8_0 and x is Q8_0
1391 *
1392 * This is the quantized equivalent of gemv_q8_0, but takes pre-quantized
1393 * input in Q8_0 format. Used for parity testing with llama.cpp.
1394 * ============================================================================ */
1395
1396/**
1397 * @brief Matrix-vector multiply with Q8_0 weights and Q8_0 input
1398 *
1399 * @param y Output vector [M]
1400 * @param W Weight matrix in Q8_0 format [M x K]
1401 * @param x_q8 Input vector in Q8_0 format [K]
1402 * @param M Number of output rows
1403 * @param K Number of columns (must be multiple of 32)
1404 */
1405void gemv_q8_0_q8_0(float *y,
1406 const void *W,
1407 const void *x_q8,
1408 int M, int K)
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}
1420
1421/*
1422 * Four-output Q8_0 dot provider.
1423 *
1424 * Each output keeps the same eight-lane accumulator and horizontal reduction
1425 * as vec_dot_q8_0_q8_0(), but four independent rows share the activation load
1426 * and expose enough independent VNNI chains to hide dot-product latency.
1427 */
1428void gemv_q8_0_q8_0_x4(float *y,
1429 const void *W,
1430 const void *x_q8,
1431 int M, int K)
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}
1493
1494/*
1495 * Two-token by four-output Q8_0 provider.
1496 *
1497 * Each output retains the certified eight-lane accumulator, K traversal, FMA
1498 * order, and horizontal reduction used by gemv_q8_0_q8_0_x4(). Processing two
1499 * independent token rows together reuses every weight load without coupling
1500 * their reductions.
1501 */
1503 int ldc,
1504 const void *W,
1505 const void *A_q8,
1506 int M, int N, int K)
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}
1612
1614 const void *W,
1615 const void *A_q8,
1616 int M, int N, int K)
1617{
1618 gemm_q8_0_q8_0_m2n4_strided(C, N, W, A_q8, M, N, K);
1619}
1620
1621/* ============================================================================
1622 * PARALLEL VERSIONS (for thread pool orchestration)
1623 *
1624 * These receive ith (thread index) and nth (total threads) from the
1625 * thread pool. OpenMP / pthreads live in the orchestration layer, NOT here.
1626 * ============================================================================ */
1627
1628/**
1629 * @brief Parallel reference GEMV for Q8_0 x Q8_0
1630 */
1632 const void *W,
1633 const void *x_q8,
1634 int M, int K,
1635 int ith, int nth)
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}
1656
1657/**
1658 * @brief Parallel SIMD GEMV for Q8_0 x Q8_0 with prefetching
1659 *
1660 * Each thread processes rows [r0, r1) where r0 = ith * ceil(M/nth).
1661 * Prefetches upcoming weight rows to hide memory latency.
1662 */
1664 const void *W,
1665 const void *x_q8,
1666 int M, int K,
1667 int ith, int nth)
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}
1712
1713/**
1714 * @brief Parallel SIMD GEMV for Q8_0 weights x FP32 input with prefetching
1715 */
1717 const void *W,
1718 const float *x,
1719 int M, int K,
1720 int ith, int nth)
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}
CPU feature detection and dispatch macros.
static int ck_env_truthy_or_qwen3vl_ocr_profile(const char *name)
Quantization block structures for weight-only quantization.
#define CK_FP16_TO_FP32(x)
#define CK_FP32_TO_FP16(x)
#define QK8_0
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.
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)
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 quantize_batch_q8_0(const float *x, void *vy, int num_rows, int k)
Batch quantize FP32 to Q8_0 format (row-major output)
static int ck_q8_0_q8_0_debug_ref(void)
static int ck_nearest_int_q8_0(float fval)
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 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_k(const float *x, void *vy, int k)
void gemm_q8_0_backward(float *dX, const void *W, const float *dY, int M, int N, int K)
Batched backward pass.
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)
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_x4(float *y, const void *W, const void *x_q8, int M, int K)
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_nt_q8_0(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_q8_0_q8_0_m2n4(float *C, const void *W, const void *A_q8, int M, int N, int K)
static int ck_q8_0_debug_ref(void)
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 quantize_row_q8_0(const float *x, void *vy, int k)
Quantize FP32 to Q8_0 format (scalar reference)
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 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_backward_ref(float *dX, const void *W, const float *dY, int M, int K)
Backward pass: compute input gradient (scalar reference)
float dot_q8_0(const void *w_q8_0, const float *x, int K)
void gemv_q8_0_backward(float *dX, const void *W, const float *dY, int M, int K)
Auto-dispatch backward.
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)
#define C(color)
Definition show_config.c:39
int8_t qs[32]
int32_t id
Definition tokenizer.h:316