22#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
31static inline int ck_min(
int a,
int b) {
return a < b ? a : b; }
38#pragma omp parallel for schedule(static)
39 for (
int i = 0; i < M; ++i) {
40 float *c_row =
C + (size_t)i * (
size_t)N;
41 for (
int j = 0; j < N; ++j) {
48#if defined(__AVX__) && !defined(__AVX512F__)
49static inline float hsum256_ps(__m256 v) {
51 __m128 lo = _mm256_castps256_ps128(v);
52 __m128 hi = _mm256_extractf128_ps(v, 1);
53 __m128 sum128 = _mm_add_ps(lo, hi);
55 __m128 shuf = _mm_movehdup_ps(sum128);
56 __m128 sums = _mm_add_ps(sum128, shuf);
57 shuf = _mm_movehl_ps(shuf, sums);
58 sums = _mm_add_ss(sums, shuf);
59 return _mm_cvtss_f32(sums);
73#pragma omp parallel for schedule(static)
74 for (
int j = 0; j < N; ++j) {
75 const float *b_row = B + (size_t)j * (
size_t)K;
76 float sum = bias ? bias[j] : 0.0f;
78#if defined(__AVX512F__)
79 __m512 acc = _mm512_setzero_ps();
81 for (; k <= K - 16; k += 16) {
82 __m512 a_vec = _mm512_loadu_ps(A + k);
83 __m512 b_vec = _mm512_loadu_ps(b_row + k);
84 acc = _mm512_fmadd_ps(a_vec, b_vec, acc);
86 sum += _mm512_reduce_add_ps(acc);
88 sum += A[k] * b_row[k];
91 __m256 acc = _mm256_setzero_ps();
93 for (; k <= K - 8; k += 8) {
94 __m256 a_vec = _mm256_loadu_ps(A + k);
95 __m256 b_vec = _mm256_loadu_ps(b_row + k);
96 acc = _mm256_add_ps(acc, _mm256_mul_ps(a_vec, b_vec));
98 sum += hsum256_ps(acc);
100 sum += A[k] * b_row[k];
103 for (
int k = 0; k < K; ++k) {
104 sum += A[k] * b_row[k];
118 for (
int i = 0; i < M; i++) {
119 for (
int j = 0; j < N; j++) {
120 double sum = bias ? (double)bias[j] : 0.0;
121 for (
int k = 0; k < K; k++) {
122 sum += (double)A[i * K + k] * (
double)B[j * K + k];
124 C[i * N + j] = (float)sum;
135 for (
int i = 0; i < M; i++) {
136 for (
int j = 0; j < N; j++) {
137 float sum = bias ? bias[j] : 0.0f;
138 for (
int k = 0; k < K; k++) {
139 sum += A[i * K + k] * B[j * K + k];
154 int bias_before_reduction;
155} ck_gemm_nt_fp32_exact_args_t;
159 const ck_gemm_nt_fp32_exact_args_t *args =
160 (
const ck_gemm_nt_fp32_exact_args_t *)opaque;
161 for (
int i = begin; i <
end; ++i) {
162 const float *a_row = args->A + (size_t)i * (
size_t)args->K;
163 float *c_row = args->C + (size_t)i * (
size_t)args->N;
164 for (
int j = 0; j < args->N; ++j) {
165 const float *b_row = args->B + (size_t)j * (
size_t)args->K;
166 float sum = args->bias_before_reduction && args->bias
167 ? args->bias[j] : 0.0f;
168 for (
int k = 0; k < args->K; ++k) {
169 sum += a_row[k] * b_row[k];
171 if (!args->bias_before_reduction && args->bias) {
172 sum += args->bias[j];
185 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0)
return;
186 ck_gemm_nt_fp32_exact_args_t args = {
187 .A = A, .B = B, .bias = bias, .C =
C,
188 .M = M, .N = N, .K = K,
193 (
size_t)M * (
size_t)N <= 4096) {
198 if (active > M) active = M;
199 int grain = M / (active * 4);
200 if (grain < 1) grain = 1;
216#pragma omp parallel for
217 for (
int i = 0; i < M; i++) {
218 for (
int j = 0; j < N; j++) {
220 for (
int k = 0; k < K; k++) {
221 sum += A[i * K + k] * B[j * K + k];
223 float bias_val = bias ? bias[j] : 0.0f;
224 C[i * N + j] = sum + bias_val;
240#if defined(__AVX512F__)
241#pragma omp parallel for
242 for (
int i = 0; i < M; i++) {
243 for (
int j = 0; j < N; j++) {
244 __m512 sum_vec = _mm512_setzero_ps();
246 for (k = 0; k <= K - 16; k += 16) {
247 __m512 a_vec = _mm512_loadu_ps(&A[i * K + k]);
248 __m512 b_vec = _mm512_loadu_ps(&B[j * K + k]);
249 sum_vec = _mm512_fmadd_ps(a_vec, b_vec, sum_vec);
251 float sum = _mm512_reduce_add_ps(sum_vec);
253 sum += A[i * K + k] * B[j * K + k];
255 float bias_val = bias ? bias[j] : 0.0f;
256 C[i * N + j] = sum + bias_val;
259#elif defined(__AVX__)
261#pragma omp parallel for
262 for (
int i = 0; i < M; i++) {
263 for (
int j = 0; j < N; j++) {
264 __m256 sum_vec = _mm256_setzero_ps();
266 for (k = 0; k <= K - 8; k += 8) {
267 __m256 a_vec = _mm256_loadu_ps(&A[i * K + k]);
268 __m256 b_vec = _mm256_loadu_ps(&B[j * K + k]);
269 __m256 prod = _mm256_mul_ps(a_vec, b_vec);
270 sum_vec = _mm256_add_ps(sum_vec, prod);
272 float sum = hsum256_ps(sum_vec);
274 sum += A[i * K + k] * B[j * K + k];
276 float bias_val = bias ? bias[j] : 0.0f;
277 C[i * N + j] = sum + bias_val;
296#if defined(__AVX512F__)
297 const int block_size = 64;
298#pragma omp parallel for
299 for (
int i = 0; i < M; i++) {
300 for (
int j = 0; j < N; j++) {
301 C[i * N + j] = bias ? bias[j] : 0.0f;
304#pragma omp parallel for collapse(3)
305 for (
int ii = 0; ii < M; ii += block_size) {
306 for (
int jj = 0; jj < N; jj += block_size) {
307 for (
int kk = 0; kk < K; kk += block_size) {
308 int i_end =
ck_min(ii + block_size, M);
309 int j_end =
ck_min(jj + block_size, N);
310 int k_end =
ck_min(kk + block_size, K);
312 for (
int i = ii; i < i_end; i++) {
313 for (
int j = jj; j < j_end; j++) {
314 __m512 sum_vec = _mm512_setzero_ps();
316 for (k = kk; k <= k_end - 16; k += 16) {
317 __m512 a_vec = _mm512_loadu_ps(&A[i * K + k]);
318 __m512 b_vec = _mm512_loadu_ps(&B[j * K + k]);
319 sum_vec = _mm512_fmadd_ps(a_vec, b_vec, sum_vec);
321 float partial_sum = _mm512_reduce_add_ps(sum_vec);
322 for (; k < k_end; k++) {
323 partial_sum += A[i * K + k] * B[j * K + k];
326 C[i * N + j] += partial_sum;
332#elif defined(__AVX__)
334 const int block_size = 32;
335#pragma omp parallel for
336 for (
int i = 0; i < M; i++) {
337 for (
int j = 0; j < N; j++) {
338 C[i * N + j] = bias ? bias[j] : 0.0f;
341#pragma omp parallel for collapse(3)
342 for (
int ii = 0; ii < M; ii += block_size) {
343 for (
int jj = 0; jj < N; jj += block_size) {
344 for (
int kk = 0; kk < K; kk += block_size) {
345 int i_end =
ck_min(ii + block_size, M);
346 int j_end =
ck_min(jj + block_size, N);
347 int k_end =
ck_min(kk + block_size, K);
349 for (
int i = ii; i < i_end; i++) {
350 for (
int j = jj; j < j_end; j++) {
351 __m256 sum_vec = _mm256_setzero_ps();
353 for (k = kk; k <= k_end - 8; k += 8) {
354 __m256 a_vec = _mm256_loadu_ps(&A[i * K + k]);
355 __m256 b_vec = _mm256_loadu_ps(&B[j * K + k]);
356 __m256 prod = _mm256_mul_ps(a_vec, b_vec);
357 sum_vec = _mm256_add_ps(sum_vec, prod);
359 float partial_sum = hsum256_ps(sum_vec);
360 for (; k < k_end; k++) {
361 partial_sum += A[i * K + k] * B[j * K + k];
364 C[i * N + j] += partial_sum;
387 for (
int i = 0; i < M; i++) {
388 for (
int j = 0; j < N; j++) {
389 double sum = bias ? (double)bias[j] : 0.0;
390 for (
int k = 0; k < K; k++) {
391 sum += (double)A[i * K + k] * (
double)B[k * N + j];
393 C[i * N + j] = (float)sum;
408#pragma omp parallel for
409 for (
int i = 0; i < M; i++) {
410 for (
int j = 0; j < N; j++) {
411 float sum = bias ? bias[j] : 0.0f;
412 for (
int k = 0; k < K; k++) {
413 sum += A[i * K + k] * B[k * N + j];
430#if defined(__AVX512F__)
433#pragma omp parallel for
434 for (
int i = 0; i < M; i++) {
437 for (; j <= N - 16; j += 16) {
438 __m512 sum_vec = bias ? _mm512_loadu_ps(&bias[j]) : _mm512_setzero_ps();
439 for (
int k = 0; k < K; k++) {
440 __m512 a_broadcast = _mm512_set1_ps(A[i * K + k]);
441 __m512 b_vec = _mm512_loadu_ps(&B[k * N + j]);
442 sum_vec = _mm512_fmadd_ps(a_broadcast, b_vec, sum_vec);
444 _mm512_storeu_ps(&
C[i * N + j], sum_vec);
448 float sum = bias ? bias[j] : 0.0f;
449 for (
int k = 0; k < K; k++) {
450 sum += A[i * K + k] * B[k * N + j];
455#elif defined(__AVX__)
457#pragma omp parallel for
458 for (
int i = 0; i < M; i++) {
460 for (; j <= N - 8; j += 8) {
461 __m256 sum_vec = bias ? _mm256_loadu_ps(&bias[j]) : _mm256_setzero_ps();
462 for (
int k = 0; k < K; k++) {
463 __m256 a_broadcast = _mm256_set1_ps(A[i * K + k]);
464 __m256 b_vec = _mm256_loadu_ps(&B[k * N + j]);
465 __m256 prod = _mm256_mul_ps(a_broadcast, b_vec);
466 sum_vec = _mm256_add_ps(sum_vec, prod);
468 _mm256_storeu_ps(&
C[i * N + j], sum_vec);
471 float sum = bias ? bias[j] : 0.0f;
472 for (
int k = 0; k < K; k++) {
473 sum += A[i * K + k] * B[k * N + j];
488#if defined(__GNUC__) && !defined(__INTEL_LLVM_COMPILER)
490#elif defined(__GNUC__)
503#if defined(__AVX512F__)
506#pragma omp parallel for
507 for (
int i = 0; i < M; i++) {
510 for (; j <= N - 16; j += 16) {
511 __m512 sum_vec = bias ? _mm512_loadu_ps(&bias[j]) : _mm512_setzero_ps();
512 for (
int k = 0; k < K; k++) {
513 __m512 a_broadcast = _mm512_set1_ps(A[i * K + k]);
514 __m512 b_vec = _mm512_loadu_ps(&B[k * N + j]);
515 sum_vec = _mm512_fmadd_ps(a_broadcast, b_vec, sum_vec);
517 _mm512_storeu_ps(&
C[i * N + j], sum_vec);
521 float sum = bias ? bias[j] : 0.0f;
522 for (
int k = 0; k < K; k++) {
523 sum += A[i * K + k] * B[k * N + j];
528#elif defined(__AVX__)
530#pragma omp parallel for
531 for (
int i = 0; i < M; i++) {
533 for (; j <= N - 8; j += 8) {
534 __m256 sum_vec = bias ? _mm256_loadu_ps(&bias[j]) : _mm256_setzero_ps();
535 for (
int k = 0; k < K; k++) {
536 __m256 a_broadcast = _mm256_set1_ps(A[i * K + k]);
537 __m256 b_vec = _mm256_loadu_ps(&B[k * N + j]);
538 __m256 prod = _mm256_mul_ps(a_broadcast, b_vec);
539 sum_vec = _mm256_add_ps(sum_vec, prod);
541 _mm256_storeu_ps(&
C[i * N + j], sum_vec);
544 float sum = bias ? bias[j] : 0.0f;
545 for (
int k = 0; k < K; k++) {
546 sum += A[i * K + k] * B[k * N + j];
558 static int cached = -1;
563 const char *v = getenv(
"CK_GEMM_NN_IMPL");
567 if (strcmp(v,
"probe") == 0 || strcmp(v,
"dup") == 0 || strcmp(v,
"1") == 0) {
600#if defined(__AVX512F__)
601 const int block_size = 64;
602#elif defined(__AVX__)
603 const int block_size = 32;
605 const int block_size = 32;
608#pragma omp parallel for
609 for (
int i = 0; i < M; i++) {
610 for (
int j = 0; j < N; j++) {
611 C[i * N + j] = bias ? bias[j] : 0.0f;
615#pragma omp parallel for
616 for (
int ii = 0; ii < M; ii += block_size) {
617 for (
int kk = 0; kk < K; kk += block_size) {
618 for (
int jj = 0; jj < N; jj += block_size) {
619 int i_end =
ck_min(ii + block_size, M);
620 int k_end =
ck_min(kk + block_size, K);
621 int j_end =
ck_min(jj + block_size, N);
623 for (
int i = ii; i < i_end; i++) {
624 for (
int k = kk; k < k_end; k++) {
625 float a_val = A[i * K + k];
626#if defined(__AVX512F__)
627 __m512 a_broadcast = _mm512_set1_ps(a_val);
629 for (j = jj; j <= j_end - 16; j += 16) {
630 __m512 b_vec = _mm512_loadu_ps(&B[k * N + j]);
631 __m512 c_vec = _mm512_loadu_ps(&
C[i * N + j]);
632 c_vec = _mm512_fmadd_ps(a_broadcast, b_vec, c_vec);
633 _mm512_storeu_ps(&
C[i * N + j], c_vec);
635 for (; j < j_end; j++) {
636 C[i * N + j] += a_val * B[k * N + j];
638#elif defined(__AVX__)
639 __m256 a_broadcast = _mm256_set1_ps(a_val);
641 for (j = jj; j <= j_end - 8; j += 8) {
642 __m256 b_vec = _mm256_loadu_ps(&B[k * N + j]);
643 __m256 c_vec = _mm256_loadu_ps(&
C[i * N + j]);
644 __m256 prod = _mm256_mul_ps(a_broadcast, b_vec);
645 c_vec = _mm256_add_ps(c_vec, prod);
646 _mm256_storeu_ps(&
C[i * N + j], c_vec);
648 for (; j < j_end; j++) {
649 C[i * N + j] += a_val * B[k * N + j];
652 for (
int j = jj; j < j_end; j++) {
653 C[i * N + j] += a_val * B[k * N + j];
675 for (
int i = 0; i < M; i++) {
676 for (
int j = 0; j < N; j++) {
677 double sum = bias ? (double)bias[j] : 0.0;
678 for (
int k = 0; k < K; k++) {
680 sum += (double)A[k * M + i] * (
double)B[k * N + j];
682 C[i * N + j] = (float)sum;
697#pragma omp parallel for
698 for (
int i = 0; i < M; i++) {
699 for (
int j = 0; j < N; j++) {
700 float sum = bias ? bias[j] : 0.0f;
701 for (
int k = 0; k < K; k++) {
702 sum += A[k * M + i] * B[k * N + j];
719#if defined(__AVX512F__)
721#pragma omp parallel for
722 for (
int i = 0; i < M; i++) {
724 for (; j <= N - 16; j += 16) {
725 __m512 sum_vec = bias ? _mm512_loadu_ps(&bias[j]) : _mm512_setzero_ps();
726 for (
int k = 0; k < K; k++) {
727 __m512 a_broadcast = _mm512_set1_ps(A[k * M + i]);
728 __m512 b_vec = _mm512_loadu_ps(&B[k * N + j]);
729 sum_vec = _mm512_fmadd_ps(a_broadcast, b_vec, sum_vec);
731 _mm512_storeu_ps(&
C[i * N + j], sum_vec);
734 float sum = bias ? bias[j] : 0.0f;
735 for (
int k = 0; k < K; k++) {
736 sum += A[k * M + i] * B[k * N + j];
741#elif defined(__AVX__)
743#pragma omp parallel for
744 for (
int i = 0; i < M; i++) {
746 for (; j <= N - 8; j += 8) {
747 __m256 sum_vec = bias ? _mm256_loadu_ps(&bias[j]) : _mm256_setzero_ps();
748 for (
int k = 0; k < K; k++) {
749 __m256 a_broadcast = _mm256_set1_ps(A[k * M + i]);
750 __m256 b_vec = _mm256_loadu_ps(&B[k * N + j]);
751 __m256 prod = _mm256_mul_ps(a_broadcast, b_vec);
752 sum_vec = _mm256_add_ps(sum_vec, prod);
754 _mm256_storeu_ps(&
C[i * N + j], sum_vec);
757 float sum = bias ? bias[j] : 0.0f;
758 for (
int k = 0; k < K; k++) {
759 sum += A[k * M + i] * B[k * N + j];
779#if defined(__AVX512F__)
780 const int block_size = 64;
781#elif defined(__AVX__)
782 const int block_size = 32;
784 const int block_size = 32;
787#pragma omp parallel for
788 for (
int i = 0; i < M; i++) {
789 for (
int j = 0; j < N; j++) {
790 C[i * N + j] = bias ? bias[j] : 0.0f;
794#pragma omp parallel for
795 for (
int ii = 0; ii < M; ii += block_size) {
796 for (
int kk = 0; kk < K; kk += block_size) {
797 for (
int jj = 0; jj < N; jj += block_size) {
798 int i_end =
ck_min(ii + block_size, M);
799 int k_end =
ck_min(kk + block_size, K);
800 int j_end =
ck_min(jj + block_size, N);
802 for (
int k = kk; k < k_end; k++) {
803 for (
int i = ii; i < i_end; i++) {
804 float a_val = A[k * M + i];
805#if defined(__AVX512F__)
806 __m512 a_broadcast = _mm512_set1_ps(a_val);
808 for (j = jj; j <= j_end - 16; j += 16) {
809 __m512 b_vec = _mm512_loadu_ps(&B[k * N + j]);
810 __m512 c_vec = _mm512_loadu_ps(&
C[i * N + j]);
811 c_vec = _mm512_fmadd_ps(a_broadcast, b_vec, c_vec);
812 _mm512_storeu_ps(&
C[i * N + j], c_vec);
814 for (; j < j_end; j++) {
815 C[i * N + j] += a_val * B[k * N + j];
817#elif defined(__AVX__)
818 __m256 a_broadcast = _mm256_set1_ps(a_val);
820 for (j = jj; j <= j_end - 8; j += 8) {
821 __m256 b_vec = _mm256_loadu_ps(&B[k * N + j]);
822 __m256 c_vec = _mm256_loadu_ps(&
C[i * N + j]);
823 __m256 prod = _mm256_mul_ps(a_broadcast, b_vec);
824 c_vec = _mm256_add_ps(c_vec, prod);
825 _mm256_storeu_ps(&
C[i * N + j], c_vec);
827 for (; j < j_end; j++) {
828 C[i * N + j] += a_val * B[k * N + j];
831 for (
int j = jj; j < j_end; j++) {
832 C[i * N + j] += a_val * B[k * N + j];
866 if (M == 1 && (
size_t)N * (
size_t)K >= 65536) {
877 if (M >= 32 && N >= 32 && K >= 32) {
882#if defined(__AVX512F__)
883 const int block_size = 64;
884#elif defined(__AVX__)
885 const int block_size = 32;
887 const int block_size = 32;
889 for (
int i = 0; i < M; i++) {
890 for (
int j = 0; j < N; j++) {
891 C[i * N + j] = bias ? bias[j] : 0.0f;
894 for (
int ii = 0; ii < M; ii += block_size) {
895 for (
int jj = 0; jj < N; jj += block_size) {
896 for (
int kk = 0; kk < K; kk += block_size) {
897 int i_end =
ck_min(ii + block_size, M);
898 int j_end =
ck_min(jj + block_size, N);
899 int k_end =
ck_min(kk + block_size, K);
901 for (
int i = ii; i < i_end; i++) {
902 for (
int j = jj; j < j_end; j++) {
903#if defined(__AVX512F__)
904 __m512 sum_vec = _mm512_setzero_ps();
906 for (k = kk; k <= k_end - 16; k += 16) {
907 __m512 a_vec = _mm512_loadu_ps(&A[i * K + k]);
908 __m512 b_vec = _mm512_loadu_ps(&B[j * K + k]);
909 sum_vec = _mm512_fmadd_ps(a_vec, b_vec, sum_vec);
911 float partial_sum = _mm512_reduce_add_ps(sum_vec);
912 for (; k < k_end; k++) {
913 partial_sum += A[i * K + k] * B[j * K + k];
915#elif defined(__AVX__)
916 __m256 sum_vec = _mm256_setzero_ps();
918 for (k = kk; k <= k_end - 8; k += 8) {
919 __m256 a_vec = _mm256_loadu_ps(&A[i * K + k]);
920 __m256 b_vec = _mm256_loadu_ps(&B[j * K + k]);
921 __m256 prod = _mm256_mul_ps(a_vec, b_vec);
922 sum_vec = _mm256_add_ps(sum_vec, prod);
924 float partial_sum = hsum256_ps(sum_vec);
925 for (; k < k_end; k++) {
926 partial_sum += A[i * K + k] * B[j * K + k];
929 float partial_sum = 0.0f;
930 for (
int k = kk; k < k_end; k++) {
931 partial_sum += A[i * K + k] * B[j * K + k];
934 C[i * N + j] += partial_sum;
953#if defined(__AVX__) && !defined(__AVX512F__)
954static inline float ck_hsum256_llamafile(__m256 value)
956 __m128 sum = _mm_add_ps(
957 _mm256_extractf128_ps(value, 1),
958 _mm256_castps256_ps128(value));
959 sum = _mm_add_ps(sum, _mm_movehl_ps(sum, sum));
960 sum = _mm_add_ss(sum, _mm_movehdup_ps(sum));
961 return _mm_cvtss_f32(sum);
966 const float *A,
const float *B,
const float *bias,
float *
C,
967 int M,
int N,
int K,
int index)
969 const int row = index / N;
970 const int col = index - row * N;
971 const float *a = A + (size_t)row * (
size_t)K;
972 const float *b = B + (size_t)col * (
size_t)K;
976#if defined(__AVX512F__)
978 __m512 acc = _mm512_setzero_ps();
979 for (; k + 16 <= K; k += 16) {
980 acc = _mm512_fmadd_ps(
981 _mm512_loadu_ps(a + k), _mm512_loadu_ps(b + k), acc);
983 sum = _mm512_reduce_add_ps(acc);
986 _mm512_setzero_ps(), _mm512_setzero_ps(),
987 _mm512_setzero_ps(), _mm512_setzero_ps()
989 for (; k + 64 <= K; k += 64) {
990 for (
int lane = 0; lane < 4; ++lane) {
991 acc[lane] = _mm512_fmadd_ps(
992 _mm512_loadu_ps(a + k + lane * 16),
993 _mm512_loadu_ps(b + k + lane * 16),
997 acc[0] = _mm512_add_ps(acc[0], acc[2]);
998 acc[1] = _mm512_add_ps(acc[1], acc[3]);
999 acc[0] = _mm512_add_ps(acc[0], acc[1]);
1000 sum = _mm512_reduce_add_ps(acc[0]);
1002#elif defined(__AVX__)
1004 __m256 acc = _mm256_setzero_ps();
1005 for (; k + 8 <= K; k += 8) {
1007 acc = _mm256_fmadd_ps(
1008 _mm256_loadu_ps(a + k), _mm256_loadu_ps(b + k), acc);
1010 acc = _mm256_add_ps(
1012 _mm256_loadu_ps(a + k), _mm256_loadu_ps(b + k)));
1015 sum = ck_hsum256_llamafile(acc);
1018 _mm256_setzero_ps(), _mm256_setzero_ps(),
1019 _mm256_setzero_ps(), _mm256_setzero_ps()
1021 for (; k + 32 <= K; k += 32) {
1022 for (
int lane = 0; lane < 4; ++lane) {
1024 acc[lane] = _mm256_fmadd_ps(
1025 _mm256_loadu_ps(a + k + lane * 8),
1026 _mm256_loadu_ps(b + k + lane * 8),
1029 acc[lane] = _mm256_add_ps(
1030 acc[lane], _mm256_mul_ps(
1031 _mm256_loadu_ps(a + k + lane * 8),
1032 _mm256_loadu_ps(b + k + lane * 8)));
1036 acc[0] = _mm256_add_ps(acc[0], acc[2]);
1037 acc[1] = _mm256_add_ps(acc[1], acc[3]);
1038 acc[0] = _mm256_add_ps(acc[0], acc[1]);
1039 sum = hsum256_ps(acc[0]);
1042 for (; k < K; ++k) {
1045 C[(size_t)row * (
size_t)N + (size_t)col] =
1046 bias ? sum + bias[col] : sum;
1050 const float *A,
const float *B,
const float *bias,
float *
C,
1051 int M,
int N,
int K,
int output_begin,
int output_end)
1053 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0)
return;
1054 const int total = M * N;
1055 if (output_begin < 0) output_begin = 0;
1056 if (output_end > total) output_end = total;
1057 for (
int index = output_begin; index < output_end; ++index) {
1066 int M,
int N,
int K)
1068 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
1072#pragma omp parallel for schedule(static) if ((size_t)M * (size_t)N >= 96)
1073 for (
int index = 0; index < M * N; ++index) {
Persistent pthread thread pool for CK-Engine inference.
void ck_threadpool_parallel_for_n(ck_threadpool_t *pool, int active_threads, int begin, int end, int grain_size, ck_range_fn_t fn, void *args)
ck_threadpool_t * ck_threadpool_global(void)
int ck_threadpool_n_threads(const ck_threadpool_t *pool)
void gemm_microkernel(const float *A, const float *B, float *C, int M, int N, int K, int B_transposed)
int ck_get_num_threads(void)
int ck_strict_parity_enabled(void)
static void gemm_naive_serial_double(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_fp32_exact_parallel_dispatch(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static int ck_min(int a, int b)
void gemm_naive_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_f32_llama_production_output_range(const float *A, const float *B, const float *bias, float *C, int M, int N, int K, int output_begin, int output_end)
static int ck_gemm_nn_impl_probe_enabled(void)
static void gemm_naive_serial_float(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static void gemm_nn_serial_double(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nn_avx512_probe(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nn_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static void gemm_nt_matvec_parallel(const float *A, const float *B, const float *bias, float *C, int N, int K)
static void ck_gemm_nt_f32_llama_production_output(const float *A, const float *B, const float *bias, float *C, int M, int N, int K, int index)
void gemm_tn_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nn_avx512(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static void gemm_tn_serial_double(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_avx512_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nn_simd(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static void ck_gemm_nt_fp32_exact_rows(int begin, int end, void *opaque)
void gemm_fine_grained_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static void ck_gemm_add_bias(float *C, const float *bias, int M, int N)
void gemm_tn_blocked(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_f32_llama_production(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_blocked_serial(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nn_blocked(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_tn_avx512(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)