36#if defined(__AVX512F__)
40#if defined(__linux__) && defined(__AMX_TILE__)
41#include <sys/syscall.h>
58static inline int ck_min_i(
int a,
int b) {
return a < b ? a : b; }
65static void gemm_bf16_scalar(
const uint16_t *A,
71 for (
int i = 0; i < M; ++i) {
72 for (
int j = 0; j < N; ++j) {
74 const size_t a_row = (size_t)i * (
size_t)K;
75 const size_t b_row = (size_t)j * (
size_t)K;
76 for (
int k = 0; k < K; ++k) {
83#if defined(__AVX512F__)
94static inline __m512 bf16_dot16(__m256i a_bf16, __m256i b_bf16, __m512 acc)
96 __m512 a_fp32 = bf16x16_to_fp32(a_bf16);
97 __m512 b_fp32 = bf16x16_to_fp32(b_bf16);
98 return _mm512_fmadd_ps(a_fp32, b_fp32, acc);
105static void gemm_bf16_avx512(
const uint16_t *A,
107 const uint16_t *bias,
111 #pragma omp parallel for schedule(dynamic)
112 for (
int i = 0; i < M; ++i) {
113 const uint16_t *a_row = A + (size_t)i * K;
115 for (
int j = 0; j < N; ++j) {
116 const uint16_t *b_row = B + (size_t)j * K;
119 __m512 sum_vec = _mm512_setzero_ps();
123 for (; k <= K - 16; k += 16) {
124 __m256i a_bf16 = _mm256_loadu_si256((
const __m256i *)(a_row + k));
125 __m256i b_bf16 = _mm256_loadu_si256((
const __m256i *)(b_row + k));
126 sum_vec = bf16_dot16(a_bf16, b_bf16, sum_vec);
130 float sum = _mm512_reduce_add_ps(sum_vec);
151static void gemm_bf16_blocked_avx512(
const uint16_t *A,
153 const uint16_t *bias,
158 #pragma omp parallel for
159 for (
int i = 0; i < M; ++i) {
160 for (
int j = 0; j < N; ++j) {
167 #pragma omp parallel for collapse(2) schedule(dynamic)
168 for (
int ii = 0; ii < M; ii +=
BLK_M) {
169 for (
int jj = 0; jj < N; jj +=
BLK_N) {
175 for (
int i = 0; i <
BLK_M; ++i) {
176 for (
int j = 0; j <
BLK_N; ++j) {
182 for (
int kk = 0; kk < K; kk +=
BLK_K) {
185 for (
int i = ii; i < i_end; ++i) {
186 const uint16_t *a_row = A + (size_t)i * K;
187 int local_i = i - ii;
189 for (
int j = jj; j < j_end; ++j) {
190 const uint16_t *b_row = B + (size_t)j * K;
191 int local_j = j - jj;
193 __m512 sum_vec = _mm512_setzero_ps();
196 for (; k <= k_end - 16; k += 16) {
197 __m256i a_bf16 = _mm256_loadu_si256((
const __m256i *)(a_row + k));
198 __m256i b_bf16 = _mm256_loadu_si256((
const __m256i *)(b_row + k));
199 sum_vec = bf16_dot16(a_bf16, b_bf16, sum_vec);
202 float partial = _mm512_reduce_add_ps(sum_vec);
203 for (; k < k_end; ++k) {
207 acc[local_i][local_j] += partial;
213 for (
int i = ii; i < i_end; ++i) {
214 for (
int j = jj; j < j_end; ++j) {
216 float new_val = old_val + acc[i - ii][j - jj];
229#if defined(__AVX512BF16__) && defined(__AVX512VL__)
232static inline __m512bh load_bf16x32(
const uint16_t *ptr)
234 return (__m512bh)_mm512_loadu_si512((
const __m512i *)ptr);
237#if defined(__AMX_TILE__) && defined(__AMX_BF16__)
239#ifndef ARCH_REQ_XCOMP_PERM
240#define ARCH_REQ_XCOMP_PERM 0x1023
242#ifndef XFEATURE_XTILE_DATA
243#define XFEATURE_XTILE_DATA 18
246typedef struct ck_amx_tile_config {
249 uint8_t reserved_0[14];
254_Static_assert(
sizeof(ck_amx_tile_config) == 64,
255 "AMX tile configuration must occupy exactly 64 bytes");
257static int ck_amx_request_xtile_data(
void)
259#if defined(__linux__)
260 static int state = 0;
267 long rc = syscall(SYS_arch_prctl, ARCH_REQ_XCOMP_PERM, XFEATURE_XTILE_DATA);
268 state = (rc == 0) ? 1 : -1;
275static void ck_amx_config_bf16_16x16x32(
void)
277 ck_amx_tile_config cfg;
278 memset(&cfg, 0,
sizeof(cfg));
287 for (
int tile = 3; tile <= 5; ++tile) {
289 cfg.colsb[tile] = 64;
292 _tile_loadconfig(&cfg);
295static void ck_amx_config_bf16_16x16_kblock(
int k_block)
297 ck_amx_tile_config cfg;
298 memset(&cfg, 0,
sizeof(cfg));
301 cfg.colsb[0] = (uint16_t)(k_block * (
int)
sizeof(uint16_t));
302 cfg.rows[1] = (uint8_t)(k_block / 2);
308 for (
int tile = 3; tile <= 5; ++tile) {
310 cfg.colsb[tile] = 64;
314 __asm__
volatile(
"" : :
"m"(cfg) :
"memory");
315 _tile_loadconfig(&cfg);
318static void ck_pack_bf16_ktile_pairs_16x16(uint16_t *dst,
324 for (
int kp = 0; kp < 16; ++kp) {
325 const int k0 = k + kp * 2;
326 for (
int nn = 0; nn < 16; ++nn) {
327 dst[(size_t)kp * 32u + (
size_t)nn * 2u + 0u] =
328 B[(size_t)(j + nn) * (size_t)K + (
size_t)k0];
329 dst[(size_t)kp * 32u + (
size_t)nn * 2u + 1u] =
330 B[(size_t)(j + nn) * (size_t)K + (
size_t)(k0 + 1)];
335static void gemm_bf16_fp32out_amx(
const uint16_t *A,
341 ck_amx_config_bf16_16x16x32();
343 uint16_t b_tile[16 * 32];
345 for (
int i = 0; i < M; i += 16) {
346 for (
int j = 0; j < N; j += 16) {
349 for (
int k = 0; k < K; k += 32) {
350 ck_pack_bf16_ktile_pairs_16x16(b_tile, B, K, j, k);
351 _tile_loadd(0, A + (
size_t)i * (
size_t)K + (
size_t)k, K * (
int)
sizeof(uint16_t));
352 _tile_loadd(1, b_tile, 32 * (
int)
sizeof(uint16_t));
353 _tile_dpbf16ps(2, 0, 1);
356 _tile_stored(2,
C + (
size_t)i * (
size_t)N + (
size_t)j, N * (
int)
sizeof(
float));
359 for (
int ii = 0; ii < 16; ++ii) {
360 float *c_row =
C + (size_t)(i + ii) * (size_t)N + (
size_t)j;
361 for (
int jj = 0; jj < 16; ++jj) {
362 c_row[jj] += bias[j + jj];
372#define HAVE_AMX_BF16 1
374#define HAVE_AMX_BF16 0
377static void gemm_bf16_native(
const uint16_t *A,
379 const uint16_t *bias,
383 #pragma omp parallel for schedule(dynamic)
384 for (
int i = 0; i < M; ++i) {
385 for (
int j = 0; j < N; ++j) {
387 __m512 sum_vec = _mm512_setzero_ps();
391 for (; k <= K - 32; k += 32) {
392 __m512bh a_vec = load_bf16x32(A + (
size_t)i * K + k);
393 __m512bh b_vec = load_bf16x32(B + (
size_t)j * K + k);
394 sum_vec = _mm512_dpbf16_ps(sum_vec, a_vec, b_vec);
397 float sum = _mm512_reduce_add_ps(sum_vec);
414#define HAVE_NATIVE_BF16 1
416#define HAVE_NATIVE_BF16 0
426 const uint16_t *bias,
430 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
436 gemm_bf16_native(A, B, bias,
C, M, N, K);
437#elif defined(__AVX512F__)
440 gemm_bf16_blocked_avx512(A, B, bias,
C, M, N, K);
442 gemm_bf16_avx512(A, B, bias,
C, M, N, K);
446 gemm_bf16_scalar(A, B, bias,
C, M, N, K);
459 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
465 const char *amx_env = getenv(
"CK_BF16_AMX");
466 if (amx_env && amx_env[0] ==
'1' &&
467 (M % 16) == 0 && (N % 16) == 0 && (K % 32) == 0 &&
468 M >= 16 && N >= 16 && K >= 32 && ck_amx_request_xtile_data()) {
469 gemm_bf16_fp32out_amx(A, B, bias,
C, M, N, K);
474 #pragma omp parallel for schedule(dynamic)
475 for (
int i = 0; i < M; ++i) {
476 const uint16_t *a_row = A + (size_t)i * K;
479 for (; j + 4 <= N; j += 4) {
480 const uint16_t *b0 = B + (size_t)(j + 0) * K;
481 const uint16_t *b1 = B + (size_t)(j + 1) * K;
482 const uint16_t *b2 = B + (size_t)(j + 2) * K;
483 const uint16_t *b3 = B + (size_t)(j + 3) * K;
484 __m512 acc0 = _mm512_setzero_ps();
485 __m512 acc1 = _mm512_setzero_ps();
486 __m512 acc2 = _mm512_setzero_ps();
487 __m512 acc3 = _mm512_setzero_ps();
490 for (; k <= K - 32; k += 32) {
491 const __m512bh a_vec = load_bf16x32(a_row + k);
492 acc0 = _mm512_dpbf16_ps(acc0, a_vec, load_bf16x32(b0 + k));
493 acc1 = _mm512_dpbf16_ps(acc1, a_vec, load_bf16x32(b1 + k));
494 acc2 = _mm512_dpbf16_ps(acc2, a_vec, load_bf16x32(b2 + k));
495 acc3 = _mm512_dpbf16_ps(acc3, a_vec, load_bf16x32(b3 + k));
498 float s0 = _mm512_reduce_add_ps(acc0);
499 float s1 = _mm512_reduce_add_ps(acc1);
500 float s2 = _mm512_reduce_add_ps(acc2);
501 float s3 = _mm512_reduce_add_ps(acc3);
515 C[(size_t)i * N + (j + 0)] = s0;
516 C[(size_t)i * N + (j + 1)] = s1;
517 C[(size_t)i * N + (j + 2)] = s2;
518 C[(size_t)i * N + (j + 3)] = s3;
522 const uint16_t *b_row = B + (size_t)j * K;
523 __m512 sum_vec = _mm512_setzero_ps();
526 for (; k <= K - 32; k += 32) {
527 const __m512bh a_vec = load_bf16x32(a_row + k);
528 const __m512bh b_vec = load_bf16x32(b_row + k);
529 sum_vec = _mm512_dpbf16_ps(sum_vec, a_vec, b_vec);
532 float sum = _mm512_reduce_add_ps(sum_vec);
539 C[(size_t)i * N + j] = sum;
542#elif defined(__AVX512F__)
543 #pragma omp parallel for schedule(dynamic)
544 for (
int i = 0; i < M; ++i) {
545 const uint16_t *a_row = A + (size_t)i * K;
547 for (
int j = 0; j < N; ++j) {
548 const uint16_t *b_row = B + (size_t)j * K;
550 __m512 sum_vec = _mm512_setzero_ps();
553 for (; k <= K - 16; k += 16) {
554 __m256i a_bf16 = _mm256_loadu_si256((
const __m256i *)(a_row + k));
555 __m256i b_bf16 = _mm256_loadu_si256((
const __m256i *)(b_row + k));
556 sum_vec = bf16_dot16(a_bf16, b_bf16, sum_vec);
559 float sum = _mm512_reduce_add_ps(sum_vec);
569 C[(size_t)i * N + j] = sum;
573 for (
int i = 0; i < M; ++i) {
574 for (
int j = 0; j < N; ++j) {
575 float sum = bias ? bias[j] : 0.0f;
576 for (
int k = 0; k < K; ++k) {
580 C[(size_t)i * N + j] = sum;
604 int row_begin,
int row_end)
606 if (!y || !w || !x || M <= 0 || K <= 0 ||
607 row_begin < 0 || row_begin >= row_end || row_end > M) {
611 for (
int i = row_begin; i < row_end; ++i) {
612 const uint16_t *w_row = w + (size_t)i * (
size_t)K;
614 for (
int k = 0; k < K; ++k) {
628 y, (
const uint16_t *)W, x, M, K, 0, M);
637} ck_gemv_bf16_args_t;
641 const ck_gemv_bf16_args_t *args = (
const ck_gemv_bf16_args_t *)opaque;
643 args->y, args->w, args->x, args->M, args->K, begin,
end);
653 (
size_t)M * (
size_t)K <= 65536) {
658 ck_gemv_bf16_args_t args = {
659 .y = y, .w = (
const uint16_t *)W, .x = x, .M = M, .K = K,
662 if (active > M) active = M;
663 int grain = M / (active * 4);
664 if (grain < 1) grain = 1;
673 int row_begin,
int row_end)
676 for (
int row = row_begin; row < row_end; ++row) {
683 const ck_gemv_bf16_args_t *args = (
const ck_gemv_bf16_args_t *)opaque;
685 args->y, args->w, args->x, args->M, args->K, begin,
end);
695 (
size_t)M * (
size_t)K <= 65536) {
697 y, (
const uint16_t *)W, x, M, K, 0, M);
701 ck_gemv_bf16_args_t args = {
702 .y = y, .w = (
const uint16_t *)W, .x = x, .M = M, .K = K,
705 if (active > M) active = M;
706 int grain = M / (active * 4);
707 if (grain < 1) grain = 1;
725 int row_begin,
int row_end)
727 const uint16_t *w = (
const uint16_t *)B;
728 if (!A || !w || !
C || M <= 0 || N <= 0 || K <= 0 ||
729 row_begin < 0 || row_begin >= row_end || row_end > M) {
733 for (
int i = row_begin; i < row_end; ++i) {
734 const float *a_row = A + (size_t)i * (
size_t)K;
735 float *c_row =
C + (size_t)i * (
size_t)N;
736 for (
int j = 0; j < N; ++j) {
737 const uint16_t *w_row = w + (size_t)j * (
size_t)K;
738 float sum = bias ? bias[j] : 0.0f;
739 for (
int k = 0; k < K; ++k) {
766} ck_gemm_nt_bf16_exact_args_t;
770 const ck_gemm_nt_bf16_exact_args_t *args =
771 (
const ck_gemm_nt_bf16_exact_args_t *)opaque;
773 args->A, args->B, args->bias, args->C,
774 args->M, args->N, args->K, begin,
end);
784 const char *disabled = getenv(
"CK_DISABLE_BF16_GEMM_PARALLEL_PREFILL");
785 if ((disabled && disabled[0] && strcmp(disabled,
"0") != 0) ||
787 (
size_t)M * (
size_t)N <= 4096) {
792 ck_gemm_nt_bf16_exact_args_t args = {
793 .A = A, .B = B, .bias = bias, .C =
C, .M = M, .N = N, .K = K,
796 if (active > M) active = M;
797 int grain = M / (active * 4);
798 if (grain < 1) grain = 1;
810 const uint16_t *bias,
814 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
818#if defined(__AVX512F__)
819 #pragma omp parallel for
820 for (
int i = 0; i < M; ++i) {
823 for (; j <= N - 16; j += 16) {
824 __m512 b_vec = bias ? bf16x16_to_fp32(_mm256_loadu_si256((
const __m256i *)(bias + j)))
825 : _mm512_setzero_ps();
826 __m256i out = fp32x16_to_bf16(b_vec);
827 _mm256_storeu_si256((__m256i *)(
C + (
size_t)i * N + j), out);
835 for (
int k = 0; k < K; ++k) {
837 __m512 a_broadcast = _mm512_set1_ps(a_val);
840 for (; j <= N - 16; j += 16) {
841 __m256i b_bf16 = _mm256_loadu_si256((
const __m256i *)(B + (
size_t)k * N + j));
842 __m512 b_fp32 = bf16x16_to_fp32(b_bf16);
844 __m256i c_bf16 = _mm256_loadu_si256((
const __m256i *)(
C + (
size_t)i * N + j));
845 __m512 c_fp32 = bf16x16_to_fp32(c_bf16);
847 c_fp32 = _mm512_fmadd_ps(a_broadcast, b_fp32, c_fp32);
849 __m256i c_out = fp32x16_to_bf16(c_fp32);
850 _mm256_storeu_si256((__m256i *)(
C + (
size_t)i * N + j), c_out);
861 for (
int i = 0; i < M; ++i) {
862 for (
int j = 0; j < N; ++j) {
864 for (
int k = 0; k < K; ++k) {
877 const uint16_t *bias,
881 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
889#if defined(__AVX512F__)
891 #pragma omp parallel for
892 for (
int i = 0; i < M; ++i) {
893 for (
int j = 0; j < N; ++j) {
900 #pragma omp parallel for
901 for (
int i = 0; i < M; ++i) {
902 for (
int j = 0; j < N; ++j) {
903 __m512 sum_vec = _mm512_setzero_ps();
906 for (; k <= K - 16; k += 16) {
908 __m512 a_fp32 = _mm512_setzero_ps();
909 for (
int kk = 0; kk < 16; ++kk) {
911 a_fp32 = _mm512_mask_mov_ps(a_fp32, 1 << kk, _mm512_set1_ps(val));
915 __m512 b_fp32 = _mm512_setzero_ps();
916 for (
int kk = 0; kk < 16; ++kk) {
918 b_fp32 = _mm512_mask_mov_ps(b_fp32, 1 << kk, _mm512_set1_ps(val));
921 sum_vec = _mm512_fmadd_ps(a_fp32, b_fp32, sum_vec);
924 float sum = _mm512_reduce_add_ps(sum_vec);
936 for (
int i = 0; i < M; ++i) {
937 for (
int j = 0; j < N; ++j) {
939 for (
int k = 0; k < K; ++k) {
967} ck_gemm_bf16_native_args_t;
971 ck_gemm_bf16_native_args_t *args = (ck_gemm_bf16_native_args_t *)opaque;
972 const int N = args->N;
973 const int K = args->K;
974 enum { ROW_TILE = 4 };
975 uint16_t *a_bf16 = (uint16_t *)alloca(
976 (
size_t)ROW_TILE * (size_t)K *
sizeof(uint16_t));
978 for (
int row0 = ith * ROW_TILE; row0 < args->M; row0 += nth * ROW_TILE) {
979 const int rows = args->M - row0 < ROW_TILE ? args->M - row0 : ROW_TILE;
980 for (
int r = 0; r < rows; ++r) {
981 const float *src = args->A + (size_t)(row0 + r) * (size_t)K;
982 uint16_t *ar = a_bf16 + (size_t)r * (
size_t)K;
988 for (; j + 4 <= N; j += 4) {
989 const uint16_t *b0 = args->B + (size_t)(j + 0) * K;
990 const uint16_t *b1 = args->B + (size_t)(j + 1) * K;
991 const uint16_t *b2 = args->B + (size_t)(j + 2) * K;
992 const uint16_t *b3 = args->B + (size_t)(j + 3) * K;
993 __m512 acc[ROW_TILE][4];
994 for (
int r = 0; r < rows; ++r) {
995 for (
int lane = 0; lane < 4; ++lane) acc[r][lane] = _mm512_setzero_ps();
998 for (; k <= K - 32; k += 32) {
999 const __m512bh bv[4] = {
1000 load_bf16x32(b0 + k), load_bf16x32(b1 + k),
1001 load_bf16x32(b2 + k), load_bf16x32(b3 + k)
1003 for (
int r = 0; r < rows; ++r) {
1005 load_bf16x32(a_bf16 + (
size_t)r * (
size_t)K + k);
1006 for (
int lane = 0; lane < 4; ++lane) {
1007 acc[r][lane] = _mm512_dpbf16_ps(acc[r][lane], av, bv[lane]);
1011 for (
int r = 0; r < rows; ++r) {
1013 for (
int lane = 0; lane < 4; ++lane) {
1014 sums[lane] = _mm512_reduce_add_ps(acc[r][lane]);
1016 const uint16_t *ar = a_bf16 + (size_t)r * (
size_t)K;
1017 for (
int tail = k; tail < K; ++tail) {
1024 float *dst = args->C + (size_t)(row0 + r) * (size_t)N;
1025 for (
int lane = 0; lane < 4; ++lane) {
1026 if (args->bias) sums[lane] += args->bias[j + lane];
1031 for (; j < N; ++j) {
1032 const uint16_t *b = args->B + (size_t)j * K;
1033 __m512 acc[ROW_TILE];
1034 for (
int r = 0; r < rows; ++r) acc[r] = _mm512_setzero_ps();
1036 for (; k <= K - 32; k += 32) {
1037 const __m512bh bv = load_bf16x32(b + k);
1038 for (
int r = 0; r < rows; ++r) {
1039 acc[r] = _mm512_dpbf16_ps(
1040 acc[r], load_bf16x32(a_bf16 + (
size_t)r * (
size_t)K + k), bv);
1043 for (
int r = 0; r < rows; ++r) {
1044 const uint16_t *ar = a_bf16 + (size_t)r * (
size_t)K;
1045 float sum = _mm512_reduce_add_ps(acc[r]);
1046 for (
int tail = k; tail < K; ++tail) {
1049 if (args->bias) sum += args->bias[j];
1050 args->C[(size_t)(row0 + r) * (size_t)N + j] =
1055 for (
int r = 0; r < rows; ++r) {
1056 const uint16_t *ar = a_bf16 + (size_t)r * (
size_t)K;
1057 float *dst = args->C + (size_t)(row0 + r) * (size_t)N;
1058 for (
int j = 0; j < N; ++j) {
1059 const uint16_t *b = args->B + (size_t)j * K;
1060 float sum = args->bias ? args->bias[j] : 0.0f;
1061 for (
int k = 0; k < K; ++k) {
1075 int M,
int N,
int K)
1077 const uint16_t *weights = (
const uint16_t *)B;
1078 if (!A || !weights || !
C || M <= 0 || N <= 0 || K <= 0)
return;
1080 ck_gemm_bf16_native_args_t args = {
1081 .A = A, .B = weights, .bias = bias, .C =
C, .M = M, .N = N, .K = K
1085 if (active > M) active = M;
1086 if (active > 24) active = 24;
1087 if (!pool || active <= 1 || (
size_t)M * (
size_t)N <= 4096) {
1098} ck_bf16_convert_args_t;
1102 ck_bf16_convert_args_t *args = (ck_bf16_convert_args_t *)opaque;
1103 const size_t begin = args->count * (size_t)ith / (
size_t)nth;
1104 const size_t end = args->count * (size_t)(ith + 1) / (size_t)nth;
1105 for (
size_t i = begin; i <
end; ++i) args->dst[i] =
float_to_bf16(args->src[i]);
1117} ck_gemm_bf16_amx_args_t;
1122 ck_gemm_bf16_amx_args_t *args = (ck_gemm_bf16_amx_args_t *)opaque;
1123 if (!ck_amx_request_xtile_data()) {
1124 __atomic_store_n(&args->failed, 1, __ATOMIC_RELAXED);
1127 ck_amx_config_bf16_16x16x32();
1128 uint16_t b_tile[16 * 32];
1129 const int mb = args->M / 16;
1130 const int nb = args->N / 16;
1131 const int m_groups = (mb + 3) / 4;
1132 const int jobs = m_groups * nb;
1134 for (
int job = ith; job < jobs; job += nth) {
1135 const int m_group = job / nb;
1136 const int j = (job % nb) * 16;
1137 const int group_blocks = (mb - m_group * 4 < 4) ? mb - m_group * 4 : 4;
1139 if (group_blocks > 1) _tile_zero(3);
1140 if (group_blocks > 2) _tile_zero(4);
1141 if (group_blocks > 3) _tile_zero(5);
1142 for (
int k = 0; k < args->K; k += 32) {
1143 ck_pack_bf16_ktile_pairs_16x16(b_tile, args->B, args->K, j, k);
1144 _tile_loadd(1, b_tile, 32 * (
int)
sizeof(uint16_t));
1145 for (
int g = 0; g < group_blocks; ++g) {
1146 const int i = (m_group * 4 + g) * 16;
1147 _tile_loadd(0, args->A + (
size_t)i * args->K + k,
1148 args->K * (
int)
sizeof(uint16_t));
1150 case 0: _tile_dpbf16ps(2, 0, 1);
break;
1151 case 1: _tile_dpbf16ps(3, 0, 1);
break;
1152 case 2: _tile_dpbf16ps(4, 0, 1);
break;
1153 default: _tile_dpbf16ps(5, 0, 1);
break;
1157 for (
int g = 0; g < group_blocks; ++g) {
1158 const int i = (m_group * 4 + g) * 16;
1159 float *tile_dst = args->C + (size_t)i * args->N + j;
1160 const int tile_stride = args->N * (int)
sizeof(
float);
1162 case 0: _tile_stored(2, tile_dst, tile_stride);
break;
1163 case 1: _tile_stored(3, tile_dst, tile_stride);
break;
1164 case 2: _tile_stored(4, tile_dst, tile_stride);
break;
1165 default: _tile_stored(5, tile_dst, tile_stride);
break;
1168 for (
int ii = 0; ii < 16; ++ii) {
1169 float *row = args->C + (size_t)(i + ii) * args->N + j;
1170 for (
int jj = 0; jj < 16; ++jj) row[jj] += args->bias[j + jj];
1177 (void)ith; (void)nth; (void)opaque;
1184} ck_bf16_round_args_t;
1188 ck_bf16_round_args_t *args = (ck_bf16_round_args_t *)opaque;
1189 const size_t begin = args->count * (size_t)ith / (
size_t)nth;
1190 const size_t end = args->count * (size_t)(ith + 1) / (size_t)nth;
1191 for (
size_t i = begin; i <
end; ++i) {
1199 return ck_amx_request_xtile_data();
1208 int M,
int N,
int K,
1212 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0 ||
1213 (M % 16) != 0 || (N % 16) != 0 || (K % 2) != 0 ||
1214 !ck_amx_request_xtile_data()) {
1218 int k_block = K < 32 ? K : 32;
1219 while (k_block > 2 && K % k_block != 0) k_block -= 2;
1220 ck_amx_config_bf16_16x16_kblock(k_block);
1221 uint16_t b_tile[16 * 32];
1222 for (
int i = 0; i < M; i += 16) {
1223 for (
int j = 0; j < N; j += 16) {
1225 _tile_loadd(2,
C + (
size_t)i * (
size_t)N + (
size_t)j,
1226 N * (
int)
sizeof(
float));
1230 for (
int k = 0; k < K; k += k_block) {
1231 memset(b_tile, 0,
sizeof(b_tile));
1232 for (
int kp = 0; kp < k_block / 2; ++kp) {
1233 const int k0 = k + kp * 2;
1234 for (
int nn = 0; nn < 16; ++nn) {
1235 b_tile[(size_t)kp * 32u + (
size_t)nn * 2u] =
1236 B[(size_t)(j + nn) * (size_t)K + (
size_t)k0];
1237 b_tile[(size_t)kp * 32u + (
size_t)nn * 2u + 1u] =
1238 B[(size_t)(j + nn) * (size_t)K + (
size_t)k0 + 1u];
1241 _tile_loadd(0, A + (
size_t)i * (
size_t)K + (
size_t)k,
1242 K * (
int)
sizeof(uint16_t));
1243 _tile_loadd(1, b_tile, 32 * (
int)
sizeof(uint16_t));
1244 _tile_dpbf16ps(2, 0, 1);
1246 _tile_stored(2,
C + (
size_t)i * (
size_t)N + (
size_t)j,
1247 N * (
int)
sizeof(
float));
1253 (void)A; (void)B; (void)
C; (void)M; (void)N; (void)K; (void)accumulate;
1262 int M,
int N,
int K,
1264 size_t a_bf16_bytes)
1267 if (!A || !B || !
C || M < 16 || N < 16 || K < 32 ||
1268 (M % 16) != 0 || (N % 16) != 0 || (K % 32) != 0) {
1270 "HARD KERNEL CONTRACT FAULT: AMX BF16 GEMM requires non-null buffers "
1271 "and M%%16=N%%16=K%%32=0 (M=%d N=%d K=%d)\n",
1275 const size_t input_count = (size_t)M * K;
1276 if (!a_bf16 || input_count > SIZE_MAX /
sizeof(uint16_t) ||
1277 a_bf16_bytes < input_count *
sizeof(uint16_t)) {
1278 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: AMX BF16 activation workspace is too small\n");
1283 if (active > 24) active = 24;
1284 ck_bf16_convert_args_t convert = {.src=A, .dst=a_bf16, .count=input_count};
1287 ck_gemm_bf16_amx_args_t gemm = {
1288 .A=a_bf16, .B=(
const uint16_t *)B, .bias=bias, .
C=
C,
1289 .M=M, .N=N, .K=K, .failed=0
1293 if (__atomic_load_n(&gemm.failed, __ATOMIC_RELAXED)) {
1294 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: AMX tile permission request failed\n");
1297 ck_bf16_round_args_t round = {.values=
C, .count=(size_t)M * N};
1302 (void)A; (void)B; (void)bias; (void)
C; (void)M; (void)N; (void)K;
1303 (void)a_bf16; (void)a_bf16_bytes;
1305 "HARD KERNEL CONTRACT FAULT: gemm_nt_bf16_amx_bf16_storage was selected "
1306 "without AMX BF16 support\n");
1315 int M,
int N,
int K)
1317 size_t input_count = 0;
1318 if (M > 0 && K > 0 && (
size_t)M <= SIZE_MAX / (
size_t)K) {
1319 input_count = (size_t)M * (
size_t)K;
1321 uint16_t *workspace = input_count > 0 && input_count <= SIZE_MAX /
sizeof(uint16_t)
1322 ? (uint16_t *)malloc(input_count *
sizeof(uint16_t))
1325 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: AMX BF16 compatibility workspace allocation failed\n");
1329 A, B, bias,
C, M, N, K, workspace, input_count *
sizeof(uint16_t));
1334 const float *A,
const void *B,
const float *bias,
float *
C,
1335 int M,
int N,
int K, uint16_t *a_bf16,
size_t a_bf16_bytes)
1337 const int amx_shape = M >= 16 && N >= 16 && K >= 32 &&
1338 (M % 16) == 0 && (N % 16) == 0 && (K % 32) == 0;
1341 A, B, bias,
C, M, N, K, a_bf16, a_bf16_bytes);
1351 int M,
int N,
int K)
1353 const int amx_shape = M >= 16 && N >= 16 && K >= 32 &&
1354 (M % 16) == 0 && (N % 16) == 0 && (K % 32) == 0;
1363static dnnl_engine_t ck_pytorch_brgemm_engine;
1364static dnnl_stream_t ck_pytorch_brgemm_stream;
1365static int ck_pytorch_brgemm_init_status = -1;
1366static pthread_once_t ck_pytorch_brgemm_once = PTHREAD_ONCE_INIT;
1367static pthread_mutex_t ck_pytorch_brgemm_lock = PTHREAD_MUTEX_INITIALIZER;
1368static const dnnl_version_t *ck_pytorch_brgemm_version;
1370static void ck_pytorch_brgemm_init(
void)
1372 const dnnl_version_t *version = dnnl_version();
1373 if (!version || version->cpu_runtime != DNNL_RUNTIME_OMP || !version->hash) {
1375 "HARD KERNEL CONTRACT FAULT: PyTorch oneDNN BF16 BRGEMM requires "
1376 "an identity-bearing OpenMP oneDNN runtime "
1377 "(found %d.%d.%d runtime=%u hash=%s)\n",
1378 version ? version->major : -1,
1379 version ? version->minor : -1,
1380 version ? version->patch : -1,
1381 version ? version->cpu_runtime : 0,
1382 version && version->hash ? version->hash :
"<missing>");
1385 ck_pytorch_brgemm_version = version;
1386 if (dnnl_engine_create(&ck_pytorch_brgemm_engine, dnnl_cpu, 0) != dnnl_success)
return;
1387 if (dnnl_stream_create(&ck_pytorch_brgemm_stream, ck_pytorch_brgemm_engine,
1388 dnnl_stream_default_flags) != dnnl_success) {
1389 dnnl_engine_destroy(ck_pytorch_brgemm_engine);
1390 ck_pytorch_brgemm_engine = NULL;
1393 ck_pytorch_brgemm_init_status = 0;
1396static void ck_pytorch_brgemm_fault(
const char *message,
int M,
int N,
int K)
1398 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: PyTorch oneDNN BF16 BRGEMM %s "
1399 "(M=%d N=%d K=%d)\n", message, M, N, K);
1403static void ck_pytorch_brgemm_require_version(
int major,
int minor,
int patch,
1404 const char *source_hash,
1405 const char *provider,
1406 int M,
int N,
int K)
1408 pthread_once(&ck_pytorch_brgemm_once, ck_pytorch_brgemm_init);
1409 if (ck_pytorch_brgemm_init_status != 0 || !ck_pytorch_brgemm_version) {
1410 ck_pytorch_brgemm_fault(
"could not initialize oneDNN", M, N, K);
1412 if (ck_pytorch_brgemm_version->major != major ||
1413 ck_pytorch_brgemm_version->minor != minor ||
1414 ck_pytorch_brgemm_version->patch != patch ||
1415 strcmp(ck_pytorch_brgemm_version->hash, source_hash) != 0) {
1417 "HARD KERNEL CONTRACT FAULT: %s requires oneDNN %d.%d.%d "
1418 "OpenMP at %s (found %d.%d.%d runtime=%u hash=%s)\n",
1419 provider, major, minor, patch, source_hash,
1420 ck_pytorch_brgemm_version->major,
1421 ck_pytorch_brgemm_version->minor,
1422 ck_pytorch_brgemm_version->patch,
1423 ck_pytorch_brgemm_version->cpu_runtime,
1424 ck_pytorch_brgemm_version->hash);
1431 const float *A,
const void *B,
const float *bias,
float *
C,
1432 int M,
int N,
int K)
1435 if (!A || !B || !
C || M <= 0 || N <= 0 || K <= 0) {
1436 ck_pytorch_brgemm_fault(
"received an invalid tensor contract", M, N, K);
1438 if (ck_pytorch_brgemm_init_status != 0) {
1439 ck_pytorch_brgemm_fault(
"could not initialize oneDNN", M, N, K);
1442 const size_t input_count = (size_t)M * (
size_t)K;
1443 const size_t output_count = (size_t)M * (
size_t)N;
1444 uint16_t *input_bf16 = (uint16_t *)malloc(input_count *
sizeof(*input_bf16));
1445 uint16_t *output_bf16 = (uint16_t *)malloc(output_count *
sizeof(*output_bf16));
1446 uint16_t *bias_bf16 = bias ? (uint16_t *)malloc((
size_t)N *
sizeof(*bias_bf16)) : NULL;
1447 if (!input_bf16 || !output_bf16 || (bias && !bias_bf16)) {
1451 ck_pytorch_brgemm_fault(
"workspace allocation failed", M, N, K);
1453 for (
size_t i = 0; i < input_count; ++i) input_bf16[i] =
float_to_bf16(A[i]);
1455 for (
int j = 0; j < N; ++j) bias_bf16[j] =
float_to_bf16(bias[j]);
1456 for (
int i = 0; i < M; ++i) {
1457 memcpy(output_bf16 + (
size_t)i * (
size_t)N,
1458 bias_bf16, (
size_t)N *
sizeof(*bias_bf16));
1462 dnnl_memory_desc_t src_md = NULL, weights_md = NULL, dst_md = NULL;
1463 dnnl_primitive_attr_t attr = NULL;
1464 dnnl_post_ops_t post_ops = NULL;
1465 dnnl_primitive_desc_t primitive_desc = NULL;
1466 dnnl_primitive_t primitive = NULL;
1467 dnnl_memory_t src_mem = NULL, weights_mem = NULL, dst_mem = NULL;
1468 dnnl_dims_t src_dims = {M, K};
1469 dnnl_dims_t weights_dims = {K, N};
1470 dnnl_dims_t dst_dims = {M, N};
1471 dnnl_dims_t src_strides = {K, 1};
1472 dnnl_dims_t weights_strides = {1, K};
1473 dnnl_dims_t dst_strides = {N, 1};
1474 dnnl_status_t status = dnnl_success;
1476#define CK_DNNL(call) do { status = (call); if (status != dnnl_success) goto cleanup; } while (0)
1477 pthread_mutex_lock(&ck_pytorch_brgemm_lock);
1478 CK_DNNL(dnnl_memory_desc_create_with_strides(&src_md, 2, src_dims, dnnl_bf16, src_strides));
1479 CK_DNNL(dnnl_memory_desc_create_with_strides(
1480 &weights_md, 2, weights_dims, dnnl_bf16, weights_strides));
1481 CK_DNNL(dnnl_memory_desc_create_with_strides(&dst_md, 2, dst_dims, dnnl_bf16, dst_strides));
1483 CK_DNNL(dnnl_primitive_attr_create(&attr));
1484 CK_DNNL(dnnl_post_ops_create(&post_ops));
1485 CK_DNNL(dnnl_post_ops_append_sum(post_ops, 1.0f, 0, dnnl_bf16));
1486 CK_DNNL(dnnl_primitive_attr_set_post_ops(attr, post_ops));
1488 CK_DNNL(dnnl_matmul_primitive_desc_create(
1489 &primitive_desc, ck_pytorch_brgemm_engine, src_md, weights_md, NULL, dst_md, attr));
1490 CK_DNNL(dnnl_primitive_create(&primitive, primitive_desc));
1491 CK_DNNL(dnnl_memory_create(&src_mem, src_md, ck_pytorch_brgemm_engine, input_bf16));
1492 CK_DNNL(dnnl_memory_create(
1493 &weights_mem, weights_md, ck_pytorch_brgemm_engine, (
void *)B));
1494 CK_DNNL(dnnl_memory_create(&dst_mem, dst_md, ck_pytorch_brgemm_engine, output_bf16));
1495 dnnl_exec_arg_t args[] = {
1496 {DNNL_ARG_SRC, src_mem},
1497 {DNNL_ARG_WEIGHTS, weights_mem},
1498 {DNNL_ARG_DST, dst_mem},
1500 CK_DNNL(dnnl_primitive_execute(
1501 primitive, ck_pytorch_brgemm_stream, (
int)(
sizeof(args) /
sizeof(args[0])), args));
1502 CK_DNNL(dnnl_stream_wait(ck_pytorch_brgemm_stream));
1505 if (dst_mem) dnnl_memory_destroy(dst_mem);
1506 if (weights_mem) dnnl_memory_destroy(weights_mem);
1507 if (src_mem) dnnl_memory_destroy(src_mem);
1508 if (primitive) dnnl_primitive_destroy(primitive);
1509 if (primitive_desc) dnnl_primitive_desc_destroy(primitive_desc);
1510 if (dst_md) dnnl_memory_desc_destroy(dst_md);
1511 if (weights_md) dnnl_memory_desc_destroy(weights_md);
1512 if (src_md) dnnl_memory_desc_destroy(src_md);
1513 if (post_ops) dnnl_post_ops_destroy(post_ops);
1514 if (attr) dnnl_primitive_attr_destroy(attr);
1515 pthread_mutex_unlock(&ck_pytorch_brgemm_lock);
1518 if (status != dnnl_success) {
1522 ck_pytorch_brgemm_fault(
"execution failed", M, N, K);
1524 for (
size_t i = 0; i < output_count; ++i)
C[i] =
bf16_to_float(output_bf16[i]);
1529 (void)A; (void)B; (void)bias; (void)
C; (void)M; (void)N; (void)K;
1530 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: PyTorch oneDNN BF16 BRGEMM was "
1531 "selected without USE_ONEDNN=1\n");
1540 int M,
int N,
int K)
1543 ck_pytorch_brgemm_require_version(
1544 3, 7, 1,
"8d263e693366ef8db40acc569cc7d8edf644556d",
1545 "gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage", M, N, K);
1548 A, B, bias,
C, M, N, K);
1552 const float *A,
const void *B,
const float *bias,
float *
C,
1553 int M,
int N,
int K)
1556 ck_pytorch_brgemm_require_version(
1557 3, 12, 0,
"80afa71049cd69a3df32adcccb623b12cd7baa22",
1558 "gemm_nt_bf16_pytorch_onednn_3_12_brgemm_bf16_storage", M, N, K);
1561 A, B, bias,
C, M, N, K);
1565 const float *input,
const void *weights,
const float *bias,
float *output,
1566 int batch,
int out_channels,
int in_channels,
int temporal,
1567 int patch_h,
int patch_w)
1570 if (!input || !weights || !bias || !output || batch <= 0 ||
1571 out_channels <= 0 || in_channels <= 0 || temporal <= 0 ||
1572 patch_h <= 0 || patch_w <= 0) {
1573 ck_pytorch_brgemm_fault(
"invalid Conv3D patch contract", batch,
1574 out_channels, in_channels * temporal * patch_h * patch_w);
1576 pthread_once(&ck_pytorch_brgemm_once, ck_pytorch_brgemm_init);
1577 if (ck_pytorch_brgemm_init_status != 0) {
1578 ck_pytorch_brgemm_fault(
"could not initialize oneDNN Conv3D", batch,
1579 out_channels, in_channels * temporal * patch_h * patch_w);
1581 ck_pytorch_brgemm_require_version(
1582 3, 7, 1,
"8d263e693366ef8db40acc569cc7d8edf644556d",
1583 "patch_projection_bf16_pytorch_onednn_conv3d_storage",
1584 batch, out_channels, in_channels * temporal * patch_h * patch_w);
1586 const size_t input_count = (size_t)batch * (
size_t)in_channels *
1587 (size_t)temporal * (
size_t)patch_h * (size_t)patch_w;
1588 const size_t output_count = (size_t)batch * (
size_t)out_channels;
1589 uint16_t *input_bf16 = (uint16_t *)malloc(input_count *
sizeof(*input_bf16));
1590 uint16_t *bias_bf16 = (uint16_t *)malloc((
size_t)out_channels *
sizeof(*bias_bf16));
1591 uint16_t *output_bf16 = (uint16_t *)malloc(output_count *
sizeof(*output_bf16));
1592 if (!input_bf16 || !bias_bf16 || !output_bf16) {
1596 ck_pytorch_brgemm_fault(
"Conv3D workspace allocation failed", batch,
1597 out_channels, in_channels * temporal * patch_h * patch_w);
1599 for (
size_t i = 0; i < input_count; ++i) input_bf16[i] =
float_to_bf16(input[i]);
1600 for (
int i = 0; i < out_channels; ++i) bias_bf16[i] =
float_to_bf16(bias[i]);
1602 dnnl_dims_t src_dims = {batch, in_channels, temporal, patch_h, patch_w};
1603 dnnl_dims_t weight_dims = {
1604 out_channels, in_channels, temporal, patch_h, patch_w};
1605 dnnl_dims_t bias_dims = {out_channels};
1606 dnnl_dims_t dst_dims = {batch, out_channels, 1, 1, 1};
1607 dnnl_dims_t strides = {temporal, patch_h, patch_w};
1608 dnnl_dims_t dilates = {0, 0, 0};
1609 dnnl_dims_t padding = {0, 0, 0};
1611 dnnl_memory_desc_t user_src_md = NULL, user_weight_md = NULL;
1612 dnnl_memory_desc_t bias_md = NULL, user_dst_md = NULL;
1613 dnnl_memory_desc_t any_src_md = NULL, any_weight_md = NULL, any_dst_md = NULL;
1614 dnnl_primitive_desc_t conv_pd = NULL;
1615 dnnl_primitive_t conv = NULL;
1616 dnnl_primitive_desc_t reorder_pd = NULL;
1617 dnnl_primitive_t reorder = NULL;
1618 dnnl_exec_arg_t reorder_args[2];
1619 dnnl_memory_t user_src = NULL, user_weight = NULL, bias_mem = NULL, user_dst = NULL;
1620 dnnl_memory_t conv_src = NULL, conv_weight = NULL, conv_dst = NULL;
1621 dnnl_status_t status = dnnl_success;
1623#define CK_DNNL_CONV(call) do { status = (call); if (status != dnnl_success) goto cleanup_conv; } while (0)
1624 pthread_mutex_lock(&ck_pytorch_brgemm_lock);
1625 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1626 &user_src_md, 5, src_dims, dnnl_bf16, dnnl_ncdhw));
1627 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1628 &user_weight_md, 5, weight_dims, dnnl_bf16, dnnl_oidhw));
1629 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1630 &bias_md, 1, bias_dims, dnnl_bf16, dnnl_x));
1631 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1632 &user_dst_md, 5, dst_dims, dnnl_bf16, dnnl_ncdhw));
1633 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1634 &any_src_md, 5, src_dims, dnnl_bf16, dnnl_format_tag_any));
1635 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1636 &any_weight_md, 5, weight_dims, dnnl_bf16, dnnl_format_tag_any));
1637 CK_DNNL_CONV(dnnl_memory_desc_create_with_tag(
1638 &any_dst_md, 5, dst_dims, dnnl_bf16, dnnl_format_tag_any));
1639 CK_DNNL_CONV(dnnl_convolution_forward_primitive_desc_create(
1640 &conv_pd, ck_pytorch_brgemm_engine, dnnl_forward_training,
1641 dnnl_convolution_direct, any_src_md, any_weight_md, bias_md, any_dst_md,
1642 strides, dilates, padding, padding, NULL));
1644 const_dnnl_memory_desc_t conv_src_md =
1645 dnnl_primitive_desc_query_md(conv_pd, dnnl_query_src_md, 0);
1646 const_dnnl_memory_desc_t conv_weight_md =
1647 dnnl_primitive_desc_query_md(conv_pd, dnnl_query_weights_md, 0);
1648 const_dnnl_memory_desc_t conv_dst_md =
1649 dnnl_primitive_desc_query_md(conv_pd, dnnl_query_dst_md, 0);
1650 CK_DNNL_CONV(dnnl_memory_create(
1651 &user_src, user_src_md, ck_pytorch_brgemm_engine, input_bf16));
1652 CK_DNNL_CONV(dnnl_memory_create(
1653 &user_weight, user_weight_md, ck_pytorch_brgemm_engine, (
void *)weights));
1654 CK_DNNL_CONV(dnnl_memory_create(
1655 &bias_mem, bias_md, ck_pytorch_brgemm_engine, bias_bf16));
1656 CK_DNNL_CONV(dnnl_memory_create(
1657 &user_dst, user_dst_md, ck_pytorch_brgemm_engine, output_bf16));
1658 CK_DNNL_CONV(dnnl_memory_create(
1659 &conv_src, conv_src_md, ck_pytorch_brgemm_engine, DNNL_MEMORY_ALLOCATE));
1660 CK_DNNL_CONV(dnnl_memory_create(
1661 &conv_weight, conv_weight_md, ck_pytorch_brgemm_engine, DNNL_MEMORY_ALLOCATE));
1662 CK_DNNL_CONV(dnnl_memory_create(
1663 &conv_dst, conv_dst_md, ck_pytorch_brgemm_engine, DNNL_MEMORY_ALLOCATE));
1665 CK_DNNL_CONV(dnnl_reorder_primitive_desc_create(
1666 &reorder_pd, user_src_md, ck_pytorch_brgemm_engine,
1667 conv_src_md, ck_pytorch_brgemm_engine, NULL));
1668 CK_DNNL_CONV(dnnl_primitive_create(&reorder, reorder_pd));
1669 reorder_args[0] = (dnnl_exec_arg_t){DNNL_ARG_FROM, user_src};
1670 reorder_args[1] = (dnnl_exec_arg_t){DNNL_ARG_TO, conv_src};
1671 CK_DNNL_CONV(dnnl_primitive_execute(
1672 reorder, ck_pytorch_brgemm_stream, 2, reorder_args));
1673 dnnl_primitive_destroy(reorder); reorder = NULL;
1674 dnnl_primitive_desc_destroy(reorder_pd); reorder_pd = NULL;
1676 CK_DNNL_CONV(dnnl_reorder_primitive_desc_create(
1677 &reorder_pd, user_weight_md, ck_pytorch_brgemm_engine,
1678 conv_weight_md, ck_pytorch_brgemm_engine, NULL));
1679 CK_DNNL_CONV(dnnl_primitive_create(&reorder, reorder_pd));
1680 reorder_args[0] = (dnnl_exec_arg_t){DNNL_ARG_FROM, user_weight};
1681 reorder_args[1] = (dnnl_exec_arg_t){DNNL_ARG_TO, conv_weight};
1682 CK_DNNL_CONV(dnnl_primitive_execute(
1683 reorder, ck_pytorch_brgemm_stream, 2, reorder_args));
1684 dnnl_primitive_destroy(reorder); reorder = NULL;
1685 dnnl_primitive_desc_destroy(reorder_pd); reorder_pd = NULL;
1687 CK_DNNL_CONV(dnnl_primitive_create(&conv, conv_pd));
1688 dnnl_exec_arg_t conv_args[] = {
1689 {DNNL_ARG_SRC, conv_src},
1690 {DNNL_ARG_WEIGHTS, conv_weight},
1691 {DNNL_ARG_BIAS, bias_mem},
1692 {DNNL_ARG_DST, conv_dst},
1694 CK_DNNL_CONV(dnnl_primitive_execute(
1695 conv, ck_pytorch_brgemm_stream,
1696 (
int)(
sizeof(conv_args) /
sizeof(conv_args[0])), conv_args));
1698 CK_DNNL_CONV(dnnl_reorder_primitive_desc_create(
1699 &reorder_pd, conv_dst_md, ck_pytorch_brgemm_engine,
1700 user_dst_md, ck_pytorch_brgemm_engine, NULL));
1701 CK_DNNL_CONV(dnnl_primitive_create(&reorder, reorder_pd));
1702 reorder_args[0] = (dnnl_exec_arg_t){DNNL_ARG_FROM, conv_dst};
1703 reorder_args[1] = (dnnl_exec_arg_t){DNNL_ARG_TO, user_dst};
1704 CK_DNNL_CONV(dnnl_primitive_execute(
1705 reorder, ck_pytorch_brgemm_stream, 2, reorder_args));
1706 CK_DNNL_CONV(dnnl_stream_wait(ck_pytorch_brgemm_stream));
1709 if (reorder) dnnl_primitive_destroy(reorder);
1710 if (reorder_pd) dnnl_primitive_desc_destroy(reorder_pd);
1711 if (conv) dnnl_primitive_destroy(conv);
1712 if (conv_dst) dnnl_memory_destroy(conv_dst);
1713 if (conv_weight) dnnl_memory_destroy(conv_weight);
1714 if (conv_src) dnnl_memory_destroy(conv_src);
1715 if (user_dst) dnnl_memory_destroy(user_dst);
1716 if (bias_mem) dnnl_memory_destroy(bias_mem);
1717 if (user_weight) dnnl_memory_destroy(user_weight);
1718 if (user_src) dnnl_memory_destroy(user_src);
1719 if (conv_pd) dnnl_primitive_desc_destroy(conv_pd);
1720 if (any_dst_md) dnnl_memory_desc_destroy(any_dst_md);
1721 if (any_weight_md) dnnl_memory_desc_destroy(any_weight_md);
1722 if (any_src_md) dnnl_memory_desc_destroy(any_src_md);
1723 if (user_dst_md) dnnl_memory_desc_destroy(user_dst_md);
1724 if (bias_md) dnnl_memory_desc_destroy(bias_md);
1725 if (user_weight_md) dnnl_memory_desc_destroy(user_weight_md);
1726 if (user_src_md) dnnl_memory_desc_destroy(user_src_md);
1727 pthread_mutex_unlock(&ck_pytorch_brgemm_lock);
1730 if (status != dnnl_success) {
1734 ck_pytorch_brgemm_fault(
"oneDNN Conv3D execution failed", batch,
1735 out_channels, in_channels * temporal * patch_h * patch_w);
1737 for (
size_t i = 0; i < output_count; ++i) output[i] =
bf16_to_float(output_bf16[i]);
1742 (void)input; (void)weights; (void)bias; (void)output; (void)batch;
1743 (void)out_channels; (void)in_channels; (void)temporal; (void)patch_h; (void)patch_w;
1744 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: PyTorch oneDNN BF16 Conv3D "
1745 "was selected without USE_ONEDNN=1\n");
1751 const float *image,
const void *weights_t0,
const void *weights_t1,
1752 const float *bias,
float *output,
int channels,
int image_h,
int image_w,
1753 int patch_size,
int out_channels,
int merge_size)
1756 if (!image || !weights_t0 || !weights_t1 || !bias || !output ||
1757 channels <= 0 || image_h <= 0 || image_w <= 0 || patch_size <= 0 ||
1758 out_channels <= 0 || merge_size <= 0 || image_h % patch_size != 0 ||
1759 image_w % patch_size != 0) {
1760 ck_pytorch_brgemm_fault(
"invalid image patch projection contract",
1761 image_h, image_w, patch_size);
1763 const int grid_h = image_h / patch_size;
1764 const int grid_w = image_w / patch_size;
1765 if (grid_h % merge_size != 0 || grid_w % merge_size != 0) {
1766 ck_pytorch_brgemm_fault(
"patch grid is not merge-tile aligned",
1767 grid_h, grid_w, merge_size);
1769 const int batch = grid_h * grid_w;
1770 const int temporal = 2;
1771 const int half_k = channels * patch_size * patch_size;
1772 const int full_k = temporal * half_k;
1773 float *patches = (
float *)malloc((
size_t)batch * (size_t)full_k *
sizeof(*patches));
1774 uint16_t *weights = (uint16_t *)malloc(
1775 (
size_t)out_channels * (size_t)full_k *
sizeof(*weights));
1776 if (!patches || !weights) {
1779 ck_pytorch_brgemm_fault(
"image patch projection workspace allocation failed",
1780 batch, out_channels, full_k);
1783 for (
int tok = 0; tok < batch; ++tok) {
1784 const int tiles_per_row = grid_w / merge_size;
1785 const int tile_area = merge_size * merge_size;
1786 const int tile = tok / tile_area;
1787 const int within = tok % tile_area;
1788 const int patch_y = (tile / tiles_per_row) * merge_size + within / merge_size;
1789 const int patch_x = (tile % tiles_per_row) * merge_size + within % merge_size;
1790 float *dst = patches + (size_t)tok * (
size_t)full_k;
1791 for (
int c = 0; c < channels; ++c) {
1792 for (
int t = 0; t < temporal; ++t) {
1793 for (
int py = 0; py < patch_size; ++py) {
1794 const float *src = image +
1795 ((size_t)c * (
size_t)image_h +
1796 (size_t)(patch_y * patch_size + py)) * (
size_t)image_w +
1797 (size_t)(patch_x * patch_size);
1798 memcpy(dst, src, (
size_t)patch_size *
sizeof(*dst));
1805 const uint16_t *w0 = (
const uint16_t *)weights_t0;
1806 const uint16_t *w1 = (
const uint16_t *)weights_t1;
1807 for (
int n = 0; n < out_channels; ++n) {
1808 uint16_t *dst = weights + (size_t)n * (
size_t)full_k;
1809 for (
int c = 0; c < channels; ++c) {
1810 const size_t channel_offset =
1811 (size_t)n * (
size_t)half_k +
1812 (size_t)c * (
size_t)patch_size * (size_t)patch_size;
1813 const size_t plane_bytes =
1814 (size_t)patch_size * (
size_t)patch_size *
sizeof(*dst);
1815 memcpy(dst, w0 + channel_offset, plane_bytes);
1816 dst += patch_size * patch_size;
1817 memcpy(dst, w1 + channel_offset, plane_bytes);
1818 dst += patch_size * patch_size;
1823 patches, weights, bias, output, batch, out_channels, channels,
1824 temporal, patch_size, patch_size);
1828 (void)image; (void)weights_t0; (void)weights_t1; (void)bias; (void)output;
1829 (void)channels; (void)image_h; (void)image_w; (void)patch_size;
1830 (void)out_channels; (void)merge_size;
1831 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: PyTorch oneDNN BF16 image "
1832 "patch projection was selected without USE_ONEDNN=1\n");
1839 const uint16_t *weights_t0;
1840 const uint16_t *weights_t1;
1851} ck_patch_projection_bf16_native_args_t;
1854 int ith,
int nth,
void *opaque)
1856 ck_patch_projection_bf16_native_args_t *args =
1857 (ck_patch_projection_bf16_native_args_t *)opaque;
1858 const int begin = args->batch * ith / nth;
1859 const int end = args->batch * (ith + 1) / nth;
1860 const int patch_area = args->patch_size * args->patch_size;
1861 const int half_k = args->channels * patch_area;
1862 const int tiles_per_row = args->grid_w / args->merge_size;
1863 const int tile_area = args->merge_size * args->merge_size;
1865 for (
int tok = begin; tok <
end; ++tok) {
1866 const int tile = tok / tile_area;
1867 const int within = tok % tile_area;
1869 (tile / tiles_per_row) * args->merge_size + within / args->merge_size;
1871 (tile % tiles_per_row) * args->merge_size + within % args->merge_size;
1872 for (
int n = 0; n < args->out_channels; ++n) {
1873 float sum = args->bias
1876#if defined(__AVX512BF16__) && defined(__AVX512VL__)
1877 if (args->patch_size == 16) {
1878 __m256 acc = _mm256_setzero_ps();
1879 for (
int c = 0; c < args->channels; ++c) {
1880 for (
int t = 0; t < 2; ++t) {
1881 const uint16_t *weights = t == 0
1882 ? args->weights_t0 : args->weights_t1;
1883 const uint16_t *weight_plane = weights +
1884 (size_t)n * (
size_t)half_k +
1885 (size_t)c * (
size_t)patch_area;
1886 for (
int py = 0; py < 16; ++py) {
1887 const float *src = args->image +
1888 ((size_t)c * (
size_t)args->image_h +
1889 (size_t)(patch_y * 16 + py)) *
1890 (
size_t)args->image_w +
1891 (size_t)(patch_x * 16);
1892 const __m256bh image_bf16 = _mm256_cvtne2ps_pbh(
1893 _mm256_loadu_ps(src + 8), _mm256_loadu_ps(src));
1894 const __m256bh weight_bf16 = (__m256bh)_mm256_loadu_si256(
1895 (
const __m256i *)(weight_plane + (size_t)py * 16u));
1896 acc = _mm256_dpbf16_ps(acc, image_bf16, weight_bf16);
1901 _mm256_storeu_ps(lanes, acc);
1902 const float sum01 = lanes[0] + lanes[1];
1903 const float sum23 = lanes[2] + lanes[3];
1904 const float sum45 = lanes[4] + lanes[5];
1905 const float sum67 = lanes[6] + lanes[7];
1906 sum += (sum01 + sum23) + (sum45 + sum67);
1910 for (
int c = 0; c < args->channels; ++c) {
1911 for (
int t = 0; t < 2; ++t) {
1912 const uint16_t *weights = t == 0
1913 ? args->weights_t0 : args->weights_t1;
1914 const uint16_t *weight_plane = weights +
1915 (size_t)n * (
size_t)half_k +
1916 (size_t)c * (
size_t)patch_area;
1917 for (
int py = 0; py < args->patch_size; ++py) {
1918 const float *src = args->image +
1919 ((size_t)c * (
size_t)args->image_h +
1920 (size_t)(patch_y * args->patch_size + py)) *
1921 (
size_t)args->image_w +
1922 (size_t)(patch_x * args->patch_size);
1923 for (
int px = 0; px < args->patch_size; ++px) {
1926 weight_plane[(
size_t)py *
1927 (
size_t)args->patch_size + (
size_t)px]);
1928 sum += value * weight;
1934 args->output[(size_t)tok * (
size_t)args->out_channels + (size_t)n] =
1941 const float *image,
const void *weights_t0,
const void *weights_t1,
1942 const float *bias,
float *output,
int channels,
int image_h,
int image_w,
1943 int patch_size,
int out_channels,
int merge_size)
1945 if (!image || !weights_t0 || !weights_t1 || !output || channels <= 0 ||
1946 image_h <= 0 || image_w <= 0 || patch_size <= 0 || out_channels <= 0 ||
1947 merge_size <= 0 || image_h % patch_size != 0 ||
1948 image_w % patch_size != 0) {
1949 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: invalid native BF16 image patch projection\n");
1952 const int grid_h = image_h / patch_size;
1953 const int grid_w = image_w / patch_size;
1954 if (grid_h % merge_size != 0 || grid_w % merge_size != 0) {
1955 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: native BF16 patch grid is not merge aligned\n");
1958 ck_patch_projection_bf16_native_args_t args = {
1960 .weights_t0 = (
const uint16_t *)weights_t0,
1961 .weights_t1 = (
const uint16_t *)weights_t1,
1964 .channels = channels,
1967 .patch_size = patch_size,
1968 .out_channels = out_channels,
1969 .merge_size = merge_size,
1971 .batch = grid_h * grid_w,
1975 if (active > args.batch) active = args.batch;
1976 if (active > 24) active = 24;
1977 if (pool && active > 1) {
1989 int M,
int N,
int K,
1990 int row_begin,
int row_end)
1993 for (
int row = row_begin; row < row_end; ++row) {
1994 float *dst =
C + (size_t)row * (
size_t)N;
1995 for (
int col = 0; col < N; ++col) {
2003 const ck_gemm_nt_bf16_exact_args_t *args =
2004 (
const ck_gemm_nt_bf16_exact_args_t *)opaque;
2006 args->A, args->B, args->bias, args->C,
2007 args->M, args->N, args->K, begin,
end);
2014 int M,
int N,
int K)
2018 (
size_t)M * (
size_t)N <= 4096) {
2023 ck_gemm_nt_bf16_exact_args_t args = {
2024 .A = A, .B = B, .bias = bias, .C =
C, .M = M, .N = N, .K = K,
2027 if (active > M) active = M;
2028 int grain = M / (active * 4);
2029 if (grain < 1) grain = 1;
2038 int M,
int N,
int K)
2044 const uint16_t *input,
2045 const uint16_t *weight,
2053 if (!d_output || !input || !weight || tokens <= 0 || in_dim <= 0 || out_dim <= 0) {
2058 for (
int t = 0; t < tokens; ++t) {
2059 for (
int i = 0; i < in_dim; ++i) {
2061 for (
int o = 0; o < out_dim; ++o) {
2062 const float dy =
bf16_to_float(d_output[(
size_t)t * (
size_t)out_dim + (
size_t)o]);
2063 const float w =
bf16_to_float(weight[(
size_t)o * (
size_t)in_dim + (
size_t)i]);
2066 d_input[(size_t)t * (
size_t)in_dim + (size_t)i] = sum;
2072 for (
int o = 0; o < out_dim; ++o) {
2073 for (
int i = 0; i < in_dim; ++i) {
2075 for (
int t = 0; t < tokens; ++t) {
2076 const float dy =
bf16_to_float(d_output[(
size_t)t * (
size_t)out_dim + (
size_t)o]);
2077 const float x =
bf16_to_float(input[(
size_t)t * (
size_t)in_dim + (
size_t)i]);
2080 d_weight[(size_t)o * (
size_t)in_dim + (size_t)i] = sum;
2086 for (
int o = 0; o < out_dim; ++o) {
2088 for (
int t = 0; t < tokens; ++t) {
2089 sum +=
bf16_to_float(d_output[(
size_t)t * (
size_t)out_dim + (
size_t)o]);
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
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)
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)
void gemm_nt_bf16_pytorch_onednn_3_12_brgemm_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_tn_bf16(const uint16_t *A, const uint16_t *B, const uint16_t *bias, uint16_t *C, int M, int N, int K)
void gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
static void ck_gemv_bf16_rows(int begin, int end, void *opaque)
void gemv_bf16_bf16_storage(float *y, const void *W, const float *x, int M, int K)
int ck_gemm_bf16_amx_available(void)
void gemm_nt_bf16_bf16_storage_parallel_dispatch(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_bf16_native_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_bf16_bf16_storage_row_range(const float *A, const void *B, const float *bias, float *C, int M, int N, int K, int row_begin, int row_end)
void gemm_nt_bf16_amx_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void patch_projection_bf16_pytorch_onednn_conv3d_storage(const float *input, const void *weights, const float *bias, float *output, int batch, int out_channels, int in_channels, int temporal, int patch_h, int patch_w)
void gemm_bf16_fp32out(const uint16_t *A, const uint16_t *B, const float *bias, float *C, int M, int N, int K)
void gemv_bf16_bf16_storage_parallel_dispatch(float *y, const void *W, const float *x, int M, int K)
void gemm_blocked_serial_bf16(const uint16_t *A, const uint16_t *B, const uint16_t *bias, uint16_t *C, int M, int N, int K)
void patch_projection_image_bf16_native_storage(const float *image, const void *weights_t0, const void *weights_t1, const float *bias, float *output, int channels, int image_h, int image_w, int patch_size, int out_channels, int merge_size)
int ck_gemm_bf16_fp32out_amx_raw(const uint16_t *A, const uint16_t *B, float *C, int M, int N, int K, int accumulate)
void gemm_nt_bf16(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_bf16_parallel_dispatch(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
static void ck_bf16_convert_work(int ith, int nth, void *opaque)
static int ck_min_i(int a, int b)
void gemv_bf16(float *y, const void *W, const float *x, int M, int K)
void gemm_nn_bf16(const uint16_t *A, const uint16_t *B, const uint16_t *bias, uint16_t *C, int M, int N, int K)
static void ck_gemm_bf16_native_work(int ith, int nth, void *opaque)
static void ck_gemm_nt_bf16_exact_rows(int begin, int end, void *opaque)
static void ck_bf16_round_work(int ith, int nth, void *opaque)
static void gemm_nt_bf16_pytorch_onednn_brgemm_bf16_storage_impl(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_backward_bf16_mixed(const uint16_t *d_output, const uint16_t *input, const uint16_t *weight, float *d_input, float *d_weight, float *d_bias, int tokens, int in_dim, int out_dim)
static void gemv_bf16_row_range(float *y, const uint16_t *w, const float *x, int M, int K, int row_begin, int row_end)
void gemv_bf16_parallel_dispatch(float *y, const void *W, const float *x, int M, int K)
static void ck_gemm_nt_bf16_storage_exact_rows(int begin, int end, void *opaque)
static void ck_patch_projection_bf16_native_work(int ith, int nth, void *opaque)
static void ck_gemm_bf16_amx_work(int ith, int nth, void *opaque)
void patch_projection_image_bf16_pytorch_onednn_conv3d_storage(const float *image, const void *weights_t0, const void *weights_t1, const float *bias, float *output, int channels, int image_h, int image_w, int patch_size, int out_channels, int merge_size)
void gemm_nt_bf16_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
void gemm_nt_bf16_prefill_shape_safe_bf16_storage(const float *A, const void *B, const float *bias, float *C, int M, int N, int K)
static void gemv_bf16_bf16_storage_row_range(float *y, const uint16_t *w, const float *x, int M, int K, int row_begin, int row_end)
void gemm_nt_bf16_prefill_shape_safe_bf16_storage_workspace(const float *A, const void *B, const float *bias, float *C, int M, int N, int K, uint16_t *a_bf16, size_t a_bf16_bytes)
void gemm_nt_bf16_row_range(const float *A, const void *B, const float *bias, float *C, int M, int N, int K, int row_begin, int row_end)
void gemm_nt_bf16_amx_bf16_storage_workspace(const float *A, const void *B, const float *bias, float *C, int M, int N, int K, uint16_t *a_bf16, size_t a_bf16_bytes)
static void ck_gemv_bf16_storage_rows(int begin, int end, void *opaque)
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)