← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_q5_k.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels_q5_k.c
3 * @brief GEMM/GEMV kernels with Q5_K quantized weights
4 *
5 * CK-ENGINE KERNEL RULES:
6 * =======================
7 * 1. NO malloc/free - memory via bump allocator, pointers passed in
8 * 2. NO OpenMP - parallelization at orchestrator/codegen layer
9 * 3. API must define: inputs, outputs, workspace, and memory layouts
10 * 4. Pure computation - deterministic, no side effects
11 *
12 * After changes: make test && make llamacpp-parity-full
13 *
14 * Implements matrix multiplication where:
15 * - Activations (input): FP32 (quantized internally to Q8_K for dot path)
16 * - Weights: Q5_K (5-bit super-block quant)
17 * - Output: FP32
18 *
19 * Q5_K Format (256 weights per super-block):
20 * - d: FP16 super-block scale
21 * - dmin: FP16 super-block minimum
22 * - scales[12]: 8 sub-block scales + 8 sub-block mins (6 bits each, packed)
23 * - qh[32]: high bits for 256 weights (1 bit each)
24 * - qs[128]: low 4 bits for 256 weights (4 bits each)
25 *
26 * Total: 2 + 2 + 12 + 32 + 128 = 176 bytes per 256 weights = 5.5 bits/weight
27 *
28 * Dequantization formula (matches llama.cpp):
29 * w = d * scale * q - dmin * mins
30 * where q = qs_val | (qh_bit << 4) = 5-bit value [0, 31]
31 */
32
33#include <stdint.h>
34#include <stddef.h>
35#include <string.h>
36#include <stdlib.h>
37#include "ckernel_quant.h"
38
39/* Include SIMD headers based on available extensions */
40#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE4_1__)
41#include <immintrin.h>
42#endif
43
44/* Q5_K constants */
45#define QK_K 256
46#define CK_Q5K_STACK_Q8_BLOCKS 128
47
49{
50 static int cached = -1;
51 if (cached < 0) {
52 const char *env = getenv("CK_DEBUG_Q5K_FP32_FALLBACK");
53 cached = (env && env[0] && env[0] != '0') ? 1 : 0;
54 }
55 return cached;
56}
57
59{
60 static int cached = -1;
61 if (cached < 0) {
62 const char *env = getenv("CK_DEBUG_Q5K_GENERIC_DOT");
63 cached = (env && env[0] && env[0] != '0') ? 1 : 0;
64 }
65 return cached;
66}
67
68/* Q5_K block definition is required by this kernel file.
69 * Keep a local ggml-compatible layout to decouple from shared headers. */
70typedef struct {
71 ck_half d;
72 ck_half dmin;
73 uint8_t scales[K_SCALE_SIZE];
74 uint8_t qh[QK_K / 8];
75 uint8_t qs[QK_K / 2];
76} block_q5_K;
77
78/* Load-time representation used by the optional prepared prefill provider.
79 * It expands only integer metadata: FP16 super-block scales are retained
80 * verbatim, while sub-block scales/mins and 5-bit codes become byte-addressable.
81 * The dot-product and FP32 reduction order remain unchanged. */
82typedef struct {
83 ck_half d;
84 ck_half dmin;
85 uint8_t scales[8];
86 uint8_t mins[8];
87 uint8_t qs[QK_K];
88} block_q5_K_prepared;
89
90_Static_assert(sizeof(block_q5_K_prepared) == 276,
91 "Q5_K prepared-size contract changed");
92
93/* Unpack 8 per-subblock scales and mins from packed Q5_K scale bytes.
94 * This mirrors the packing contract used by llama.cpp. */
95static inline void unpack_q5_k_scales(const uint8_t *scales,
96 uint8_t *sc,
97 uint8_t *m) {
98 sc[0] = scales[0] & 0x3F;
99 sc[1] = scales[1] & 0x3F;
100 sc[2] = scales[2] & 0x3F;
101 sc[3] = scales[3] & 0x3F;
102
103 m[0] = scales[4] & 0x3F;
104 m[1] = scales[5] & 0x3F;
105 m[2] = scales[6] & 0x3F;
106 m[3] = scales[7] & 0x3F;
107
108 sc[4] = (scales[8] & 0x0F) | ((scales[0] >> 6) << 4);
109 sc[5] = (scales[9] & 0x0F) | ((scales[1] >> 6) << 4);
110 sc[6] = (scales[10] & 0x0F) | ((scales[2] >> 6) << 4);
111 sc[7] = (scales[11] & 0x0F) | ((scales[3] >> 6) << 4);
112
113 m[4] = (scales[8] >> 4) | ((scales[4] >> 6) << 4);
114 m[5] = (scales[9] >> 4) | ((scales[5] >> 6) << 4);
115 m[6] = (scales[10] >> 4) | ((scales[6] >> 6) << 4);
116 m[7] = (scales[11] >> 4) | ((scales[7] >> 6) << 4);
117}
118
119static inline uint8_t q5_k_quant_value(const block_q5_K *block, int subblock, int i) {
120 const uint8_t *ql = block->qs + (subblock / 2) * 32;
121 const uint8_t low = (subblock & 1) ? (uint8_t)(ql[i] >> 4) : (uint8_t)(ql[i] & 0x0F);
122 const uint8_t high = (block->qh[i] & (uint8_t)(1u << subblock)) ? 16u : 0u;
123 return (uint8_t)(low | high);
124}
125
127{
128 return sizeof(block_q5_K_prepared);
129}
130
131void ck_q5_k_prepare_weight(const void *src, void *dst, int N, int K)
132{
133 if (!src || !dst || N <= 0 || K <= 0 || (K % QK_K) != 0) return;
134 const block_q5_K *input = (const block_q5_K *)src;
135 block_q5_K_prepared *output = (block_q5_K_prepared *)dst;
136 const size_t blocks = (size_t)N * (size_t)(K / QK_K);
137 for (size_t b = 0; b < blocks; ++b) {
138 output[b].d = input[b].d;
139 output[b].dmin = input[b].dmin;
140 unpack_q5_k_scales(input[b].scales, output[b].scales, output[b].mins);
141 for (int sb = 0; sb < 8; ++sb) {
142 for (int i = 0; i < 32; ++i) {
143 output[b].qs[sb * 32 + i] = q5_k_quant_value(&input[b], sb, i);
144 }
145 }
146 }
147}
148
149/* quantize_row_q8_k() is implemented in gemm_kernels_q4k_q8k.c */
150void quantize_row_q8_k(const float *x, void *vy, int k);
151
152#if defined(__AVX2__)
153static inline __m256i ck_q5k_scale_shuffle_avx2(int i)
154{
155 static const uint8_t k_shuffle[256] = {
156 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1, 0, 1,
157 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3, 2, 3,
158 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5, 4, 5,
159 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7, 6, 7,
160 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9, 8, 9,
161 10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,10,11,
162 12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,12,13,
163 14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15,14,15
164 };
165 return _mm256_loadu_si256((const __m256i *)(const void *)(k_shuffle + 32 * i));
166}
167
168static inline __m256i ck_mm256_set_m128i(__m128i hi, __m128i lo)
169{
170 return _mm256_inserti128_si256(_mm256_castsi128_si256(lo), hi, 1);
171}
172
173static inline float ck_q5k_hsum256_ps(__m256 v)
174{
175 __m128 sum = _mm256_extractf128_ps(v, 1);
176 sum = _mm_add_ps(sum, _mm256_castps256_ps128(v));
177 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
178 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
179 return _mm_cvtss_f32(sum);
180}
181
182static float dot_q5_k_q8_k_row_avx2(const block_q5_K *w, const block_q8_K *x, int nb) {
183 static const uint32_t kmask1 = 0x3f3f3f3fU;
184 static const uint32_t kmask2 = 0x0f0f0f0fU;
185 static const uint32_t kmask3 = 0x03030303U;
186
187 const __m256i m4 = _mm256_set1_epi8(0x0f);
188 const __m128i mzero = _mm_setzero_si128();
189 const __m256i mone = _mm256_set1_epi8(1);
190
191 uint32_t utmp[4] = {0, 0, 0, 0};
192 __m256 acc = _mm256_setzero_ps();
193 float summs = 0.0f;
194
195 for (int b = 0; b < nb; ++b) {
196 const block_q5_K *wb = &w[b];
197 const block_q8_K *xb = &x[b];
198 const uint8_t *q5 = wb->qs;
199 const int8_t *q8 = xb->qs;
200
201 const float d = CK_FP16_TO_FP32(wb->d) * xb->d;
202 const float dmin = -CK_FP16_TO_FP32(wb->dmin) * xb->d;
203
204 memcpy(utmp, wb->scales, 12);
205 utmp[3] = ((utmp[2] >> 4) & kmask2) | (((utmp[1] >> 6) & kmask3) << 4);
206 const uint32_t uaux = utmp[1] & kmask1;
207 utmp[1] = (utmp[2] & kmask2) | (((utmp[0] >> 6) & kmask3) << 4);
208 utmp[2] = uaux;
209 utmp[0] &= kmask1;
210
211 const __m256i mins_and_scales =
212 _mm256_cvtepu8_epi16(_mm_set_epi32((int)utmp[3], (int)utmp[2], (int)utmp[1], (int)utmp[0]));
213
214 const __m256i q8sums = _mm256_loadu_si256((const __m256i *)(const void *)xb->bsums);
215 const __m128i q8s = _mm_hadd_epi16(_mm256_extracti128_si256(q8sums, 0),
216 _mm256_extracti128_si256(q8sums, 1));
217 const __m128i prod = _mm_madd_epi16(_mm256_extracti128_si256(mins_and_scales, 1), q8s);
218 const __m128i hsum = _mm_hadd_epi32(_mm_hadd_epi32(prod, mzero), mzero);
219 summs += dmin * (float)_mm_extract_epi32(hsum, 0);
220
221 const __m128i sc128 = _mm256_extracti128_si256(mins_and_scales, 0);
222 const __m256i scales = ck_mm256_set_m128i(sc128, sc128);
223 const __m256i hbits = _mm256_loadu_si256((const __m256i *)(const void *)wb->qh);
224 __m256i hmask = mone;
225 __m256i sumi = _mm256_setzero_si256();
226 int bit = 0;
227
228 for (int j = 0; j < QK_K / 64; ++j) {
229 const __m256i scale_0 = _mm256_shuffle_epi8(scales, ck_q5k_scale_shuffle_avx2(2 * j + 0));
230 const __m256i scale_1 = _mm256_shuffle_epi8(scales, ck_q5k_scale_shuffle_avx2(2 * j + 1));
231
232 const __m256i q5bits = _mm256_loadu_si256((const __m256i *)(const void *)q5);
233 q5 += 32;
234
235 const __m256i q5l_0 = _mm256_and_si256(q5bits, m4);
236 const __m256i q5h_0 = _mm256_slli_epi16(_mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
237 const __m256i q5_0 = _mm256_add_epi8(q5l_0, q5h_0);
238 hmask = _mm256_slli_epi16(hmask, 1);
239
240 const __m256i q5l_1 = _mm256_and_si256(_mm256_srli_epi16(q5bits, 4), m4);
241 const __m256i q5h_1 = _mm256_slli_epi16(_mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
242 const __m256i q5_1 = _mm256_add_epi8(q5l_1, q5h_1);
243 hmask = _mm256_slli_epi16(hmask, 1);
244
245 const __m256i q8_0 = _mm256_loadu_si256((const __m256i *)(const void *)q8);
246 q8 += 32;
247 const __m256i q8_1 = _mm256_loadu_si256((const __m256i *)(const void *)q8);
248 q8 += 32;
249
250 __m256i p16_0 = _mm256_maddubs_epi16(q5_0, q8_0);
251 __m256i p16_1 = _mm256_maddubs_epi16(q5_1, q8_1);
252 p16_0 = _mm256_madd_epi16(scale_0, p16_0);
253 p16_1 = _mm256_madd_epi16(scale_1, p16_1);
254 sumi = _mm256_add_epi32(sumi, _mm256_add_epi32(p16_0, p16_1));
255 }
256
257 acc = _mm256_fmadd_ps(
258 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi), acc);
259 }
260
261 return ck_q5k_hsum256_ps(acc) + summs;
262}
263
264static void dot_q5_k_q8_k_rows4_avx2(
265 const block_q5_K *w,
266 const block_q8_K *const x[4],
267 int rows,
268 int nb,
269 float out[4])
270{
271 static const uint32_t kmask1 = 0x3f3f3f3fU;
272 static const uint32_t kmask2 = 0x0f0f0f0fU;
273 static const uint32_t kmask3 = 0x03030303U;
274 const __m256i m4 = _mm256_set1_epi8(0x0f);
275 const __m128i mzero = _mm_setzero_si128();
276 const __m256i mone = _mm256_set1_epi8(1);
277 __m256 acc[4] = {
278 _mm256_setzero_ps(), _mm256_setzero_ps(),
279 _mm256_setzero_ps(), _mm256_setzero_ps(),
280 };
281 float summs[4] = {0.0f, 0.0f, 0.0f, 0.0f};
282
283 for (int b = 0; b < nb; ++b) {
284 const block_q5_K *wb = &w[b];
285 uint32_t utmp[4] = {0, 0, 0, 0};
286 memcpy(utmp, wb->scales, 12);
287 utmp[3] = ((utmp[2] >> 4) & kmask2) |
288 (((utmp[1] >> 6) & kmask3) << 4);
289 const uint32_t uaux = utmp[1] & kmask1;
290 utmp[1] = (utmp[2] & kmask2) |
291 (((utmp[0] >> 6) & kmask3) << 4);
292 utmp[2] = uaux;
293 utmp[0] &= kmask1;
294
295 const __m256i mins_and_scales = _mm256_cvtepu8_epi16(
296 _mm_set_epi32((int)utmp[3], (int)utmp[2],
297 (int)utmp[1], (int)utmp[0]));
298 const __m128i sc128 =
299 _mm256_extracti128_si256(mins_and_scales, 0);
300 const __m256i scales = ck_mm256_set_m128i(sc128, sc128);
301 const __m256i hbits = _mm256_loadu_si256(
302 (const __m256i *)(const void *)wb->qh);
303 __m256i sumi[4] = {
304 _mm256_setzero_si256(), _mm256_setzero_si256(),
305 _mm256_setzero_si256(), _mm256_setzero_si256(),
306 };
307
308 for (int row = 0; row < rows; ++row) {
309 const block_q8_K *xb = &x[row][b];
310 const __m256i q8sums = _mm256_loadu_si256(
311 (const __m256i *)(const void *)xb->bsums);
312 const __m128i q8s = _mm_hadd_epi16(
313 _mm256_extracti128_si256(q8sums, 0),
314 _mm256_extracti128_si256(q8sums, 1));
315 const __m128i prod = _mm_madd_epi16(
316 _mm256_extracti128_si256(mins_and_scales, 1), q8s);
317 const __m128i hsum = _mm_hadd_epi32(
318 _mm_hadd_epi32(prod, mzero), mzero);
319 const float dmin = -CK_FP16_TO_FP32(wb->dmin) * xb->d;
320 summs[row] += dmin * (float)_mm_extract_epi32(hsum, 0);
321 }
322
323 const uint8_t *q5 = wb->qs;
324 __m256i hmask = mone;
325 int bit = 0;
326 for (int j = 0; j < QK_K / 64; ++j) {
327 const __m256i scale_0 = _mm256_shuffle_epi8(
328 scales, ck_q5k_scale_shuffle_avx2(2 * j));
329 const __m256i scale_1 = _mm256_shuffle_epi8(
330 scales, ck_q5k_scale_shuffle_avx2(2 * j + 1));
331 const __m256i q5bits = _mm256_loadu_si256(
332 (const __m256i *)(const void *)q5);
333 q5 += 32;
334 const __m256i q5l_0 = _mm256_and_si256(q5bits, m4);
335 const __m256i q5h_0 = _mm256_slli_epi16(
336 _mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
337 const __m256i q5_0 = _mm256_add_epi8(q5l_0, q5h_0);
338 hmask = _mm256_slli_epi16(hmask, 1);
339 const __m256i q5l_1 = _mm256_and_si256(
340 _mm256_srli_epi16(q5bits, 4), m4);
341 const __m256i q5h_1 = _mm256_slli_epi16(
342 _mm256_srli_epi16(_mm256_and_si256(hbits, hmask), bit++), 4);
343 const __m256i q5_1 = _mm256_add_epi8(q5l_1, q5h_1);
344 hmask = _mm256_slli_epi16(hmask, 1);
345
346 for (int row = 0; row < rows; ++row) {
347 const block_q8_K *xb = &x[row][b];
348 const __m256i q8_0 = _mm256_loadu_si256(
349 (const __m256i *)(const void *)&xb->qs[j * 64]);
350 const __m256i q8_1 = _mm256_loadu_si256(
351 (const __m256i *)(const void *)&xb->qs[j * 64 + 32]);
352 __m256i p16_0 = _mm256_maddubs_epi16(q5_0, q8_0);
353 __m256i p16_1 = _mm256_maddubs_epi16(q5_1, q8_1);
354 p16_0 = _mm256_madd_epi16(scale_0, p16_0);
355 p16_1 = _mm256_madd_epi16(scale_1, p16_1);
356 sumi[row] = _mm256_add_epi32(
357 sumi[row], _mm256_add_epi32(p16_0, p16_1));
358 }
359 }
360
361 for (int row = 0; row < rows; ++row) {
362 const float d = CK_FP16_TO_FP32(wb->d) * x[row][b].d;
363 acc[row] = _mm256_fmadd_ps(
364 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi[row]), acc[row]);
365 }
366 }
367
368 for (int row = 0; row < rows; ++row) {
369 out[row] = ck_q5k_hsum256_ps(acc[row]) + summs[row];
370 }
371}
372
373static float dot_q5_k_prepared_q8_k_row_avx2(
374 const block_q5_K_prepared *w, const block_q8_K *x, int nb)
375{
376 const __m128i mzero = _mm_setzero_si128();
377 __m256 acc = _mm256_setzero_ps();
378 float summs = 0.0f;
379
380 for (int b = 0; b < nb; ++b) {
381 const block_q5_K_prepared *wb = &w[b];
382 const block_q8_K *xb = &x[b];
383 const float d = CK_FP16_TO_FP32(wb->d) * xb->d;
384 const float dmin = -CK_FP16_TO_FP32(wb->dmin) * xb->d;
385
386 const __m128i mins8 = _mm_loadl_epi64((const __m128i *)(const void *)wb->mins);
387 const __m256i mins16 = _mm256_cvtepu8_epi16(mins8);
388 const __m256i q8sums = _mm256_loadu_si256((const __m256i *)(const void *)xb->bsums);
389 const __m128i q8s = _mm_hadd_epi16(
390 _mm256_extracti128_si256(q8sums, 0),
391 _mm256_extracti128_si256(q8sums, 1));
392 const __m128i prod = _mm_madd_epi16(
393 _mm256_castsi256_si128(mins16), q8s);
394 const __m128i hsum = _mm_hadd_epi32(_mm_hadd_epi32(prod, mzero), mzero);
395 summs += dmin * (float)_mm_extract_epi32(hsum, 0);
396
397 __m256i sumi = _mm256_setzero_si256();
398 for (int sb = 0; sb < 8; ++sb) {
399 const __m256i q5 = _mm256_loadu_si256(
400 (const __m256i *)(const void *)(wb->qs + sb * 32));
401 const __m256i q8 = _mm256_loadu_si256(
402 (const __m256i *)(const void *)(xb->qs + sb * 32));
403 __m256i p16 = _mm256_maddubs_epi16(q5, q8);
404 const __m256i scale = _mm256_set1_epi16((int16_t)wb->scales[sb]);
405 p16 = _mm256_madd_epi16(scale, p16);
406 sumi = _mm256_add_epi32(sumi, p16);
407 }
408 acc = _mm256_fmadd_ps(
409 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi), acc);
410 }
411 return ck_q5k_hsum256_ps(acc) + summs;
412}
413
414static void dot_q5_k_prepared_q8_k_m4_avx2(
415 const block_q5_K_prepared *w,
416 const block_q8_K *const x[4],
417 int rows, int nb, float out[4])
418{
419 const __m128i mzero = _mm_setzero_si128();
420 __m256 acc[4] = {
421 _mm256_setzero_ps(), _mm256_setzero_ps(),
422 _mm256_setzero_ps(), _mm256_setzero_ps(),
423 };
424 float summs[4] = {0.0f, 0.0f, 0.0f, 0.0f};
425
426 for (int b = 0; b < nb; ++b) {
427 const block_q5_K_prepared *wb = &w[b];
428 const __m128i mins8 = _mm_loadl_epi64((const __m128i *)(const void *)wb->mins);
429 const __m256i mins16 = _mm256_cvtepu8_epi16(mins8);
430 __m256i sumi[4] = {
431 _mm256_setzero_si256(), _mm256_setzero_si256(),
432 _mm256_setzero_si256(), _mm256_setzero_si256(),
433 };
434
435 for (int r = 0; r < rows; ++r) {
436 const block_q8_K *xb = &x[r][b];
437 const __m256i q8sums = _mm256_loadu_si256(
438 (const __m256i *)(const void *)xb->bsums);
439 const __m128i q8s = _mm_hadd_epi16(
440 _mm256_extracti128_si256(q8sums, 0),
441 _mm256_extracti128_si256(q8sums, 1));
442 const __m128i prod = _mm_madd_epi16(
443 _mm256_castsi256_si128(mins16), q8s);
444 const __m128i hsum = _mm_hadd_epi32(
445 _mm_hadd_epi32(prod, mzero), mzero);
446 const float dmin = -CK_FP16_TO_FP32(wb->dmin) * xb->d;
447 summs[r] += dmin * (float)_mm_extract_epi32(hsum, 0);
448 }
449
450 for (int sb = 0; sb < 8; ++sb) {
451 const __m256i q5 = _mm256_loadu_si256(
452 (const __m256i *)(const void *)(wb->qs + sb * 32));
453 const __m256i scale = _mm256_set1_epi16((int16_t)wb->scales[sb]);
454 for (int r = 0; r < rows; ++r) {
455 const __m256i q8 = _mm256_loadu_si256(
456 (const __m256i *)(const void *)(x[r][b].qs + sb * 32));
457 __m256i p16 = _mm256_maddubs_epi16(q5, q8);
458 p16 = _mm256_madd_epi16(scale, p16);
459 sumi[r] = _mm256_add_epi32(sumi[r], p16);
460 }
461 }
462 for (int r = 0; r < rows; ++r) {
463 const float d = CK_FP16_TO_FP32(wb->d) * x[r][b].d;
464 acc[r] = _mm256_fmadd_ps(
465 _mm256_set1_ps(d), _mm256_cvtepi32_ps(sumi[r]), acc[r]);
466 }
467 }
468 for (int r = 0; r < rows; ++r) {
469 out[r] = ck_q5k_hsum256_ps(acc[r]) + summs[r];
470 }
471}
472#endif
473
474/* Llama-compatible dot path: Q5_K weights x Q8_K activations for a full row.
475 * Keep the eight lane sums live across all blocks, matching ggml's generic
476 * Q5_K/Q8_K reduction order. Reducing each block to a scalar first is close,
477 * but can move borderline logits in long decode parity tests. */
478static float dot_q5_k_q8_k_row(const block_q5_K *w, const block_q8_K *x, int nb) {
479#if defined(__AVX2__)
481 return dot_q5_k_q8_k_row_avx2(w, x, nb);
482 }
483#endif
484
485 static const uint32_t kmask1 = 0x3f3f3f3fU;
486 static const uint32_t kmask2 = 0x0f0f0f0fU;
487 static const uint32_t kmask3 = 0x03030303U;
488
489 uint32_t utmp[4] = {0, 0, 0, 0};
490 const uint8_t *scales = (const uint8_t *)&utmp[0];
491 const uint8_t *mins = (const uint8_t *)&utmp[2];
492
493 int8_t aux8[QK_K];
494 int16_t aux16[8];
495 float sums[8] = {0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f, 0.0f};
496 int32_t aux32[8];
497
498 float sumf = 0.0f;
499 for (int b = 0; b < nb; ++b) {
500 const block_q5_K *wb = &w[b];
501 const block_q8_K *xb = &x[b];
502 const uint8_t *q4 = wb->qs;
503 const uint8_t *hm = wb->qh;
504 int8_t *a = aux8;
505 uint8_t m = 1;
506 memset(aux32, 0, sizeof(aux32));
507
508 for (int j = 0; j < QK_K / 64; ++j) {
509 for (int l = 0; l < 32; ++l) a[l] = (int8_t)(q4[l] & 0xF);
510 for (int l = 0; l < 32; ++l) a[l] += (hm[l] & m ? 16 : 0);
511 a += 32;
512 m <<= 1;
513
514 for (int l = 0; l < 32; ++l) a[l] = (int8_t)(q4[l] >> 4);
515 for (int l = 0; l < 32; ++l) a[l] += (hm[l] & m ? 16 : 0);
516 a += 32;
517 m <<= 1;
518
519 q4 += 32;
520 }
521
522 memcpy(utmp, wb->scales, 12);
523 utmp[3] = ((utmp[2] >> 4) & kmask2) | (((utmp[1] >> 6) & kmask3) << 4);
524 const uint32_t uaux = utmp[1] & kmask1;
525 utmp[1] = (utmp[2] & kmask2) | (((utmp[0] >> 6) & kmask3) << 4);
526 utmp[2] = uaux;
527 utmp[0] &= kmask1;
528
529 int sumi = 0;
530 for (int j = 0; j < QK_K / 16; ++j) {
531 sumi += (int)xb->bsums[j] * (int)mins[j / 2];
532 }
533
534 a = aux8;
535 const int8_t *q8 = xb->qs;
536 int is = 0;
537 for (int j = 0; j < QK_K / 32; ++j) {
538 const int32_t scale = (int32_t)scales[is++];
539
540 for (int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
541 for (int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
542 q8 += 8; a += 8;
543
544 for (int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
545 for (int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
546 q8 += 8; a += 8;
547
548 for (int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
549 for (int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
550 q8 += 8; a += 8;
551
552 for (int l = 0; l < 8; ++l) aux16[l] = (int16_t)(q8[l] * a[l]);
553 for (int l = 0; l < 8; ++l) aux32[l] += scale * aux16[l];
554 q8 += 8; a += 8;
555 }
556
557 const float d = CK_FP16_TO_FP32(wb->d) * xb->d;
558 for (int l = 0; l < 8; ++l) {
559 sums[l] += d * (float)aux32[l];
560 }
561 const float dmin = CK_FP16_TO_FP32(wb->dmin) * xb->d;
562 sumf -= dmin * (float)sumi;
563 }
564
565 for (int l = 0; l < 8; ++l) {
566 sumf += sums[l];
567 }
568 return sumf;
569}
570
571/* FP32 fallback for oversized K (very rare for current models). */
572static void gemv_q5_k_ref_fp32(float *y, const void *W, const float *x, int M, int K)
573{
574 const block_q5_K *blocks = (const block_q5_K *)W;
575 const int blocks_per_row = K / QK_K;
576
577 for (int m = 0; m < M; m++) {
578 const float *x_row = x;
579 float sum = 0.0f;
580
581 for (int b = 0; b < blocks_per_row; b++) {
582 const block_q5_K *block = &blocks[m * blocks_per_row + b];
583 const float d = CK_FP16_TO_FP32(block->d);
584 const float dmin = CK_FP16_TO_FP32(block->dmin);
585 uint8_t sc_arr[8], m_arr[8];
586 unpack_q5_k_scales(block->scales, sc_arr, m_arr);
587
588 for (int sb = 0; sb < 8; sb++) {
589 const float d_sub = d * (float)sc_arr[sb];
590 const float m_sub = dmin * (float)m_arr[sb];
591
592 for (int i = 0; i < 32; i++) {
593 const uint8_t q = q5_k_quant_value(block, sb, i);
594 sum += (d_sub * (float)q - m_sub) * x_row[b * QK_K + sb * 32 + i];
595 }
596 }
597 }
598
599 y[m] = sum;
600 }
601}
602
603static void gemm_nt_q5_k_ref_fp32(const float *A,
604 const void *B,
605 const float *bias,
606 float *C,
607 int M, int N, int K)
608{
609 const block_q5_K *blocks = (const block_q5_K *)B;
610 const int blocks_per_col = K / QK_K;
611
612 for (int m = 0; m < M; m++) {
613 const float *a_row = &A[m * K];
614
615 for (int n = 0; n < N; n++) {
616 float sum = 0.0f;
617 const block_q5_K *w_row = &blocks[n * blocks_per_col];
618 for (int b = 0; b < blocks_per_col; b++) {
619 const block_q5_K *block = &w_row[b];
620 const float d = CK_FP16_TO_FP32(block->d);
621 const float dmin = CK_FP16_TO_FP32(block->dmin);
622 uint8_t sc_arr[8], m_arr[8];
623 unpack_q5_k_scales(block->scales, sc_arr, m_arr);
624
625 for (int sb = 0; sb < 8; sb++) {
626 const float d_sub = d * (float)sc_arr[sb];
627 const float m_sub = dmin * (float)m_arr[sb];
628
629 for (int i = 0; i < 32; i++) {
630 const uint8_t q = q5_k_quant_value(block, sb, i);
631 sum += (d_sub * (float)q - m_sub) * a_row[b * QK_K + sb * 32 + i];
632 }
633 }
634 }
635
636 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
637 }
638 }
639}
640
641/* ============================================================================
642 * Q5_K x Q8_K Kernels (explicit contract)
643 *
644 * WHY THESE EXIST:
645 * llama.cpp's Q5_K matmul contract is "Q5_K weights x Q8_K activations".
646 * The activation quantization is part of the numerical contract, not just
647 * an optimization. If we accidentally do FP32 activation dot here, we can
648 * get large parity drift at attn_proj/mlp_down while tests still pass if
649 * they compare against FP32-dequant references.
650 *
651 * These entry points make the contract explicit in code:
652 * - gemv_q5_k_q8_k(): decode-style matrix-vector (single token)
653 * - gemm_nt_q5_k_q8_k(): prefill-style matrix-matrix
654 * ============================================================================ */
655
656void gemv_q5_k_q8_k_ref(float *y,
657 const void *W,
658 const void *x_q8,
659 int M, int K)
660{
661 if (!y || !W || !x_q8 || M <= 0 || K <= 0) {
662 return;
663 }
664 if (K % QK_K != 0) {
665 return;
666 }
667
668 const block_q5_K *blocks = (const block_q5_K *)W;
669 const block_q8_K *x = (const block_q8_K *)x_q8;
670 const int blocks_per_row = K / QK_K;
671
672 for (int m = 0; m < M; ++m) {
673 const block_q5_K *w_row = &blocks[m * blocks_per_row];
674 y[m] = dot_q5_k_q8_k_row(w_row, x, blocks_per_row);
675 }
676}
677
678void gemm_nt_q5_k_q8_k_ref(const void *A_q8,
679 const void *B,
680 const float *bias,
681 float *C,
682 int M, int N, int K)
683{
684 if (!A_q8 || !B || !C || M <= 0 || N <= 0 || K <= 0) {
685 return;
686 }
687 if (K % QK_K != 0) {
688 return;
689 }
690
691 const block_q8_K *A = (const block_q8_K *)A_q8;
692 const block_q5_K *W = (const block_q5_K *)B;
693 const int blocks_per_row = K / QK_K;
694
695 for (int m = 0; m < M; ++m) {
696 const block_q8_K *a_row = &A[m * blocks_per_row];
697 for (int n = 0; n < N; ++n) {
698 const block_q5_K *w_row = &W[n * blocks_per_row];
699 const float sum = dot_q5_k_q8_k_row(w_row, a_row, blocks_per_row);
700 C[m * N + n] = sum + (bias ? bias[n] : 0.0f);
701 }
702 }
703}
704
705void gemm_nt_q5_k_prepared(const float *A,
706 const void *B_prepared,
707 const float *bias,
708 float *C,
709 int M, int N, int K)
710{
711#if !defined(__AVX2__)
712 (void)A; (void)B_prepared; (void)bias; (void)C;
713 (void)M; (void)N; (void)K;
714#else
715 if (!A || !B_prepared || !C || M <= 0 || N <= 0 || K <= 0 ||
716 (K % QK_K) != 0) return;
717 const int blocks_per_row = K / QK_K;
718 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) return;
719 const block_q5_K_prepared *W = (const block_q5_K_prepared *)B_prepared;
720 for (int m = 0; m < M; ++m) {
721 block_q8_K a_q8[blocks_per_row];
722 quantize_row_q8_k(A + (size_t)m * K, a_q8, K);
723 for (int n = 0; n < N; ++n) {
724 const float sum = dot_q5_k_prepared_q8_k_row_avx2(
725 W + (size_t)n * blocks_per_row, a_q8, blocks_per_row);
726 C[(size_t)m * N + n] = sum + (bias ? bias[n] : 0.0f);
727 }
728 }
729#endif
730}
731
732void gemm_nt_q5_k_prepared_m4(const float *A,
733 const void *B_prepared,
734 const float *bias,
735 float *C,
736 int M, int N, int K)
737{
738#if !defined(__AVX2__)
739 (void)A; (void)B_prepared; (void)bias; (void)C;
740 (void)M; (void)N; (void)K;
741#else
742 if (!A || !B_prepared || !C || M <= 0 || N <= 0 || K <= 0 ||
743 (K % QK_K) != 0) return;
744 const int blocks_per_row = K / QK_K;
745 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) return;
746 const block_q5_K_prepared *W = (const block_q5_K_prepared *)B_prepared;
747 for (int m = 0; m < M; m += 4) {
748 const int rows = M - m < 4 ? M - m : 4;
749 block_q8_K a_q8[4][blocks_per_row];
750 const block_q8_K *row_ptrs[4] = {
751 a_q8[0], a_q8[1], a_q8[2], a_q8[3],
752 };
753 for (int r = 0; r < rows; ++r) {
754 quantize_row_q8_k(A + (size_t)(m + r) * K, a_q8[r], K);
755 }
756 for (int n = 0; n < N; ++n) {
757 float sums[4];
758 dot_q5_k_prepared_q8_k_m4_avx2(
759 W + (size_t)n * blocks_per_row,
760 row_ptrs, rows, blocks_per_row, sums);
761 for (int r = 0; r < rows; ++r) {
762 C[(size_t)(m + r) * N + n] = sums[r] + (bias ? bias[n] : 0.0f);
763 }
764 }
765 }
766#endif
767}
768
770 const void *B_prepared,
771 const float *bias,
772 float *C,
773 int M, int N, int K,
774 int n_begin, int n_end)
775{
776#if !defined(__AVX2__)
777 (void)A_q8; (void)B_prepared; (void)bias; (void)C;
778 (void)M; (void)N; (void)K; (void)n_begin; (void)n_end;
779#else
780 if (!A_q8 || !B_prepared || !C || M <= 0 || N <= 0 || K <= 0 ||
781 (K % QK_K) != 0 || n_begin < 0 || n_end > N ||
782 n_begin >= n_end) return;
783 const int blocks_per_row = K / QK_K;
784 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) return;
785 const block_q8_K *A = (const block_q8_K *)A_q8;
786 const block_q5_K_prepared *W =
787 (const block_q5_K_prepared *)B_prepared;
788
789 for (int m = 0; m < M; m += 4) {
790 const int rows = M - m < 4 ? M - m : 4;
791 const block_q8_K *row_ptrs[4] = {
792 A + (size_t)(m + 0) * blocks_per_row,
793 A + (size_t)(m + (rows > 1 ? 1 : 0)) * blocks_per_row,
794 A + (size_t)(m + (rows > 2 ? 2 : 0)) * blocks_per_row,
795 A + (size_t)(m + (rows > 3 ? 3 : 0)) * blocks_per_row,
796 };
797 for (int n = n_begin; n < n_end; ++n) {
798 float sums[4];
799 dot_q5_k_prepared_q8_k_m4_avx2(
800 W + (size_t)n * blocks_per_row,
801 row_ptrs, rows, blocks_per_row, sums);
802 for (int r = 0; r < rows; ++r) {
803 C[(size_t)(m + r) * N + n] =
804 sums[r] + (bias ? bias[n] : 0.0f);
805 }
806 }
807 }
808#endif
809}
810
811/* ============================================================================
812 * FP32 adapter path (keeps existing call sites stable)
813 *
814 * Existing generated code and orchestration call gemv_q5_k/gemm_nt_q5_k with
815 * FP32 activations. These adapter functions quantize activations to Q8_K and
816 * then call the explicit Q5_K x Q8_K kernels above.
817 * ============================================================================ */
818
819void gemv_q5_k_ref(float *y, const void *W, const float *x, int M, int K)
820{
821 if (!y || !W || !x || M <= 0 || K <= 0) {
822 return;
823 }
825 gemv_q5_k_ref_fp32(y, W, x, M, K);
826 return;
827 }
828 if (K % QK_K != 0) {
829 gemv_q5_k_ref_fp32(y, W, x, M, K);
830 return;
831 }
832
833 const block_q5_K *blocks = (const block_q5_K *)W;
834 const int blocks_per_row = K / QK_K;
835 if (blocks_per_row > CK_Q5K_STACK_Q8_BLOCKS) {
836 gemv_q5_k_ref_fp32(y, W, x, M, K);
837 return;
838 }
839
841 /* Q8_K bytes are part of the numerical ABI. Use the shared provider,
842 * whose FP-contraction policy is validated against llama.cpp. */
843 quantize_row_q8_k(x, x_q8, K);
844 gemv_q5_k_q8_k_ref(y, blocks, x_q8, M, K);
845}
846
847/* ============================================================================
848 * GEMM NT Reference: C = A @ B^T + bias
849 * - A: FP32 activation matrix [M, K] (quantized internally to Q8_K per row)
850 * - B: Q5_K weight matrix [N, K] (stored transposed, accessed as [N, K])
851 * - bias: Optional FP32 bias [N]
852 * - C: FP32 output matrix [M, N]
853 * ============================================================================ */
854
855void gemm_nt_q5_k_ref(const float *A,
856 const void *B,
857 const float *bias,
858 float *C,
859 int M, int N, int K)
860{
861 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0) {
862 return;
863 }
865 gemm_nt_q5_k_ref_fp32(A, B, bias, C, M, N, K);
866 return;
867 }
868 if (K % QK_K != 0) {
869 gemm_nt_q5_k_ref_fp32(A, B, bias, C, M, N, K);
870 return;
871 }
872
873 const block_q5_K *blocks = (const block_q5_K *)B;
874 const int blocks_per_col = K / QK_K;
875 if (blocks_per_col > CK_Q5K_STACK_Q8_BLOCKS) {
876 gemm_nt_q5_k_ref_fp32(A, B, bias, C, M, N, K);
877 return;
878 }
879
880 for (int m = 0; m < M; ++m) {
881 const float *a_row = &A[m * K];
883 quantize_row_q8_k(a_row, a_q8, K);
884 gemm_nt_q5_k_q8_k_ref(a_q8, blocks, bias, &C[m * N], 1, N, K);
885 }
886}
887
888/* ============================================================================
889 * Dispatch wrappers - select best available implementation
890 * ============================================================================ */
891
892void gemv_q5_k_q8_k(float *y,
893 const void *W,
894 const void *x_q8,
895 int M, int K)
896{
897#if defined(__AVX512F__)
898 /* TODO: AVX-512 implementation */
899 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
900#elif defined(__AVX2__)
901 /* TODO: AVX-2 implementation */
902 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
903#elif defined(__AVX__)
904 /* TODO: AVX implementation */
905 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
906#elif defined(__SSE4_1__)
907 /* TODO: SSE4.1 implementation */
908 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
909#else
910 gemv_q5_k_q8_k_ref(y, W, x_q8, M, K);
911#endif
912}
913
914void gemm_nt_q5_k_q8_k(const void *A_q8,
915 const void *B,
916 const float *bias,
917 float *C,
918 int M, int N, int K)
919{
920#if defined(__AVX512F__)
921 /* TODO: AVX-512 implementation */
922 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
923#elif defined(__AVX2__)
924 /* TODO: AVX-2 implementation */
925 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
926#elif defined(__AVX__)
927 /* TODO: AVX implementation */
928 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
929#elif defined(__SSE4_1__)
930 /* TODO: SSE4.1 implementation */
931 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
932#else
933 gemm_nt_q5_k_q8_k_ref(A_q8, B, bias, C, M, N, K);
934#endif
935}
936
938 int output_stride,
939 const void *weights,
940 const void *const input_rows[4],
941 int rows,
942 int output_dim,
943 int input_dim)
944{
945 if (!output || !weights || !input_rows || rows <= 0 || rows > 4 ||
946 output_stride < output_dim || output_dim <= 0 || input_dim <= 0 ||
947 (input_dim % QK_K) != 0) {
948 return;
949 }
950 for (int row = 0; row < rows; ++row) {
951 if (!input_rows[row]) return;
952 }
953
954#if defined(__AVX2__)
955 const block_q5_K *blocks = (const block_q5_K *)weights;
956 const int blocks_per_row = input_dim / QK_K;
957 const block_q8_K *inputs[4] = {
958 (const block_q8_K *)input_rows[0],
959 (const block_q8_K *)input_rows[rows > 1 ? 1 : 0],
960 (const block_q8_K *)input_rows[rows > 2 ? 2 : 0],
961 (const block_q8_K *)input_rows[rows > 3 ? 3 : 0],
962 };
963 for (int n = 0; n < output_dim; ++n) {
964 float values[4];
965 dot_q5_k_q8_k_rows4_avx2(
966 blocks + (size_t)n * (size_t)blocks_per_row,
967 inputs, rows, blocks_per_row, values);
968 for (int row = 0; row < rows; ++row) {
969 output[(size_t)row * (size_t)output_stride + (size_t)n] =
970 values[row];
971 }
972 }
973#else
974 for (int row = 0; row < rows; ++row) {
976 output + (size_t)row * (size_t)output_stride,
977 weights, input_rows[row], output_dim, input_dim);
978 }
979#endif
980}
981
982void gemv_q5_k(float *y, const void *W, const float *x, int M, int K)
983{
984#if defined(__AVX512F__)
985 /* TODO: AVX-512 implementation */
986 gemv_q5_k_ref(y, W, x, M, K);
987#elif defined(__AVX2__)
988 /* TODO: AVX-2 implementation */
989 gemv_q5_k_ref(y, W, x, M, K);
990#elif defined(__AVX__)
991 /* TODO: AVX implementation */
992 gemv_q5_k_ref(y, W, x, M, K);
993#elif defined(__SSE4_1__)
994 /* TODO: SSE4.1 implementation */
995 gemv_q5_k_ref(y, W, x, M, K);
996#else
997 gemv_q5_k_ref(y, W, x, M, K);
998#endif
999}
1000
1001void gemm_nt_q5_k(const float *A,
1002 const void *B,
1003 const float *bias,
1004 float *C,
1005 int M, int N, int K)
1006{
1007#if defined(__AVX512F__)
1008 /* TODO: AVX-512 implementation */
1009 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1010#elif defined(__AVX2__)
1011 /* TODO: AVX-2 implementation */
1012 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1013#elif defined(__AVX__)
1014 /* TODO: AVX implementation */
1015 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1016#elif defined(__SSE4_1__)
1017 /* TODO: SSE4.1 implementation */
1018 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1019#else
1020 gemm_nt_q5_k_ref(A, B, bias, C, M, N, K);
1021#endif
1022}
Quantization block structures for weight-only quantization.
#define K_SCALE_SIZE
uint16_t ck_half
#define CK_FP16_TO_FP32(x)
static void gemm_nt_q5_k_ref_fp32(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
static int ck_q5k_debug_generic_dot(void)
void gemv_q5_k_ref(float *y, const void *W, const float *x, int M, int K)
void ck_q5_k_prepare_weight(const void *src, void *dst, int N, int K)
void gemm_q5_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)
static uint8_t q5_k_quant_value(const block_q5_K *block, int subblock, int i)
void gemv_q5_k_q8_k(float *y, const void *W, const void *x_q8, int M, int K)
void gemm_nt_q5_k_prepared_m4(const float *A, const void *B_prepared, const float *bias, float *C, int M, int N, int K)
void quantize_row_q8_k(const float *x, void *vy, int k)
void gemv_q5_k(float *y, const void *W, const float *x, int M, int K)
static int ck_q5k_debug_fp32_fallback(void)
#define CK_Q5K_STACK_Q8_BLOCKS
void gemm_nt_q5_k_ref(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_k_q8_k(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_k_q8_k_ref(const void *A_q8, const void *B, const float *bias, float *C, int M, int N, int K)
static float dot_q5_k_q8_k_row(const block_q5_K *w, const block_q8_K *x, int nb)
static void gemv_q5_k_ref_fp32(float *y, const void *W, const float *x, int M, int K)
size_t ck_q5_k_prepared_block_size(void)
static void unpack_q5_k_scales(const uint8_t *scales, uint8_t *sc, uint8_t *m)
void gemm_nt_q5_k_prepared(const float *A, const void *B_prepared, const float *bias, float *C, int M, int N, int K)
void gemm_nt_q5_k_prepared_q8_m4_nrange(const void *A_q8, const void *B_prepared, const float *bias, float *C, int M, int N, int K, int n_begin, int n_end)
void gemm_nt_q5_k(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemv_q5_k_q8_k_ref(float *y, const void *W, const void *x_q8, int M, int K)
#define QK_K
#define C(color)
Definition show_config.c:39
int8_t qs[256]
int16_t bsums[256/16]