← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels_f16.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels_f16.c
3 * @brief GEMM kernels with FP16 (half-precision) 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 * - Weights: FP16 (IEEE half-precision, used by vision encoders)
16 * - Activations: FP32
17 * - Output: FP32
18 *
19 * Used for multimodal projection layers (mmproj-*.gguf files).
20 */
21
22#ifndef _GNU_SOURCE
23#define _GNU_SOURCE
24#endif
25
26#include <stdint.h>
27#include <stddef.h>
28#include <stdbool.h>
29#include "ckernel_quant.h" /* For ck_fp16_to_fp32 */
30#include "ckernel_engine.h"
31#include "ck_threadpool.h"
32#include "ggml_runtime_compat.h"
33
34#include <dlfcn.h>
35#include <stdlib.h>
36#include <string.h>
37
38#if defined(__AVX512F__) || defined(__AVX__) || defined(__F16C__)
39#include <immintrin.h>
40#endif
41
42typedef struct ggml_context *(*ck_f16_ggml_init_fn)(struct ggml_init_params);
43typedef void (*ck_f16_ggml_free_fn)(struct ggml_context *);
44typedef struct ggml_tensor *(*ck_f16_ggml_new_tensor_2d_fn)(struct ggml_context *, enum ggml_type, int64_t, int64_t);
45typedef struct ggml_tensor *(*ck_f16_ggml_mul_mat_fn)(struct ggml_context *, struct ggml_tensor *, struct ggml_tensor *);
46typedef struct ggml_cgraph *(*ck_f16_ggml_new_graph_fn)(struct ggml_context *);
47typedef void (*ck_f16_ggml_build_forward_expand_fn)(struct ggml_cgraph *, struct ggml_tensor *);
48typedef enum ggml_status (*ck_f16_ggml_graph_compute_with_ctx_fn)(struct ggml_context *, struct ggml_cgraph *, int);
49typedef void (*ck_f16_ggml_cpu_init_fn)(void);
50typedef void *(*ck_f16_ggml_get_data_fn)(const struct ggml_tensor *);
51typedef float *(*ck_f16_ggml_get_data_f32_fn)(const struct ggml_tensor *);
52
54{
55 static int tried = 0;
56 static ck_f16_ggml_init_fn fn = NULL;
57 if (!tried) {
58 tried = 1;
59 fn = (ck_f16_ggml_init_fn) dlsym(RTLD_DEFAULT, "ggml_init");
60 }
61 return fn;
62}
63
65{
66 static int tried = 0;
67 static ck_f16_ggml_free_fn fn = NULL;
68 if (!tried) {
69 tried = 1;
70 fn = (ck_f16_ggml_free_fn) dlsym(RTLD_DEFAULT, "ggml_free");
71 }
72 return fn;
73}
74
76{
77 static int tried = 0;
78 static ck_f16_ggml_new_tensor_2d_fn fn = NULL;
79 if (!tried) {
80 tried = 1;
81 fn = (ck_f16_ggml_new_tensor_2d_fn) dlsym(RTLD_DEFAULT, "ggml_new_tensor_2d");
82 }
83 return fn;
84}
85
87{
88 static int tried = 0;
89 static ck_f16_ggml_mul_mat_fn fn = NULL;
90 if (!tried) {
91 tried = 1;
92 fn = (ck_f16_ggml_mul_mat_fn) dlsym(RTLD_DEFAULT, "ggml_mul_mat");
93 }
94 return fn;
95}
96
98{
99 static int tried = 0;
100 static ck_f16_ggml_new_graph_fn fn = NULL;
101 if (!tried) {
102 tried = 1;
103 fn = (ck_f16_ggml_new_graph_fn) dlsym(RTLD_DEFAULT, "ggml_new_graph");
104 }
105 return fn;
106}
107
109{
110 static int tried = 0;
112 if (!tried) {
113 tried = 1;
114 fn = (ck_f16_ggml_build_forward_expand_fn) dlsym(RTLD_DEFAULT, "ggml_build_forward_expand");
115 }
116 return fn;
117}
118
120{
121 static int tried = 0;
123 if (!tried) {
124 tried = 1;
125 fn = (ck_f16_ggml_graph_compute_with_ctx_fn) dlsym(RTLD_DEFAULT, "ggml_graph_compute_with_ctx");
126 }
127 return fn;
128}
129
131{
132 static int tried = 0;
133 static ck_f16_ggml_cpu_init_fn fn = NULL;
134 if (!tried) {
135 tried = 1;
136 fn = (ck_f16_ggml_cpu_init_fn) dlsym(RTLD_DEFAULT, "ggml_cpu_init");
137 }
138 return fn;
139}
140
142{
143 static int tried = 0;
144 static ck_f16_ggml_get_data_fn fn = NULL;
145 if (!tried) {
146 tried = 1;
147 fn = (ck_f16_ggml_get_data_fn) dlsym(RTLD_DEFAULT, "ggml_get_data");
148 }
149 return fn;
150}
151
153{
154 static int tried = 0;
155 static ck_f16_ggml_get_data_f32_fn fn = NULL;
156 if (!tried) {
157 tried = 1;
158 fn = (ck_f16_ggml_get_data_f32_fn) dlsym(RTLD_DEFAULT, "ggml_get_data_f32");
159 }
160 return fn;
161}
162
163static int gemm_nt_f16_ggml_strict(const float *A,
164 const void *B,
165 const float *bias,
166 float *C,
167 int M,
168 int N,
169 int K)
170{
181
182 if (!ggml_cpu_init_fn || !ggml_init_fn || !ggml_free_fn || !ggml_new_tensor_2d_fn ||
183 !ggml_mul_mat_fn || !ggml_new_graph_fn || !ggml_build_forward_expand_fn ||
184 !ggml_graph_compute_with_ctx_fn || !ggml_get_data_fn || !ggml_get_data_f32_fn) {
185 return 0;
186 }
187
188 ggml_cpu_init_fn();
189
190 const size_t output_bytes = (size_t) M * (size_t) N * sizeof(float);
191 const size_t mem_size = ((size_t) 128 * 1024 * 1024) + output_bytes;
192
193 struct ggml_init_params params = {
195 .mem_buffer = NULL,
196 .no_alloc = false,
197 };
198 struct ggml_context *ctx = ggml_init_fn(params);
199 if (!ctx) {
200 return 0;
201 }
202
203 int ok = 0;
204 struct ggml_tensor *w = ggml_new_tensor_2d_fn(ctx, GGML_TYPE_F16, K, N);
205 struct ggml_tensor *x = ggml_new_tensor_2d_fn(ctx, GGML_TYPE_F32, K, M);
206 if (!w || !x) {
207 ggml_free_fn(ctx);
208 return 0;
209 }
210
211 {
212 void *w_data = ggml_get_data_fn(w);
213 void *x_data = ggml_get_data_fn(x);
214 if (!w_data || !x_data) {
215 ggml_free_fn(ctx);
216 return 0;
217 }
218 memcpy(w_data, B, (size_t) K * (size_t) N * sizeof(uint16_t));
219 memcpy(x_data, A, (size_t) K * (size_t) M * sizeof(float));
220 }
221
222 struct ggml_tensor *y = ggml_mul_mat_fn(ctx, w, x);
223 if (!y) {
224 ggml_free_fn(ctx);
225 return 0;
226 }
227
228 struct ggml_cgraph *gf = ggml_new_graph_fn(ctx);
229 if (!gf) {
230 ggml_free_fn(ctx);
231 return 0;
232 }
233 ggml_build_forward_expand_fn(gf, y);
234 if (ggml_graph_compute_with_ctx_fn(ctx, gf, 1) != GGML_STATUS_SUCCESS) {
235 ggml_free_fn(ctx);
236 return 0;
237 }
238
239 {
240 const float *src = ggml_get_data_f32_fn(y);
241 if (!src) {
242 ggml_free_fn(ctx);
243 return 0;
244 }
245 for (int m = 0; m < M; ++m) {
246 memcpy(C + (size_t) m * (size_t) N,
247 src + (size_t) m * (size_t) N,
248 (size_t) N * sizeof(float));
249 if (bias) {
250 for (int n = 0; n < N; ++n) {
251 C[(size_t) m * (size_t) N + (size_t) n] += bias[n];
252 }
253 }
254 }
255 }
256
257 ok = 1;
258 ggml_free_fn(ctx);
259 return ok;
260}
261
262int ck_gemm_nt_f16_ggml_oracle(const float *A,
263 const void *B,
264 const float *bias,
265 float *C,
266 int M,
267 int N,
268 int K)
269{
270 return gemm_nt_f16_ggml_strict(A, B, bias, C, M, N, K);
271}
272
273/* ============================================================================
274 * FP16 Conversion Utilities (if not using F16C)
275 * ============================================================================ */
276
277#ifndef __F16C__
278/* Software FP16 to FP32 conversion (already in ggml_quants.h) */
279#define fp16_to_fp32(x) ggml_fp16_to_fp32(x)
280#define fp32_to_fp16(x) ggml_fp32_to_fp16(x)
281#else
282/* Hardware F16C support */
283#include <immintrin.h>
284static inline float fp16_to_fp32(uint16_t h) {
285 return _cvtsh_ss(h);
286}
287static inline uint16_t fp32_to_fp16(float f) {
288 return _cvtss_sh(f, 0);
289}
290#endif
291
292static inline void ck_f32_to_f16_row_local(uint16_t *dst, const float *src, int n)
293{
294#ifdef __AVX512F__
295 int i = 0;
296 const int n16 = (n / 16) * 16;
297 for (; i < n16; i += 16) {
298 const __m512 v = _mm512_loadu_ps(src + i);
299 const __m256i h = _mm512_cvtps_ph(v, 0);
300 _mm256_storeu_si256((__m256i *)(dst + i), h);
301 }
302 for (; i < n; ++i) {
303 dst[i] = fp32_to_fp16(src[i]);
304 }
305#elif defined(__F16C__) && defined(__AVX__)
306 int i = 0;
307 const int n8 = (n / 8) * 8;
308 for (; i < n8; i += 8) {
309 const __m256 v = _mm256_loadu_ps(src + i);
310 const __m128i h = _mm256_cvtps_ph(v, 0);
311 _mm_storeu_si128((__m128i *)(dst + i), h);
312 }
313 for (; i < n; ++i) {
314 dst[i] = fp32_to_fp16(src[i]);
315 }
316#else
317 for (int i = 0; i < n; ++i) {
318 dst[i] = fp32_to_fp16(src[i]);
319 }
320#endif
321}
322
323#if defined(__AVX512F__) && defined(__F16C__)
324static inline float ck_dot_f16_f16_avx512(const uint16_t *w,
325 const uint16_t *x,
326 int k)
327{
328 int i = 0;
329 const int k16 = (k / 16) * 16;
330 __m512 acc = _mm512_setzero_ps();
331 for (; i < k16; i += 16) {
332 const __m256i wh = _mm256_loadu_si256((const __m256i *)(w + i));
333 const __m256i xh = _mm256_loadu_si256((const __m256i *)(x + i));
334 const __m512 wf = _mm512_cvtph_ps(wh);
335 const __m512 xf = _mm512_cvtph_ps(xh);
336#ifdef __FMA__
337 acc = _mm512_fmadd_ps(wf, xf, acc);
338#else
339 acc = _mm512_add_ps(acc, _mm512_mul_ps(wf, xf));
340#endif
341 }
342 float sum = _mm512_reduce_add_ps(acc);
343 for (; i < k; ++i) {
344 sum += fp16_to_fp32(w[i]) * fp16_to_fp32(x[i]);
345 }
346 return sum;
347}
348
349static inline void ck_dot_f16_f16_avx512_4(const uint16_t *w,
350 const uint16_t *x,
351 int k,
352 float sums[4])
353{
354 int i = 0;
355 const int k16 = (k / 16) * 16;
356 __m512 acc0 = _mm512_setzero_ps();
357 __m512 acc1 = _mm512_setzero_ps();
358 __m512 acc2 = _mm512_setzero_ps();
359 __m512 acc3 = _mm512_setzero_ps();
360 const uint16_t *w1 = w + k;
361 const uint16_t *w2 = w1 + k;
362 const uint16_t *w3 = w2 + k;
363
364 for (; i < k16; i += 16) {
365 const __m512 xf = _mm512_cvtph_ps(
366 _mm256_loadu_si256((const __m256i *)(x + i)));
367 const __m512 wf0 = _mm512_cvtph_ps(
368 _mm256_loadu_si256((const __m256i *)(w + i)));
369 const __m512 wf1 = _mm512_cvtph_ps(
370 _mm256_loadu_si256((const __m256i *)(w1 + i)));
371 const __m512 wf2 = _mm512_cvtph_ps(
372 _mm256_loadu_si256((const __m256i *)(w2 + i)));
373 const __m512 wf3 = _mm512_cvtph_ps(
374 _mm256_loadu_si256((const __m256i *)(w3 + i)));
375#ifdef __FMA__
376 acc0 = _mm512_fmadd_ps(wf0, xf, acc0);
377 acc1 = _mm512_fmadd_ps(wf1, xf, acc1);
378 acc2 = _mm512_fmadd_ps(wf2, xf, acc2);
379 acc3 = _mm512_fmadd_ps(wf3, xf, acc3);
380#else
381 acc0 = _mm512_add_ps(acc0, _mm512_mul_ps(wf0, xf));
382 acc1 = _mm512_add_ps(acc1, _mm512_mul_ps(wf1, xf));
383 acc2 = _mm512_add_ps(acc2, _mm512_mul_ps(wf2, xf));
384 acc3 = _mm512_add_ps(acc3, _mm512_mul_ps(wf3, xf));
385#endif
386 }
387
388 sums[0] = _mm512_reduce_add_ps(acc0);
389 sums[1] = _mm512_reduce_add_ps(acc1);
390 sums[2] = _mm512_reduce_add_ps(acc2);
391 sums[3] = _mm512_reduce_add_ps(acc3);
392 for (; i < k; ++i) {
393 const float xv = fp16_to_fp32(x[i]);
394 sums[0] += fp16_to_fp32(w[i]) * xv;
395 sums[1] += fp16_to_fp32(w1[i]) * xv;
396 sums[2] += fp16_to_fp32(w2[i]) * xv;
397 sums[3] += fp16_to_fp32(w3[i]) * xv;
398 }
399}
400#endif
401
402#if defined(__F16C__) && defined(__AVX__)
403static inline float ck_hsum256_ps(__m256 v)
404{
405 __m128 sum = _mm_add_ps(_mm256_extractf128_ps(v, 1),
406 _mm256_castps256_ps128(v));
407 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
408 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
409 return _mm_cvtss_f32(sum);
410}
411
412static inline float ck_dot_f16_f16_avx(const uint16_t *w, const uint16_t *x, int k)
413{
414 int i = 0;
415 const int k8 = (k / 8) * 8;
416 __m256 acc = _mm256_setzero_ps();
417 for (; i < k8; i += 8) {
418 const __m128i wh = _mm_loadu_si128((const __m128i *)(w + i));
419 const __m128i xh = _mm_loadu_si128((const __m128i *)(x + i));
420 const __m256 wf = _mm256_cvtph_ps(wh);
421 const __m256 xf = _mm256_cvtph_ps(xh);
422#ifdef __FMA__
423 acc = _mm256_fmadd_ps(wf, xf, acc);
424#else
425 acc = _mm256_add_ps(acc, _mm256_mul_ps(wf, xf));
426#endif
427 }
428 float sum = ck_hsum256_ps(acc);
429 for (; i < k; ++i) {
430 sum += fp16_to_fp32(w[i]) * fp16_to_fp32(x[i]);
431 }
432 return sum;
433}
434
435static inline void ck_dot_f16_f16_avx4(const uint16_t *w,
436 const uint16_t *x,
437 int k,
438 float sums[4])
439{
440 int i = 0;
441 const int k8 = (k / 8) * 8;
442 __m256 acc0 = _mm256_setzero_ps();
443 __m256 acc1 = _mm256_setzero_ps();
444 __m256 acc2 = _mm256_setzero_ps();
445 __m256 acc3 = _mm256_setzero_ps();
446 const uint16_t *w1 = w + k;
447 const uint16_t *w2 = w1 + k;
448 const uint16_t *w3 = w2 + k;
449
450 for (; i < k8; i += 8) {
451 const __m128i xh = _mm_loadu_si128((const __m128i *)(x + i));
452 const __m256 xf = _mm256_cvtph_ps(xh);
453 const __m256 wf0 = _mm256_cvtph_ps(
454 _mm_loadu_si128((const __m128i *)(w + i)));
455 const __m256 wf1 = _mm256_cvtph_ps(
456 _mm_loadu_si128((const __m128i *)(w1 + i)));
457 const __m256 wf2 = _mm256_cvtph_ps(
458 _mm_loadu_si128((const __m128i *)(w2 + i)));
459 const __m256 wf3 = _mm256_cvtph_ps(
460 _mm_loadu_si128((const __m128i *)(w3 + i)));
461#ifdef __FMA__
462 acc0 = _mm256_fmadd_ps(wf0, xf, acc0);
463 acc1 = _mm256_fmadd_ps(wf1, xf, acc1);
464 acc2 = _mm256_fmadd_ps(wf2, xf, acc2);
465 acc3 = _mm256_fmadd_ps(wf3, xf, acc3);
466#else
467 acc0 = _mm256_add_ps(acc0, _mm256_mul_ps(wf0, xf));
468 acc1 = _mm256_add_ps(acc1, _mm256_mul_ps(wf1, xf));
469 acc2 = _mm256_add_ps(acc2, _mm256_mul_ps(wf2, xf));
470 acc3 = _mm256_add_ps(acc3, _mm256_mul_ps(wf3, xf));
471#endif
472 }
473
474 sums[0] = ck_hsum256_ps(acc0);
475 sums[1] = ck_hsum256_ps(acc1);
476 sums[2] = ck_hsum256_ps(acc2);
477 sums[3] = ck_hsum256_ps(acc3);
478 for (; i < k; ++i) {
479 const float xv = fp16_to_fp32(x[i]);
480 sums[0] += fp16_to_fp32(w[i]) * xv;
481 sums[1] += fp16_to_fp32(w1[i]) * xv;
482 sums[2] += fp16_to_fp32(w2[i]) * xv;
483 sums[3] += fp16_to_fp32(w3[i]) * xv;
484 }
485}
486
487static inline void ck_dot_f16_f16_avx_m4n2(const uint16_t *w0,
488 const uint16_t *w1,
489 const uint16_t *x0,
490 const uint16_t *x1,
491 const uint16_t *x2,
492 const uint16_t *x3,
493 int k,
494 float sums[4][2])
495{
496 int i = 0;
497 const int k8 = (k / 8) * 8;
498 __m256 acc00 = _mm256_setzero_ps();
499 __m256 acc01 = _mm256_setzero_ps();
500 __m256 acc10 = _mm256_setzero_ps();
501 __m256 acc11 = _mm256_setzero_ps();
502 __m256 acc20 = _mm256_setzero_ps();
503 __m256 acc21 = _mm256_setzero_ps();
504 __m256 acc30 = _mm256_setzero_ps();
505 __m256 acc31 = _mm256_setzero_ps();
506
507 for (; i < k8; i += 8) {
508 const __m256 wf0 = _mm256_cvtph_ps(
509 _mm_loadu_si128((const __m128i *)(w0 + i)));
510 const __m256 wf1 = _mm256_cvtph_ps(
511 _mm_loadu_si128((const __m128i *)(w1 + i)));
512 const __m256 xf0 = _mm256_cvtph_ps(
513 _mm_loadu_si128((const __m128i *)(x0 + i)));
514 const __m256 xf1 = _mm256_cvtph_ps(
515 _mm_loadu_si128((const __m128i *)(x1 + i)));
516 const __m256 xf2 = _mm256_cvtph_ps(
517 _mm_loadu_si128((const __m128i *)(x2 + i)));
518 const __m256 xf3 = _mm256_cvtph_ps(
519 _mm_loadu_si128((const __m128i *)(x3 + i)));
520#ifdef __FMA__
521 acc00 = _mm256_fmadd_ps(wf0, xf0, acc00);
522 acc01 = _mm256_fmadd_ps(wf1, xf0, acc01);
523 acc10 = _mm256_fmadd_ps(wf0, xf1, acc10);
524 acc11 = _mm256_fmadd_ps(wf1, xf1, acc11);
525 acc20 = _mm256_fmadd_ps(wf0, xf2, acc20);
526 acc21 = _mm256_fmadd_ps(wf1, xf2, acc21);
527 acc30 = _mm256_fmadd_ps(wf0, xf3, acc30);
528 acc31 = _mm256_fmadd_ps(wf1, xf3, acc31);
529#else
530 acc00 = _mm256_add_ps(acc00, _mm256_mul_ps(wf0, xf0));
531 acc01 = _mm256_add_ps(acc01, _mm256_mul_ps(wf1, xf0));
532 acc10 = _mm256_add_ps(acc10, _mm256_mul_ps(wf0, xf1));
533 acc11 = _mm256_add_ps(acc11, _mm256_mul_ps(wf1, xf1));
534 acc20 = _mm256_add_ps(acc20, _mm256_mul_ps(wf0, xf2));
535 acc21 = _mm256_add_ps(acc21, _mm256_mul_ps(wf1, xf2));
536 acc30 = _mm256_add_ps(acc30, _mm256_mul_ps(wf0, xf3));
537 acc31 = _mm256_add_ps(acc31, _mm256_mul_ps(wf1, xf3));
538#endif
539 }
540
541 sums[0][0] = ck_hsum256_ps(acc00);
542 sums[0][1] = ck_hsum256_ps(acc01);
543 sums[1][0] = ck_hsum256_ps(acc10);
544 sums[1][1] = ck_hsum256_ps(acc11);
545 sums[2][0] = ck_hsum256_ps(acc20);
546 sums[2][1] = ck_hsum256_ps(acc21);
547 sums[3][0] = ck_hsum256_ps(acc30);
548 sums[3][1] = ck_hsum256_ps(acc31);
549 for (; i < k; ++i) {
550 const float wv0 = fp16_to_fp32(w0[i]);
551 const float wv1 = fp16_to_fp32(w1[i]);
552 const float xv0 = fp16_to_fp32(x0[i]);
553 const float xv1 = fp16_to_fp32(x1[i]);
554 const float xv2 = fp16_to_fp32(x2[i]);
555 const float xv3 = fp16_to_fp32(x3[i]);
556 sums[0][0] += wv0 * xv0;
557 sums[0][1] += wv1 * xv0;
558 sums[1][0] += wv0 * xv1;
559 sums[1][1] += wv1 * xv1;
560 sums[2][0] += wv0 * xv2;
561 sums[2][1] += wv1 * xv2;
562 sums[3][0] += wv0 * xv3;
563 sums[3][1] += wv1 * xv3;
564 }
565}
566#endif
567
568static inline float ck_dot_f16_f16_local(const uint16_t *w, const uint16_t *x, int k)
569{
570#if defined(__AVX512F__) && defined(__F16C__)
571 return ck_dot_f16_f16_avx512(w, x, k);
572#elif defined(__F16C__) && defined(__AVX__)
573 return ck_dot_f16_f16_avx(w, x, k);
574#else
575 float sum = 0.0f;
576 for (int i = 0; i < k; ++i) {
577 sum += fp16_to_fp32(w[i]) * fp16_to_fp32(x[i]);
578 }
579 return sum;
580#endif
581}
582
584{
585#if defined(__AVX512F__) && defined(__F16C__)
586 return 16;
587#elif defined(__F16C__) && defined(__AVX__)
588 return 8;
589#else
590 return 1;
591#endif
592}
593
594/* ============================================================================
595 * GEMV: y = W @ x (W is FP16, x and y are FP32)
596 * ============================================================================ */
597
598/**
599 * @brief Matrix-vector multiply with FP16 weights (scalar reference)
600 *
601 * @param y Output vector [M]
602 * @param W Weight matrix in FP16 [M x K]
603 * @param x Input vector [K]
604 * @param M Number of output rows
605 * @param K Number of columns
606 */
607void gemv_f16_ref(float *y,
608 const uint16_t *W,
609 const float *x,
610 int M, int K)
611{
612 for (int row = 0; row < M; row++) {
613 float sum = 0.0f;
614 const uint16_t *w_row = &W[row * K];
615
616 for (int k = 0; k < K; k++) {
617 float w = fp16_to_fp32(w_row[k]);
618 sum += w * x[k];
619 }
620
621 y[row] = sum;
622 }
623}
624
625#ifdef __AVX512F__
626/**
627 * @brief Matrix-vector multiply with FP16 weights (AVX-512)
628 *
629 * Converts FP16 to FP32 in registers using VCVTPH2PS.
630 */
631void gemv_f16_avx512(float *y,
632 const uint16_t *W,
633 const float *x,
634 int M, int K)
635{
636 const int K16 = K / 16 * 16;
637
638 for (int row = 0; row < M; row++) {
639 __m512 acc = _mm512_setzero_ps();
640 const uint16_t *w_row = &W[row * K];
641
642 /* Process 16 elements at a time */
643 for (int k = 0; k < K16; k += 16) {
644 /* Load 16 x FP16 weights */
645 __m256i w_f16 = _mm256_loadu_si256((const __m256i *)&w_row[k]);
646
647 /* Convert FP16 to FP32 */
648 __m512 w_f32 = _mm512_cvtph_ps(w_f16);
649
650 /* Load 16 x FP32 inputs */
651 __m512 x_vec = _mm512_loadu_ps(&x[k]);
652
653 /* FMA */
654 acc = _mm512_fmadd_ps(w_f32, x_vec, acc);
655 }
656
657 /* Horizontal sum */
658 float sum = _mm512_reduce_add_ps(acc);
659
660 /* Handle remainder */
661 for (int k = K16; k < K; k++) {
662 sum += fp16_to_fp32(w_row[k]) * x[k];
663 }
664
665 y[row] = sum;
666 }
667}
668#endif /* __AVX512F__ */
669
670/**
671 * @brief Auto-dispatch GEMV based on available SIMD
672 */
673void gemv_f16(float *y,
674 const uint16_t *W,
675 const float *x,
676 int M, int K)
677{
678#ifdef __AVX512F__
679 gemv_f16_avx512(y, W, x, M, K);
680#else
681 gemv_f16_ref(y, W, x, M, K);
682#endif
683}
684
685/* ============================================================================
686 * GEMM: Y = W @ X (W is FP16, X and Y are FP32)
687 * ============================================================================ */
688
689/**
690 * @brief Matrix-matrix multiply with FP16 weights (scalar reference)
691 *
692 * @param Y Output matrix [M x N]
693 * @param W Weight matrix in FP16 [M x K]
694 * @param X Input matrix [K x N]
695 * @param M Number of output rows
696 * @param N Batch size
697 * @param K Hidden dimension
698 */
699void gemm_f16_ref(float *Y,
700 const uint16_t *W,
701 const float *X,
702 int M, int N, int K)
703{
704 for (int n = 0; n < N; n++) {
705 gemv_f16_ref(&Y[n * M], W, &X[n * K], M, K);
706 }
707}
708
709#ifdef __AVX512F__
710/**
711 * @brief Matrix-matrix multiply with FP16 weights (AVX-512)
712 */
713void gemm_f16_avx512(float *Y,
714 const uint16_t *W,
715 const float *X,
716 int M, int N, int K)
717{
718 const int K16 = K / 16 * 16;
719
720 for (int row = 0; row < M; row++) {
721 const uint16_t *w_row = &W[row * K];
722
723 /* Pre-convert weight row to FP32 in cache-sized chunks */
724 /* For now, convert on-the-fly per batch element */
725
726 for (int n = 0; n < N; n++) {
727 __m512 acc = _mm512_setzero_ps();
728 const float *x_col = &X[n * K];
729
730 for (int k = 0; k < K16; k += 16) {
731 __m256i w_f16 = _mm256_loadu_si256((const __m256i *)&w_row[k]);
732 __m512 w_f32 = _mm512_cvtph_ps(w_f16);
733 __m512 x_vec = _mm512_loadu_ps(&x_col[k]);
734 acc = _mm512_fmadd_ps(w_f32, x_vec, acc);
735 }
736
737 float sum = _mm512_reduce_add_ps(acc);
738
739 for (int k = K16; k < K; k++) {
740 sum += fp16_to_fp32(w_row[k]) * x_col[k];
741 }
742
743 Y[n * M + row] = sum;
744 }
745 }
746}
747#endif /* __AVX512F__ */
748
749/**
750 * @brief Auto-dispatch GEMM based on available SIMD
751 */
752void gemm_f16(float *Y,
753 const uint16_t *W,
754 const float *X,
755 int M, int N, int K)
756{
757#ifdef __AVX512F__
758 gemm_f16_avx512(Y, W, X, M, N, K);
759#else
760 gemm_f16_ref(Y, W, X, M, N, K);
761#endif
762}
763
764static void gemm_f16_input_fp16_ref(float *Y,
765 const uint16_t *W,
766 const float *X,
767 int M, int N, int K)
768{
769#pragma omp parallel for schedule(static) if(N > 1)
770 for (int n = 0; n < N; ++n) {
771 const float *x_row = &X[(size_t)n * (size_t)K];
772 uint16_t x_f16[K];
773
774 ck_f32_to_f16_row_local(x_f16, x_row, K);
775
776 int row = 0;
777#if defined(__AVX512F__) && defined(__F16C__)
778 for (; row + 3 < M; row += 4) {
779 float sums[4];
780 ck_dot_f16_f16_avx512_4(
781 &W[(size_t)row * (size_t)K], x_f16, K, sums);
782 Y[(size_t)n * (size_t)M + (size_t)row] = sums[0];
783 Y[(size_t)n * (size_t)M + (size_t)row + 1] = sums[1];
784 Y[(size_t)n * (size_t)M + (size_t)row + 2] = sums[2];
785 Y[(size_t)n * (size_t)M + (size_t)row + 3] = sums[3];
786 }
787#elif defined(__F16C__) && defined(__AVX__)
788 for (; row + 3 < M; row += 4) {
789 float sums[4];
790 ck_dot_f16_f16_avx4(&W[(size_t)row * (size_t)K], x_f16, K, sums);
791 Y[(size_t)n * (size_t)M + (size_t)row] = sums[0];
792 Y[(size_t)n * (size_t)M + (size_t)row + 1] = sums[1];
793 Y[(size_t)n * (size_t)M + (size_t)row + 2] = sums[2];
794 Y[(size_t)n * (size_t)M + (size_t)row + 3] = sums[3];
795 }
796#endif
797 for (; row < M; ++row) {
798 const uint16_t *w_row = &W[(size_t)row * (size_t)K];
799 const float sum = ck_dot_f16_f16_local(w_row, x_f16, K);
800 Y[(size_t)n * (size_t)M + (size_t)row] = sum;
801 }
802 }
803}
804
805typedef struct {
806 float *Y;
807 const uint16_t *W;
808 const float *X;
809 int M;
810 int N;
811 int K;
812} ck_gemm_f16_input_fp16_args_t;
813
814#if defined(__F16C__) && defined(__AVX__) && !defined(__AVX512F__)
815static int ck_gemm_f16_m4n2_enabled(void);
816#endif
817
818static void ck_gemm_f16_input_fp16_work(int ith, int nth, void *opaque)
819{
820 ck_gemm_f16_input_fp16_args_t *args = (ck_gemm_f16_input_fp16_args_t *) opaque;
821 const int M = args->M;
822 const int N = args->N;
823 const int K = args->K;
824
825 const int token_groups = (N + 3) / 4;
826 for (int group = ith; group < token_groups; group += nth) {
827 const int n0 = group * 4;
828 const int token_count = N - n0 < 4 ? N - n0 : 4;
829#if defined(__F16C__) && defined(__AVX__) && !defined(__AVX512F__)
830 if (token_count == 4 && ck_gemm_f16_m4n2_enabled()) {
831 uint16_t x_f16[4][K];
832 for (int t = 0; t < 4; ++t) {
834 x_f16[t], args->X + (size_t)(n0 + t) * (size_t)K, K);
835 }
836
837 int row = 0;
838 for (; row + 1 < M; row += 2) {
839 float sums[4][2];
840 const uint16_t *w0 = args->W + (size_t)row * (size_t)K;
841 ck_dot_f16_f16_avx_m4n2(
842 w0, w0 + K,
843 x_f16[0], x_f16[1], x_f16[2], x_f16[3], K, sums);
844 for (int t = 0; t < 4; ++t) {
845 float *out = args->Y + (size_t)(n0 + t) * (size_t)M + (size_t)row;
846 out[0] = sums[t][0];
847 out[1] = sums[t][1];
848 }
849 }
850 for (; row < M; ++row) {
851 const uint16_t *w = args->W + (size_t)row * (size_t)K;
852 for (int t = 0; t < 4; ++t) {
853 args->Y[(size_t)(n0 + t) * (size_t)M + (size_t)row] =
854 ck_dot_f16_f16_local(w, x_f16[t], K);
855 }
856 }
857 continue;
858 }
859#endif
860 for (int n = n0; n < n0 + token_count; ++n) {
861 const float *x_row = args->X + (size_t)n * (size_t)K;
862 uint16_t x_f16[K];
863
864 ck_f32_to_f16_row_local(x_f16, x_row, K);
865
866 int row = 0;
867#if defined(__AVX512F__) && defined(__F16C__)
868 for (; row + 3 < M; row += 4) {
869 float sums[4];
870 ck_dot_f16_f16_avx512_4(
871 args->W + (size_t)row * (size_t)K, x_f16, K, sums);
872 args->Y[(size_t)n * (size_t)M + (size_t)row] = sums[0];
873 args->Y[(size_t)n * (size_t)M + (size_t)row + 1] = sums[1];
874 args->Y[(size_t)n * (size_t)M + (size_t)row + 2] = sums[2];
875 args->Y[(size_t)n * (size_t)M + (size_t)row + 3] = sums[3];
876 }
877#elif defined(__F16C__) && defined(__AVX__)
878 for (; row + 3 < M; row += 4) {
879 float sums[4];
880 ck_dot_f16_f16_avx4(
881 args->W + (size_t)row * (size_t)K, x_f16, K, sums);
882 args->Y[(size_t)n * (size_t)M + (size_t)row] = sums[0];
883 args->Y[(size_t)n * (size_t)M + (size_t)row + 1] = sums[1];
884 args->Y[(size_t)n * (size_t)M + (size_t)row + 2] = sums[2];
885 args->Y[(size_t)n * (size_t)M + (size_t)row + 3] = sums[3];
886 }
887#endif
888 for (; row < M; ++row) {
889 const uint16_t *w_row = args->W + (size_t)row * (size_t)K;
890 const float sum = ck_dot_f16_f16_local(w_row, x_f16, K);
891 args->Y[(size_t)n * (size_t)M + (size_t)row] = sum;
892 }
893 }
894 }
895}
896
897static int ck_gemm_f16_threadpool_enabled(int M, int N, int K)
898{
899 const char *disable = getenv("CK_DISABLE_F16_GEMM_THREADPOOL");
900 if (disable && disable[0] && strcmp(disable, "0") != 0) return 0;
901 if (M < 256 || N < 16 || K < 256) return 0;
902 return 1;
903}
904
905#if defined(__F16C__) && defined(__AVX__) && !defined(__AVX512F__)
906static int ck_gemm_f16_m4n2_enabled(void)
907{
908 const char *disable = getenv("CK_DISABLE_F16_GEMM_M4N2");
909 return !(disable && disable[0] && strcmp(disable, "0") != 0);
910}
911#endif
912
913static int ck_gemm_f16_pick_active_threads(const ck_threadpool_t *pool, int M, int N, int K)
914{
915 int nth = pool ? ck_threadpool_n_threads(pool) : 1;
916 if (nth <= 1) return 1;
917
918 const char *cap_env = getenv("CK_F16_GEMM_THREAD_CAP");
919 int cap = cap_env && cap_env[0] ? atoi(cap_env) : 24;
920 if (cap < 1) cap = 1;
921 if (cap > nth) cap = nth;
922
923 int active = N;
924 if (M >= 1024 && K >= 1024 && active < 8) active = 8;
925 if (active > cap) active = cap;
926 if (active > nth) active = nth;
927 return active < 1 ? 1 : active;
928}
929
931 const uint16_t *W,
932 const float *X,
933 int M, int N, int K)
934{
935 if (!ck_gemm_f16_threadpool_enabled(M, N, K)) {
936 return 0;
937 }
938
939 ck_threadpool_t *pool = ck_threadpool_global();
940 const int active = ck_gemm_f16_pick_active_threads(pool, M, N, K);
941 if (!pool || active <= 1) {
942 return 0;
943 }
944
945 ck_gemm_f16_input_fp16_args_t args = {
946 .Y = Y,
947 .W = W,
948 .X = X,
949 .M = M,
950 .N = N,
951 .K = K,
952 };
954 return 1;
955}
956
957/**
958 * @brief NT GEMM wrapper for FP16 weights with the engine's standard ABI.
959 *
960 * Contract:
961 * A: [M, K] fp32 activation matrix
962 * B: [N, K] fp16 weight matrix stored row-major (transposed layout)
963 * C: [M, N] fp32 output matrix
964 *
965 * This wrapper follows llama.cpp's CPU F16 mul_mat contract: activation rows
966 * are rounded to FP16 first, then the dot runs as F16 x F16 with FP32 output
967 * accumulation. The lower-level gemm_f16() helper remains the direct F16-weight
968 * x FP32-activation operator for generic use.
969 */
970void gemm_nt_f16(const float *A,
971 const void *B,
972 const float *bias,
973 float *C,
974 int M, int N, int K)
975{
977 gemm_nt_f16_ggml_strict(A, B, bias, C, M, N, K)) {
978 return;
979 }
980
981 if (!gemm_f16_input_fp16_threadpool(C, (const uint16_t *)B, A, N, M, K)) {
982 gemm_f16_input_fp16_ref(C, (const uint16_t *)B, A, N, M, K);
983 }
984
985 if (!bias) {
986 return;
987 }
988
989#pragma omp parallel for schedule(static) if(M > 1)
990 for (int i = 0; i < M; ++i) {
991 float *c_row = C + (size_t)i * (size_t)N;
992 for (int j = 0; j < N; ++j) {
993 c_row[j] += bias[j];
994 }
995 }
996}
997
998void gemm_nt_f16_clipped(const float *A,
999 const void *B,
1000 const float *bias,
1001 const float *input_min,
1002 const float *input_max,
1003 const float *output_min,
1004 const float *output_max,
1005 float *C,
1006 int M, int N, int K)
1007{
1008 const float in_min = input_min ? input_min[0] : -3.4028234663852886e38f;
1009 const float in_max = input_max ? input_max[0] : 3.4028234663852886e38f;
1010 const float out_min = output_min ? output_min[0] : -3.4028234663852886e38f;
1011 const float out_max = output_max ? output_max[0] : 3.4028234663852886e38f;
1012 const uint16_t *W = (const uint16_t *)B;
1013
1014#pragma omp parallel for schedule(static) if(M > 1)
1015 for (int m = 0; m < M; ++m) {
1016 const float *a_row = A + (size_t)m * (size_t)K;
1017 uint16_t a_f16[K];
1018
1019 for (int k = 0; k < K; ++k) {
1020 float x = a_row[k];
1021 if (x < in_min) x = in_min;
1022 if (x > in_max) x = in_max;
1023 a_f16[k] = fp32_to_fp16(x);
1024 }
1025
1026 float *c_row = C + (size_t)m * (size_t)N;
1027 for (int n = 0; n < N; ++n) {
1028 const uint16_t *w_row = W + (size_t)n * (size_t)K;
1029 float sum = bias ? bias[n] : 0.0f;
1030 for (int k = 0; k < K; ++k) {
1031 sum += fp16_to_fp32(w_row[k]) * fp16_to_fp32(a_f16[k]);
1032 }
1033 if (sum < out_min) sum = out_min;
1034 if (sum > out_max) sum = out_max;
1035 c_row[n] = sum;
1036 }
1037 }
1038}
1039
1040/* ============================================================================
1041 * FP16 Tensor Conversion Utilities
1042 * ============================================================================ */
1043
1044/**
1045 * @brief Convert FP16 tensor to FP32
1046 */
1047void convert_f16_to_f32(float *dst, const uint16_t *src, size_t count)
1048{
1049#ifdef __AVX512F__
1050 const size_t count16 = count / 16 * 16;
1051
1052 for (size_t i = 0; i < count16; i += 16) {
1053 __m256i f16 = _mm256_loadu_si256((const __m256i *)&src[i]);
1054 __m512 f32 = _mm512_cvtph_ps(f16);
1055 _mm512_storeu_ps(&dst[i], f32);
1056 }
1057
1058 for (size_t i = count16; i < count; i++) {
1059 dst[i] = fp16_to_fp32(src[i]);
1060 }
1061#else
1062 for (size_t i = 0; i < count; i++) {
1063 dst[i] = fp16_to_fp32(src[i]);
1064 }
1065#endif
1066}
1067
1068/**
1069 * @brief Convert FP32 tensor to FP16
1070 */
1071void convert_f32_to_f16(uint16_t *dst, const float *src, size_t count)
1072{
1073#ifdef __AVX512F__
1074 const size_t count16 = count / 16 * 16;
1075
1076 for (size_t i = 0; i < count16; i += 16) {
1077 __m512 f32 = _mm512_loadu_ps(&src[i]);
1078 __m256i f16 = _mm512_cvtps_ph(f32, 0);
1079 _mm256_storeu_si256((__m256i *)&dst[i], f16);
1080 }
1081
1082 for (size_t i = count16; i < count; i++) {
1083 dst[i] = fp32_to_fp16(src[i]);
1084 }
1085#else
1086 for (size_t i = 0; i < count; i++) {
1087 dst[i] = fp32_to_fp16(src[i]);
1088 }
1089#endif
1090}
1091
1092/* ============================================================================
1093 * Backward Pass: Gradient w.r.t. Input
1094 *
1095 * Given: dL/dY (gradient of loss w.r.t. output)
1096 * Compute: dL/dX = W^T @ dL/dY
1097 *
1098 * For F16 weights, we convert to FP32 on-the-fly during backprop.
1099 * ============================================================================ */
1100
1101/**
1102 * @brief Backward pass: compute input gradient (scalar reference)
1103 *
1104 * @param dX Output gradient w.r.t. input [K]
1105 * @param W Weight matrix in FP16 format [M x K]
1106 * @param dY Gradient w.r.t. output [M]
1107 * @param M Number of output rows
1108 * @param K Number of columns (input dimension)
1109 */
1111 const uint16_t *W,
1112 const float *dY,
1113 int M, int K)
1114{
1115 /* Zero output gradient */
1116 for (int k = 0; k < K; k++) {
1117 dX[k] = 0.0f;
1118 }
1119
1120 /* Accumulate: dX += W^T @ dY */
1121 for (int row = 0; row < M; row++) {
1122 const float dy = dY[row];
1123 const uint16_t *w_row = &W[row * K];
1124
1125 for (int k = 0; k < K; k++) {
1126 float w = fp16_to_fp32(w_row[k]);
1127 dX[k] += w * dy;
1128 }
1129 }
1130}
1131
1132#ifdef __AVX512F__
1133/**
1134 * @brief Backward pass with AVX-512
1135 */
1136void gemv_f16_backward_avx512(float *dX,
1137 const uint16_t *W,
1138 const float *dY,
1139 int M, int K)
1140{
1141 const int K16 = K / 16 * 16;
1142
1143 /* Zero output */
1144 for (int k = 0; k < K16; k += 16) {
1145 _mm512_storeu_ps(&dX[k], _mm512_setzero_ps());
1146 }
1147 for (int k = K16; k < K; k++) {
1148 dX[k] = 0.0f;
1149 }
1150
1151 for (int row = 0; row < M; row++) {
1152 const __m512 vdy = _mm512_set1_ps(dY[row]);
1153 const uint16_t *w_row = &W[row * K];
1154
1155 for (int k = 0; k < K16; k += 16) {
1156 /* Load and convert F16 weights */
1157 __m256i w_f16 = _mm256_loadu_si256((const __m256i *)&w_row[k]);
1158 __m512 w_f32 = _mm512_cvtph_ps(w_f16);
1159
1160 /* Compute gradient */
1161 __m512 grad = _mm512_mul_ps(w_f32, vdy);
1162
1163 /* Accumulate */
1164 __m512 dx_cur = _mm512_loadu_ps(&dX[k]);
1165 _mm512_storeu_ps(&dX[k], _mm512_add_ps(dx_cur, grad));
1166 }
1167
1168 /* Remainder */
1169 for (int k = K16; k < K; k++) {
1170 dX[k] += fp16_to_fp32(w_row[k]) * dY[row];
1171 }
1172 }
1173}
1174#endif
1175
1176/**
1177 * @brief Auto-dispatch backward
1178 */
1179void gemv_f16_backward(float *dX,
1180 const uint16_t *W,
1181 const float *dY,
1182 int M, int K)
1183{
1184#ifdef __AVX512F__
1185 gemv_f16_backward_avx512(dX, W, dY, M, K);
1186#else
1187 gemv_f16_backward_ref(dX, W, dY, M, K);
1188#endif
1189}
1190
1191/**
1192 * @brief Batched backward pass
1193 */
1194void gemm_f16_backward(float *dX,
1195 const uint16_t *W,
1196 const float *dY,
1197 int M, int N, int K)
1198{
1199 for (int n = 0; n < N; n++) {
1200 gemv_f16_backward(&dX[n * K], W, &dY[n * M], M, K);
1201 }
1202}
1203
1204/* ============================================================================
1205 * Dot Product Utility
1206 * ============================================================================ */
1207
1208float dot_f16(const uint16_t *w_f16, const float *x, int K)
1209{
1210 float result;
1211 gemv_f16(&result, w_f16, x, 1, K);
1212 return result;
1213}
#define RTLD_DEFAULT
Persistent pthread thread pool for CK-Engine inference.
void ck_threadpool_dispatch_n(ck_threadpool_t *pool, int active_threads, ck_work_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
int ck_strict_parity_enabled(void)
Quantization block structures for weight-only quantization.
static ck_f16_ggml_init_fn ck_f16_resolve_ggml_init(void)
static ck_f16_ggml_graph_compute_with_ctx_fn ck_f16_resolve_ggml_graph_compute_with_ctx(void)
void gemm_f16_ref(float *Y, const uint16_t *W, const float *X, int M, int N, int K)
Matrix-matrix multiply with FP16 weights (scalar reference)
struct ggml_cgraph *(* ck_f16_ggml_new_graph_fn)(struct ggml_context *)
int ck_gemm_nt_f16_ggml_oracle(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_f16_backward(float *dX, const uint16_t *W, const float *dY, int M, int N, int K)
Batched backward pass.
static float ck_dot_f16_f16_local(const uint16_t *w, const uint16_t *x, int k)
void gemm_f16(float *Y, const uint16_t *W, const float *X, int M, int N, int K)
Auto-dispatch GEMM based on available SIMD.
struct ggml_tensor *(* ck_f16_ggml_mul_mat_fn)(struct ggml_context *, struct ggml_tensor *, struct ggml_tensor *)
static void gemm_f16_input_fp16_ref(float *Y, const uint16_t *W, const float *X, int M, int N, int K)
void gemm_nt_f16(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
NT GEMM wrapper for FP16 weights with the engine's standard ABI.
void gemv_f16(float *y, const uint16_t *W, const float *x, int M, int K)
Auto-dispatch GEMV based on available SIMD.
void(* ck_f16_ggml_build_forward_expand_fn)(struct ggml_cgraph *, struct ggml_tensor *)
void convert_f16_to_f32(float *dst, const uint16_t *src, size_t count)
Convert FP16 tensor to FP32.
void *(* ck_f16_ggml_get_data_fn)(const struct ggml_tensor *)
static int ck_gemm_f16_pick_active_threads(const ck_threadpool_t *pool, int M, int N, int K)
#define fp16_to_fp32(x)
void gemv_f16_backward_ref(float *dX, const uint16_t *W, const float *dY, int M, int K)
Backward pass: compute input gradient (scalar reference)
float dot_f16(const uint16_t *w_f16, const float *x, int K)
int ck_gemm_nt_f16_simd_lanes(void)
static int gemm_nt_f16_ggml_strict(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void convert_f32_to_f16(uint16_t *dst, const float *src, size_t count)
Convert FP32 tensor to FP16.
enum ggml_status(* ck_f16_ggml_graph_compute_with_ctx_fn)(struct ggml_context *, struct ggml_cgraph *, int)
void(* ck_f16_ggml_free_fn)(struct ggml_context *)
static int gemm_f16_input_fp16_threadpool(float *Y, const uint16_t *W, const float *X, int M, int N, int K)
void(* ck_f16_ggml_cpu_init_fn)(void)
void gemm_nt_f16_clipped(const float *A, const void *B, const float *bias, const float *input_min, const float *input_max, const float *output_min, const float *output_max, float *C, int M, int N, int K)
static ck_f16_ggml_cpu_init_fn ck_f16_resolve_ggml_cpu_init(void)
static ck_f16_ggml_build_forward_expand_fn ck_f16_resolve_ggml_build_forward_expand(void)
float *(* ck_f16_ggml_get_data_f32_fn)(const struct ggml_tensor *)
static void ck_gemm_f16_input_fp16_work(int ith, int nth, void *opaque)
struct ggml_tensor *(* ck_f16_ggml_new_tensor_2d_fn)(struct ggml_context *, enum ggml_type, int64_t, int64_t)
struct ggml_context *(* ck_f16_ggml_init_fn)(struct ggml_init_params)
static ck_f16_ggml_get_data_f32_fn ck_f16_resolve_ggml_get_data_f32(void)
static ck_f16_ggml_new_tensor_2d_fn ck_f16_resolve_ggml_new_tensor_2d(void)
static ck_f16_ggml_mul_mat_fn ck_f16_resolve_ggml_mul_mat(void)
static ck_f16_ggml_get_data_fn ck_f16_resolve_ggml_get_data(void)
void gemv_f16_backward(float *dX, const uint16_t *W, const float *dY, int M, int K)
Auto-dispatch backward.
static ck_f16_ggml_new_graph_fn ck_f16_resolve_ggml_new_graph(void)
static void ck_f32_to_f16_row_local(uint16_t *dst, const float *src, int n)
void gemv_f16_ref(float *y, const uint16_t *W, const float *x, int M, int K)
Matrix-vector multiply with FP16 weights (scalar reference)
static ck_f16_ggml_free_fn ck_f16_resolve_ggml_free(void)
#define fp32_to_fp16(x)
static int ck_gemm_f16_threadpool_enabled(int M, int N, int K)
@ GGML_STATUS_SUCCESS
@ GGML_TYPE_F32
@ GGML_TYPE_F16
#define C(color)
Definition show_config.c:39