← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q5_1_q8_1.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels_q5_1_q8_1.c
3 * @brief Q5_1 x Q8_1 contract kernels used for ggml parity (Gemma-sensitive path)
4 */
5
6#include <stdint.h>
7#include <string.h>
8#include <math.h>
9
10#include "ckernel_quant.h"
11
12#if defined(__AVX2__)
13#include <immintrin.h>
14#endif
15
16#define CK_Q51_STACK_Q8_BLOCKS 256
17
18/* Q8_1 is used only by this contract kernel path. Keep a local definition so
19 * this file does not depend on global header churn.
20 * Layout matches ggml: fp16 d, fp16 s, 32 int8 quants. */
21#ifndef QK8_1
22#define QK8_1 32
23#endif
24
25typedef struct {
26 ck_half d;
27 ck_half s;
28 int8_t qs[QK8_1];
29} block_q8_1;
30
31#if defined(__AVX2__)
32static const uint32_t ck_q51_high_nibble_lut[16] = {
33 0x00000000U, 0x00000010U, 0x00001000U, 0x00001010U,
34 0x00100000U, 0x00100010U, 0x00101000U, 0x00101010U,
35 0x10000000U, 0x10000010U, 0x10001000U, 0x10001010U,
36 0x10100000U, 0x10100010U, 0x10101000U, 0x10101010U,
37};
38
39static inline uint32_t ck_q51_high_nibble_word(uint32_t bits4) {
40 return ck_q51_high_nibble_lut[bits4 & 0x0fU];
41}
42
43static inline __m128i ck_q51_high_bits_16_avx2(uint32_t bits16) {
44 const uint32_t w0 = ck_q51_high_nibble_word(bits16);
45 const uint32_t w1 = ck_q51_high_nibble_word(bits16 >> 4);
46 const uint32_t w2 = ck_q51_high_nibble_word(bits16 >> 8);
47 const uint32_t w3 = ck_q51_high_nibble_word(bits16 >> 12);
48 return _mm_set_epi32((int)w3, (int)w2, (int)w1, (int)w0);
49}
50
51static inline int ck_q51_hsum256_epi32(__m256i v) {
52 const __m128i lo = _mm256_castsi256_si128(v);
53 const __m128i hi = _mm256_extracti128_si256(v, 1);
54 __m128i sum = _mm_add_epi32(lo, hi);
55 sum = _mm_hadd_epi32(sum, sum);
56 sum = _mm_hadd_epi32(sum, sum);
57 return _mm_cvtsi128_si32(sum);
58}
59
60#endif
61
62/* Quantize one FP32 row to Q8_1 blocks (ggml-compatible scalar path). */
63static void quantize_row_q8_1_scalar(const float *x, block_q8_1 *y, int k) {
64 const int nb = k / QK8_1;
65 for (int b = 0; b < nb; ++b) {
66 const float *xb = x + (size_t)b * QK8_1;
67 float amax = 0.0f;
68 for (int j = 0; j < QK8_1; ++j) {
69 float av = xb[j] >= 0.0f ? xb[j] : -xb[j];
70 if (av > amax) amax = av;
71 }
72
73 const float d = amax / 127.0f;
74 const float id = (d != 0.0f) ? (1.0f / d) : 0.0f;
75 y[b].d = CK_FP32_TO_FP16(d);
76
77 int sum = 0;
78 for (int j = 0; j < QK8_1; ++j) {
79 int q = (int)roundf(xb[j] * id);
80 y[b].qs[j] = (int8_t)q;
81 sum += q;
82 }
83 y[b].s = CK_FP32_TO_FP16((float)sum * d);
84 }
85}
86
87#if defined(__AVX2__)
88static inline __m256i dot_q5_1_q8_1_block_sumi_avx2(const block_q5_1 *w,
89 const block_q8_1 *x) {
90 uint32_t qh;
91 memcpy(&qh, w->qh, sizeof(qh));
92
93 const __m128i qpacked = _mm_loadu_si128((const __m128i *)(const void *)w->qs);
94 const __m128i low_mask = _mm_set1_epi8(0x0f);
95 const __m128i qlo = _mm_or_si128(_mm_and_si128(qpacked, low_mask),
96 ck_q51_high_bits_16_avx2(qh));
97 const __m128i qhi = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(qpacked, 4), low_mask),
98 ck_q51_high_bits_16_avx2(qh >> 16));
99 const __m256i q5 = _mm256_inserti128_si256(_mm256_castsi128_si256(qlo), qhi, 1);
100 const __m256i q8 = _mm256_loadu_si256((const __m256i *)(const void *)x->qs);
101#if defined(__AVXVNNI__)
102 return _mm256_dpbusd_epi32(_mm256_setzero_si256(), q5, q8);
103#else
104 const __m256i prod16 = _mm256_maddubs_epi16(q5, q8);
105 return _mm256_madd_epi16(prod16, _mm256_set1_epi16(1));
106#endif
107}
108
109/* One 32-element block dot: Q5_1(weights) x Q8_1(activations), AVX2.
110 * The high-bit placement intentionally mirrors dot_q5_1_q8_1_block().
111 */
112static float dot_q5_1_q8_1_block_avx2(const block_q5_1 *w, const block_q8_1 *x) {
113 const __m256i sumi = dot_q5_1_q8_1_block_sumi_avx2(w, x);
114
115 const float wd = CK_FP16_TO_FP32(w->d);
116 const float wm = CK_FP16_TO_FP32(w->m);
117 const float xd = CK_FP16_TO_FP32(x->d);
118 const float xs = CK_FP16_TO_FP32(x->s);
119 return (wd * xd) * (float)ck_q51_hsum256_epi32(sumi) + wm * xs;
120}
121
122static inline void dot_q5_1_q8_1_block_m4_avx2(
123 const block_q5_1 *w,
124 const block_q8_1 *x0,
125 const block_q8_1 *x1,
126 const block_q8_1 *x2,
127 const block_q8_1 *x3,
128 float out[4]) {
129 uint32_t qh;
130 memcpy(&qh, w->qh, sizeof(qh));
131
132 const __m128i qpacked = _mm_loadu_si128((const __m128i *)(const void *)w->qs);
133 const __m128i low_mask = _mm_set1_epi8(0x0f);
134 const __m128i qlo = _mm_or_si128(_mm_and_si128(qpacked, low_mask),
135 ck_q51_high_bits_16_avx2(qh));
136 const __m128i qhi = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(qpacked, 4), low_mask),
137 ck_q51_high_bits_16_avx2(qh >> 16));
138 const __m256i q5 = _mm256_inserti128_si256(_mm256_castsi128_si256(qlo), qhi, 1);
139 const block_q8_1 *rows[4] = {x0, x1, x2, x3};
140 const float wd = CK_FP16_TO_FP32(w->d);
141 const float wm = CK_FP16_TO_FP32(w->m);
142
143 for (int row = 0; row < 4; ++row) {
144 const __m256i q8 = _mm256_loadu_si256(
145 (const __m256i *)(const void *)rows[row]->qs);
146#if defined(__AVXVNNI__)
147 const __m256i sumi = _mm256_dpbusd_epi32(_mm256_setzero_si256(), q5, q8);
148#else
149 const __m256i prod16 = _mm256_maddubs_epi16(q5, q8);
150 const __m256i sumi = _mm256_madd_epi16(prod16, _mm256_set1_epi16(1));
151#endif
152 const float xd = CK_FP16_TO_FP32(rows[row]->d);
153 const float xs = CK_FP16_TO_FP32(rows[row]->s);
154 out[row] = (wd * xd) * (float)ck_q51_hsum256_epi32(sumi) + wm * xs;
155 }
156}
157
158static inline void dot_q5_1_q8_1_block_m8_avx2(
159 const block_q5_1 *w,
160 const block_q8_1 *const rows[8],
161 float out[8]) {
162 uint32_t qh;
163 memcpy(&qh, w->qh, sizeof(qh));
164
165 const __m128i qpacked = _mm_loadu_si128((const __m128i *)(const void *)w->qs);
166 const __m128i low_mask = _mm_set1_epi8(0x0f);
167 const __m128i qlo = _mm_or_si128(_mm_and_si128(qpacked, low_mask),
168 ck_q51_high_bits_16_avx2(qh));
169 const __m128i qhi = _mm_or_si128(_mm_and_si128(_mm_srli_epi16(qpacked, 4), low_mask),
170 ck_q51_high_bits_16_avx2(qh >> 16));
171 const __m256i q5 = _mm256_inserti128_si256(_mm256_castsi128_si256(qlo), qhi, 1);
172 const float wd = CK_FP16_TO_FP32(w->d);
173 const float wm = CK_FP16_TO_FP32(w->m);
174
175 for (int row = 0; row < 8; ++row) {
176 const __m256i q8 = _mm256_loadu_si256(
177 (const __m256i *)(const void *)rows[row]->qs);
178#if defined(__AVXVNNI__)
179 const __m256i sumi = _mm256_dpbusd_epi32(_mm256_setzero_si256(), q5, q8);
180#else
181 const __m256i prod16 = _mm256_maddubs_epi16(q5, q8);
182 const __m256i sumi = _mm256_madd_epi16(prod16, _mm256_set1_epi16(1));
183#endif
184 const float xd = CK_FP16_TO_FP32(rows[row]->d);
185 const float xs = CK_FP16_TO_FP32(rows[row]->s);
186 out[row] = (wd * xd) * (float)ck_q51_hsum256_epi32(sumi) + wm * xs;
187 }
188}
189
190#endif
191
192/* One 32-element block dot: Q5_1(weights) x Q8_1(activations). */
193static float dot_q5_1_q8_1_block(const block_q5_1 *w, const block_q8_1 *x) {
194#if defined(__AVX2__)
195 return dot_q5_1_q8_1_block_avx2(w, x);
196#else
197 uint32_t qh;
198 memcpy(&qh, w->qh, sizeof(qh));
199
200 int sumi0 = 0;
201 int sumi1 = 0;
202 for (int j = 0; j < QK5_1 / 2; ++j) {
203 const uint8_t xh0 = (uint8_t)(((qh >> (j + 0)) << 4) & 0x10);
204 const uint8_t xh1 = (uint8_t)(((qh >> (j + 12)) ) & 0x10);
205 const int32_t q0 = (int32_t)((w->qs[j] & 0x0F) | xh0);
206 const int32_t q1 = (int32_t)((w->qs[j] >> 4) | xh1);
207 sumi0 += q0 * (int32_t)x->qs[j];
208 sumi1 += q1 * (int32_t)x->qs[j + QK5_1 / 2];
209 }
210
211 const float wd = CK_FP16_TO_FP32(w->d);
212 const float wm = CK_FP16_TO_FP32(w->m);
213 const float xd = CK_FP16_TO_FP32(x->d);
214 const float xs = CK_FP16_TO_FP32(x->s);
215 return (wd * xd) * (float)(sumi0 + sumi1) + wm * xs;
216#endif
217}
218
219void gemv_q5_1_q8_1_ref(float *y,
220 const void *W,
221 const void *x_q8,
222 int M, int K)
223{
224 if (!y || !W || !x_q8 || M <= 0 || K <= 0 || (K % QK5_1) != 0) {
225 return;
226 }
227
228 const block_q5_1 *blocks = (const block_q5_1 *)W;
229 const block_q8_1 *x = (const block_q8_1 *)x_q8;
230 const int blocks_per_row = K / QK5_1;
231
232 for (int row = 0; row < M; ++row) {
233 const block_q5_1 *w_row = &blocks[row * blocks_per_row];
234 float sum = 0.0f;
235 for (int b = 0; b < blocks_per_row; ++b) {
236 sum += dot_q5_1_q8_1_block(&w_row[b], &x[b]);
237 }
238 y[row] = sum;
239 }
240}
241
242void gemm_nt_q5_1_q8_1_ref(const void *A_q8,
243 const void *B,
244 const float *bias,
245 float *C,
246 int M, int N, int K)
247{
248 if (!A_q8 || !B || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK5_1) != 0) {
249 return;
250 }
251
252 const block_q8_1 *A = (const block_q8_1 *)A_q8;
253 const block_q5_1 *W = (const block_q5_1 *)B;
254 const int blocks_per_row = K / QK5_1;
255
256 for (int m = 0; m < M; ++m) {
257 const block_q8_1 *a_row = &A[m * blocks_per_row];
258 for (int n = 0; n < N; ++n) {
259 const block_q5_1 *w_row = &W[n * blocks_per_row];
260 float sum = 0.0f;
261 for (int b = 0; b < blocks_per_row; ++b) {
262 sum += dot_q5_1_q8_1_block(&w_row[b], &a_row[b]);
263 }
264 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
265 }
266 }
267}
268
269void gemv_q5_1_q8_1(float *y,
270 const void *W,
271 const float *x,
272 int M,
273 int K)
274{
275 if (!y || !W || !x || M <= 0 || K <= 0 || (K % QK5_1) != 0) {
276 return;
277 }
278
279 const int blocks_per_row = K / QK5_1;
280 if (blocks_per_row > CK_Q51_STACK_Q8_BLOCKS) {
281 return;
282 }
283
284 block_q8_1 x_q8[CK_Q51_STACK_Q8_BLOCKS];
285 quantize_row_q8_1_scalar(x, x_q8, K);
286 gemv_q5_1_q8_1_ref(y, W, x_q8, M, K);
287}
288
289void gemm_nt_q5_1_q8_1(const float *A,
290 const void *B,
291 const float *bias,
292 float *C,
293 int M,
294 int N,
295 int K)
296{
297 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK5_1) != 0) {
298 return;
299 }
300
301 const int blocks_per_row = K / QK5_1;
302 if (blocks_per_row > CK_Q51_STACK_Q8_BLOCKS) {
303 return;
304 }
305
306 const block_q5_1 *W = (const block_q5_1 *)B;
307
308 for (int m = 0; m < M; ++m) {
309 block_q8_1 a_q8[CK_Q51_STACK_Q8_BLOCKS];
310 quantize_row_q8_1_scalar(&A[m * K], a_q8, K);
311 float *c_row = &C[(size_t)m * (size_t)N];
312
313 int n = 0;
314 for (; n + 7 < N; n += 8) {
315 const block_q5_1 *w0 = &W[(size_t)(n + 0) * (size_t)blocks_per_row];
316 const block_q5_1 *w1 = &W[(size_t)(n + 1) * (size_t)blocks_per_row];
317 const block_q5_1 *w2 = &W[(size_t)(n + 2) * (size_t)blocks_per_row];
318 const block_q5_1 *w3 = &W[(size_t)(n + 3) * (size_t)blocks_per_row];
319 const block_q5_1 *w4 = &W[(size_t)(n + 4) * (size_t)blocks_per_row];
320 const block_q5_1 *w5 = &W[(size_t)(n + 5) * (size_t)blocks_per_row];
321 const block_q5_1 *w6 = &W[(size_t)(n + 6) * (size_t)blocks_per_row];
322 const block_q5_1 *w7 = &W[(size_t)(n + 7) * (size_t)blocks_per_row];
323 float s0 = 0.0f;
324 float s1 = 0.0f;
325 float s2 = 0.0f;
326 float s3 = 0.0f;
327 float s4 = 0.0f;
328 float s5 = 0.0f;
329 float s6 = 0.0f;
330 float s7 = 0.0f;
331
332 for (int b = 0; b < blocks_per_row; ++b) {
333 const block_q8_1 *x = &a_q8[b];
334 s0 += dot_q5_1_q8_1_block(&w0[b], x);
335 s1 += dot_q5_1_q8_1_block(&w1[b], x);
336 s2 += dot_q5_1_q8_1_block(&w2[b], x);
337 s3 += dot_q5_1_q8_1_block(&w3[b], x);
338 s4 += dot_q5_1_q8_1_block(&w4[b], x);
339 s5 += dot_q5_1_q8_1_block(&w5[b], x);
340 s6 += dot_q5_1_q8_1_block(&w6[b], x);
341 s7 += dot_q5_1_q8_1_block(&w7[b], x);
342 }
343
344 c_row[n + 0] = s0 + (bias ? bias[n + 0] : 0.0f);
345 c_row[n + 1] = s1 + (bias ? bias[n + 1] : 0.0f);
346 c_row[n + 2] = s2 + (bias ? bias[n + 2] : 0.0f);
347 c_row[n + 3] = s3 + (bias ? bias[n + 3] : 0.0f);
348 c_row[n + 4] = s4 + (bias ? bias[n + 4] : 0.0f);
349 c_row[n + 5] = s5 + (bias ? bias[n + 5] : 0.0f);
350 c_row[n + 6] = s6 + (bias ? bias[n + 6] : 0.0f);
351 c_row[n + 7] = s7 + (bias ? bias[n + 7] : 0.0f);
352 }
353
354 for (; n < N; ++n) {
355 const block_q5_1 *w_row = &W[(size_t)n * (size_t)blocks_per_row];
356 float sum = 0.0f;
357 for (int b = 0; b < blocks_per_row; ++b) {
358 sum += dot_q5_1_q8_1_block(&w_row[b], &a_q8[b]);
359 }
360 c_row[n] = sum + (bias ? bias[n] : 0.0f);
361 }
362 }
363}
364
365void gemm_nt_q5_1_q8_1_m4(const float *A,
366 const void *B,
367 const float *bias,
368 float *C,
369 int M,
370 int N,
371 int K)
372{
373#if defined(__AVX2__)
374 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK5_1) != 0) {
375 return;
376 }
377
378 const int blocks_per_row = K / QK5_1;
379 if (blocks_per_row > CK_Q51_STACK_Q8_BLOCKS) {
380 gemm_nt_q5_1_q8_1(A, B, bias, C, M, N, K);
381 return;
382 }
383
384 const block_q5_1 *weights = (const block_q5_1 *)B;
385 int m = 0;
386 for (; m + 3 < M; m += 4) {
387 block_q8_1 activation_q8[4][CK_Q51_STACK_Q8_BLOCKS];
388 for (int row = 0; row < 4; ++row) {
390 &A[(size_t)(m + row) * (size_t)K], activation_q8[row], K);
391 }
392
393 for (int n = 0; n < N; ++n) {
394 const block_q5_1 *weight_row =
395 &weights[(size_t)n * (size_t)blocks_per_row];
396 float sums[4] = {0.0f, 0.0f, 0.0f, 0.0f};
397 for (int block = 0; block < blocks_per_row; ++block) {
398 float partial[4];
399 dot_q5_1_q8_1_block_m4_avx2(
400 &weight_row[block],
401 &activation_q8[0][block], &activation_q8[1][block],
402 &activation_q8[2][block], &activation_q8[3][block], partial);
403 sums[0] += partial[0];
404 sums[1] += partial[1];
405 sums[2] += partial[2];
406 sums[3] += partial[3];
407 }
408 const float add = bias ? bias[n] : 0.0f;
409 C[(size_t)(m + 0) * (size_t)N + n] = sums[0] + add;
410 C[(size_t)(m + 1) * (size_t)N + n] = sums[1] + add;
411 C[(size_t)(m + 2) * (size_t)N + n] = sums[2] + add;
412 C[(size_t)(m + 3) * (size_t)N + n] = sums[3] + add;
413 }
414 }
415
416 if (m < M) {
418 A + (size_t)m * (size_t)K, B, bias,
419 C + (size_t)m * (size_t)N, M - m, N, K);
420 }
421#else
422 gemm_nt_q5_1_q8_1(A, B, bias, C, M, N, K);
423#endif
424}
425
426void gemm_nt_q5_1_q8_1_m8(const float *A,
427 const void *B,
428 const float *bias,
429 float *C,
430 int M,
431 int N,
432 int K)
433{
434#if defined(__AVX2__)
435 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK5_1) != 0) {
436 return;
437 }
438
439 const int blocks_per_row = K / QK5_1;
440 if (blocks_per_row > CK_Q51_STACK_Q8_BLOCKS) {
441 gemm_nt_q5_1_q8_1_m4(A, B, bias, C, M, N, K);
442 return;
443 }
444
445 const block_q5_1 *weights = (const block_q5_1 *)B;
446 int m = 0;
447 for (; m + 7 < M; m += 8) {
448 block_q8_1 activation_q8[8][CK_Q51_STACK_Q8_BLOCKS];
449 for (int row = 0; row < 8; ++row) {
451 &A[(size_t)(m + row) * (size_t)K], activation_q8[row], K);
452 }
453
454 for (int n = 0; n < N; ++n) {
455 const block_q5_1 *weight_row =
456 &weights[(size_t)n * (size_t)blocks_per_row];
457 float sums[8] = {0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f};
458 for (int block = 0; block < blocks_per_row; ++block) {
459 const block_q8_1 *rows[8] = {
460 &activation_q8[0][block], &activation_q8[1][block],
461 &activation_q8[2][block], &activation_q8[3][block],
462 &activation_q8[4][block], &activation_q8[5][block],
463 &activation_q8[6][block], &activation_q8[7][block],
464 };
465 float partial[8];
466 dot_q5_1_q8_1_block_m8_avx2(&weight_row[block], rows, partial);
467 for (int row = 0; row < 8; ++row) {
468 sums[row] += partial[row];
469 }
470 }
471 const float add = bias ? bias[n] : 0.0f;
472 for (int row = 0; row < 8; ++row) {
473 C[(size_t)(m + row) * (size_t)N + n] = sums[row] + add;
474 }
475 }
476 }
477
478 if (m < M) {
480 A + (size_t)m * (size_t)K, B, bias,
481 C + (size_t)m * (size_t)N, M - m, N, K);
482 }
483#else
484 gemm_nt_q5_1_q8_1_m4(A, B, bias, C, M, N, K);
485#endif
486}
Quantization block structures for weight-only quantization.
uint16_t ck_half
#define CK_FP16_TO_FP32(x)
#define QK5_1
#define CK_FP32_TO_FP16(x)
void gemm_nt_q5_1_q8_1_m4(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q5_1_q8_1_ref(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_nt_q5_1_q8_1_ref(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
static void quantize_row_q8_1_scalar(const float *x, block_q8_1 *y, int k)
#define QK8_1
void gemm_nt_q5_1_q8_1(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_1_q8_1_m8(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q5_1_q8_1(float *y, const void *W, const float *x, int M, int K)
#define CK_Q51_STACK_Q8_BLOCKS
static float dot_q5_1_q8_1_block(const block_q5_1 *w, const block_q8_1 *x)
#define C(color)
Definition show_config.c:39
uint8_t qs[32/2]
uint8_t qh[4]