38#if defined(__AVX512F__) || defined(__AVX__) || defined(__F16C__)
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 *);
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 *);
110 static int tried = 0;
121 static int tried = 0;
132 static int tried = 0;
143 static int tried = 0;
154 static int tried = 0;
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) {
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;
198 struct ggml_context *ctx = ggml_init_fn(params);
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);
212 void *w_data = ggml_get_data_fn(w);
213 void *x_data = ggml_get_data_fn(x);
214 if (!w_data || !x_data) {
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));
222 struct ggml_tensor *y = ggml_mul_mat_fn(ctx, w, x);
228 struct ggml_cgraph *gf = ggml_new_graph_fn(ctx);
233 ggml_build_forward_expand_fn(gf, y);
240 const float *src = ggml_get_data_f32_fn(y);
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));
250 for (
int n = 0; n < N; ++n) {
251 C[(size_t) m * (
size_t) N + (size_t) n] += bias[n];
279#define fp16_to_fp32(x) ggml_fp16_to_fp32(x)
280#define fp32_to_fp16(x) ggml_fp32_to_fp16(x)
283#include <immintrin.h>
288 return _cvtss_sh(f, 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);
305#elif defined(__F16C__) && defined(__AVX__)
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);
317 for (
int i = 0; i < n; ++i) {
323#if defined(__AVX512F__) && defined(__F16C__)
324static inline float ck_dot_f16_f16_avx512(
const uint16_t *w,
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);
337 acc = _mm512_fmadd_ps(wf, xf, acc);
339 acc = _mm512_add_ps(acc, _mm512_mul_ps(wf, xf));
342 float sum = _mm512_reduce_add_ps(acc);
349static inline void ck_dot_f16_f16_avx512_4(
const uint16_t *w,
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;
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)));
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);
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));
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);
402#if defined(__F16C__) && defined(__AVX__)
403static inline float ck_hsum256_ps(__m256 v)
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);
412static inline float ck_dot_f16_f16_avx(
const uint16_t *w,
const uint16_t *x,
int k)
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);
423 acc = _mm256_fmadd_ps(wf, xf, acc);
425 acc = _mm256_add_ps(acc, _mm256_mul_ps(wf, xf));
428 float sum = ck_hsum256_ps(acc);
435static inline void ck_dot_f16_f16_avx4(
const uint16_t *w,
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;
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)));
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);
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));
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);
487static inline void ck_dot_f16_f16_avx_m4n2(
const uint16_t *w0,
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();
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)));
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);
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));
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);
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;
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);
576 for (
int i = 0; i < k; ++i) {
585#if defined(__AVX512F__) && defined(__F16C__)
587#elif defined(__F16C__) && defined(__AVX__)
612 for (
int row = 0; row < M; row++) {
614 const uint16_t *w_row = &W[row * K];
616 for (
int k = 0; k < K; k++) {
631void gemv_f16_avx512(
float *y,
636 const int K16 = K / 16 * 16;
638 for (
int row = 0; row < M; row++) {
639 __m512 acc = _mm512_setzero_ps();
640 const uint16_t *w_row = &W[row * K];
643 for (
int k = 0; k < K16; k += 16) {
645 __m256i w_f16 = _mm256_loadu_si256((
const __m256i *)&w_row[k]);
648 __m512 w_f32 = _mm512_cvtph_ps(w_f16);
651 __m512 x_vec = _mm512_loadu_ps(&x[k]);
654 acc = _mm512_fmadd_ps(w_f32, x_vec, acc);
658 float sum = _mm512_reduce_add_ps(acc);
661 for (
int k = K16; k < K; k++) {
679 gemv_f16_avx512(y, W, x, M, K);
704 for (
int n = 0; n < N; n++) {
713void gemm_f16_avx512(
float *Y,
718 const int K16 = K / 16 * 16;
720 for (
int row = 0; row < M; row++) {
721 const uint16_t *w_row = &W[row * K];
726 for (
int n = 0; n < N; n++) {
727 __m512 acc = _mm512_setzero_ps();
728 const float *x_col = &X[n * K];
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);
737 float sum = _mm512_reduce_add_ps(acc);
739 for (
int k = K16; k < K; k++) {
743 Y[n * M + row] = sum;
758 gemm_f16_avx512(Y, W, X, M, N, K);
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];
777#if defined(__AVX512F__) && defined(__F16C__)
778 for (; row + 3 < M; row += 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];
787#elif defined(__F16C__) && defined(__AVX__)
788 for (; row + 3 < M; row += 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];
797 for (; row < M; ++row) {
798 const uint16_t *w_row = &W[(size_t)row * (
size_t)K];
800 Y[(size_t)n * (
size_t)M + (size_t)row] = sum;
812} ck_gemm_f16_input_fp16_args_t;
814#if defined(__F16C__) && defined(__AVX__) && !defined(__AVX512F__)
815static int ck_gemm_f16_m4n2_enabled(
void);
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;
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);
838 for (; row + 1 < M; row += 2) {
840 const uint16_t *w0 = args->W + (size_t)row * (
size_t)K;
841 ck_dot_f16_f16_avx_m4n2(
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;
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] =
860 for (
int n = n0; n < n0 + token_count; ++n) {
861 const float *x_row = args->X + (size_t)n * (
size_t)K;
867#if defined(__AVX512F__) && defined(__F16C__)
868 for (; row + 3 < M; row += 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];
877#elif defined(__F16C__) && defined(__AVX__)
878 for (; row + 3 < M; row += 4) {
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];
888 for (; row < M; ++row) {
889 const uint16_t *w_row = args->W + (size_t)row * (
size_t)K;
891 args->Y[(size_t)n * (
size_t)M + (size_t)row] = sum;
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;
905#if defined(__F16C__) && defined(__AVX__) && !defined(__AVX512F__)
906static int ck_gemm_f16_m4n2_enabled(
void)
908 const char *disable = getenv(
"CK_DISABLE_F16_GEMM_M4N2");
909 return !(disable && disable[0] && strcmp(disable,
"0") != 0);
916 if (nth <= 1)
return 1;
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;
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;
941 if (!pool || active <= 1) {
945 ck_gemm_f16_input_fp16_args_t args = {
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) {
1001 const float *input_min,
1002 const float *input_max,
1003 const float *output_min,
1004 const float *output_max,
1006 int M,
int N,
int K)
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;
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;
1019 for (
int k = 0; k < K; ++k) {
1021 if (x < in_min) x = in_min;
1022 if (x > in_max) x = in_max;
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) {
1033 if (sum < out_min) sum = out_min;
1034 if (sum > out_max) sum = out_max;
1050 const size_t count16 = count / 16 * 16;
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);
1058 for (
size_t i = count16; i < count; i++) {
1062 for (
size_t i = 0; i < count; i++) {
1074 const size_t count16 = count / 16 * 16;
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);
1082 for (
size_t i = count16; i < count; i++) {
1086 for (
size_t i = 0; i < count; i++) {
1116 for (
int k = 0; k < K; k++) {
1121 for (
int row = 0; row < M; row++) {
1122 const float dy = dY[row];
1123 const uint16_t *w_row = &W[row * K];
1125 for (
int k = 0; k < K; k++) {
1136void gemv_f16_backward_avx512(
float *dX,
1141 const int K16 = K / 16 * 16;
1144 for (
int k = 0; k < K16; k += 16) {
1145 _mm512_storeu_ps(&dX[k], _mm512_setzero_ps());
1147 for (
int k = K16; k < K; k++) {
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];
1155 for (
int k = 0; k < K16; k += 16) {
1157 __m256i w_f16 = _mm256_loadu_si256((
const __m256i *)&w_row[k]);
1158 __m512 w_f32 = _mm512_cvtph_ps(w_f16);
1161 __m512 grad = _mm512_mul_ps(w_f32, vdy);
1164 __m512 dx_cur = _mm512_loadu_ps(&dX[k]);
1165 _mm512_storeu_ps(&dX[k], _mm512_add_ps(dx_cur, grad));
1169 for (
int k = K16; k < K; k++) {
1185 gemv_f16_backward_avx512(dX, W, dY, M, K);
1197 int M,
int N,
int K)
1199 for (
int n = 0; n < N; n++) {
1208float dot_f16(
const uint16_t *w_f16,
const float *x,
int K)
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)
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)
static int ck_gemm_f16_threadpool_enabled(int M, int N, int K)