← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q4k_q8k_vnni.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels_q4k_q8k_vnni.c
3 * @brief VNNI Q4_K x Q8_K matvec kernel (inference only)
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 * The canonical providers require AVX2. The x8 output-interleaved provider
15 * uses 256-bit VNNI from AVX-VNNI or AVX-512 VNNI+VL. The x16 provider is a
16 * separate AVX-512 VNNI diagnostic path; production promotion is sweep-gated.
17 *
18 * Packed-meta status:
19 * Packed layouts are internal providers selected by the declared production
20 * dispatcher. Weight-identity caches own their lifetime until runtime shutdown;
21 * the kernel map records layout and ISA requirements. Canonical GGUF-layout
22 * Q4_K remains the parity fallback for unsupported ISAs and uncovered shapes.
23 */
24
25#include <stddef.h>
26#include <stdint.h>
27#include <math.h>
28#include <stdlib.h>
29#include <string.h>
30
31#include "ckernel_engine.h"
32#include "ck_threadpool.h"
33#include "ckernel_quant.h"
34#include "ck_speed_profiles.h"
35
37{
38 static int cached = -1;
39 if (cached < 0) {
40 const char *env = getenv("CK_Q4K_X16_CHUNK4");
41 cached = env ? ck_env_value_truthy(env) : 1;
42 }
43 return cached;
44}
45
46#if (defined(__AVX512VNNI__) && defined(__AVX512VL__)) || defined(__AVX2__)
47#include <immintrin.h>
48#endif
49
50#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
51#define CK_HAS_AVX_VNNI_256 1
52static inline __m256i ck_dpbusd_i32x8(__m256i acc, __m256i lhs, __m256i rhs)
53{
54 return _mm256_dpbusd_epi32(acc, lhs, rhs);
55}
56#elif defined(__AVXVNNI__)
57#define CK_HAS_AVX_VNNI_256 1
58static inline __m256i ck_dpbusd_i32x8(__m256i acc, __m256i lhs, __m256i rhs)
59{
60 return _mm256_dpbusd_avx_epi32(acc, lhs, rhs);
61}
62#endif
63
64#if defined(__AVX512VNNI__) && defined(__AVX512BW__)
65#define CK_HAS_AVX512_VNNI_512 1
66#endif
67
68void gemv_q4_k_q8_k_ref(float *y,
69 const void *W,
70 const void *x_q8,
71 int M, int K);
72
73void gemv_q4_k_q8_k_avx2(float *y,
74 const void *W,
75 const void *x_q8,
76 int M, int K);
77
78#if (defined(__AVX512VNNI__) && defined(__AVX512VL__)) || defined(__AVX2__)
79static inline int32_t hsum256_epi32(__m256i v) {
80 __m128i lo = _mm256_castsi256_si128(v);
81 __m128i hi = _mm256_extracti128_si256(v, 1);
82 __m128i sum = _mm_add_epi32(lo, hi);
83 sum = _mm_add_epi32(sum, _mm_shuffle_epi32(sum, 0x4e));
84 sum = _mm_add_epi32(sum, _mm_shuffle_epi32(sum, 0xb1));
85 return _mm_cvtsi128_si32(sum);
86}
87#endif
88
89#if defined(CK_HAS_AVX_VNNI_256)
90static inline int32_t dot_q4_k_q8_k_32_vnni(const uint8_t *q4_packed_32,
91 const int8_t *q8_32,
92 int high_nibble) {
93 const __m256i q8_bytes = _mm256_loadu_si256((const __m256i *)q8_32);
94 const __m256i packed = _mm256_loadu_si256((const __m256i *)q4_packed_32);
95 const __m256i mask4 = _mm256_set1_epi8(0x0F);
96 const __m256i q4_bytes = high_nibble
97 ? _mm256_and_si256(_mm256_srli_epi16(packed, 4), mask4)
98 : _mm256_and_si256(packed, mask4);
99 __m256i acc = _mm256_setzero_si256();
100 acc = ck_dpbusd_i32x8(acc, q4_bytes, q8_bytes);
101 return hsum256_epi32(acc);
102}
103
104static inline int32_t dot_q4_k_q8_k_32_vnni_q8v(const uint8_t *q4_packed_32,
105 __m256i q8_bytes,
106 int high_nibble) {
107 const __m256i packed = _mm256_loadu_si256((const __m256i *)q4_packed_32);
108 const __m256i mask4 = _mm256_set1_epi8(0x0F);
109 const __m256i q4_bytes = high_nibble
110 ? _mm256_and_si256(_mm256_srli_epi16(packed, 4), mask4)
111 : _mm256_and_si256(packed, mask4);
112 __m256i acc = _mm256_setzero_si256();
113 acc = ck_dpbusd_i32x8(acc, q4_bytes, q8_bytes);
114 return hsum256_epi32(acc);
115}
116
117static inline __m256i q4_k_unpack_32_vnni_bytes(__m256i packed, int high_nibble) {
118 const __m256i mask4 = _mm256_set1_epi8(0x0F);
119 return high_nibble
120 ? _mm256_and_si256(_mm256_srli_epi16(packed, 4), mask4)
121 : _mm256_and_si256(packed, mask4);
122}
123
124static inline int32_t dot_q4_k_q8_k_32_vnni_q4v_q8v(__m256i q4_bytes, __m256i q8_bytes) {
125 __m256i acc = _mm256_setzero_si256();
126 acc = ck_dpbusd_i32x8(acc, q4_bytes, q8_bytes);
127 return hsum256_epi32(acc);
128}
129
130
131static inline float dot_q4_k_q8_k_vnni_block(const block_q4_K *w,
132 const block_q8_K *x) {
133 uint8_t sc[8], m_val[8];
134 unpack_q4_k_scales(w->scales, sc, m_val);
135
136 const float d = CK_FP16_TO_FP32(w->d) * x->d;
137 const float dmin = CK_FP16_TO_FP32(w->dmin) * x->d;
138 float sumf = 0.0f;
139
140 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
141 const uint8_t *qs = &w->qs[q_offset];
142 const int8_t *q8_lo = &x->qs[j];
143 const int8_t *q8_hi = &x->qs[j + 32];
144 if (j + 128 < QK_K) {
145 __builtin_prefetch(&w->qs[q_offset + 64], 0, 1);
146 __builtin_prefetch(&x->qs[j + 128], 0, 1);
147 }
148
149 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni(qs, q8_lo, 0);
150 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni(qs, q8_hi, 1);
151 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
152 (int32_t)x->bsums[j / 16 + 1];
153 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
154 (int32_t)x->bsums[(j + 32) / 16 + 1];
155
156 sumf += d * (float)sc[is] * (float)sum_lo;
157 sumf -= dmin * (float)m_val[is] * (float)bsum_lo;
158 sumf += d * (float)sc[is + 1] * (float)sum_hi;
159 sumf -= dmin * (float)m_val[is + 1] * (float)bsum_hi;
160 }
161
162 return sumf;
163}
164#endif
165
166#if defined(__AVX2__)
167static inline __m256i q4_k_unpack_32_avx2_bytes(__m256i packed, int high_nibble) {
168 const __m256i mask4 = _mm256_set1_epi8(0x0F);
169 return high_nibble
170 ? _mm256_and_si256(_mm256_srli_epi16(packed, 4), mask4)
171 : _mm256_and_si256(packed, mask4);
172}
173
174static inline __m256i dot_q4_k_q8_k_32_avx2_i32x8(__m256i q4_bytes,
175 __m256i q8_bytes) {
176 const __m256i pair_sums_i16 = _mm256_maddubs_epi16(q4_bytes, q8_bytes);
177 const __m256i ones = _mm256_set1_epi16(1);
178 return _mm256_madd_epi16(pair_sums_i16, ones);
179}
180
181static inline int32_t dot_q4_k_q8_k_32_avx2_q4v_q8v(__m256i q4_bytes,
182 __m256i q8_bytes) {
183 const __m256i sums_i32 = dot_q4_k_q8_k_32_avx2_i32x8(q4_bytes, q8_bytes);
184 return hsum256_epi32(sums_i32);
185}
186
187static inline __m256i dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(__m256i q4_bytes,
188 __m256i q8_bytes,
189 uint8_t scale) {
190 const __m256i pair_sums_i16 = _mm256_maddubs_epi16(q4_bytes, q8_bytes);
191 const __m256i scale_i16 = _mm256_set1_epi16((int16_t)scale);
192 return _mm256_madd_epi16(pair_sums_i16, scale_i16);
193}
194
195#endif
196
197
198typedef struct {
199 ck_half d;
200 ck_half dmin;
201 uint8_t sc[8];
202 uint8_t m[8];
203 uint8_t qs[QK_K];
204} block_q4_K_packed_u8;
205
206typedef struct {
207 ck_half d;
208 ck_half dmin;
209 uint8_t sc[8];
210 uint8_t m[8];
211 uint8_t qs[QK_K / 2];
212} block_q4_K_packed_meta;
213
214typedef struct {
215 ck_half d[8];
216 ck_half dmin[8];
217 uint8_t sc[8][8];
218 uint8_t m[8][8];
219 uint8_t qs[8][QK_K / 2];
220 uint8_t active;
221 uint8_t reserved[7];
222} block_q4_K_packed_meta_x8;
223
224typedef struct {
225 ck_half d[16];
226 ck_half dmin[16];
227 uint8_t sc[16][8];
228 uint8_t m[16][8];
229 uint8_t qs[16][QK_K / 2];
230 uint8_t active;
231 uint8_t reserved[15];
232} block_q4_K_packed_meta_x16;
233
234typedef struct {
235 ck_half d[16];
236 ck_half dmin[16];
237 uint8_t sc[16][8];
238 uint8_t m[16][8];
239 uint8_t qs[16][QK_K];
240 uint8_t active;
241 uint8_t reserved[15];
242} block_q4_K_packed_u8_x16;
243
244typedef struct {
245 ck_half d[8];
246 ck_half dmin[8];
247 uint8_t sc[8][8];
248 uint8_t m[8][8];
249 uint8_t qs[(QK_K / 2) * 8];
250 uint8_t active;
251 uint8_t reserved[7];
252} block_q4_K_packed_vnni_x8;
253
254_Static_assert(sizeof(block_q4_K_packed_vnni_x8) == 1192,
255 "Q4_K VNNI x8 layout must match kernel-map prepared_bytes");
256
257typedef struct {
258 ck_half d[16];
259 ck_half dmin[16];
260 uint8_t sc[8][16];
261 uint8_t m[8][16];
262 uint8_t qs[(QK_K / 2) * 16];
263 uint8_t active;
264 uint8_t reserved[15];
265} block_q4_K_packed_vnni_x16;
266
268{
269 return sizeof(block_q4_K_packed_vnni_x8);
270}
271
273{
274#if defined(CK_HAS_AVX_VNNI_256)
275 return 1;
276#else
277 return 0;
278#endif
279}
280
282{
283 return sizeof(block_q4_K_packed_vnni_x16);
284}
285
287{
288#if defined(CK_HAS_AVX512_VNNI_512)
289 return 1;
290#else
291 return 0;
292#endif
293}
294
295void pack_q4_k_to_packed_vnni_x8(const void *src, void *dst, int N, int K)
296{
297 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
298 return;
299 }
300 const block_q4_K *in = (const block_q4_K *)src;
301 block_q4_K_packed_vnni_x8 *out = (block_q4_K_packed_vnni_x8 *)dst;
302 const int blocks_per_row = K / QK_K;
303 const int groups = (N + 7) / 8;
304 memset(out, 0,
305 (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
306
307 for (int group = 0; group < groups; ++group) {
308 const int n0 = group * 8;
309 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
310 for (int block = 0; block < blocks_per_row; ++block) {
311 block_q4_K_packed_vnni_x8 *packed =
312 out + (size_t)group * (size_t)blocks_per_row +
313 (size_t)block;
314 packed->active = (uint8_t)active;
315 for (int lane = 0; lane < active; ++lane) {
316 const block_q4_K *source =
317 in + (size_t)(n0 + lane) * (size_t)blocks_per_row +
318 (size_t)block;
319 uint8_t scales[8];
320 uint8_t mins[8];
321 packed->d[lane] = source->d;
322 packed->dmin[lane] = source->dmin;
323 unpack_q4_k_scales(source->scales, scales, mins);
324 for (int subblock = 0; subblock < 8; ++subblock) {
325 packed->sc[subblock][lane] = scales[subblock];
326 packed->m[subblock][lane] = mins[subblock];
327 }
328 for (int pair = 0; pair < QK_K / 64; ++pair) {
329 for (int segment = 0; segment < 8; ++segment) {
330 memcpy(packed->qs + (size_t)pair * 256u +
331 (size_t)segment * 32u +
332 (size_t)lane * 4u,
333 source->qs + (size_t)pair * 32u +
334 (size_t)segment * 4u,
335 4u);
336 }
337 }
338 }
339 }
340 }
341}
342
343void pack_q4_k_to_packed_vnni_x16(const void *src, void *dst, int N, int K)
344{
345 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
346 return;
347 }
348 const block_q4_K *in = (const block_q4_K *)src;
349 block_q4_K_packed_vnni_x16 *out = (block_q4_K_packed_vnni_x16 *)dst;
350 const int blocks_per_row = K / QK_K;
351 const int groups = (N + 15) / 16;
352 memset(out, 0,
353 (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
354
355 for (int group = 0; group < groups; ++group) {
356 const int n0 = group * 16;
357 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
358 for (int block = 0; block < blocks_per_row; ++block) {
359 block_q4_K_packed_vnni_x16 *packed =
360 out + (size_t)group * (size_t)blocks_per_row +
361 (size_t)block;
362 packed->active = (uint8_t)active;
363 for (int lane = 0; lane < active; ++lane) {
364 const block_q4_K *source =
365 in + (size_t)(n0 + lane) * (size_t)blocks_per_row +
366 (size_t)block;
367 uint8_t scales[8];
368 uint8_t mins[8];
369 packed->d[lane] = source->d;
370 packed->dmin[lane] = source->dmin;
371 unpack_q4_k_scales(source->scales, scales, mins);
372 for (int subblock = 0; subblock < 8; ++subblock) {
373 packed->sc[subblock][lane] = scales[subblock];
374 packed->m[subblock][lane] = mins[subblock];
375 }
376 for (int pair = 0; pair < QK_K / 64; ++pair) {
377 for (int segment = 0; segment < 8; ++segment) {
378 memcpy(packed->qs + (size_t)pair * 512u +
379 (size_t)segment * 64u +
380 (size_t)lane * 4u,
381 source->qs + (size_t)pair * 32u +
382 (size_t)segment * 4u,
383 4u);
384 }
385 }
386 }
387 }
388 }
389}
390
391#if defined(__AVX2__)
392static inline float hsum256_ps_q4k_packed(__m256 v)
393{
394 __m128 sum = _mm_add_ps(_mm256_castps256_ps128(v), _mm256_extractf128_ps(v, 1));
395 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
396 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
397 return _mm_cvtss_f32(sum);
398}
399
400/* Preserve llama.cpp's Q4_K x Q8_K AVX2 arithmetic while reading the x16
401 * packed weight layout. Packing changes addresses only; it must not collapse
402 * the eight FP32 accumulation lanes or interleave minimum terms into them. */
403static inline float dot_q4_k_packed_meta_x16_q8_k_llama_avx2(
404 const block_q4_K_packed_meta_x16 *weights,
405 int blocks_per_row,
406 int lane,
407 const block_q8_K *activation)
408{
409 __m256 acc = _mm256_setzero_ps();
410 __m128 acc_min = _mm_setzero_ps();
411 const __m256i nibble_mask = _mm256_set1_epi8(0x0F);
412
413 for (int block = 0; block < blocks_per_row; ++block) {
414 const block_q4_K_packed_meta_x16 *w = &weights[block];
415 const block_q8_K *x = &activation[block];
416 const float d = CK_FP16_TO_FP32(w->d[lane]) * x->d;
417 const float dmin = -CK_FP16_TO_FP32(w->dmin[lane]) * x->d;
418
419 const __m128i mins = _mm_loadl_epi64((const __m128i *)w->m[lane]);
420 const __m128i q8_sums = _mm_hadd_epi16(
421 _mm_loadu_si128((const __m128i *)&x->bsums[0]),
422 _mm_loadu_si128((const __m128i *)&x->bsums[8]));
423 const __m128i min_products = _mm_madd_epi16(_mm_cvtepu8_epi16(mins), q8_sums);
424 acc_min = _mm_fmadd_ps(
425 _mm_set1_ps(dmin), _mm_cvtepi32_ps(min_products), acc_min);
426
427 __m256i block_sum = _mm256_setzero_si256();
428 for (int group = 0; group < QK_K / 64; ++group) {
429 const __m256i packed = _mm256_loadu_si256(
430 (const __m256i *)&w->qs[lane][group * 32]);
431 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
432 const __m256i q4_hi = _mm256_and_si256(
433 _mm256_srli_epi16(packed, 4), nibble_mask);
434 const __m256i q8_lo = _mm256_loadu_si256(
435 (const __m256i *)&x->qs[group * 64]);
436 const __m256i q8_hi = _mm256_loadu_si256(
437 (const __m256i *)&x->qs[group * 64 + 32]);
438 __m256i lo = _mm256_maddubs_epi16(q4_lo, q8_lo);
439 __m256i hi = _mm256_maddubs_epi16(q4_hi, q8_hi);
440 lo = _mm256_madd_epi16(_mm256_set1_epi16(w->sc[lane][2 * group]), lo);
441 hi = _mm256_madd_epi16(_mm256_set1_epi16(w->sc[lane][2 * group + 1]), hi);
442 block_sum = _mm256_add_epi32(block_sum, _mm256_add_epi32(lo, hi));
443 }
444 acc = _mm256_fmadd_ps(
445 _mm256_set1_ps(d), _mm256_cvtepi32_ps(block_sum), acc);
446 }
447
448 acc_min = _mm_add_ps(acc_min, _mm_movehl_ps(acc_min, acc_min));
449 acc_min = _mm_add_ss(acc_min, _mm_movehdup_ps(acc_min));
450 return hsum256_ps_q4k_packed(acc) + _mm_cvtss_f32(acc_min);
451}
452#endif
453
455{
456 return sizeof(block_q4_K_packed_u8);
457}
458
460{
461 return sizeof(block_q4_K_packed_meta);
462}
463
465{
466 return sizeof(block_q4_K_packed_meta_x8);
467}
468
470{
471 return sizeof(block_q4_K_packed_meta_x16);
472}
473
475{
476 return sizeof(block_q4_K_packed_u8_x16);
477}
478
479void pack_q4_k_to_packed_u8(const void *src, void *dst, int N, int K)
480{
481 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
482 return;
483 }
484 const block_q4_K *in = (const block_q4_K *)src;
485 block_q4_K_packed_u8 *out = (block_q4_K_packed_u8 *)dst;
486 const int blocks_per_row = K / QK_K;
487 for (int n = 0; n < N; ++n) {
488 for (int b = 0; b < blocks_per_row; ++b) {
489 const block_q4_K *sb = in + (size_t)n * (size_t)blocks_per_row + (size_t)b;
490 block_q4_K_packed_u8 *pb = out + (size_t)n * (size_t)blocks_per_row + (size_t)b;
491 pb->d = sb->d;
492 pb->dmin = sb->dmin;
493 unpack_q4_k_scales(sb->scales, pb->sc, pb->m);
494 for (int j = 0, q_offset = 0; j < QK_K; j += 64, q_offset += 32) {
495 const uint8_t *qs = &sb->qs[q_offset];
496 for (int l = 0; l < 32; ++l) {
497 pb->qs[j + l] = qs[l] & 0x0F;
498 pb->qs[j + 32 + l] = qs[l] >> 4;
499 }
500 }
501 }
502 }
503}
504
505void pack_q4_k_to_packed_meta(const void *src, void *dst, int N, int K)
506{
507 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
508 return;
509 }
510 const block_q4_K *in = (const block_q4_K *)src;
511 block_q4_K_packed_meta *out = (block_q4_K_packed_meta *)dst;
512 const int blocks_per_row = K / QK_K;
513 for (int n = 0; n < N; ++n) {
514 for (int b = 0; b < blocks_per_row; ++b) {
515 const block_q4_K *sb = in + (size_t)n * (size_t)blocks_per_row + (size_t)b;
516 block_q4_K_packed_meta *pb = out + (size_t)n * (size_t)blocks_per_row + (size_t)b;
517 pb->d = sb->d;
518 pb->dmin = sb->dmin;
519 unpack_q4_k_scales(sb->scales, pb->sc, pb->m);
520 memcpy(pb->qs, sb->qs, sizeof(pb->qs));
521 }
522 }
523}
524
525void pack_q4_k_to_packed_meta_x8(const void *src, void *dst, int N, int K)
526{
527 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
528 return;
529 }
530 const block_q4_K *in = (const block_q4_K *)src;
531 block_q4_K_packed_meta_x8 *out = (block_q4_K_packed_meta_x8 *)dst;
532 const int blocks_per_row = K / QK_K;
533 const int groups = (N + 7) / 8;
534 memset(out, 0, (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
535
536 for (int g = 0; g < groups; ++g) {
537 const int n0 = g * 8;
538 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
539 for (int b = 0; b < blocks_per_row; ++b) {
540 block_q4_K_packed_meta_x8 *pb = out + (size_t)g * (size_t)blocks_per_row + (size_t)b;
541 pb->active = (uint8_t)active;
542 for (int lane = 0; lane < active; ++lane) {
543 const block_q4_K *sb = in + (size_t)(n0 + lane) * (size_t)blocks_per_row + (size_t)b;
544 pb->d[lane] = sb->d;
545 pb->dmin[lane] = sb->dmin;
546 unpack_q4_k_scales(sb->scales, pb->sc[lane], pb->m[lane]);
547 memcpy(pb->qs[lane], sb->qs, sizeof(pb->qs[lane]));
548 }
549 }
550 }
551}
552
553void pack_q4_k_to_packed_meta_x16(const void *src, void *dst, int N, int K)
554{
555 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
556 return;
557 }
558 const block_q4_K *in = (const block_q4_K *)src;
559 block_q4_K_packed_meta_x16 *out = (block_q4_K_packed_meta_x16 *)dst;
560 const int blocks_per_row = K / QK_K;
561 const int groups = (N + 15) / 16;
562 memset(out, 0, (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
563
564 for (int g = 0; g < groups; ++g) {
565 const int n0 = g * 16;
566 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
567 for (int b = 0; b < blocks_per_row; ++b) {
568 block_q4_K_packed_meta_x16 *pb = out + (size_t)g * (size_t)blocks_per_row + (size_t)b;
569 pb->active = (uint8_t)active;
570 for (int lane = 0; lane < active; ++lane) {
571 const block_q4_K *sb = in + (size_t)(n0 + lane) * (size_t)blocks_per_row + (size_t)b;
572 pb->d[lane] = sb->d;
573 pb->dmin[lane] = sb->dmin;
574 unpack_q4_k_scales(sb->scales, pb->sc[lane], pb->m[lane]);
575 memcpy(pb->qs[lane], sb->qs, sizeof(pb->qs[lane]));
576 }
577 }
578 }
579}
580
581void pack_q4_k_to_packed_u8_x16(const void *src, void *dst, int N, int K)
582{
583 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) {
584 return;
585 }
586 const block_q4_K *in = (const block_q4_K *)src;
587 block_q4_K_packed_u8_x16 *out = (block_q4_K_packed_u8_x16 *)dst;
588 const int blocks_per_row = K / QK_K;
589 const int groups = (N + 15) / 16;
590 memset(out, 0, (size_t)groups * (size_t)blocks_per_row * sizeof(*out));
591
592 for (int g = 0; g < groups; ++g) {
593 const int n0 = g * 16;
594 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
595 for (int b = 0; b < blocks_per_row; ++b) {
596 block_q4_K_packed_u8_x16 *pb = out + (size_t)g * (size_t)blocks_per_row + (size_t)b;
597 pb->active = (uint8_t)active;
598 for (int lane = 0; lane < active; ++lane) {
599 const block_q4_K *sb = in + (size_t)(n0 + lane) * (size_t)blocks_per_row + (size_t)b;
600 pb->d[lane] = sb->d;
601 pb->dmin[lane] = sb->dmin;
602 unpack_q4_k_scales(sb->scales, pb->sc[lane], pb->m[lane]);
603 for (int j = 0, q_offset = 0; j < QK_K; j += 64, q_offset += 32) {
604 const uint8_t *qs = &sb->qs[q_offset];
605 for (int l = 0; l < 32; ++l) {
606 pb->qs[lane][j + l] = qs[l] & 0x0F;
607 pb->qs[lane][j + 32 + l] = qs[l] >> 4;
608 }
609 }
610 }
611 }
612 }
613}
614
615#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
616static inline int32_t dot_q4_packed_u8_q8_32_vnni(const uint8_t *q4_32,
617 const int8_t *q8_32)
618{
619 const __m256i q4 = _mm256_loadu_si256((const __m256i *)q4_32);
620 const __m256i q8 = _mm256_loadu_si256((const __m256i *)q8_32);
621 __m256i acc = _mm256_setzero_si256();
622 acc = _mm256_dpbusd_epi32(acc, q4, q8);
623 return hsum256_epi32(acc);
624}
625
626static inline int32_t dot_q4_packed_u8_q8_32_vnni_q8v(const uint8_t *q4_32,
627 __m256i q8)
628{
629 const __m256i q4 = _mm256_loadu_si256((const __m256i *)q4_32);
630 __m256i acc = _mm256_setzero_si256();
631 acc = _mm256_dpbusd_epi32(acc, q4, q8);
632 return hsum256_epi32(acc);
633}
634#endif
635
636#if !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
637static inline int32_t dot_q4_packed_u8_q8_32_ref(const uint8_t *q4_32,
638 const int8_t *q8_32)
639{
640 int32_t acc = 0;
641 for (int i = 0; i < 32; ++i) {
642 acc += (int32_t)q4_32[i] * (int32_t)q8_32[i];
643 }
644 return acc;
645}
646#endif
647
648static inline float dot_q4_k_packed_u8_q8_k_block(const block_q4_K_packed_u8 *w,
649 const block_q8_K *x)
650{
651 const float d = CK_FP16_TO_FP32(w->d) * x->d;
652 const float dmin = CK_FP16_TO_FP32(w->dmin) * x->d;
653 float sumf = 0.0f;
654 for (int j = 0; j < QK_K; j += 32) {
655 const int is = j / 32;
656#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
657 const int32_t sumi = dot_q4_packed_u8_q8_32_vnni(&w->qs[j], &x->qs[j]);
658#else
659 const int32_t sumi = dot_q4_packed_u8_q8_32_ref(&w->qs[j], &x->qs[j]);
660#endif
661 const int32_t bsum = (int32_t)x->bsums[j / 16] + (int32_t)x->bsums[j / 16 + 1];
662 sumf += d * (float)w->sc[is] * (float)sumi;
663 sumf -= dmin * (float)w->m[is] * (float)bsum;
664 }
665 return sumf;
666}
667
668static inline float dot_q4_k_packed_meta_q8_k_block(const block_q4_K_packed_meta *w,
669 const block_q8_K *x)
670{
671 const float d = CK_FP16_TO_FP32(w->d) * x->d;
672 const float dmin = CK_FP16_TO_FP32(w->dmin) * x->d;
673 float sumf = 0.0f;
674 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
675 const uint8_t *qs = &w->qs[q_offset];
676 const int8_t *q8_lo = &x->qs[j];
677 const int8_t *q8_hi = &x->qs[j + 32];
678
679#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
680 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni(qs, q8_lo, 0);
681 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni(qs, q8_hi, 1);
682#else
683 int32_t sum_lo = 0;
684 int32_t sum_hi = 0;
685 for (int l = 0; l < 32; ++l) {
686 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo[l];
687 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi[l];
688 }
689#endif
690 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
691 (int32_t)x->bsums[j / 16 + 1];
692 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
693 (int32_t)x->bsums[(j + 32) / 16 + 1];
694
695 sumf += d * (float)w->sc[is] * (float)sum_lo;
696 sumf -= dmin * (float)w->m[is] * (float)bsum_lo;
697 sumf += d * (float)w->sc[is + 1] * (float)sum_hi;
698 sumf -= dmin * (float)w->m[is + 1] * (float)bsum_hi;
699 }
700 return sumf;
701}
702
703static inline void accum_q4_k_packed_meta_x8_q8_k_block(float acc[8],
704 const block_q4_K_packed_meta_x8 *w,
705 int active,
706 const block_q8_K *x)
707{
708 const float xd = x->d;
709 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
710 const int8_t *q8_lo_ptr = &x->qs[j];
711 const int8_t *q8_hi_ptr = &x->qs[j + 32];
712 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
713 (int32_t)x->bsums[j / 16 + 1];
714 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
715 (int32_t)x->bsums[(j + 32) / 16 + 1];
716
717#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
718 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
719 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
720#elif defined(__AVX2__)
721 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
722 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
723#endif
724
725 for (int lane = 0; lane < active; ++lane) {
726 const uint8_t *qs = &w->qs[lane][q_offset];
727#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
728 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
729 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
730#elif defined(__AVX2__)
731 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
732 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
733 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
734 const int32_t sum_lo = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo);
735 const int32_t sum_hi = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi);
736#else
737 int32_t sum_lo = 0;
738 int32_t sum_hi = 0;
739 for (int l = 0; l < 32; ++l) {
740 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
741 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
742 }
743#endif
744 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
745 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
746 acc[lane] += d * (float)w->sc[lane][is] * (float)sum_lo;
747 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum_lo;
748 acc[lane] += d * (float)w->sc[lane][is + 1] * (float)sum_hi;
749 acc[lane] -= dmin * (float)w->m[lane][is + 1] * (float)bsum_hi;
750 }
751 }
752}
753
754/* Match the loaded-model Q4_Kx8 contract: each pair of 32-element subblocks
755 * is reduced in integer arithmetic, followed by one FP32 value FMA and one
756 * minimum FMA. The two FP32 accumulators remain separate until the end. */
758 float acc[8], float acc_min[8],
759 const block_q4_K_packed_meta_x8 *w, int active,
760 const block_q8_K *x)
761{
762 const float xd = x->d;
763#if defined(__AVX2__)
764 float d[8] = {0};
765 float dmin[8] = {0};
766 for (int lane = 0; lane < active; ++lane) {
767 d[lane] = CK_FP16_TO_FP32(w->d[lane]);
768 dmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
769 }
770 const __m256 scale = _mm256_mul_ps(_mm256_loadu_ps(d), _mm256_set1_ps(xd));
771 const __m256 min_scale = _mm256_mul_ps(_mm256_loadu_ps(dmin), _mm256_set1_ps(xd));
772#endif
773
774 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
775 int32_t iacc[8] = {0};
776 int32_t iacc_min[8] = {0};
777 const int8_t *q8_lo_ptr = &x->qs[j];
778 const int8_t *q8_hi_ptr = &x->qs[j + 32];
779 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
780 (int32_t)x->bsums[j / 16 + 1];
781 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
782 (int32_t)x->bsums[(j + 32) / 16 + 1];
783
784#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
785 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
786 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
787#elif defined(__AVX2__)
788 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
789 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
790#endif
791 for (int lane = 0; lane < active; ++lane) {
792 const uint8_t *qs = &w->qs[lane][q_offset];
793#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
794 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
795 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
796#elif defined(__AVX2__)
797 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
798 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
799 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
800 const int32_t sum_lo = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo);
801 const int32_t sum_hi = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi);
802#else
803 int32_t sum_lo = 0;
804 int32_t sum_hi = 0;
805 for (int l = 0; l < 32; ++l) {
806 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
807 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
808 }
809#endif
810 iacc[lane] = (int32_t)w->sc[lane][is] * sum_lo +
811 (int32_t)w->sc[lane][is + 1] * sum_hi;
812 iacc_min[lane] = (int32_t)w->m[lane][is] * bsum_lo +
813 (int32_t)w->m[lane][is + 1] * bsum_hi;
814 }
815
816#if defined(__AVX2__)
817 const __m256 acc_vec = _mm256_fmadd_ps(
818 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc)),
819 scale,
820 _mm256_loadu_ps(acc));
821 const __m256 acc_min_vec = _mm256_fmadd_ps(
822 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc_min)),
823 min_scale,
824 _mm256_loadu_ps(acc_min));
825 _mm256_storeu_ps(acc, acc_vec);
826 _mm256_storeu_ps(acc_min, acc_min_vec);
827#else
828 for (int lane = 0; lane < active; ++lane) {
829 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
830 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
831 acc[lane] = fmaf((float)iacc[lane], d, acc[lane]);
832 acc_min[lane] = fmaf((float)iacc_min[lane], dmin, acc_min[lane]);
833 }
834#endif
835 }
836}
837
838/* Multi-row form of the exact repacked-matmul contract. Q4 unpacking is shared
839 * across rows, while every output retains the same ascending subblock order,
840 * separate value/minimum accumulators, and final subtraction as the scalar-row
841 * provider. Callers currently use four or eight rows to measure the reuse versus
842 * register-pressure tradeoff without changing the numerical contract. */
844 float acc[8][8], float acc_min[8][8],
845 const block_q4_K_packed_meta_x8 *w, int active,
846 const block_q8_K *x[8], int rows)
847{
848#if defined(__AVX2__)
849 float wd[8] = {0};
850 float wdmin[8] = {0};
851 for (int lane = 0; lane < active; ++lane) {
852 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
853 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
854 }
855 const __m256 weight_scale = _mm256_loadu_ps(wd);
856 const __m256 weight_min_scale = _mm256_loadu_ps(wdmin);
857
858 for (int j = 0, is = 0, q_offset = 0;
859 j < QK_K; j += 64, is += 2, q_offset += 32) {
860 int32_t iacc[8][8] = {{0}};
861 int32_t iacc_min[8][8] = {{0}};
862 __m256i q8_lo[8];
863 __m256i q8_hi[8];
864 int32_t bsum_lo[8];
865 int32_t bsum_hi[8];
866
867 for (int row = 0; row < rows; ++row) {
868 q8_lo[row] = _mm256_loadu_si256((const __m256i *)&x[row]->qs[j]);
869 q8_hi[row] = _mm256_loadu_si256((const __m256i *)&x[row]->qs[j + 32]);
870 bsum_lo[row] = (int32_t)x[row]->bsums[j / 16] +
871 (int32_t)x[row]->bsums[j / 16 + 1];
872 bsum_hi[row] = (int32_t)x[row]->bsums[(j + 32) / 16] +
873 (int32_t)x[row]->bsums[(j + 32) / 16 + 1];
874 }
875
876 for (int lane = 0; lane < active; ++lane) {
877 const __m256i packed = _mm256_loadu_si256(
878 (const __m256i *)&w->qs[lane][q_offset]);
879#if defined(CK_HAS_AVX_VNNI_256)
880 const __m256i q4_lo = q4_k_unpack_32_vnni_bytes(packed, 0);
881 const __m256i q4_hi = q4_k_unpack_32_vnni_bytes(packed, 1);
882#else
883 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
884 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
885#endif
886 const int32_t scale_lo = (int32_t)w->sc[lane][is];
887 const int32_t scale_hi = (int32_t)w->sc[lane][is + 1];
888 const int32_t min_lo = (int32_t)w->m[lane][is];
889 const int32_t min_hi = (int32_t)w->m[lane][is + 1];
890
891 for (int row = 0; row < rows; ++row) {
892#if defined(CK_HAS_AVX_VNNI_256)
893 const int32_t sum_lo =
894 dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo, q8_lo[row]);
895 const int32_t sum_hi =
896 dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi, q8_hi[row]);
897#else
898 const int32_t sum_lo =
899 dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo[row]);
900 const int32_t sum_hi =
901 dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi[row]);
902#endif
903 iacc[row][lane] = scale_lo * sum_lo + scale_hi * sum_hi;
904 iacc_min[row][lane] =
905 min_lo * bsum_lo[row] + min_hi * bsum_hi[row];
906 }
907 }
908
909 for (int row = 0; row < rows; ++row) {
910 const __m256 row_scale = _mm256_set1_ps(x[row]->d);
911 const __m256 value = _mm256_fmadd_ps(
912 _mm256_cvtepi32_ps(
913 _mm256_loadu_si256((const __m256i *)iacc[row])),
914 _mm256_mul_ps(weight_scale, row_scale),
915 _mm256_loadu_ps(acc[row]));
916 const __m256 minimum = _mm256_fmadd_ps(
917 _mm256_cvtepi32_ps(
918 _mm256_loadu_si256((const __m256i *)iacc_min[row])),
919 _mm256_mul_ps(weight_min_scale, row_scale),
920 _mm256_loadu_ps(acc_min[row]));
921 _mm256_storeu_ps(acc[row], value);
922 _mm256_storeu_ps(acc_min[row], minimum);
923 }
924 }
925#else
926 for (int row = 0; row < rows; ++row) {
928 acc[row], acc_min[row], w, active, x[row]);
929 }
930#endif
931}
932
933/* VNNI-native 4M x 8N microkernel. Q4 bytes are interleaved by output lane,
934 * allowing each vpdpbusd lane to accumulate one output column. The standard
935 * Q8_K row remains unchanged; four activation bytes are broadcast to all
936 * output lanes. Float updates preserve the accepted pairwise split-min order. */
938 float acc[4][8], float acc_min[4][8],
939 const block_q4_K_packed_vnni_x8 *w,
940 const block_q8_K *x[4], int rows)
941{
942#if defined(CK_HAS_AVX_VNNI_256)
943 float wd[8];
944 float wdmin[8];
945 for (int lane = 0; lane < 8; ++lane) {
946 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
947 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
948 }
949 const __m256 weight_scale = _mm256_loadu_ps(wd);
950 const __m256 weight_min_scale = _mm256_loadu_ps(wdmin);
951 const __m256i nibble_mask = _mm256_set1_epi8(0x0f);
952
953 for (int pair = 0; pair < QK_K / 64; ++pair) {
954 const int j = pair * 64;
955 const int is = pair * 2;
956 __m256i sum_lo[4];
957 __m256i sum_hi[4];
958 for (int row = 0; row < rows; ++row) {
959 sum_lo[row] = _mm256_setzero_si256();
960 sum_hi[row] = _mm256_setzero_si256();
961 }
962
963 for (int segment = 0; segment < 8; ++segment) {
964 const __m256i packed = _mm256_loadu_si256(
965 (const __m256i *)(w->qs + (size_t)pair * 256u +
966 (size_t)segment * 32u));
967 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
968 const __m256i q4_hi = _mm256_and_si256(
969 _mm256_srli_epi16(packed, 4), nibble_mask);
970 for (int row = 0; row < rows; ++row) {
971 int32_t q8_lo_word;
972 int32_t q8_hi_word;
973 memcpy(&q8_lo_word, x[row]->qs + j + segment * 4,
974 sizeof(q8_lo_word));
975 memcpy(&q8_hi_word, x[row]->qs + j + 32 + segment * 4,
976 sizeof(q8_hi_word));
977 sum_lo[row] = ck_dpbusd_i32x8(
978 sum_lo[row], q4_lo, _mm256_set1_epi32(q8_lo_word));
979 sum_hi[row] = ck_dpbusd_i32x8(
980 sum_hi[row], q4_hi, _mm256_set1_epi32(q8_hi_word));
981 }
982 }
983
984 const __m256i scale_lo = _mm256_cvtepu8_epi32(
985 _mm_loadl_epi64((const __m128i *)w->sc[is]));
986 const __m256i scale_hi = _mm256_cvtepu8_epi32(
987 _mm_loadl_epi64((const __m128i *)w->sc[is + 1]));
988 const __m256i min_lo = _mm256_cvtepu8_epi32(
989 _mm_loadl_epi64((const __m128i *)w->m[is]));
990 const __m256i min_hi = _mm256_cvtepu8_epi32(
991 _mm_loadl_epi64((const __m128i *)w->m[is + 1]));
992
993 for (int row = 0; row < rows; ++row) {
994 const __m256i weighted = _mm256_add_epi32(
995 _mm256_mullo_epi32(sum_lo[row], scale_lo),
996 _mm256_mullo_epi32(sum_hi[row], scale_hi));
997 const int32_t bsum_lo =
998 (int32_t)x[row]->bsums[j / 16] +
999 (int32_t)x[row]->bsums[j / 16 + 1];
1000 const int32_t bsum_hi =
1001 (int32_t)x[row]->bsums[(j + 32) / 16] +
1002 (int32_t)x[row]->bsums[(j + 32) / 16 + 1];
1003 const __m256i weighted_min = _mm256_add_epi32(
1004 _mm256_mullo_epi32(min_lo, _mm256_set1_epi32(bsum_lo)),
1005 _mm256_mullo_epi32(min_hi, _mm256_set1_epi32(bsum_hi)));
1006 const __m256 row_scale = _mm256_set1_ps(x[row]->d);
1007 const __m256 value = _mm256_fmadd_ps(
1008 _mm256_cvtepi32_ps(weighted),
1009 _mm256_mul_ps(weight_scale, row_scale),
1010 _mm256_loadu_ps(acc[row]));
1011 const __m256 minimum = _mm256_fmadd_ps(
1012 _mm256_cvtepi32_ps(weighted_min),
1013 _mm256_mul_ps(weight_min_scale, row_scale),
1014 _mm256_loadu_ps(acc_min[row]));
1015 _mm256_storeu_ps(acc[row], value);
1016 _mm256_storeu_ps(acc_min[row], minimum);
1017 }
1018 }
1019#else
1020 (void)acc;
1021 (void)acc_min;
1022 (void)w;
1023 (void)x;
1024 (void)rows;
1025#endif
1026}
1027
1028/* Output-lane x8 packing with the compact MoE provider's FP32 update order.
1029 * Integer dots are shared across eight output columns, but every lane applies
1030 * low, low-min, high, high-min terms in the same sequence as
1031 * dot_q4_k_q8_k_vnni_block_rows4(). */
1033 float block_sums[4][8],
1034 const block_q4_K_packed_vnni_x8 *w,
1035 const block_q8_K *x[4],
1036 int rows)
1037{
1038#if defined(CK_HAS_AVX_VNNI_256)
1039 float wd[8];
1040 float wdmin[8];
1041 for (int lane = 0; lane < 8; ++lane) {
1042 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
1043 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1044 }
1045 const __m256 weight_scale = _mm256_loadu_ps(wd);
1046 const __m256 weight_min_scale = _mm256_loadu_ps(wdmin);
1047 const __m256i nibble_mask = _mm256_set1_epi8(0x0f);
1048
1049 for (int pair = 0; pair < QK_K / 64; ++pair) {
1050 const int j = pair * 64;
1051 const int is = pair * 2;
1052 __m256i sum_lo[4];
1053 __m256i sum_hi[4];
1054 for (int row = 0; row < rows; ++row) {
1055 sum_lo[row] = _mm256_setzero_si256();
1056 sum_hi[row] = _mm256_setzero_si256();
1057 }
1058
1059 for (int segment = 0; segment < 8; ++segment) {
1060 const __m256i packed = _mm256_loadu_si256(
1061 (const __m256i *)(w->qs + (size_t)pair * 256u +
1062 (size_t)segment * 32u));
1063 const __m256i q4_lo = _mm256_and_si256(packed, nibble_mask);
1064 const __m256i q4_hi = _mm256_and_si256(
1065 _mm256_srli_epi16(packed, 4), nibble_mask);
1066 for (int row = 0; row < rows; ++row) {
1067 int32_t q8_lo_word;
1068 int32_t q8_hi_word;
1069 memcpy(&q8_lo_word, x[row]->qs + j + segment * 4,
1070 sizeof(q8_lo_word));
1071 memcpy(&q8_hi_word, x[row]->qs + j + 32 + segment * 4,
1072 sizeof(q8_hi_word));
1073 sum_lo[row] = ck_dpbusd_i32x8(
1074 sum_lo[row], q4_lo, _mm256_set1_epi32(q8_lo_word));
1075 sum_hi[row] = ck_dpbusd_i32x8(
1076 sum_hi[row], q4_hi, _mm256_set1_epi32(q8_hi_word));
1077 }
1078 }
1079
1080 const __m256 scale_lo = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1081 _mm_loadl_epi64((const __m128i *)w->sc[is])));
1082 const __m256 scale_hi = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1083 _mm_loadl_epi64((const __m128i *)w->sc[is + 1])));
1084 const __m256 min_lo = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1085 _mm_loadl_epi64((const __m128i *)w->m[is])));
1086 const __m256 min_hi = _mm256_cvtepi32_ps(_mm256_cvtepu8_epi32(
1087 _mm_loadl_epi64((const __m128i *)w->m[is + 1])));
1088
1089 for (int row = 0; row < rows; ++row) {
1090 const __m256 row_scale = _mm256_set1_ps(x[row]->d);
1091 const __m256 d = _mm256_mul_ps(weight_scale, row_scale);
1092 const __m256 dmin = _mm256_mul_ps(weight_min_scale, row_scale);
1093 __m256 value = _mm256_loadu_ps(block_sums[row]);
1094 value = _mm256_fmadd_ps(
1095 _mm256_mul_ps(d, scale_lo),
1096 _mm256_cvtepi32_ps(sum_lo[row]), value);
1097 const int32_t bsum_lo =
1098 (int32_t)x[row]->bsums[j / 16] +
1099 (int32_t)x[row]->bsums[j / 16 + 1];
1100 value = _mm256_fnmadd_ps(
1101 _mm256_mul_ps(dmin, min_lo),
1102 _mm256_set1_ps((float)bsum_lo), value);
1103 value = _mm256_fmadd_ps(
1104 _mm256_mul_ps(d, scale_hi),
1105 _mm256_cvtepi32_ps(sum_hi[row]), value);
1106 const int32_t bsum_hi =
1107 (int32_t)x[row]->bsums[(j + 32) / 16] +
1108 (int32_t)x[row]->bsums[(j + 32) / 16 + 1];
1109 value = _mm256_fnmadd_ps(
1110 _mm256_mul_ps(dmin, min_hi),
1111 _mm256_set1_ps((float)bsum_hi), value);
1112 _mm256_storeu_ps(block_sums[row], value);
1113 }
1114 }
1115#else
1116 (void)block_sums;
1117 (void)w;
1118 (void)x;
1119 (void)rows;
1120#endif
1121}
1122
1124 float *output,
1125 const void *weights_packed,
1126 const void *input_q8,
1127 int rows,
1128 int output_dim,
1129 int input_dim)
1130{
1131 if (!output || !weights_packed || !input_q8 || rows <= 0 || rows > 4 ||
1132 output_dim <= 0 || input_dim <= 0 || (input_dim % QK_K) != 0) {
1133 return;
1134 }
1135#if defined(CK_HAS_AVX_VNNI_256)
1136 const block_q8_K *input = (const block_q8_K *)input_q8;
1137 const block_q4_K_packed_vnni_x8 *weights =
1138 (const block_q4_K_packed_vnni_x8 *)weights_packed;
1139 const int blocks_per_row = input_dim / QK_K;
1140 const int groups = (output_dim + 7) / 8;
1141 for (int group = 0; group < groups; ++group) {
1142 const int n0 = group * 8;
1143 const int active = n0 + 8 <= output_dim ? 8 : output_dim - n0;
1144 float acc[4][8] = {{0}};
1145 for (int block = 0; block < blocks_per_row; ++block) {
1146 float block_sums[4][8] = {{0}};
1147 const block_q8_K *input_rows[4] = {NULL, NULL, NULL, NULL};
1148 for (int row = 0; row < rows; ++row) {
1149 input_rows[row] = input +
1150 (size_t)row * (size_t)blocks_per_row + (size_t)block;
1151 }
1153 block_sums,
1154 weights + (size_t)group * (size_t)blocks_per_row +
1155 (size_t)block,
1156 input_rows,
1157 rows);
1158 for (int row = 0; row < rows; ++row) {
1159 const __m256 prior = _mm256_loadu_ps(acc[row]);
1160 const __m256 current = _mm256_loadu_ps(block_sums[row]);
1161 _mm256_storeu_ps(acc[row], _mm256_add_ps(prior, current));
1162 }
1163 }
1164 for (int row = 0; row < rows; ++row) {
1165 for (int lane = 0; lane < active; ++lane) {
1166 output[(size_t)row * (size_t)output_dim +
1167 (size_t)n0 + (size_t)lane] = acc[row][lane];
1168 }
1169 }
1170 }
1171#else
1172 (void)rows;
1173 (void)output_dim;
1174 (void)input_dim;
1175#endif
1176}
1177
1179{
1180#if defined(CK_HAS_AVX_VNNI_256) && defined(__AVX512F__) && \
1181 defined(__AVX512VNNI__) && defined(__AVX512VL__)
1182 return 1;
1183#else
1184 return 0;
1185#endif
1186}
1187
1188/*
1189 * AVX-512 VNNI 16N microkernel. A scheduling tile may contain up to 16 token
1190 * rows; rows are evaluated in groups of eight so the integer dot accumulators
1191 * remain register-resident instead of spilling the entire 16M x 16N tile.
1192 * Each output lane retains the accepted pairwise split-min FP32 update order.
1193 */
1195 float acc[16][16], float acc_min[16][16],
1196 const block_q4_K_packed_vnni_x16 *w,
1197 const block_q8_K *x[16], int rows)
1198{
1199#if defined(CK_HAS_AVX512_VNNI_512)
1200 float wd[16];
1201 float wdmin[16];
1202 for (int lane = 0; lane < 16; ++lane) {
1203 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
1204 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1205 }
1206 const __m512 weight_scale = _mm512_loadu_ps(wd);
1207 const __m512 weight_min_scale = _mm512_loadu_ps(wdmin);
1208 const __m512i nibble_mask = _mm512_set1_epi8(0x0f);
1209
1210 for (int pair = 0; pair < QK_K / 64; ++pair) {
1211 const int j = pair * 64;
1212 const int is = pair * 2;
1213 const __m512i scale_lo = _mm512_cvtepu8_epi32(
1214 _mm_loadu_si128((const __m128i *)w->sc[is]));
1215 const __m512i scale_hi = _mm512_cvtepu8_epi32(
1216 _mm_loadu_si128((const __m128i *)w->sc[is + 1]));
1217 const __m512i min_lo = _mm512_cvtepu8_epi32(
1218 _mm_loadu_si128((const __m128i *)w->m[is]));
1219 const __m512i min_hi = _mm512_cvtepu8_epi32(
1220 _mm_loadu_si128((const __m128i *)w->m[is + 1]));
1221
1222 for (int row_base = 0; row_base < rows; row_base += 8) {
1223 const int row_count =
1224 row_base + 8 <= rows ? 8 : (rows - row_base);
1225 __m512i sum_lo[8];
1226 __m512i sum_hi[8];
1227 for (int row = 0; row < row_count; ++row) {
1228 sum_lo[row] = _mm512_setzero_si512();
1229 sum_hi[row] = _mm512_setzero_si512();
1230 }
1231
1232 for (int segment = 0; segment < 8; ++segment) {
1233 const __m512i packed = _mm512_loadu_si512(
1234 (const void *)(w->qs + (size_t)pair * 512u +
1235 (size_t)segment * 64u));
1236 const __m512i q4_lo =
1237 _mm512_and_si512(packed, nibble_mask);
1238 const __m512i q4_hi = _mm512_and_si512(
1239 _mm512_srli_epi16(packed, 4), nibble_mask);
1240 for (int row = 0; row < row_count; ++row) {
1241 int32_t q8_lo_word;
1242 int32_t q8_hi_word;
1243 const block_q8_K *activation = x[row_base + row];
1244 memcpy(&q8_lo_word,
1245 activation->qs + j + segment * 4,
1246 sizeof(q8_lo_word));
1247 memcpy(&q8_hi_word,
1248 activation->qs + j + 32 + segment * 4,
1249 sizeof(q8_hi_word));
1250 sum_lo[row] = _mm512_dpbusd_epi32(
1251 sum_lo[row], q4_lo,
1252 _mm512_set1_epi32(q8_lo_word));
1253 sum_hi[row] = _mm512_dpbusd_epi32(
1254 sum_hi[row], q4_hi,
1255 _mm512_set1_epi32(q8_hi_word));
1256 }
1257 }
1258
1259 for (int row = 0; row < row_count; ++row) {
1260 const block_q8_K *activation = x[row_base + row];
1261 const __m512i weighted = _mm512_add_epi32(
1262 _mm512_mullo_epi32(sum_lo[row], scale_lo),
1263 _mm512_mullo_epi32(sum_hi[row], scale_hi));
1264 const int32_t bsum_lo =
1265 (int32_t)activation->bsums[j / 16] +
1266 (int32_t)activation->bsums[j / 16 + 1];
1267 const int32_t bsum_hi =
1268 (int32_t)activation->bsums[(j + 32) / 16] +
1269 (int32_t)activation->bsums[(j + 32) / 16 + 1];
1270 const __m512i weighted_min = _mm512_add_epi32(
1271 _mm512_mullo_epi32(
1272 min_lo, _mm512_set1_epi32(bsum_lo)),
1273 _mm512_mullo_epi32(
1274 min_hi, _mm512_set1_epi32(bsum_hi)));
1275 const __m512 row_scale =
1276 _mm512_set1_ps(activation->d);
1277 const int output_row = row_base + row;
1278 const __m512 value = _mm512_fmadd_ps(
1279 _mm512_cvtepi32_ps(weighted),
1280 _mm512_mul_ps(weight_scale, row_scale),
1281 _mm512_loadu_ps(acc[output_row]));
1282 const __m512 minimum = _mm512_fmadd_ps(
1283 _mm512_cvtepi32_ps(weighted_min),
1284 _mm512_mul_ps(weight_min_scale, row_scale),
1285 _mm512_loadu_ps(acc_min[output_row]));
1286 _mm512_storeu_ps(acc[output_row], value);
1287 _mm512_storeu_ps(acc_min[output_row], minimum);
1288 }
1289 }
1290 }
1291#else
1292 (void)acc;
1293 (void)acc_min;
1294 (void)w;
1295 (void)x;
1296 (void)rows;
1297#endif
1298}
1299
1301 float acc[16], float acc_min[16],
1302 const block_q4_K_packed_vnni_x16 *w,
1303 const block_q8_K *x)
1304{
1305#if defined(CK_HAS_AVX512_VNNI_512)
1306 const __m512i nibble_mask = _mm512_set1_epi8(0x0f);
1307 __m512i iacc = _mm512_setzero_si512();
1308 __m512i iacc_min = _mm512_setzero_si512();
1309 for (int pair = 0; pair < QK_K / 64; ++pair) {
1310 const int j = pair * 64;
1311 const int is = pair * 2;
1312 __m512i sum_lo = _mm512_setzero_si512();
1313 __m512i sum_hi = _mm512_setzero_si512();
1314 for (int segment = 0; segment < 8; ++segment) {
1315 const __m512i packed = _mm512_loadu_si512(
1316 (const void *)(w->qs + (size_t)pair * 512u +
1317 (size_t)segment * 64u));
1318 const __m512i q4_lo =
1319 _mm512_and_si512(packed, nibble_mask);
1320 const __m512i q4_hi = _mm512_and_si512(
1321 _mm512_srli_epi16(packed, 4), nibble_mask);
1322 int32_t q8_lo_word;
1323 int32_t q8_hi_word;
1324 memcpy(&q8_lo_word, x->qs + j + segment * 4,
1325 sizeof(q8_lo_word));
1326 memcpy(&q8_hi_word, x->qs + j + 32 + segment * 4,
1327 sizeof(q8_hi_word));
1328 sum_lo = _mm512_dpbusd_epi32(
1329 sum_lo, q4_lo, _mm512_set1_epi32(q8_lo_word));
1330 sum_hi = _mm512_dpbusd_epi32(
1331 sum_hi, q4_hi, _mm512_set1_epi32(q8_hi_word));
1332 }
1333
1334 const __m512i scale_lo = _mm512_cvtepu8_epi32(
1335 _mm_loadu_si128((const __m128i *)w->sc[is]));
1336 const __m512i scale_hi = _mm512_cvtepu8_epi32(
1337 _mm_loadu_si128((const __m128i *)w->sc[is + 1]));
1338 iacc = _mm512_add_epi32(
1339 iacc,
1340 _mm512_add_epi32(
1341 _mm512_mullo_epi32(sum_lo, scale_lo),
1342 _mm512_mullo_epi32(sum_hi, scale_hi)));
1343
1344 const int32_t bsum_lo =
1345 (int32_t)x->bsums[j / 16] +
1346 (int32_t)x->bsums[j / 16 + 1];
1347 const int32_t bsum_hi =
1348 (int32_t)x->bsums[(j + 32) / 16] +
1349 (int32_t)x->bsums[(j + 32) / 16 + 1];
1350 const __m512i min_lo = _mm512_cvtepu8_epi32(
1351 _mm_loadu_si128((const __m128i *)w->m[is]));
1352 const __m512i min_hi = _mm512_cvtepu8_epi32(
1353 _mm_loadu_si128((const __m128i *)w->m[is + 1]));
1354 iacc_min = _mm512_add_epi32(
1355 iacc_min,
1356 _mm512_add_epi32(
1357 _mm512_mullo_epi32(
1358 min_lo, _mm512_set1_epi32(bsum_lo)),
1359 _mm512_mullo_epi32(
1360 min_hi, _mm512_set1_epi32(bsum_hi))));
1361 }
1362
1363 float wd[16];
1364 float wdmin[16];
1365 for (int lane = 0; lane < 16; ++lane) {
1366 wd[lane] = CK_FP16_TO_FP32(w->d[lane]);
1367 wdmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1368 }
1369 const __m512 xd = _mm512_set1_ps(x->d);
1370 _mm512_storeu_ps(
1371 acc,
1372 _mm512_fmadd_ps(
1373 _mm512_cvtepi32_ps(iacc),
1374 _mm512_mul_ps(_mm512_loadu_ps(wd), xd),
1375 _mm512_loadu_ps(acc)));
1376 _mm512_storeu_ps(
1377 acc_min,
1378 _mm512_fmadd_ps(
1379 _mm512_cvtepi32_ps(iacc_min),
1380 _mm512_mul_ps(_mm512_loadu_ps(wdmin), xd),
1381 _mm512_loadu_ps(acc_min)));
1382#else
1383 (void)acc;
1384 (void)acc_min;
1385 (void)w;
1386 (void)x;
1387#endif
1388}
1389
1390/* The repacked GEMV provider has a distinct reduction boundary from GEMM:
1391 * all four 64-element pairs in one Q4_K block are combined in int32 before
1392 * one value FMA and one minimum FMA update the FP32 accumulators. */
1394 float acc[8], float acc_min[8],
1395 const block_q4_K_packed_meta_x8 *w, int active,
1396 const block_q8_K *x)
1397{
1398 int32_t iacc[8] = {0};
1399 int32_t iacc_min[8] = {0};
1400 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1401 const int8_t *q8_lo_ptr = &x->qs[j];
1402 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1403 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1404 (int32_t)x->bsums[j / 16 + 1];
1405 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1406 (int32_t)x->bsums[(j + 32) / 16 + 1];
1407#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1408 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1409 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1410#elif defined(__AVX2__)
1411 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1412 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1413#endif
1414 for (int lane = 0; lane < active; ++lane) {
1415 const uint8_t *qs = &w->qs[lane][q_offset];
1416#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1417 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
1418 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
1419#elif defined(__AVX2__)
1420 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1421 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
1422 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
1423 const int32_t sum_lo = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_lo, q8_lo);
1424 const int32_t sum_hi = dot_q4_k_q8_k_32_avx2_q4v_q8v(q4_hi, q8_hi);
1425#else
1426 int32_t sum_lo = 0;
1427 int32_t sum_hi = 0;
1428 for (int l = 0; l < 32; ++l) {
1429 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
1430 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
1431 }
1432#endif
1433 iacc[lane] += (int32_t)w->sc[lane][is] * sum_lo +
1434 (int32_t)w->sc[lane][is + 1] * sum_hi;
1435 iacc_min[lane] += (int32_t)w->m[lane][is] * bsum_lo +
1436 (int32_t)w->m[lane][is + 1] * bsum_hi;
1437 }
1438 }
1439
1440#if defined(__AVX2__)
1441 float d[8] = {0};
1442 float dmin[8] = {0};
1443 for (int lane = 0; lane < active; ++lane) {
1444 d[lane] = CK_FP16_TO_FP32(w->d[lane]);
1445 dmin[lane] = CK_FP16_TO_FP32(w->dmin[lane]);
1446 }
1447 const __m256 xd = _mm256_set1_ps(x->d);
1448 const __m256 acc_vec = _mm256_fmadd_ps(
1449 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc)),
1450 _mm256_mul_ps(_mm256_loadu_ps(d), xd), _mm256_loadu_ps(acc));
1451 const __m256 min_vec = _mm256_fmadd_ps(
1452 _mm256_cvtepi32_ps(_mm256_loadu_si256((const __m256i *)iacc_min)),
1453 _mm256_mul_ps(_mm256_loadu_ps(dmin), xd), _mm256_loadu_ps(acc_min));
1454 _mm256_storeu_ps(acc, acc_vec);
1455 _mm256_storeu_ps(acc_min, min_vec);
1456#else
1457 for (int lane = 0; lane < active; ++lane) {
1458 const float d = CK_FP16_TO_FP32(w->d[lane]) * x->d;
1459 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * x->d;
1460 acc[lane] = fmaf((float)iacc[lane], d, acc[lane]);
1461 acc_min[lane] = fmaf((float)iacc_min[lane], dmin, acc_min[lane]);
1462 }
1463#endif
1464}
1465
1466
1467
1468
1469static inline void accum_q4_k_packed_meta_x8_q8_k_block_mreuse(float acc[8][8],
1470 const block_q4_K_packed_meta_x8 *w,
1471 int active,
1472 const block_q8_K *A,
1473 int blocks_per_vec,
1474 int block_index,
1475 int m0,
1476 int m_count)
1477{
1478#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1479 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1480 for (int lane0 = 0; lane0 < active; lane0 += 4) {
1481 const int lanes = (lane0 + 4 <= active) ? 4 : (active - lane0);
1482 __m256i q4_lo[4];
1483 __m256i q4_hi[4];
1484 float wd[4], wdmin[4], sc_lo[4], sc_hi[4], min_lo[4], min_hi[4];
1485
1486 for (int l = 0; l < lanes; ++l) {
1487 const int lane = lane0 + l;
1488 const uint8_t *qs = &w->qs[lane][q_offset];
1489 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1490 q4_lo[l] = q4_k_unpack_32_vnni_bytes(packed, 0);
1491 q4_hi[l] = q4_k_unpack_32_vnni_bytes(packed, 1);
1492 wd[l] = CK_FP16_TO_FP32(w->d[lane]);
1493 wdmin[l] = CK_FP16_TO_FP32(w->dmin[lane]);
1494 sc_lo[l] = (float)w->sc[lane][is];
1495 sc_hi[l] = (float)w->sc[lane][is + 1];
1496 min_lo[l] = (float)w->m[lane][is];
1497 min_hi[l] = (float)w->m[lane][is + 1];
1498 }
1499
1500 for (int mt = 0; mt < m_count; ++mt) {
1501 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1502 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1503 (int32_t)x->bsums[j / 16 + 1];
1504 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1505 (int32_t)x->bsums[(j + 32) / 16 + 1];
1506 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)&x->qs[j]);
1507 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)&x->qs[j + 32]);
1508 const float xd = x->d;
1509 for (int l = 0; l < lanes; ++l) {
1510 const int lane = lane0 + l;
1511 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo[l], q8_lo);
1512 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi[l], q8_hi);
1513 const float d = wd[l] * xd;
1514 const float dmin = wdmin[l] * xd;
1515 acc[mt][lane] += d * sc_lo[l] * (float)sum_lo;
1516 acc[mt][lane] -= dmin * min_lo[l] * (float)bsum_lo;
1517 acc[mt][lane] += d * sc_hi[l] * (float)sum_hi;
1518 acc[mt][lane] -= dmin * min_hi[l] * (float)bsum_hi;
1519 }
1520 }
1521 }
1522 }
1523#else
1524 for (int mt = 0; mt < m_count; ++mt) {
1525 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1526 accum_q4_k_packed_meta_x8_q8_k_block(acc[mt], w, active, x);
1527 }
1528#endif
1529}
1530
1531
1532static inline void accum_q4_k_packed_u8_x16_q8_k_block(float acc[16],
1533 const block_q4_K_packed_u8_x16 *w,
1534 int active,
1535 const block_q8_K *x)
1536{
1537 const float xd = x->d;
1538 for (int j = 0; j < QK_K; j += 32) {
1539 const int is = j / 32;
1540 const int8_t *q8_ptr = &x->qs[j];
1541 const int32_t bsum = (int32_t)x->bsums[j / 16] + (int32_t)x->bsums[j / 16 + 1];
1542#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1543 const __m256i q8 = _mm256_loadu_si256((const __m256i *)q8_ptr);
1544#endif
1545 for (int lane = 0; lane < active; ++lane) {
1546 const uint8_t *q4 = &w->qs[lane][j];
1547#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1548 const int32_t sumi = dot_q4_packed_u8_q8_32_vnni_q8v(q4, q8);
1549#else
1550 const int32_t sumi = dot_q4_packed_u8_q8_32_ref(q4, q8_ptr);
1551#endif
1552 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
1553 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
1554 acc[lane] += d * (float)w->sc[lane][is] * (float)sumi;
1555 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum;
1556 }
1557 }
1558}
1559
1560
1561static inline void accum_q4_k_packed_meta_x16_q8_k_block_mreuse(float acc[8][16],
1562 const block_q4_K_packed_meta_x16 *w,
1563 int active,
1564 const block_q8_K *A,
1565 int blocks_per_vec,
1566 int block_index,
1567 int m0,
1568 int m_count)
1569{
1570 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1571 for (int lane = 0; lane < active; ++lane) {
1572 const uint8_t *qs = &w->qs[lane][q_offset];
1573#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1574 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1575 const __m256i q4_lo = q4_k_unpack_32_vnni_bytes(packed, 0);
1576 const __m256i q4_hi = q4_k_unpack_32_vnni_bytes(packed, 1);
1577#elif defined(__AVX2__)
1578 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1579 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
1580 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
1581#endif
1582 const float wd = CK_FP16_TO_FP32(w->d[lane]);
1583 const float wdmin = CK_FP16_TO_FP32(w->dmin[lane]);
1584#if !defined(__AVX2__) || (defined(__AVX512VNNI__) && defined(__AVX512VL__))
1585 const float sc_lo = (float)w->sc[lane][is];
1586 const float sc_hi = (float)w->sc[lane][is + 1];
1587#endif
1588 const float min_lo = (float)w->m[lane][is];
1589 const float min_hi = (float)w->m[lane][is + 1];
1590
1591 for (int mt = 0; mt < m_count; ++mt) {
1592 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1593 const int8_t *q8_lo_ptr = &x->qs[j];
1594 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1595 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1596 (int32_t)x->bsums[j / 16 + 1];
1597 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1598 (int32_t)x->bsums[(j + 32) / 16 + 1];
1599#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1600 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1601 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1602 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo, q8_lo);
1603 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi, q8_hi);
1604#elif defined(__AVX2__)
1605 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1606 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1607 const __m256i sum_lo_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_lo, q8_lo, w->sc[lane][is]);
1608 const __m256i sum_hi_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_hi, q8_hi, w->sc[lane][is + 1]);
1609 const int32_t sum_scaled = hsum256_epi32(_mm256_add_epi32(sum_lo_v, sum_hi_v));
1610#else
1611 int32_t sum_lo = 0;
1612 int32_t sum_hi = 0;
1613 for (int l = 0; l < 32; ++l) {
1614 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
1615 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
1616 }
1617#endif
1618 const float xd = x->d;
1619 const float d = wd * xd;
1620 const float dmin = wdmin * xd;
1621#if defined(__AVX2__) && !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
1622 acc[mt][lane] += d * (float)sum_scaled;
1623 acc[mt][lane] -= dmin * min_lo * (float)bsum_lo;
1624 acc[mt][lane] -= dmin * min_hi * (float)bsum_hi;
1625#else
1626 acc[mt][lane] += d * sc_lo * (float)sum_lo;
1627 acc[mt][lane] -= dmin * min_lo * (float)bsum_lo;
1628 acc[mt][lane] += d * sc_hi * (float)sum_hi;
1629 acc[mt][lane] -= dmin * min_hi * (float)bsum_hi;
1630#endif
1631 }
1632 }
1633 }
1634}
1635
1637 const block_q4_K_packed_meta_x16 *w,
1638 int active,
1639 const block_q8_K *A,
1640 int blocks_per_vec,
1641 int block_index,
1642 int m0,
1643 int m_count)
1644{
1645#if (defined(__AVX512VNNI__) && defined(__AVX512VL__)) || defined(__AVX2__)
1646 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1647 for (int lane0 = 0; lane0 < active; lane0 += 4) {
1648 const int lanes = (lane0 + 4 <= active) ? 4 : (active - lane0);
1649 __m256i q4_lo[4];
1650 __m256i q4_hi[4];
1651 float wd[4], wdmin[4], sc_lo[4], sc_hi[4], min_lo[4], min_hi[4];
1652
1653 for (int l = 0; l < lanes; ++l) {
1654 const int lane = lane0 + l;
1655 const uint8_t *qs = &w->qs[lane][q_offset];
1656 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1657#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1658 q4_lo[l] = q4_k_unpack_32_vnni_bytes(packed, 0);
1659 q4_hi[l] = q4_k_unpack_32_vnni_bytes(packed, 1);
1660#else
1661 q4_lo[l] = q4_k_unpack_32_avx2_bytes(packed, 0);
1662 q4_hi[l] = q4_k_unpack_32_avx2_bytes(packed, 1);
1663#endif
1664 wd[l] = CK_FP16_TO_FP32(w->d[lane]);
1665 wdmin[l] = CK_FP16_TO_FP32(w->dmin[lane]);
1666 sc_lo[l] = (float)w->sc[lane][is];
1667 sc_hi[l] = (float)w->sc[lane][is + 1];
1668 min_lo[l] = (float)w->m[lane][is];
1669 min_hi[l] = (float)w->m[lane][is + 1];
1670 }
1671
1672 for (int mt = 0; mt < m_count; ++mt) {
1673 const block_q8_K *x = A + (size_t)(m0 + mt) * (size_t)blocks_per_vec + (size_t)block_index;
1674 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1675 (int32_t)x->bsums[j / 16 + 1];
1676 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1677 (int32_t)x->bsums[(j + 32) / 16 + 1];
1678 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)&x->qs[j]);
1679 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)&x->qs[j + 32]);
1680 const float xd = x->d;
1681 for (int l = 0; l < lanes; ++l) {
1682 const int lane = lane0 + l;
1683#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1684 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo[l], q8_lo);
1685 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi[l], q8_hi);
1686#else
1687 const __m256i sum_lo_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_lo[l], q8_lo, w->sc[lane][is]);
1688 const __m256i sum_hi_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_hi[l], q8_hi, w->sc[lane][is + 1]);
1689 const int32_t sum_scaled = hsum256_epi32(_mm256_add_epi32(sum_lo_v, sum_hi_v));
1690#endif
1691 const float d = wd[l] * xd;
1692 const float dmin = wdmin[l] * xd;
1693#if defined(__AVX2__) && !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
1694 acc[mt][lane] += d * (float)sum_scaled;
1695 acc[mt][lane] -= dmin * min_lo[l] * (float)bsum_lo;
1696 acc[mt][lane] -= dmin * min_hi[l] * (float)bsum_hi;
1697#else
1698 acc[mt][lane] += d * sc_lo[l] * (float)sum_lo;
1699 acc[mt][lane] -= dmin * min_lo[l] * (float)bsum_lo;
1700 acc[mt][lane] += d * sc_hi[l] * (float)sum_hi;
1701 acc[mt][lane] -= dmin * min_hi[l] * (float)bsum_hi;
1702#endif
1703 }
1704 }
1705 }
1706 }
1707#else
1708 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc, w, active, A, blocks_per_vec, block_index, m0, m_count);
1709#endif
1710}
1711
1712static inline void accum_q4_k_packed_meta_x16_q8_k_block(float acc[16],
1713 const block_q4_K_packed_meta_x16 *w,
1714 int active,
1715 const block_q8_K *x)
1716{
1717 const float xd = x->d;
1718 for (int j = 0, is = 0, q_offset = 0; j < QK_K; j += 64, is += 2, q_offset += 32) {
1719 const int8_t *q8_lo_ptr = &x->qs[j];
1720 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1721 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1722 (int32_t)x->bsums[j / 16 + 1];
1723 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1724 (int32_t)x->bsums[(j + 32) / 16 + 1];
1725
1726#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1727 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1728 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1729#elif defined(__AVX2__)
1730 const __m256i q8_lo = _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1731 const __m256i q8_hi = _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
1732#endif
1733
1734 for (int lane = 0; lane < active; ++lane) {
1735 const uint8_t *qs = &w->qs[lane][q_offset];
1736#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1737 const int32_t sum_lo = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
1738 const int32_t sum_hi = dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
1739#elif defined(__AVX2__)
1740 const __m256i packed = _mm256_loadu_si256((const __m256i *)qs);
1741 const __m256i q4_lo = q4_k_unpack_32_avx2_bytes(packed, 0);
1742 const __m256i q4_hi = q4_k_unpack_32_avx2_bytes(packed, 1);
1743 const __m256i sum_lo_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_lo, q8_lo, w->sc[lane][is]);
1744 const __m256i sum_hi_v = dot_q4_k_q8_k_32_avx2_q4v_q8v_scaled_i32x8(q4_hi, q8_hi, w->sc[lane][is + 1]);
1745 const int32_t sum_scaled = hsum256_epi32(_mm256_add_epi32(sum_lo_v, sum_hi_v));
1746#else
1747 int32_t sum_lo = 0;
1748 int32_t sum_hi = 0;
1749 for (int l = 0; l < 32; ++l) {
1750 sum_lo += (int32_t)(qs[l] & 0x0F) * (int32_t)q8_lo_ptr[l];
1751 sum_hi += (int32_t)(qs[l] >> 4) * (int32_t)q8_hi_ptr[l];
1752 }
1753#endif
1754 const float d = CK_FP16_TO_FP32(w->d[lane]) * xd;
1755 const float dmin = CK_FP16_TO_FP32(w->dmin[lane]) * xd;
1756#if defined(__AVX2__) && !(defined(__AVX512VNNI__) && defined(__AVX512VL__))
1757 acc[lane] += d * (float)sum_scaled;
1758 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum_lo;
1759 acc[lane] -= dmin * (float)w->m[lane][is + 1] * (float)bsum_hi;
1760#else
1761 acc[lane] += d * (float)w->sc[lane][is] * (float)sum_lo;
1762 acc[lane] -= dmin * (float)w->m[lane][is] * (float)bsum_lo;
1763 acc[lane] += d * (float)w->sc[lane][is + 1] * (float)sum_hi;
1764 acc[lane] -= dmin * (float)w->m[lane][is + 1] * (float)bsum_hi;
1765#endif
1766 }
1767 }
1768}
1769
1770void gemm_nt_q4_k_packed_u8_q8_k(const void *A_q8,
1771 const void *B_packed,
1772 const float *bias,
1773 float *C,
1774 int M, int N, int K)
1775{
1776 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1777 return;
1778 }
1779 const block_q8_K *A = (const block_q8_K *)A_q8;
1780 const block_q4_K_packed_u8 *W = (const block_q4_K_packed_u8 *)B_packed;
1781 const int blocks_per_vec = K / QK_K;
1782 const int blocks_per_row = K / QK_K;
1783 for (int m = 0; m < M; ++m) {
1784 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1785 float *c_row = C + (size_t)m * (size_t)N;
1786 for (int n = 0; n < N; ++n) {
1787 const block_q4_K_packed_u8 *w_row = W + (size_t)n * (size_t)blocks_per_row;
1788 float sum = bias ? bias[n] : 0.0f;
1789 for (int b = 0; b < blocks_per_row; ++b) {
1790 sum += dot_q4_k_packed_u8_q8_k_block(&w_row[b], &a_row[b]);
1791 }
1792 c_row[n] = sum;
1793 }
1794 }
1795}
1796
1797
1798void gemm_nt_q4_k_packed_meta_q8_k(const void *A_q8,
1799 const void *B_packed,
1800 const float *bias,
1801 float *C,
1802 int M, int N, int K)
1803{
1804 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1805 return;
1806 }
1807 const block_q8_K *A = (const block_q8_K *)A_q8;
1808 const block_q4_K_packed_meta *W = (const block_q4_K_packed_meta *)B_packed;
1809 const int blocks_per_vec = K / QK_K;
1810 const int blocks_per_row = K / QK_K;
1811 for (int m = 0; m < M; ++m) {
1812 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1813 float *c_row = C + (size_t)m * (size_t)N;
1814 for (int n = 0; n < N; ++n) {
1815 const block_q4_K_packed_meta *w_row = W + (size_t)n * (size_t)blocks_per_row;
1816 float sum = bias ? bias[n] : 0.0f;
1817 for (int b = 0; b < blocks_per_row; ++b) {
1818 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
1819 }
1820 c_row[n] = sum;
1821 }
1822 }
1823}
1824
1826 const void *B_packed,
1827 const float *bias,
1828 float *C,
1829 int M, int N, int K,
1830 int m0, int m1, int n0, int n1)
1831{
1832 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1833 return;
1834 }
1835 if (m0 < 0) m0 = 0;
1836 if (n0 < 0) n0 = 0;
1837 if (m1 > M) m1 = M;
1838 if (n1 > N) n1 = N;
1839 if (m0 >= m1 || n0 >= n1) {
1840 return;
1841 }
1842
1843 const block_q8_K *A = (const block_q8_K *)A_q8;
1844 const block_q4_K_packed_meta *W = (const block_q4_K_packed_meta *)B_packed;
1845 const int blocks_per_vec = K / QK_K;
1846 const int blocks_per_row = K / QK_K;
1847
1848 for (int m = m0; m < m1; ++m) {
1849 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1850 float *c_row = C + (size_t)m * (size_t)N;
1851 for (int n = n0; n < n1; ++n) {
1852 const block_q4_K_packed_meta *w_row = W + (size_t)n * (size_t)blocks_per_row;
1853 float sum = bias ? bias[n] : 0.0f;
1854 for (int b = 0; b < blocks_per_row; ++b) {
1855 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
1856 }
1857 c_row[n] = sum;
1858 }
1859 }
1860}
1861
1863 const void *B_packed_x8,
1864 const float *bias,
1865 float *C,
1866 int M, int N, int K)
1867{
1868 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1869 return;
1870 }
1871 const block_q8_K *A = (const block_q8_K *)A_q8;
1872 const block_q4_K_packed_meta_x8 *W = (const block_q4_K_packed_meta_x8 *)B_packed_x8;
1873 const int blocks_per_vec = K / QK_K;
1874 const int blocks_per_row = K / QK_K;
1875 const int groups = (N + 7) / 8;
1876
1877 for (int m = 0; m < M; ++m) {
1878 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_vec;
1879 float *c_row = C + (size_t)m * (size_t)N;
1880 for (int g = 0; g < groups; ++g) {
1881 const int n0 = g * 8;
1882 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
1883 float acc[8];
1884 for (int lane = 0; lane < active; ++lane) {
1885 acc[lane] = bias ? bias[n0 + lane] : 0.0f;
1886 }
1887 for (int b = 0; b < blocks_per_row; ++b) {
1888 const block_q4_K_packed_meta_x8 *w_group =
1889 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
1890 accum_q4_k_packed_meta_x8_q8_k_block(acc, w_group, active, &a_row[b]);
1891 }
1892 for (int lane = 0; lane < active; ++lane) {
1893 c_row[n0 + lane] = acc[lane];
1894 }
1895 }
1896 }
1897}
1898
1900 const void *A_q8, const void *B_packed_x8, const float *bias, float *C,
1901 int M, int N, int K)
1902{
1903 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
1904 return;
1905 }
1906 const block_q8_K *A = (const block_q8_K *)A_q8;
1907 const block_q4_K_packed_meta_x8 *W = (const block_q4_K_packed_meta_x8 *)B_packed_x8;
1908 const int blocks_per_row = K / QK_K;
1909 const int groups = (N + 7) / 8;
1910
1911 for (int m = 0; m < M; ++m) {
1912 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_row;
1913 float *c_row = C + (size_t)m * (size_t)N;
1914 for (int g = 0; g < groups; ++g) {
1915 const int n0 = g * 8;
1916 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
1917 float acc[8] = {0};
1918 float acc_min[8] = {0};
1919 for (int b = 0; b < blocks_per_row; ++b) {
1920 const block_q4_K_packed_meta_x8 *w_group =
1921 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
1923 acc, acc_min, w_group, active, &a_row[b]);
1924 }
1925 float values[8];
1926#if defined(__AVX2__)
1927 _mm256_storeu_ps(values, _mm256_sub_ps(_mm256_loadu_ps(acc), _mm256_loadu_ps(acc_min)));
1928#else
1929 for (int lane = 0; lane < active; ++lane) {
1930 values[lane] = acc[lane] - acc_min[lane];
1931 }
1932#endif
1933 for (int lane = 0; lane < active; ++lane) {
1934 float value = values[lane];
1935 if (bias) {
1936 value += bias[n0 + lane];
1937 }
1938 c_row[n0 + lane] = value;
1939 }
1940 }
1941 }
1942}
1943
1945 const void *A_q8, const void *B_packed_x8, const float *bias, float *C,
1946 int M, int N, int K)
1947{
1948#if defined(__AVX512F__) && defined(__AVX512BW__) && defined(__AVX512DQ__)
1949 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
1950 (K % QK_K) != 0 || (N % 16) != 0) {
1951 return;
1952 }
1953 const block_q8_K *A = (const block_q8_K *)A_q8;
1954 const block_q4_K_packed_meta_x8 *W =
1955 (const block_q4_K_packed_meta_x8 *)B_packed_x8;
1956 const int blocks_per_row = K / QK_K;
1957
1958 for (int row = 0; row < M; ++row) {
1959 const block_q8_K *a_row = A + (size_t)row * (size_t)blocks_per_row;
1960 float *c_row = C + (size_t)row * (size_t)N;
1961 for (int n0 = 0; n0 < N; n0 += 16) {
1962 const int group0 = n0 / 8;
1963 __m512 acc = _mm512_setzero_ps();
1964 __m512 acc_min = _mm512_setzero_ps();
1965
1966 for (int b = 0; b < blocks_per_row; ++b) {
1967 const block_q4_K_packed_meta_x8 *w0 =
1968 W + (size_t)group0 * (size_t)blocks_per_row + (size_t)b;
1969 const block_q4_K_packed_meta_x8 *w1 =
1970 W + (size_t)(group0 + 1) * (size_t)blocks_per_row + (size_t)b;
1971 const block_q8_K *x = &a_row[b];
1972 float d[16];
1973 float dmin[16];
1974 for (int lane = 0; lane < 16; ++lane) {
1975 const block_q4_K_packed_meta_x8 *w = lane < 8 ? w0 : w1;
1976 const int wl = lane & 7;
1977 d[lane] = CK_FP16_TO_FP32(w->d[wl]);
1978 dmin[lane] = CK_FP16_TO_FP32(w->dmin[wl]);
1979 }
1980 const __m512 scale =
1981 _mm512_mul_ps(_mm512_loadu_ps(d), _mm512_set1_ps(x->d));
1982 const __m512 min_scale =
1983 _mm512_mul_ps(_mm512_loadu_ps(dmin), _mm512_set1_ps(x->d));
1984
1985 for (int j = 0, is = 0, q_offset = 0;
1986 j < QK_K;
1987 j += 64, is += 2, q_offset += 32) {
1988 int32_t iacc[16];
1989 int32_t iacc_min[16];
1990 const int8_t *q8_lo_ptr = &x->qs[j];
1991 const int8_t *q8_hi_ptr = &x->qs[j + 32];
1992 const int32_t bsum_lo = (int32_t)x->bsums[j / 16] +
1993 (int32_t)x->bsums[j / 16 + 1];
1994 const int32_t bsum_hi = (int32_t)x->bsums[(j + 32) / 16] +
1995 (int32_t)x->bsums[(j + 32) / 16 + 1];
1996#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
1997 const __m256i q8_lo =
1998 _mm256_loadu_si256((const __m256i *)q8_lo_ptr);
1999 const __m256i q8_hi =
2000 _mm256_loadu_si256((const __m256i *)q8_hi_ptr);
2001#endif
2002 for (int lane = 0; lane < 16; ++lane) {
2003 const block_q4_K_packed_meta_x8 *w = lane < 8 ? w0 : w1;
2004 const int wl = lane & 7;
2005 const uint8_t *qs = &w->qs[wl][q_offset];
2006#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
2007 const int32_t sum_lo =
2008 dot_q4_k_q8_k_32_vnni_q8v(qs, q8_lo, 0);
2009 const int32_t sum_hi =
2010 dot_q4_k_q8_k_32_vnni_q8v(qs, q8_hi, 1);
2011#else
2012 int32_t sum_lo = 0;
2013 int32_t sum_hi = 0;
2014 for (int i = 0; i < 32; ++i) {
2015 sum_lo += (int32_t)(qs[i] & 0x0F) *
2016 (int32_t)q8_lo_ptr[i];
2017 sum_hi += (int32_t)(qs[i] >> 4) *
2018 (int32_t)q8_hi_ptr[i];
2019 }
2020#endif
2021 /* llama.cpp's AVX-512 q4_K_8x8 provider keeps each
2022 * 32-value dot in an int16 lane through PMADDUBSW and
2023 * wrapping VPADDW operations, then widens while
2024 * applying the two sub-block scales. */
2025 const int16_t packed_sum_lo = (int16_t)sum_lo;
2026 const int16_t packed_sum_hi = (int16_t)sum_hi;
2027 iacc[lane] = (int32_t)w->sc[wl][is] * (int32_t)packed_sum_lo +
2028 (int32_t)w->sc[wl][is + 1] * (int32_t)packed_sum_hi;
2029 iacc_min[lane] = (int32_t)w->m[wl][is] * bsum_lo +
2030 (int32_t)w->m[wl][is + 1] * bsum_hi;
2031 }
2032 acc = _mm512_fmadd_ps(
2033 _mm512_cvtepi32_ps(_mm512_loadu_si512(iacc)),
2034 scale,
2035 acc);
2036 acc_min = _mm512_fmadd_ps(
2037 _mm512_cvtepi32_ps(_mm512_loadu_si512(iacc_min)),
2038 min_scale,
2039 acc_min);
2040 }
2041 }
2042 __m512 value = _mm512_sub_ps(acc, acc_min);
2043 if (bias) {
2044 value = _mm512_add_ps(value, _mm512_loadu_ps(bias + n0));
2045 }
2046 _mm512_storeu_ps(c_row + n0, value);
2047 }
2048 }
2049#else
2051 A_q8, B_packed_x8, bias, C, M, N, K);
2052#endif
2053}
2054
2056 const void *A_q8, const void *B_packed_x8, const float *bias, float *C,
2057 int M, int N, int K)
2058{
2059 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2060 return;
2061 }
2062 const block_q8_K *A = (const block_q8_K *)A_q8;
2063 const block_q4_K_packed_meta_x8 *W = (const block_q4_K_packed_meta_x8 *)B_packed_x8;
2064 const int blocks_per_row = K / QK_K;
2065 const int groups = (N + 7) / 8;
2066 for (int m = 0; m < M; ++m) {
2067 const block_q8_K *a_row = A + (size_t)m * (size_t)blocks_per_row;
2068 float *c_row = C + (size_t)m * (size_t)N;
2069 for (int g = 0; g < groups; ++g) {
2070 const int n0 = g * 8;
2071 const int active = (n0 + 8 <= N) ? 8 : (N - n0);
2072 float acc[8] = {0};
2073 float acc_min[8] = {0};
2074 for (int b = 0; b < blocks_per_row; ++b) {
2075 const block_q4_K_packed_meta_x8 *w_group =
2076 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
2078 acc, acc_min, w_group, active, &a_row[b]);
2079 }
2080 float values[8];
2081#if defined(__AVX2__)
2082 _mm256_storeu_ps(values, _mm256_sub_ps(
2083 _mm256_loadu_ps(acc), _mm256_loadu_ps(acc_min)));
2084#else
2085 for (int lane = 0; lane < active; ++lane) values[lane] = acc[lane] - acc_min[lane];
2086#endif
2087 for (int lane = 0; lane < active; ++lane) {
2088 c_row[n0 + lane] = values[lane] + (bias ? bias[n0 + lane] : 0.0f);
2089 }
2090 }
2091 }
2092}
2093
2095 const void *A_q8, const void *B_packed_x16,
2096 const float *bias, float *C,
2097 int M, int N, int K)
2098{
2099 if (!A_q8 || !B_packed_x16 || !C ||
2100 M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0 ||
2102 return;
2103 }
2104 const block_q8_K *A = (const block_q8_K *)A_q8;
2105 const block_q4_K_packed_vnni_x16 *W =
2106 (const block_q4_K_packed_vnni_x16 *)B_packed_x16;
2107 const int blocks_per_row = K / QK_K;
2108 const int groups = (N + 15) / 16;
2109 for (int m = 0; m < M; ++m) {
2110 const block_q8_K *a_row =
2111 A + (size_t)m * (size_t)blocks_per_row;
2112 float *c_row = C + (size_t)m * (size_t)N;
2113 for (int g = 0; g < groups; ++g) {
2114 const int n0 = g * 16;
2115 const int active = (n0 + 16 <= N) ? 16 : (N - n0);
2116 float acc[16] = {0};
2117 float acc_min[16] = {0};
2118 for (int b = 0; b < blocks_per_row; ++b) {
2119 const block_q4_K_packed_vnni_x16 *w_group =
2120 W + (size_t)g * (size_t)blocks_per_row + (size_t)b;
2122 acc, acc_min, w_group, &a_row[b]);
2123 }
2124 float values[16];
2125#if defined(CK_HAS_AVX512_VNNI_512)
2126 _mm512_storeu_ps(values, _mm512_sub_ps(
2127 _mm512_loadu_ps(acc), _mm512_loadu_ps(acc_min)));
2128#else
2129 for (int lane = 0; lane < active; ++lane) {
2130 values[lane] = acc[lane] - acc_min[lane];
2131 }
2132#endif
2133 for (int lane = 0; lane < active; ++lane) {
2134 c_row[n0 + lane] =
2135 values[lane] + (bias ? bias[n0 + lane] : 0.0f);
2136 }
2137 }
2138 }
2139}
2140
2141
2142typedef struct {
2143 const block_q8_K *A;
2144 const block_q4_K_packed_meta *W;
2145 const float *bias;
2146 float *C;
2147 int M;
2148 int N;
2149 int K;
2150 int blocks_per_vec;
2151 int blocks_per_row;
2152} gemm_q4_packed_meta_thread_work_t;
2153
2154typedef struct {
2155 const block_q8_K *A;
2156 const block_q4_K_packed_meta_x8 *W;
2157 const float *bias;
2158 float *C;
2159 int M;
2160 int N;
2161 int K;
2162 int blocks_per_vec;
2163 int blocks_per_row;
2164 int groups;
2165 int tile_m;
2166 int jobs;
2167} gemm_q4_packed_meta_x8_thread_work_t;
2168
2169typedef struct {
2170 const block_q8_K *A;
2171 const block_q4_K_packed_vnni_x8 *W;
2172 const float *bias;
2173 float *C;
2174 int M;
2175 int N;
2176 int blocks_per_row;
2177 int groups;
2178} gemm_q4_packed_vnni_x8_thread_work_t;
2179
2180typedef struct {
2181 const block_q8_K *A;
2182 const block_q4_K_packed_vnni_x16 *W;
2183 const float *bias;
2184 float *C;
2185 int M;
2186 int N;
2187 int blocks_per_row;
2188 int groups;
2189} gemm_q4_packed_vnni_x16_thread_work_t;
2190
2191typedef struct {
2192 const block_q8_K *A;
2193 const block_q4_K_packed_meta_x16 *W;
2194 const float *bias;
2195 float *C;
2196 int M;
2197 int N;
2198 int K;
2199 int blocks_per_vec;
2200 int blocks_per_row;
2201 int groups;
2202 int tile_m;
2203 int jobs;
2204} gemm_q4_packed_meta_x16_thread_work_t;
2205
2206typedef struct {
2207 const block_q8_K *A;
2208 const block_q4_K_packed_u8_x16 *W;
2209 const float *bias;
2210 float *C;
2211 int M;
2212 int N;
2213 int K;
2214 int blocks_per_vec;
2215 int blocks_per_row;
2216 int groups;
2217 int tile_m;
2218 int jobs;
2219} gemm_q4_packed_u8_x16_thread_work_t;
2220
2221static void gemm_q4_packed_meta_thread_fn(int ith, int nth, void *args)
2222{
2223 gemm_q4_packed_meta_thread_work_t *a = (gemm_q4_packed_meta_thread_work_t *)args;
2224 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2225 return;
2226 }
2227 const int dm = (a->M + nth - 1) / nth;
2228 const int m0 = dm * ith;
2229 const int m1 = (m0 + dm < a->M) ? (m0 + dm) : a->M;
2230 if (m0 >= a->M) {
2231 return;
2232 }
2233 for (int m = m0; m < m1; ++m) {
2234 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2235 float *c_row = a->C + (size_t)m * (size_t)a->N;
2236 for (int n = 0; n < a->N; ++n) {
2237 const block_q4_K_packed_meta *w_row = a->W + (size_t)n * (size_t)a->blocks_per_row;
2238 float sum = a->bias ? a->bias[n] : 0.0f;
2239 for (int b = 0; b < a->blocks_per_row; ++b) {
2240 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
2241 }
2242 c_row[n] = sum;
2243 }
2244 }
2245}
2246
2248 const void *B_packed,
2249 const float *bias,
2250 float *C,
2251 int M, int N, int K,
2252 int active_threads)
2253{
2254 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2255 return;
2256 }
2257 ck_threadpool_t *pool = ck_threadpool_global();
2258 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2259 if (active_threads <= 0 || active_threads > pool_threads) {
2260 active_threads = pool_threads;
2261 }
2262 if (active_threads > M) {
2263 active_threads = M;
2264 }
2265 if (active_threads <= 1) {
2266 gemm_nt_q4_k_packed_meta_q8_k(A_q8, B_packed, bias, C, M, N, K);
2267 return;
2268 }
2269 gemm_q4_packed_meta_thread_work_t work = {
2270 .A = (const block_q8_K *)A_q8,
2271 .W = (const block_q4_K_packed_meta *)B_packed,
2272 .bias = bias,
2273 .C = C,
2274 .M = M,
2275 .N = N,
2276 .K = K,
2277 .blocks_per_vec = K / QK_K,
2278 .blocks_per_row = K / QK_K,
2279 };
2280 ck_threadpool_dispatch_n(pool, active_threads, gemm_q4_packed_meta_thread_fn, &work);
2281}
2282
2283
2284static void gemm_q4_packed_meta_nsplit_thread_fn(int ith, int nth, void *args)
2285{
2286 gemm_q4_packed_meta_thread_work_t *a = (gemm_q4_packed_meta_thread_work_t *)args;
2287 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2288 return;
2289 }
2290 const int dn = (a->N + nth - 1) / nth;
2291 const int n0 = dn * ith;
2292 const int n1 = (n0 + dn < a->N) ? (n0 + dn) : a->N;
2293 if (n0 >= a->N) {
2294 return;
2295 }
2296 for (int m = 0; m < a->M; ++m) {
2297 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2298 float *c_row = a->C + (size_t)m * (size_t)a->N;
2299 for (int n = n0; n < n1; ++n) {
2300 const block_q4_K_packed_meta *w_row = a->W + (size_t)n * (size_t)a->blocks_per_row;
2301 float sum = a->bias ? a->bias[n] : 0.0f;
2302 for (int b = 0; b < a->blocks_per_row; ++b) {
2303 sum += dot_q4_k_packed_meta_q8_k_block(&w_row[b], &a_row[b]);
2304 }
2305 c_row[n] = sum;
2306 }
2307 }
2308}
2309
2311 const void *B_packed,
2312 const float *bias,
2313 float *C,
2314 int M, int N, int K,
2315 int active_threads)
2316{
2317 if (!A_q8 || !B_packed || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2318 return;
2319 }
2320 ck_threadpool_t *pool = ck_threadpool_global();
2321 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2322 if (active_threads <= 0 || active_threads > pool_threads) {
2323 active_threads = pool_threads;
2324 }
2325 if (active_threads > N) {
2326 active_threads = N;
2327 }
2328 if (active_threads <= 1) {
2329 gemm_nt_q4_k_packed_meta_q8_k(A_q8, B_packed, bias, C, M, N, K);
2330 return;
2331 }
2332 gemm_q4_packed_meta_thread_work_t work = {
2333 .A = (const block_q8_K *)A_q8,
2334 .W = (const block_q4_K_packed_meta *)B_packed,
2335 .bias = bias,
2336 .C = C,
2337 .M = M,
2338 .N = N,
2339 .K = K,
2340 .blocks_per_vec = K / QK_K,
2341 .blocks_per_row = K / QK_K,
2342 };
2344}
2345
2346static void gemm_q4_packed_meta_x8_nsplit_thread_fn(int ith, int nth, void *args)
2347{
2348 gemm_q4_packed_meta_x8_thread_work_t *a = (gemm_q4_packed_meta_x8_thread_work_t *)args;
2349 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2350 return;
2351 }
2352 const int dg = (a->groups + nth - 1) / nth;
2353 const int g0 = dg * ith;
2354 const int g1 = (g0 + dg < a->groups) ? (g0 + dg) : a->groups;
2355 if (g0 >= a->groups) {
2356 return;
2357 }
2358
2359 for (int m = 0; m < a->M; ++m) {
2360 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2361 float *c_row = a->C + (size_t)m * (size_t)a->N;
2362 for (int g = g0; g < g1; ++g) {
2363 const int n0 = g * 8;
2364 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2365 float acc[8];
2366 for (int lane = 0; lane < active; ++lane) {
2367 acc[lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
2368 }
2369 for (int b = 0; b < a->blocks_per_row; ++b) {
2370 const block_q4_K_packed_meta_x8 *w_group =
2371 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2372 accum_q4_k_packed_meta_x8_q8_k_block(acc, w_group, active, &a_row[b]);
2373 }
2374 for (int lane = 0; lane < active; ++lane) {
2375 c_row[n0 + lane] = acc[lane];
2376 }
2377 }
2378 }
2379}
2380
2382 const void *B_packed_x8,
2383 const float *bias,
2384 float *C,
2385 int M, int N, int K,
2386 int active_threads)
2387{
2388 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2389 return;
2390 }
2391 ck_threadpool_t *pool = ck_threadpool_global();
2392 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2393 const int groups = (N + 7) / 8;
2394 if (active_threads <= 0 || active_threads > pool_threads) {
2395 active_threads = pool_threads;
2396 }
2397 if (active_threads > groups) {
2398 active_threads = groups;
2399 }
2400 if (active_threads <= 1) {
2401 gemm_nt_q4_k_packed_meta_x8_q8_k(A_q8, B_packed_x8, bias, C, M, N, K);
2402 return;
2403 }
2404 gemm_q4_packed_meta_x8_thread_work_t work = {
2405 .A = (const block_q8_K *)A_q8,
2406 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2407 .bias = bias,
2408 .C = C,
2409 .M = M,
2410 .N = N,
2411 .K = K,
2412 .blocks_per_vec = K / QK_K,
2413 .blocks_per_row = K / QK_K,
2414 .groups = groups,
2415 };
2417}
2418
2419static void gemm_q4_packed_meta_x8_mtile_thread_fn(int ith, int nth, void *args)
2420{
2421 gemm_q4_packed_meta_x8_thread_work_t *a = (gemm_q4_packed_meta_x8_thread_work_t *)args;
2422 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2423 return;
2424 }
2425 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
2426 if (tile_m > 8) tile_m = 8;
2427 const int mt = (a->M + tile_m - 1) / tile_m;
2428 const int total = mt * a->groups;
2429
2430 for (int job = ith; job < total; job += nth) {
2431 const int g = job / mt;
2432 const int tm = job - g * mt;
2433 const int m0 = tm * tile_m;
2434 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
2435 const int n0 = g * 8;
2436 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2437 if (m0 >= a->M || g >= a->groups) {
2438 continue;
2439 }
2440
2441 float acc[8][8];
2442 for (int mt_lane = 0; mt_lane < m1 - m0; ++mt_lane) {
2443 for (int lane = 0; lane < active; ++lane) {
2444 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
2445 }
2446 }
2447
2448 for (int b = 0; b < a->blocks_per_row; ++b) {
2449 const block_q4_K_packed_meta_x8 *w_group =
2450 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2451 for (int m = m0; m < m1; ++m) {
2452 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
2453 accum_q4_k_packed_meta_x8_q8_k_block(acc[m - m0], w_group, active, &a_row[b]);
2454 }
2455 }
2456
2457 for (int m = m0; m < m1; ++m) {
2458 float *c_row = a->C + (size_t)m * (size_t)a->N;
2459 for (int lane = 0; lane < active; ++lane) {
2460 c_row[n0 + lane] = acc[m - m0][lane];
2461 }
2462 }
2463 }
2464}
2465
2466
2467static void gemm_q4_packed_meta_x8_mreuse_thread_fn(int ith, int nth, void *args)
2468{
2469 gemm_q4_packed_meta_x8_thread_work_t *a = (gemm_q4_packed_meta_x8_thread_work_t *)args;
2470 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2471 return;
2472 }
2473 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
2474 if (tile_m > 8) tile_m = 8;
2475 const int mt = (a->M + tile_m - 1) / tile_m;
2476 const int total = mt * a->groups;
2477
2478 for (int job = ith; job < total; job += nth) {
2479 const int g = job / mt;
2480 const int tm = job - g * mt;
2481 const int m0 = tm * tile_m;
2482 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
2483 const int m_count = m1 - m0;
2484 const int n0 = g * 8;
2485 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2486 if (m0 >= a->M || g >= a->groups || m_count <= 0) {
2487 continue;
2488 }
2489
2490 float acc[8][8];
2491 for (int mt_lane = 0; mt_lane < m_count; ++mt_lane) {
2492 for (int lane = 0; lane < active; ++lane) {
2493 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
2494 }
2495 }
2496
2497 for (int b = 0; b < a->blocks_per_row; ++b) {
2498 const block_q4_K_packed_meta_x8 *w_group =
2499 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2500 accum_q4_k_packed_meta_x8_q8_k_block_mreuse(acc, w_group, active, a->A,
2501 a->blocks_per_vec, b, m0, m_count);
2502 }
2503
2504 for (int m = m0; m < m1; ++m) {
2505 float *c_row = a->C + (size_t)m * (size_t)a->N;
2506 for (int lane = 0; lane < active; ++lane) {
2507 c_row[n0 + lane] = acc[m - m0][lane];
2508 }
2509 }
2510 }
2511}
2512
2513/* Reorder independent output work only. Each output keeps ascending K-block
2514 * traversal, separate value/minimum accumulators, and one final subtraction,
2515 * which is the llama.cpp repacked-matmul numerical contract. */
2517 int ith, int nth, void *args)
2518{
2519 gemm_q4_packed_meta_x8_thread_work_t *a =
2520 (gemm_q4_packed_meta_x8_thread_work_t *)args;
2521 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2522 return;
2523 }
2524
2525 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
2526 if (tile_m > 8) tile_m = 8;
2527 const int mt = (a->M + tile_m - 1) / tile_m;
2528 const int total = mt * a->groups;
2529
2530 for (int job = ith; job < total; job += nth) {
2531 const int g = job / mt;
2532 const int tm = job - g * mt;
2533 const int m0 = tm * tile_m;
2534 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
2535 const int m_count = m1 - m0;
2536 const int n0 = g * 8;
2537 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2538 if (m_count <= 0 || g >= a->groups) {
2539 continue;
2540 }
2541
2542 float acc[8][8] = {{0}};
2543 float acc_min[8][8] = {{0}};
2544 for (int b = 0; b < a->blocks_per_row; ++b) {
2545 const block_q4_K_packed_meta_x8 *w_group =
2546 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2547 for (int m = m0; m < m1; ++m) {
2548 const block_q8_K *a_row =
2549 a->A + (size_t)m * (size_t)a->blocks_per_vec;
2551 acc[m - m0], acc_min[m - m0], w_group, active, &a_row[b]);
2552 }
2553 }
2554
2555 for (int m = m0; m < m1; ++m) {
2556 float values[8];
2557#if defined(__AVX2__)
2558 _mm256_storeu_ps(values,
2559 _mm256_sub_ps(_mm256_loadu_ps(acc[m - m0]),
2560 _mm256_loadu_ps(acc_min[m - m0])));
2561#else
2562 for (int lane = 0; lane < active; ++lane) {
2563 values[lane] = acc[m - m0][lane] - acc_min[m - m0][lane];
2564 }
2565#endif
2566 float *c_row = a->C + (size_t)m * (size_t)a->N;
2567 for (int lane = 0; lane < active; ++lane) {
2568 float value = values[lane];
2569 if (a->bias) value += a->bias[n0 + lane];
2570 c_row[n0 + lane] = value;
2571 }
2572 }
2573 }
2574}
2575
2577 int ith, int nth, void *args)
2578{
2579 gemm_q4_packed_meta_x8_thread_work_t *a =
2580 (gemm_q4_packed_meta_x8_thread_work_t *)args;
2581 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2582 return;
2583 }
2584
2585 const int row_tiles = (a->M + 3) / 4;
2586 const int total = row_tiles * a->groups;
2587 for (int job = ith; job < total; job += nth) {
2588 const int g = job / row_tiles;
2589 const int row_tile = job - g * row_tiles;
2590 const int m0 = row_tile * 4;
2591 const int rows = (m0 + 4 <= a->M) ? 4 : (a->M - m0);
2592 const int n0 = g * 8;
2593 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2594 if (rows <= 0 || active <= 0 || g >= a->groups) {
2595 continue;
2596 }
2597
2598 float acc[8][8] = {{0}};
2599 float acc_min[8][8] = {{0}};
2600 for (int b = 0; b < a->blocks_per_row; ++b) {
2601 const block_q4_K_packed_meta_x8 *w_group =
2602 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2603 const block_q8_K *x[8] = {NULL};
2604 for (int row = 0; row < rows; ++row) {
2605 x[row] = a->A + (size_t)(m0 + row) *
2606 (size_t)a->blocks_per_vec + (size_t)b;
2607 }
2609 acc, acc_min, w_group, active, x, rows);
2610 }
2611
2612 for (int row = 0; row < rows; ++row) {
2613 float values[8];
2614#if defined(__AVX2__)
2615 _mm256_storeu_ps(values,
2616 _mm256_sub_ps(_mm256_loadu_ps(acc[row]),
2617 _mm256_loadu_ps(acc_min[row])));
2618#else
2619 for (int lane = 0; lane < active; ++lane) {
2620 values[lane] = acc[row][lane] - acc_min[row][lane];
2621 }
2622#endif
2623 float *c_row = a->C + (size_t)(m0 + row) * (size_t)a->N;
2624 for (int lane = 0; lane < active; ++lane) {
2625 float value = values[lane];
2626 if (a->bias) value += a->bias[n0 + lane];
2627 c_row[n0 + lane] = value;
2628 }
2629 }
2630 }
2631}
2632
2634 int ith, int nth, void *args)
2635{
2636 gemm_q4_packed_meta_x8_thread_work_t *a =
2637 (gemm_q4_packed_meta_x8_thread_work_t *)args;
2638 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2639 return;
2640 }
2641
2642 const int row_tiles = (a->M + 7) / 8;
2643 const int total = row_tiles * a->groups;
2644 for (int job = ith; job < total; job += nth) {
2645 const int g = job / row_tiles;
2646 const int row_tile = job - g * row_tiles;
2647 const int m0 = row_tile * 8;
2648 const int rows = (m0 + 8 <= a->M) ? 8 : (a->M - m0);
2649 const int n0 = g * 8;
2650 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2651 if (rows <= 0 || active <= 0 || g >= a->groups) {
2652 continue;
2653 }
2654
2655 float acc[8][8] = {{0}};
2656 float acc_min[8][8] = {{0}};
2657 for (int b = 0; b < a->blocks_per_row; ++b) {
2658 const block_q4_K_packed_meta_x8 *w_group =
2659 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
2660 const block_q8_K *x[8] = {NULL};
2661 for (int row = 0; row < rows; ++row) {
2662 x[row] = a->A + (size_t)(m0 + row) *
2663 (size_t)a->blocks_per_vec + (size_t)b;
2664 }
2666 acc, acc_min, w_group, active, x, rows);
2667 }
2668
2669 for (int row = 0; row < rows; ++row) {
2670 float values[8];
2671#if defined(__AVX2__)
2672 _mm256_storeu_ps(values,
2673 _mm256_sub_ps(_mm256_loadu_ps(acc[row]),
2674 _mm256_loadu_ps(acc_min[row])));
2675#else
2676 for (int lane = 0; lane < active; ++lane) {
2677 values[lane] = acc[row][lane] - acc_min[row][lane];
2678 }
2679#endif
2680 float *c_row = a->C + (size_t)(m0 + row) * (size_t)a->N;
2681 for (int lane = 0; lane < active; ++lane) {
2682 float value = values[lane];
2683 if (a->bias) value += a->bias[n0 + lane];
2684 c_row[n0 + lane] = value;
2685 }
2686 }
2687 }
2688}
2689
2691 gemm_q4_packed_vnni_x8_thread_work_t *a,
2692 int job,
2693 int row_tiles)
2694{
2695 const int group = job / row_tiles;
2696 const int row_tile = job - group * row_tiles;
2697 const int m0 = row_tile * 4;
2698 const int rows = (m0 + 4 <= a->M) ? 4 : (a->M - m0);
2699 const int n0 = group * 8;
2700 const int active = (n0 + 8 <= a->N) ? 8 : (a->N - n0);
2701 if (rows <= 0 || active <= 0 || group >= a->groups) {
2702 return;
2703 }
2704
2705 float acc[4][8] = {{0}};
2706 float acc_min[4][8] = {{0}};
2707 for (int block = 0; block < a->blocks_per_row; ++block) {
2708 const block_q4_K_packed_vnni_x8 *weights =
2709 a->W + (size_t)group * (size_t)a->blocks_per_row +
2710 (size_t)block;
2711 const block_q8_K *x[4] = {NULL};
2712 for (int row = 0; row < rows; ++row) {
2713 x[row] = a->A + (size_t)(m0 + row) *
2714 (size_t)a->blocks_per_row + (size_t)block;
2715 }
2717 acc, acc_min, weights, x, rows);
2718 }
2719
2720 for (int row = 0; row < rows; ++row) {
2721 float values[8];
2722 _mm256_storeu_ps(values, _mm256_sub_ps(
2723 _mm256_loadu_ps(acc[row]),
2724 _mm256_loadu_ps(acc_min[row])));
2725 float *output = a->C + (size_t)(m0 + row) * (size_t)a->N;
2726 for (int lane = 0; lane < active; ++lane) {
2727 output[n0 + lane] = values[lane] +
2728 (a->bias ? a->bias[n0 + lane] : 0.0f);
2729 }
2730 }
2731}
2732
2734 int ith, int nth, void *args)
2735{
2736 gemm_q4_packed_vnni_x8_thread_work_t *a =
2737 (gemm_q4_packed_vnni_x8_thread_work_t *)args;
2738 if (!a || ith < 0 || nth <= 0 || ith >= nth) return;
2739
2740 const int row_tiles = (a->M + 3) / 4;
2741 const int total = row_tiles * a->groups;
2742 for (int job = ith; job < total; job += nth) {
2743 gemm_q4_packed_vnni_x8_q8k_4m_job(a, job, row_tiles);
2744 }
2745}
2746
2748 int begin, int end, void *args)
2749{
2750 gemm_q4_packed_vnni_x8_thread_work_t *a =
2751 (gemm_q4_packed_vnni_x8_thread_work_t *)args;
2752 if (!a || begin < 0 || begin >= end) return;
2753 const int row_tiles = (a->M + 3) / 4;
2754 for (int job = begin; job < end; ++job) {
2755 gemm_q4_packed_vnni_x8_q8k_4m_job(a, job, row_tiles);
2756 }
2757}
2758
2760 int ith, int nth, void *args)
2761{
2762 gemm_q4_packed_vnni_x16_thread_work_t *a =
2763 (gemm_q4_packed_vnni_x16_thread_work_t *)args;
2764 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
2765 return;
2766 }
2767
2768 const int row_tiles = (a->M + 15) / 16;
2769 const int total = row_tiles * a->groups;
2770 for (int job = ith; job < total; job += nth) {
2771 const int group = job / row_tiles;
2772 const int row_tile = job - group * row_tiles;
2773 const int m0 = row_tile * 16;
2774 const int rows = (m0 + 16 <= a->M) ? 16 : (a->M - m0);
2775 const int n0 = group * 16;
2776 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
2777 if (rows <= 0 || active <= 0 || group >= a->groups) {
2778 continue;
2779 }
2780
2781 float acc[16][16] = {{0}};
2782 float acc_min[16][16] = {{0}};
2783 for (int block = 0; block < a->blocks_per_row; ++block) {
2784 const block_q4_K_packed_vnni_x16 *weights =
2785 a->W + (size_t)group * (size_t)a->blocks_per_row +
2786 (size_t)block;
2787 const block_q8_K *x[16] = {NULL};
2788 for (int row = 0; row < rows; ++row) {
2789 x[row] = a->A + (size_t)(m0 + row) *
2790 (size_t)a->blocks_per_row + (size_t)block;
2791 }
2793 acc, acc_min, weights, x, rows);
2794 }
2795
2796 for (int row = 0; row < rows; ++row) {
2797#if defined(CK_HAS_AVX512_VNNI_512)
2798 float values[16];
2799 _mm512_storeu_ps(values, _mm512_sub_ps(
2800 _mm512_loadu_ps(acc[row]),
2801 _mm512_loadu_ps(acc_min[row])));
2802 float *output = a->C + (size_t)(m0 + row) * (size_t)a->N;
2803 for (int lane = 0; lane < active; ++lane) {
2804 output[n0 + lane] = values[lane] +
2805 (a->bias ? a->bias[n0 + lane] : 0.0f);
2806 }
2807#else
2808 (void)active;
2809#endif
2810 }
2811 }
2812}
2813
2815 const void *B_packed_x8,
2816 const float *bias,
2817 float *C,
2818 int M, int N, int K,
2819 int tile_m,
2820 int active_threads)
2821{
2822 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2823 return;
2824 }
2825 ck_threadpool_t *pool = ck_threadpool_global();
2826 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2827 const int groups = (N + 7) / 8;
2828 int tm = tile_m > 0 ? tile_m : 4;
2829 if (tm > 8) tm = 8;
2830 const int mt = (M + tm - 1) / tm;
2831 const int jobs = mt * groups;
2832 if (active_threads <= 0 || active_threads > pool_threads) {
2833 active_threads = pool_threads;
2834 }
2835 if (active_threads > jobs) {
2836 active_threads = jobs;
2837 }
2838 if (active_threads <= 1) {
2839 gemm_nt_q4_k_packed_meta_x8_q8_k(A_q8, B_packed_x8, bias, C, M, N, K);
2840 return;
2841 }
2842 gemm_q4_packed_meta_x8_thread_work_t work = {
2843 .A = (const block_q8_K *)A_q8,
2844 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2845 .bias = bias,
2846 .C = C,
2847 .M = M,
2848 .N = N,
2849 .K = K,
2850 .blocks_per_vec = K / QK_K,
2851 .blocks_per_row = K / QK_K,
2852 .groups = groups,
2853 .tile_m = tm,
2854 .jobs = jobs,
2855 };
2857}
2858
2859
2861 const void *B_packed_x8,
2862 const float *bias,
2863 float *C,
2864 int M, int N, int K,
2865 int tile_m,
2866 int active_threads)
2867{
2868 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
2869 return;
2870 }
2871 ck_threadpool_t *pool = ck_threadpool_global();
2872 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2873 const int groups = (N + 7) / 8;
2874 int tm = tile_m > 0 ? tile_m : 4;
2875 if (tm > 8) tm = 8;
2876 const int mt = (M + tm - 1) / tm;
2877 const int jobs = mt * groups;
2878 if (active_threads <= 0 || active_threads > pool_threads) {
2879 active_threads = pool_threads;
2880 }
2881 if (active_threads > jobs) {
2882 active_threads = jobs;
2883 }
2884 gemm_q4_packed_meta_x8_thread_work_t work = {
2885 .A = (const block_q8_K *)A_q8,
2886 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2887 .bias = bias,
2888 .C = C,
2889 .M = M,
2890 .N = N,
2891 .K = K,
2892 .blocks_per_vec = K / QK_K,
2893 .blocks_per_row = K / QK_K,
2894 .groups = groups,
2895 .tile_m = tm,
2896 .jobs = jobs,
2897 };
2898 if (active_threads <= 1) {
2900 return;
2901 }
2903}
2904
2906 const void *A_q8, const void *B_packed_x8, const float *bias, float *C,
2907 int M, int N, int K, int tile_m, int active_threads)
2908{
2909 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
2910 (K % QK_K) != 0) {
2911 return;
2912 }
2913 ck_threadpool_t *pool = ck_threadpool_global();
2914 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2915 const int groups = (N + 7) / 8;
2916 int tm = tile_m > 0 ? tile_m : 4;
2917 if (tm > 8) tm = 8;
2918 const int jobs = ((M + tm - 1) / tm) * groups;
2919 if (active_threads <= 0 || active_threads > pool_threads) {
2920 active_threads = pool_threads;
2921 }
2922 if (active_threads > jobs) active_threads = jobs;
2923
2924 gemm_q4_packed_meta_x8_thread_work_t work = {
2925 .A = (const block_q8_K *)A_q8,
2926 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2927 .bias = bias,
2928 .C = C,
2929 .M = M,
2930 .N = N,
2931 .K = K,
2932 .blocks_per_vec = K / QK_K,
2933 .blocks_per_row = K / QK_K,
2934 .groups = groups,
2935 .tile_m = tm,
2936 .jobs = jobs,
2937 };
2938 if (active_threads <= 1 || !pool) {
2940 return;
2941 }
2943 pool, active_threads,
2945}
2946
2948 const void *A_q8, const void *B_packed_x8, const float *bias, float *C,
2949 int M, int N, int K, int active_threads)
2950{
2951 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
2952 (K % QK_K) != 0) {
2953 return;
2954 }
2955 ck_threadpool_t *pool = ck_threadpool_global();
2956 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2957 const int groups = (N + 7) / 8;
2958 const int jobs = ((M + 3) / 4) * groups;
2959 if (active_threads <= 0 || active_threads > pool_threads) {
2960 active_threads = pool_threads;
2961 }
2962 if (active_threads > jobs) active_threads = jobs;
2963
2964 gemm_q4_packed_meta_x8_thread_work_t work = {
2965 .A = (const block_q8_K *)A_q8,
2966 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
2967 .bias = bias,
2968 .C = C,
2969 .M = M,
2970 .N = N,
2971 .K = K,
2972 .blocks_per_vec = K / QK_K,
2973 .blocks_per_row = K / QK_K,
2974 .groups = groups,
2975 .tile_m = 4,
2976 .jobs = jobs,
2977 };
2978 if (active_threads <= 1 || !pool) {
2980 return;
2981 }
2983 pool, active_threads,
2985}
2986
2988 const void *A_q8, const void *B_packed_x8, const float *bias, float *C,
2989 int M, int N, int K, int active_threads)
2990{
2991 if (!A_q8 || !B_packed_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
2992 (K % QK_K) != 0) {
2993 return;
2994 }
2995 ck_threadpool_t *pool = ck_threadpool_global();
2996 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
2997 const int groups = (N + 7) / 8;
2998 const int jobs = ((M + 7) / 8) * groups;
2999 if (active_threads <= 0 || active_threads > pool_threads) {
3000 active_threads = pool_threads;
3001 }
3002 if (active_threads > jobs) active_threads = jobs;
3003
3004 gemm_q4_packed_meta_x8_thread_work_t work = {
3005 .A = (const block_q8_K *)A_q8,
3006 .W = (const block_q4_K_packed_meta_x8 *)B_packed_x8,
3007 .bias = bias,
3008 .C = C,
3009 .M = M,
3010 .N = N,
3011 .K = K,
3012 .blocks_per_vec = K / QK_K,
3013 .blocks_per_row = K / QK_K,
3014 .groups = groups,
3015 .tile_m = 8,
3016 .jobs = jobs,
3017 };
3018 if (active_threads <= 1 || !pool) {
3020 return;
3021 }
3023 pool, active_threads,
3025}
3026
3028 const void *A_q8, const void *B_packed_vnni_x8, const float *bias,
3029 float *C, int M, int N, int K, int active_threads)
3030{
3031 if (!A_q8 || !B_packed_vnni_x8 || !C || M <= 0 || N <= 0 || K <= 0 ||
3032 (K % QK_K) != 0) {
3033 return;
3034 }
3035 ck_threadpool_t *pool = ck_threadpool_global();
3036 const int pool_threads = pool ? ck_threadpool_capacity(pool) : 1;
3037 const int groups = (N + 7) / 8;
3038 const int jobs = ((M + 3) / 4) * groups;
3039 if (active_threads <= 0 || active_threads > pool_threads) {
3040 active_threads = pool_threads;
3041 }
3042 if (active_threads > jobs) active_threads = jobs;
3043
3044 gemm_q4_packed_vnni_x8_thread_work_t work = {
3045 .A = (const block_q8_K *)A_q8,
3046 .W = (const block_q4_K_packed_vnni_x8 *)B_packed_vnni_x8,
3047 .bias = bias,
3048 .C = C,
3049 .M = M,
3050 .N = N,
3051 .blocks_per_row = K / QK_K,
3052 .groups = groups,
3053 };
3054 if (active_threads <= 1 || !pool) {
3056 return;
3057 }
3059 const int grain = 4;
3061 pool, active_threads, 0, jobs, grain,
3063 } else {
3065 pool, active_threads,
3067 }
3068}
3069
3071 const void *A_q8, const void *B_packed_vnni_x16, const float *bias,
3072 float *C, int M, int N, int K, int active_threads)
3073{
3074 if (!A_q8 || !B_packed_vnni_x16 || !C ||
3075 M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0 ||
3077 return;
3078 }
3079 ck_threadpool_t *pool = ck_threadpool_global();
3080 const int pool_threads = pool ? ck_threadpool_capacity(pool) : 1;
3081 const int groups = (N + 15) / 16;
3082 const int jobs = ((M + 15) / 16) * groups;
3083 if (active_threads <= 0 || active_threads > pool_threads) {
3084 active_threads = pool_threads;
3085 }
3086 if (active_threads > jobs) active_threads = jobs;
3087
3088 gemm_q4_packed_vnni_x16_thread_work_t work = {
3089 .A = (const block_q8_K *)A_q8,
3090 .W = (const block_q4_K_packed_vnni_x16 *)B_packed_vnni_x16,
3091 .bias = bias,
3092 .C = C,
3093 .M = M,
3094 .N = N,
3095 .blocks_per_row = K / QK_K,
3096 .groups = groups,
3097 };
3098 if (active_threads <= 1 || !pool) {
3100 return;
3101 }
3103 pool, active_threads,
3105}
3106
3107
3108static void gemm_q4_packed_meta_x16_mtile_thread_fn(int ith, int nth, void *args)
3109{
3110 gemm_q4_packed_meta_x16_thread_work_t *a = (gemm_q4_packed_meta_x16_thread_work_t *)args;
3111 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
3112 return;
3113 }
3114 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3115 if (tile_m > 8) tile_m = 8;
3116 const int mt = (a->M + tile_m - 1) / tile_m;
3117 const int total = mt * a->groups;
3118
3119 for (int job = ith; job < total; job += nth) {
3120 const int g = job / mt;
3121 const int tm = job - g * mt;
3122 const int m0 = tm * tile_m;
3123 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3124 const int n0 = g * 16;
3125 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
3126 if (m0 >= a->M || g >= a->groups) {
3127 continue;
3128 }
3129
3130 float acc[8][16];
3131 for (int mt_lane = 0; mt_lane < m1 - m0; ++mt_lane) {
3132 for (int lane = 0; lane < active; ++lane) {
3133 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
3134 }
3135 }
3136
3137 for (int b = 0; b < a->blocks_per_row; ++b) {
3138 const block_q4_K_packed_meta_x16 *w_group =
3139 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3140 for (int m = m0; m < m1; ++m) {
3141 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3142 accum_q4_k_packed_meta_x16_q8_k_block(acc[m - m0], w_group, active, &a_row[b]);
3143 }
3144 }
3145
3146 for (int m = m0; m < m1; ++m) {
3147 float *c_row = a->C + (size_t)m * (size_t)a->N;
3148 for (int lane = 0; lane < active; ++lane) {
3149 c_row[n0 + lane] = acc[m - m0][lane];
3150 }
3151 }
3152 }
3153}
3154
3155
3157 const gemm_q4_packed_meta_x16_thread_work_t *a,
3158 int job,
3159 int mt,
3160 int tile_m)
3161{
3162 const int g = job / mt;
3163 const int tm = job - g * mt;
3164 const int m0 = tm * tile_m;
3165 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3166 const int m_count = m1 - m0;
3167 const int n0 = g * 16;
3168 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
3169 if (m0 >= a->M || g >= a->groups || m_count <= 0) {
3170 return;
3171 }
3172
3173#if defined(__AVX2__)
3174 const block_q4_K_packed_meta_x16 *w_group =
3175 a->W + (size_t)g * (size_t)a->blocks_per_row;
3176 for (int m = m0; m < m1; ++m) {
3177 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3178 float *c_row = a->C + (size_t)m * (size_t)a->N;
3179 for (int lane = 0; lane < active; ++lane) {
3180 float value = dot_q4_k_packed_meta_x16_q8_k_llama_avx2(
3181 w_group, a->blocks_per_row, lane, a_row);
3182 c_row[n0 + lane] = value + (a->bias ? a->bias[n0 + lane] : 0.0f);
3183 }
3184 }
3185#else
3186 float acc[8][16];
3187 for (int mt_lane = 0; mt_lane < m_count; ++mt_lane) {
3188 for (int lane = 0; lane < active; ++lane) {
3189 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
3190 }
3191 }
3192 for (int b = 0; b < a->blocks_per_row; ++b) {
3193 const block_q4_K_packed_meta_x16 *w_block =
3194 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3196 accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(acc, w_block, active, a->A,
3197 a->blocks_per_vec, b, m0, m_count);
3198 } else {
3199 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc, w_block, active, a->A,
3200 a->blocks_per_vec, b, m0, m_count);
3201 }
3202 }
3203 for (int m = m0; m < m1; ++m) {
3204 float *c_row = a->C + (size_t)m * (size_t)a->N;
3205 for (int lane = 0; lane < active; ++lane) {
3206 c_row[n0 + lane] = acc[m - m0][lane];
3207 }
3208 }
3209#endif
3210}
3211
3212static void gemm_q4_packed_meta_x16_mreuse_thread_fn(int ith, int nth, void *args)
3213{
3214 gemm_q4_packed_meta_x16_thread_work_t *a = (gemm_q4_packed_meta_x16_thread_work_t *)args;
3215 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
3216 return;
3217 }
3218 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3219 if (tile_m > 8) tile_m = 8;
3220 const int mt = (a->M + tile_m - 1) / tile_m;
3221 const int total = mt * a->groups;
3222
3223 for (int job = ith; job < total; job += nth) {
3225 }
3226}
3227
3229 const void *B_packed_x16,
3230 const float *bias,
3231 float *C,
3232 int M, int N, int K,
3233 int tile_m,
3234 int active_threads)
3235{
3236 if (!A_q8 || !B_packed_x16 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
3237 return;
3238 }
3239 ck_threadpool_t *pool = ck_threadpool_global();
3240 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3241 const int groups = (N + 15) / 16;
3242 int tm = tile_m > 0 ? tile_m : 4;
3243 if (tm > 8) tm = 8;
3244 const int mt = (M + tm - 1) / tm;
3245 const int jobs = mt * groups;
3246 if (active_threads <= 0 || active_threads > pool_threads) {
3247 active_threads = pool_threads;
3248 }
3249 if (active_threads > jobs) {
3250 active_threads = jobs;
3251 }
3252 gemm_q4_packed_meta_x16_thread_work_t work = {
3253 .A = (const block_q8_K *)A_q8,
3254 .W = (const block_q4_K_packed_meta_x16 *)B_packed_x16,
3255 .bias = bias,
3256 .C = C,
3257 .M = M,
3258 .N = N,
3259 .K = K,
3260 .blocks_per_vec = K / QK_K,
3261 .blocks_per_row = K / QK_K,
3262 .groups = groups,
3263 .tile_m = tm,
3264 .jobs = jobs,
3265 };
3266 if (active_threads <= 1) {
3268 return;
3269 }
3271}
3272
3274 const void *B_packed_x16,
3275 const float *bias,
3276 float *C,
3277 int M, int N, int K,
3278 int tile_m,
3279 int active_threads)
3280{
3281 if (!A_q8 || !B_packed_x16 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) {
3282 return;
3283 }
3284 ck_threadpool_t *pool = ck_threadpool_global();
3285 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3286 const int groups = (N + 15) / 16;
3287 int tm = tile_m > 0 ? tile_m : 4;
3288 if (tm > 8) tm = 8;
3289 const int mt = (M + tm - 1) / tm;
3290 const int jobs = mt * groups;
3291 if (active_threads <= 0 || active_threads > pool_threads) {
3292 active_threads = pool_threads;
3293 }
3294 if (active_threads > jobs) {
3295 active_threads = jobs;
3296 }
3297 gemm_q4_packed_meta_x16_thread_work_t work = {
3298 .A = (const block_q8_K *)A_q8,
3299 .W = (const block_q4_K_packed_meta_x16 *)B_packed_x16,
3300 .bias = bias,
3301 .C = C,
3302 .M = M,
3303 .N = N,
3304 .K = K,
3305 .blocks_per_vec = K / QK_K,
3306 .blocks_per_row = K / QK_K,
3307 .groups = groups,
3308 .tile_m = tm,
3309 .jobs = jobs,
3310 };
3311 if (active_threads <= 1) {
3313 return;
3314 }
3316}
3317
3318
3319typedef struct {
3320 const block_q8_K *A;
3321 const block_q4_K_packed_meta_x16 *W;
3322 const float *bias;
3323 float *C;
3324 int M;
3325 int D;
3326 int K;
3327 int tile_m;
3328 int blocks_per_vec;
3329 int blocks_per_row;
3330 int groups_d;
3331 int jobs;
3332} gemm_q4_gateup_swiglu_x16_work_t;
3333
3334
3335static void gemm_q4_gateup_swiglu_x16_thread_fn(int ith, int nth, void *args)
3336{
3337 gemm_q4_gateup_swiglu_x16_work_t *a = (gemm_q4_gateup_swiglu_x16_work_t *)args;
3338 if (!a || ith < 0 || nth <= 0 || ith >= nth) {
3339 return;
3340 }
3341
3342 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3343 if (tile_m > 8) tile_m = 8;
3344 const int mt = (a->M + tile_m - 1) / tile_m;
3345
3346 for (int job = ith; job < a->jobs; job += nth) {
3347 const int g = job / mt;
3348 const int tm = job - g * mt;
3349 const int m0 = tm * tile_m;
3350 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3351 const int m_count = m1 - m0;
3352 const int d0 = g * 16;
3353 const int active = (d0 + 16 <= a->D) ? 16 : (a->D - d0);
3354 if (m_count <= 0 || active <= 0 || g >= a->groups_d) {
3355 continue;
3356 }
3357
3358 float acc_gate[8][16];
3359 float acc_up[8][16];
3360 for (int mt_lane = 0; mt_lane < m_count; ++mt_lane) {
3361 for (int lane = 0; lane < active; ++lane) {
3362 acc_gate[mt_lane][lane] = a->bias ? a->bias[d0 + lane] : 0.0f;
3363 acc_up[mt_lane][lane] = a->bias ? a->bias[a->D + d0 + lane] : 0.0f;
3364 }
3365 }
3366
3367 for (int b = 0; b < a->blocks_per_row; ++b) {
3368 const block_q4_K_packed_meta_x16 *w_gate =
3369 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3370 const block_q4_K_packed_meta_x16 *w_up =
3371 a->W + (size_t)(a->groups_d + g) * (size_t)a->blocks_per_row + (size_t)b;
3373 accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(acc_gate, w_gate, active, a->A,
3374 a->blocks_per_vec, b, m0, m_count);
3375 accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(acc_up, w_up, active, a->A,
3376 a->blocks_per_vec, b, m0, m_count);
3377 } else {
3378 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc_gate, w_gate, active, a->A,
3379 a->blocks_per_vec, b, m0, m_count);
3380 accum_q4_k_packed_meta_x16_q8_k_block_mreuse(acc_up, w_up, active, a->A,
3381 a->blocks_per_vec, b, m0, m_count);
3382 }
3383 }
3384
3385 for (int m = m0; m < m1; ++m) {
3386 float *c_row = a->C + (size_t)m * (size_t)a->D;
3387#if defined(__clang__) || defined(__INTEL_LLVM_COMPILER)
3388#pragma clang loop vectorize(enable) interleave(enable)
3389#elif defined(__GNUC__)
3390#pragma GCC ivdep
3391#endif
3392 for (int lane = 0; lane < active; ++lane) {
3393 const float gate = acc_gate[m - m0][lane];
3394 const float up = acc_up[m - m0][lane];
3395 c_row[d0 + lane] = (gate / (1.0f + expf(-gate))) * up;
3396 }
3397 }
3398 }
3399}
3400
3402 const void *B_packed_x16,
3403 const float *bias,
3404 float *C,
3405 int M, int D, int K,
3406 int tile_m,
3407 int active_threads)
3408{
3409 if (!A_q8 || !B_packed_x16 || !C || M <= 0 || D <= 0 || K <= 0 || (K % QK_K) != 0) {
3410 return;
3411 }
3412 ck_threadpool_t *pool = ck_threadpool_global();
3413 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3414 int tm = tile_m > 0 ? tile_m : 4;
3415 if (tm > 8) tm = 8;
3416 const int groups_d = (D + 15) / 16;
3417 const int mt = (M + tm - 1) / tm;
3418 const int jobs = mt * groups_d;
3419 if (active_threads <= 0 || active_threads > pool_threads) {
3420 active_threads = pool_threads;
3421 }
3422 if (active_threads > jobs) {
3423 active_threads = jobs;
3424 }
3425 if (active_threads < 1) {
3426 active_threads = 1;
3427 }
3428
3429 gemm_q4_gateup_swiglu_x16_work_t work = {
3430 .A = (const block_q8_K *)A_q8,
3431 .W = (const block_q4_K_packed_meta_x16 *)B_packed_x16,
3432 .bias = bias,
3433 .C = C,
3434 .M = M,
3435 .D = D,
3436 .K = K,
3437 .tile_m = tm,
3438 .blocks_per_vec = K / QK_K,
3439 .blocks_per_row = K / QK_K,
3440 .groups_d = groups_d,
3441 .jobs = jobs,
3442 };
3443 if (!pool || active_threads <= 1) {
3445 return;
3446 }
3448}
3449
3450
3451static void gemm_q4_packed_u8_x16_mtile_thread_fn(int ith, int nth, void *args)
3452{
3453 gemm_q4_packed_u8_x16_thread_work_t *a = (gemm_q4_packed_u8_x16_thread_work_t *)args;
3454 if (!a || ith < 0 || nth <= 0 || ith >= nth) return;
3455 int tile_m = a->tile_m > 0 ? a->tile_m : 4;
3456 if (tile_m > 8) tile_m = 8;
3457 const int mt = (a->M + tile_m - 1) / tile_m;
3458 const int total = mt * a->groups;
3459
3460 for (int job = ith; job < total; job += nth) {
3461 const int g = job / mt;
3462 const int tm = job - g * mt;
3463 const int m0 = tm * tile_m;
3464 const int m1 = (m0 + tile_m < a->M) ? (m0 + tile_m) : a->M;
3465 const int n0 = g * 16;
3466 const int active = (n0 + 16 <= a->N) ? 16 : (a->N - n0);
3467 if (m0 >= a->M || g >= a->groups) continue;
3468
3469 float acc[8][16];
3470 for (int mt_lane = 0; mt_lane < m1 - m0; ++mt_lane) {
3471 for (int lane = 0; lane < active; ++lane) {
3472 acc[mt_lane][lane] = a->bias ? a->bias[n0 + lane] : 0.0f;
3473 }
3474 }
3475
3476 for (int b = 0; b < a->blocks_per_row; ++b) {
3477 const block_q4_K_packed_u8_x16 *w_group =
3478 a->W + (size_t)g * (size_t)a->blocks_per_row + (size_t)b;
3479 for (int m = m0; m < m1; ++m) {
3480 const block_q8_K *a_row = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3481 accum_q4_k_packed_u8_x16_q8_k_block(acc[m - m0], w_group, active, &a_row[b]);
3482 }
3483 }
3484
3485 for (int m = m0; m < m1; ++m) {
3486 float *c_row = a->C + (size_t)m * (size_t)a->N;
3487 for (int lane = 0; lane < active; ++lane) {
3488 c_row[n0 + lane] = acc[m - m0][lane];
3489 }
3490 }
3491 }
3492}
3493
3495 const void *B_packed_u8_x16,
3496 const float *bias,
3497 float *C,
3498 int M, int N, int K,
3499 int tile_m,
3500 int active_threads)
3501{
3502 if (!A_q8 || !B_packed_u8_x16 || !C || M <= 0 || N <= 0 || K <= 0 || (K % QK_K) != 0) return;
3503 ck_threadpool_t *pool = ck_threadpool_global();
3504 const int pool_threads = pool ? ck_threadpool_n_threads(pool) : 1;
3505 const int groups = (N + 15) / 16;
3506 int tm = tile_m > 0 ? tile_m : 4;
3507 if (tm > 8) tm = 8;
3508 const int mt = (M + tm - 1) / tm;
3509 const int jobs = mt * groups;
3510 if (active_threads <= 0 || active_threads > pool_threads) active_threads = pool_threads;
3511 if (active_threads > jobs) active_threads = jobs;
3512 gemm_q4_packed_u8_x16_thread_work_t work = {
3513 .A = (const block_q8_K *)A_q8,
3514 .W = (const block_q4_K_packed_u8_x16 *)B_packed_u8_x16,
3515 .bias = bias,
3516 .C = C,
3517 .M = M,
3518 .N = N,
3519 .K = K,
3520 .blocks_per_vec = K / QK_K,
3521 .blocks_per_row = K / QK_K,
3522 .groups = groups,
3523 .tile_m = tm,
3524 .jobs = jobs,
3525 };
3526 if (active_threads <= 1) {
3528 return;
3529 }
3531}
3532
3533
3535 const void *W,
3536 const void *x_q8,
3537 int M, int K)
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}
3570
3571#if defined(CK_HAS_AVX_VNNI_256)
3572static inline void dot_q4_k_q8_k_vnni_block_rows4(
3573 const block_q4_K *w,
3574 const block_q8_K *const x[4],
3575 int rows,
3576 float out[4])
3577{
3578 uint8_t sc[8], m_val[8];
3579 unpack_q4_k_scales(w->scales, sc, m_val);
3580 for (int row = 0; row < rows; ++row) out[row] = 0.0f;
3581
3582 for (int j = 0, is = 0, q_offset = 0;
3583 j < QK_K; j += 64, is += 2, q_offset += 32) {
3584 const __m256i packed = _mm256_loadu_si256(
3585 (const __m256i *)(const void *)&w->qs[q_offset]);
3586 const __m256i q4_lo = q4_k_unpack_32_vnni_bytes(packed, 0);
3587 const __m256i q4_hi = q4_k_unpack_32_vnni_bytes(packed, 1);
3588
3589 for (int row = 0; row < rows; ++row) {
3590 const block_q8_K *xr = x[row];
3591 const __m256i q8_lo = _mm256_loadu_si256(
3592 (const __m256i *)(const void *)&xr->qs[j]);
3593 const __m256i q8_hi = _mm256_loadu_si256(
3594 (const __m256i *)(const void *)&xr->qs[j + 32]);
3595 const int32_t sum_lo =
3596 dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_lo, q8_lo);
3597 const int32_t sum_hi =
3598 dot_q4_k_q8_k_32_vnni_q4v_q8v(q4_hi, q8_hi);
3599 const int32_t bsum_lo = (int32_t)xr->bsums[j / 16] +
3600 (int32_t)xr->bsums[j / 16 + 1];
3601 const int32_t bsum_hi = (int32_t)xr->bsums[(j + 32) / 16] +
3602 (int32_t)xr->bsums[(j + 32) / 16 + 1];
3603 const float d = CK_FP16_TO_FP32(w->d) * xr->d;
3604 const float dmin = CK_FP16_TO_FP32(w->dmin) * xr->d;
3605
3606 out[row] += d * (float)sc[is] * (float)sum_lo;
3607 out[row] -= dmin * (float)m_val[is] * (float)bsum_lo;
3608 out[row] += d * (float)sc[is + 1] * (float)sum_hi;
3609 out[row] -= dmin * (float)m_val[is + 1] * (float)bsum_hi;
3610 }
3611 }
3612}
3613#endif
3614
3616 int output_stride,
3617 const void *weights,
3618 const void *const input_rows[4],
3619 int rows,
3620 int output_dim,
3621 int input_dim)
3622{
3623 if (!output || !weights || !input_rows || rows <= 0 || rows > 4 ||
3624 output_stride < output_dim || output_dim <= 0 || input_dim <= 0 ||
3625 (input_dim % QK_K) != 0) {
3626 return;
3627 }
3628 for (int row = 0; row < rows; ++row) {
3629 if (!input_rows[row]) return;
3630 }
3631
3632#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3633 const block_q4_K *blocks = (const block_q4_K *)weights;
3634 const int blocks_per_row = input_dim / QK_K;
3635 const block_q8_K *inputs[4] = {
3636 (const block_q8_K *)input_rows[0],
3637 (const block_q8_K *)input_rows[rows > 1 ? 1 : 0],
3638 (const block_q8_K *)input_rows[rows > 2 ? 2 : 0],
3639 (const block_q8_K *)input_rows[rows > 3 ? 3 : 0],
3640 };
3641 for (int n = 0; n < output_dim; ++n) {
3642 const block_q4_K *weight_row =
3643 blocks + (size_t)n * (size_t)blocks_per_row;
3644 float sums[4] = {0.0f, 0.0f, 0.0f, 0.0f};
3645 for (int block = 0; block < blocks_per_row; ++block) {
3646 const block_q8_K *block_rows[4] = {
3647 &inputs[0][block], &inputs[1][block],
3648 &inputs[2][block], &inputs[3][block],
3649 };
3650 float block_sums[4];
3651 dot_q4_k_q8_k_vnni_block_rows4(
3652 &weight_row[block], block_rows, rows, block_sums);
3653 for (int row = 0; row < rows; ++row) {
3654 sums[row] += block_sums[row];
3655 }
3656 }
3657 for (int row = 0; row < rows; ++row) {
3658 output[(size_t)row * (size_t)output_stride + (size_t)n] = sums[row];
3659 }
3660 }
3661#else
3662 for (int row = 0; row < rows; ++row) {
3664 output + (size_t)row * (size_t)output_stride,
3665 weights, input_rows[row], output_dim, input_dim);
3666 }
3667#endif
3668}
3669
3670
3671static inline float ck_q4k_silu_f32(float x)
3672{
3673 return x / (1.0f + expf(-x));
3674}
3675
3676typedef struct {
3677 const block_q8_K *A;
3678 const block_q4_K *W;
3679 const float *bias;
3680 float *C;
3681 int M;
3682 int D;
3683 int K;
3684 int blocks_per_vec;
3685 int blocks_per_row;
3686} gemm_q4_gateup_swiglu_work_t;
3687
3688static void gemm_q4_gateup_swiglu_thread_fn(int ith, int nth, void *args)
3689{
3690#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3691 gemm_q4_gateup_swiglu_work_t *a = (gemm_q4_gateup_swiglu_work_t *)args;
3692 if (!a || ith < 0 || nth <= 0 || ith >= nth) return;
3693
3694 const int dd = (a->D + nth - 1) / nth;
3695 const int d0 = dd * ith;
3696 const int d1 = (d0 + dd < a->D) ? (d0 + dd) : a->D;
3697 if (d0 >= a->D) return;
3698
3699 for (int d = d0; d < d1; ++d) {
3700 const block_q4_K *w_gate = a->W + (size_t)d * (size_t)a->blocks_per_row;
3701 const block_q4_K *w_up = a->W + (size_t)(a->D + d) * (size_t)a->blocks_per_row;
3702 const float b_gate = a->bias ? a->bias[d] : 0.0f;
3703 const float b_up = a->bias ? a->bias[a->D + d] : 0.0f;
3704 for (int m = 0; m < a->M; ++m) {
3705 const block_q8_K *x = a->A + (size_t)m * (size_t)a->blocks_per_vec;
3706 float gate = b_gate;
3707 float up = b_up;
3708 for (int b = 0; b < a->blocks_per_row; ++b) {
3709 gate += dot_q4_k_q8_k_vnni_block(&w_gate[b], &x[b]);
3710 up += dot_q4_k_q8_k_vnni_block(&w_up[b], &x[b]);
3711 }
3712 a->C[(size_t)m * (size_t)a->D + (size_t)d] = ck_q4k_silu_f32(gate) * up;
3713 }
3714 }
3715#else
3716 (void)ith;
3717 (void)nth;
3718 (void)args;
3719#endif
3720}
3721
3723 const void *B_gate_up,
3724 const float *bias,
3725 float *C,
3726 int M,
3727 int D,
3728 int K,
3729 int threads)
3730{
3731#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3732 if (!A_q8 || !B_gate_up || !C || M <= 0 || D <= 0 || K <= 0 || (K % QK_K) != 0) {
3733 return;
3734 }
3735
3736 ck_threadpool_t *pool = ck_threadpool_global();
3737 int active = threads > 0 ? threads : (pool ? ck_threadpool_n_threads(pool) : 1);
3738 if (active < 1) active = 1;
3739 if (active > D) active = D;
3740
3741 gemm_q4_gateup_swiglu_work_t work = {
3742 .A = (const block_q8_K *)A_q8,
3743 .W = (const block_q4_K *)B_gate_up,
3744 .bias = bias,
3745 .C = C,
3746 .M = M,
3747 .D = D,
3748 .K = K,
3749 .blocks_per_vec = K / QK_K,
3750 .blocks_per_row = K / QK_K,
3751 };
3752
3753 if (active <= 1 || !pool) {
3755 return;
3756 }
3758#else
3759 (void)A_q8;
3760 (void)B_gate_up;
3761 (void)bias;
3762 (void)C;
3763 (void)M;
3764 (void)D;
3765 (void)K;
3766 (void)threads;
3767#endif
3768}
3769
3770
3772 const void *W,
3773 const void *x_q8,
3774 int M, int K,
3775 int ith, int nth)
3776{
3777#if defined(__AVX512VNNI__) && defined(__AVX512VL__)
3778 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
3779 return;
3780 }
3781 if (ith < 0 || nth <= 0 || ith >= nth) {
3782 return;
3783 }
3784
3785 const int dr = (M + nth - 1) / nth;
3786 const int r0 = dr * ith;
3787 const int r1 = (r0 + dr < M) ? (r0 + dr) : M;
3788 if (r0 >= M) {
3789 return;
3790 }
3791
3792 const block_q4_K *blocks = (const block_q4_K *)W;
3793 const block_q8_K *x = (const block_q8_K *)x_q8;
3794 const int blocks_per_row = K / QK_K;
3795
3796 for (int row = r0; row < r1; ++row) {
3797 const block_q4_K *w_row = blocks + (size_t)row * (size_t)blocks_per_row;
3798 float sum = 0.0f;
3799 for (int b = 0; b < blocks_per_row; ++b) {
3800 sum += dot_q4_k_q8_k_vnni_block(&w_row[b], &x[b]);
3801 }
3802 y[row] = sum;
3803 }
3804#else
3805 (void)y;
3806 (void)W;
3807 (void)x_q8;
3808 (void)M;
3809 (void)K;
3810 (void)ith;
3811 (void)nth;
3812#endif
3813}
static int ck_env_value_truthy(const char *v)
Persistent pthread thread pool for CK-Engine inference.
int ck_threadpool_capacity(const ck_threadpool_t *pool)
void ck_threadpool_parallel_for_n(ck_threadpool_t *pool, int active_threads, int begin, int end, int grain_size, ck_range_fn_t fn, void *args)
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_gemm_dynamic_schedule_enabled(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)
int ck_strict_parity_enabled(void)
Quantization block structures for weight-only quantization.
uint16_t ck_half
#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
static void accum_q4_k_packed_meta_x8_q8_k_gemv_block(float acc[8], float acc_min[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)
void gemv_q4_k_q8_k_avx2(float *y, const void *W, const void *x_q8, int M, int K)
static void gemm_q4_packed_meta_x16_mreuse_process_job(const gemm_q4_packed_meta_x16_thread_work_t *a, int job, int mt, int tile_m)
static void accum_q4_k_packed_vnni_x16_q8_k_gemv_block(float acc[16], float acc_min[16], const block_q4_K_packed_vnni_x16 *w, const block_q8_K *x)
size_t q4_k_packed_meta_x16_block_size(void)
static void gemm_q4_packed_meta_x8_mtile_thread_fn(int ith, int nth, void *args)
static void gemm_q4_packed_vnni_x8_q8k_4m_job(gemm_q4_packed_vnni_x8_thread_work_t *a, int job, int row_tiles)
void gemv_q4_k_q8_k_parallel_vnni(float *y, const void *W, const void *x_q8, int M, int K, int ith, int nth)
static float dot_q4_k_packed_meta_q8_k_block(const block_q4_K_packed_meta *w, const block_q8_K *x)
void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mreuse(const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
static void gemm_q4_packed_vnni_x8_q8k_4m_range_fn(int begin, int end, void *args)
size_t q4_k_packed_vnni_x16_block_size(void)
static void gemm_q4_packed_vnni_x8_q8k_4m_thread_fn(int ith, int nth, void *args)
static void gemm_q4_packed_meta_x8_split_min_4m_thread_fn(int ith, int nth, void *args)
void gemm_nt_q4_k_packed_meta_x8_q8_k(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse(float acc[8][16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
static void accum_q4_k_packed_meta_x8_q8_k_superblock(float acc[8], float acc_min[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_4m(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
void pack_q4_k_to_packed_meta_x16(const void *src, void *dst, int N, int K)
static float ck_q4k_silu_f32(float x)
static void accum_q4_k_packed_meta_x8_q8_k_block(float acc[8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x)
static void dot_q4_k_packed_vnni_x8_q8_k_compact_order(float block_sums[4][8], const block_q4_K_packed_vnni_x8 *w, const block_q8_K *x[4], int rows)
size_t q4_k_packed_meta_block_size(void)
static void gemm_q4_packed_meta_x16_mtile_thread_fn(int ith, int nth, void *args)
void gemm_nt_q4_k_packed_meta_q8_k_tile(const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K, int m0, int m1, int n0, int n1)
void gemv_q4_k_q8_k_vnni(float *y, const void *W, const void *x_q8, int M, int K)
void pack_q4_k_to_packed_vnni_x8(const void *src, void *dst, int N, int K)
void gemm_nt_q4_k_packed_meta_x8_q8_k_gemv_order(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
static float dot_q4_k_packed_u8_q8_k_block(const block_q4_K_packed_u8 *w, const block_q8_K *x)
static void accum_q4_k_packed_meta_x16_q8_k_block(float acc[16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *x)
void gemm_q4_k_q8_k_packed_vnni_x8_compact_order_rows4(float *output, const void *weights_packed, const void *input_q8, int rows, int output_dim, int input_dim)
size_t q4_k_packed_meta_x8_block_size(void)
int ck_q4k_packed_vnni_x8_compact_order_available(void)
void gemm_nt_q4_k_packed_meta_x16_q8_k_threaded_mtile(const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
static void gemm_q4_packed_meta_x8_split_min_8m_thread_fn(int ith, int nth, void *args)
void pack_q4_k_to_packed_u8_x16(const void *src, void *dst, int N, int K)
static void gemm_q4_packed_meta_x8_mreuse_thread_fn(int ith, int nth, void *args)
int ck_q4k_packed_vnni_x8_available(void)
int ck_q4k_packed_vnni_x16_available(void)
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_mreuse(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
void gemm_nt_q4_k_packed_u8_q8_k(const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K)
void pack_q4_k_to_packed_vnni_x16(const void *src, void *dst, int N, int K)
void pack_q4_k_to_packed_meta(const void *src, void *dst, int N, int K)
size_t q4_k_packed_u8_x16_block_size(void)
static void gemm_q4_packed_u8_x16_mtile_thread_fn(int ith, int nth, void *args)
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mreuse(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
static void accum_q4_k_packed_meta_x16_q8_k_block_mreuse_chunk4(float acc[8][16], const block_q4_K_packed_meta_x16 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
void gemm_nt_q4_k_packed_meta_x8_q8_k_split_min_threaded_8m(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
void gemm_nt_q4_k_packed_meta_x16_gateup_swiglu_fused_vnni(const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int D, int K, int tile_m, int active_threads)
void gemm_nt_q4_k_packed_meta_q8_k(const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K)
static void accum_q4_k_packed_meta_x8_q8_k_block_mreuse(float acc[8][8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *A, int blocks_per_vec, int block_index, int m0, int m_count)
size_t q4_k_packed_vnni_x8_block_size(void)
void pack_q4_k_to_packed_u8(const void *src, void *dst, int N, int K)
static void accum_q4_k_packed_u8_x16_q8_k_block(float acc[16], const block_q4_K_packed_u8_x16 *w, int active, const block_q8_K *x)
void gemm_nt_q4_k_packed_vnni_x8_q8_k_split_min_threaded_4m(const void *A_q8, const void *B_packed_vnni_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
void gemm_nt_q4_k_packed_meta_q8_k_threaded_nsplit(const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K, int active_threads)
static void gemm_q4_packed_meta_x16_mreuse_thread_fn(int ith, int nth, void *args)
static int32_t dot_q4_packed_u8_q8_32_ref(const uint8_t *q4_32, const int8_t *q8_32)
static void gemm_q4_packed_vnni_x16_q8k_16m_thread_fn(int ith, int nth, void *args)
void gemv_q4_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)
static void gemm_q4_packed_meta_x8_split_min_mreuse_thread_fn(int ith, int nth, void *args)
void gemm_nt_q4_k_packed_meta_x8_q8_k_superblock_order(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
size_t q4_k_packed_u8_block_size(void)
static void accum_q4_k_packed_meta_x8_q8_k_superblock_rows(float acc[8][8], float acc_min[8][8], const block_q4_K_packed_meta_x8 *w, int active, const block_q8_K *x[8], int rows)
static void gemm_q4_packed_meta_nsplit_thread_fn(int ith, int nth, void *args)
void pack_q4_k_to_packed_meta_x8(const void *src, void *dst, int N, int K)
void gemm_nt_q4_k_packed_vnni_x16_q8_k_gemv_order(const void *A_q8, const void *B_packed_x16, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q4_k_packed_u8_x16_q8_k_threaded_mtile(const void *A_q8, const void *B_packed_u8_x16, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
static void gemm_q4_packed_meta_x8_nsplit_thread_fn(int ith, int nth, void *args)
static void gemm_q4_packed_meta_thread_fn(int ith, int nth, void *args)
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_mtile(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int tile_m, int active_threads)
void gemm_q4_k_q8_k_compact_rows4(float *output, int output_stride, const void *weights, const void *const input_rows[4], int rows, int output_dim, int input_dim)
void gemm_nt_q4_k_packed_meta_x8_q8_k_threaded_nsplit(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K, int active_threads)
void gemm_nt_q4_k_q8_k_gateup_swiglu_fused_vnni(const void *A_q8, const void *B_gate_up, const float *bias, float *C, int M, int D, int K, int threads)
void gemm_nt_q4_k_packed_meta_x16_q8_k_llama_order(const void *A_q8, const void *B_packed_x8, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q4_k_packed_vnni_x16_q8_k_split_min_threaded_16m(const void *A_q8, const void *B_packed_vnni_x16, const float *bias, float *C, int M, int N, int K, int active_threads)
static int ck_q4k_x16_chunk4_enabled(void)
static void accum_q4_k_packed_vnni_x16_q8_k_16m_superblock(float acc[16][16], float acc_min[16][16], const block_q4_K_packed_vnni_x16 *w, const block_q8_K *x[16], int rows)
void gemm_nt_q4_k_packed_meta_q8_k_threaded(const void *A_q8, const void *B_packed, const float *bias, float *C, int M, int N, int K, int active_threads)
static void accum_q4_k_packed_vnni_x8_q8_k_4m_superblock(float acc[4][8], float acc_min[4][8], const block_q4_K_packed_vnni_x8 *w, const block_q8_K *x[4], int rows)
static void gemm_q4_gateup_swiglu_thread_fn(int ith, int nth, void *args)
static void gemm_q4_gateup_swiglu_x16_thread_fn(int ith, int nth, void *args)
#define C(color)
Definition show_config.c:39
uint8_t scales[12]
uint8_t qs[256/2]
int8_t qs[256]
int16_t bsums[256/16]
uint32_t end
Definition utf8.c:215