← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels.c File Reference

General matrix multiply (GEMM) kernels with SIMD (SSE/AVX/AVX512) More...

#include "ckernel_engine.h"
#include "ck_threadpool.h"
#include <stdlib.h>
#include <string.h>

Go to the source code of this file.

Functions

static void ck_gemm_add_bias (float *C, const float *bias, int M, int N)
 
static int ck_gemm_nn_impl_probe_enabled (void)
 
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)
 
static void ck_gemm_nt_fp32_exact_rows (int begin, int end, void *opaque)
 
static int ck_min (int a, int b)
 
void gemm_avx512_parallel (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_fine_grained_parallel (const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
 
void gemm_naive_parallel (const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
 
static void gemm_naive_serial_double (const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
 
static void gemm_naive_serial_float (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)
 
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_blocked (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_nn_serial_double (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)
 
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_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)
 
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 void gemm_nt_matvec_parallel (const float *A, const float *B, const float *bias, float *C, 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)
 
void gemm_tn_blocked (const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
 
void gemm_tn_parallel (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)
 

Detailed Description

General matrix multiply (GEMM) kernels with SIMD (SSE/AVX/AVX512)

CK-ENGINE KERNEL RULES:

  1. NO malloc/free - memory via bump allocator, pointers passed in
  2. NO OpenMP - parallelization at orchestrator/codegen layer
  3. API must define: inputs, outputs, workspace, and memory layouts
  4. Pure computation - deterministic, no side effects

After changes: make test && make llamacpp-parity-full

LEGACY EXCEPTION: This file contains OpenMP for backward compatibility. New kernels should NOT use OpenMP internally.

GEMM: C = alpha * A @ B + beta * C (with optional bias)

Definition in file gemm_kernels.c.

Function Documentation

◆ ck_gemm_add_bias()

static void ck_gemm_add_bias ( float *  C,
const float *  bias,
int  M,
int  N 
)
inlinestatic

Definition at line 33 of file gemm_kernels.c.

34{
35 if (!bias) {
36 return;
37 }
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) {
42 c_row[j] += bias[j];
43 }
44 }
45}
#define C(color)
Definition show_config.c:39

References C.

Referenced by gemm_blocked_serial().

◆ ck_gemm_nn_impl_probe_enabled()

static int ck_gemm_nn_impl_probe_enabled ( void  )
static

Definition at line 556 of file gemm_kernels.c.

557{
558 static int cached = -1;
559 if (cached != -1) {
560 return cached;
561 }
562 cached = 0;
563 const char *v = getenv("CK_GEMM_NN_IMPL");
564 if (!v || !v[0]) {
565 return cached;
566 }
567 if (strcmp(v, "probe") == 0 || strcmp(v, "dup") == 0 || strcmp(v, "1") == 0) {
568 cached = 1;
569 }
570 return cached;
571}

Referenced by gemm_nn_simd().

◆ ck_gemm_nt_f32_llama_production_output()

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 
)
inlinestatic

Definition at line 965 of file gemm_kernels.c.

968{
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;
973 float sum = 0.0f;
974 int k = 0;
975
976#if defined(__AVX512F__)
977 if (M > 1) {
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);
982 }
983 sum = _mm512_reduce_add_ps(acc);
984 } else {
985 __m512 acc[4] = {
986 _mm512_setzero_ps(), _mm512_setzero_ps(),
987 _mm512_setzero_ps(), _mm512_setzero_ps()
988 };
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),
994 acc[lane]);
995 }
996 }
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]);
1001 }
1002#elif defined(__AVX__)
1003 if (M > 1) {
1004 __m256 acc = _mm256_setzero_ps();
1005 for (; k + 8 <= K; k += 8) {
1006#if defined(__FMA__)
1007 acc = _mm256_fmadd_ps(
1008 _mm256_loadu_ps(a + k), _mm256_loadu_ps(b + k), acc);
1009#else
1010 acc = _mm256_add_ps(
1011 acc, _mm256_mul_ps(
1012 _mm256_loadu_ps(a + k), _mm256_loadu_ps(b + k)));
1013#endif
1014 }
1015 sum = ck_hsum256_llamafile(acc);
1016 } else {
1017 __m256 acc[4] = {
1018 _mm256_setzero_ps(), _mm256_setzero_ps(),
1019 _mm256_setzero_ps(), _mm256_setzero_ps()
1020 };
1021 for (; k + 32 <= K; k += 32) {
1022 for (int lane = 0; lane < 4; ++lane) {
1023#if defined(__FMA__)
1024 acc[lane] = _mm256_fmadd_ps(
1025 _mm256_loadu_ps(a + k + lane * 8),
1026 _mm256_loadu_ps(b + k + lane * 8),
1027 acc[lane]);
1028#else
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)));
1033#endif
1034 }
1035 }
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]);
1040 }
1041#endif
1042 for (; k < K; ++k) {
1043 sum += a[k] * b[k];
1044 }
1045 C[(size_t)row * (size_t)N + (size_t)col] =
1046 bias ? sum + bias[col] : sum;
1047}

References C.

Referenced by gemm_nt_f32_llama_production(), and gemm_nt_f32_llama_production_output_range().

◆ ck_gemm_nt_fp32_exact_rows()

static void ck_gemm_nt_fp32_exact_rows ( int  begin,
int  end,
void *  opaque 
)
static

Definition at line 157 of file gemm_kernels.c.

158{
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];
170 }
171 if (!args->bias_before_reduction && args->bias) {
172 sum += args->bias[j];
173 }
174 c_row[j] = sum;
175 }
176 }
177}
uint32_t end
Definition utf8.c:215

References end.

Referenced by gemm_nt_fp32_exact_parallel_dispatch().

◆ ck_min()

static int ck_min ( int  a,
int  b 
)
inlinestatic

Definition at line 31 of file gemm_kernels.c.

31{ return a < b ? a : b; }

Referenced by gemm_blocked_serial(), gemm_fine_grained_parallel(), gemm_nn_blocked(), and gemm_tn_blocked().

◆ gemm_avx512_parallel()

void gemm_avx512_parallel ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 230 of file gemm_kernels.c.

235{
237 gemm_naive_serial_float(A, B, bias, C, M, N, K);
238 return;
239 }
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();
245 int k;
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);
250 }
251 float sum = _mm512_reduce_add_ps(sum_vec);
252 for (; k < K; k++) {
253 sum += A[i * K + k] * B[j * K + k];
254 }
255 float bias_val = bias ? bias[j] : 0.0f;
256 C[i * N + j] = sum + bias_val;
257 }
258 }
259#elif defined(__AVX__)
260 // AVX1 path: 256-bit vectors, no FMA (use mul + add)
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();
265 int k;
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);
271 }
272 float sum = hsum256_ps(sum_vec);
273 for (; k < K; k++) {
274 sum += A[i * K + k] * B[j * K + k];
275 }
276 float bias_val = bias ? bias[j] : 0.0f;
277 C[i * N + j] = sum + bias_val;
278 }
279 }
280#else
281 gemm_naive_parallel(A, B, bias, C, M, N, K);
282#endif
283}
int ck_strict_parity_enabled(void)
void gemm_naive_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
static void gemm_naive_serial_float(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)

References C, ck_strict_parity_enabled(), gemm_naive_parallel(), and gemm_naive_serial_float().

◆ gemm_blocked_serial()

void gemm_blocked_serial ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 849 of file gemm_kernels.c.

854{
855 // Ensure threads are initialized (auto-detects on first call)
856 (void)ck_get_num_threads();
857
859 gemm_naive_serial_float(A, B, bias, C, M, N, K);
860 return;
861 }
862
863 // Decode-time matvec (M=1) is extremely common and benefits from parallelism over N.
864 // Lower threshold to parallelize more ops; OpenMP overhead is ~1-2μs per barrier.
865 // For N*K >= 64K elements, parallel is worthwhile.
866 if (M == 1 && (size_t)N * (size_t)K >= 65536) {
867 gemm_nt_matvec_parallel(A, B, bias, C, N, K);
868 return;
869 }
870
871 /*
872 * Use gemm_microkernel for large matrices - it uses MKL/oneDNN when available,
873 * which is substantially faster than our hand-written SIMD kernels.
874 * B is stored as [N x K] (transposed), so we pass B_transposed=1.
875 * Note: Use threshold of 32 to avoid numerical precision issues with small matrices.
876 */
877 if (M >= 32 && N >= 32 && K >= 32) {
878 gemm_microkernel(A, B, C, M, N, K, 1); // B_transposed=1
879 ck_gemm_add_bias(C, bias, M, N);
880 return;
881 }
882#if defined(__AVX512F__)
883 const int block_size = 64;
884#elif defined(__AVX__)
885 const int block_size = 32;
886#else
887 const int block_size = 32;
888#endif
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;
892 }
893 }
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);
900
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();
905 int k;
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);
910 }
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];
914 }
915#elif defined(__AVX__)
916 __m256 sum_vec = _mm256_setzero_ps();
917 int k;
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);
923 }
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];
927 }
928#else
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];
932 }
933#endif
934 C[i * N + j] += partial_sum;
935 }
936 }
937 }
938 }
939 }
940}
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)
static int ck_min(int a, int b)
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_add_bias(float *C, const float *bias, int M, int N)

References C, ck_gemm_add_bias(), ck_get_num_threads(), ck_min(), ck_strict_parity_enabled(), gemm_microkernel(), gemm_naive_serial_float(), and gemm_nt_matvec_parallel().

Referenced by ck_attention_project_head_major(), ck_gemm_nt_quant(), ck_mlp_swiglu_forward(), ck_mlp_swiglu_forward_fused_token(), ck_qkv_project_head_major(), ck_qkv_project_head_major_token(), ck_train_gemm_work(), gemm_blocked_serial_train_parallel_dispatch(), mlp_token_parallel(), and mlp_token_parallel_exact().

◆ gemm_fine_grained_parallel()

void gemm_fine_grained_parallel ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 286 of file gemm_kernels.c.

291{
293 gemm_naive_serial_float(A, B, bias, C, M, N, K);
294 return;
295 }
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;
302 }
303 }
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);
311
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();
315 int k;
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);
320 }
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];
324 }
325#pragma omp atomic
326 C[i * N + j] += partial_sum;
327 }
328 }
329 }
330 }
331 }
332#elif defined(__AVX__)
333 // AVX1 cache-blocked version
334 const int block_size = 32; // Smaller block for L1 cache
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;
339 }
340 }
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);
348
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();
352 int k;
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);
358 }
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];
362 }
363#pragma omp atomic
364 C[i * N + j] += partial_sum;
365 }
366 }
367 }
368 }
369 }
370#else
371 gemm_naive_parallel(A, B, bias, C, M, N, K);
372#endif
373}

References C, ck_min(), ck_strict_parity_enabled(), gemm_naive_parallel(), and gemm_naive_serial_float().

◆ gemm_naive_parallel()

void gemm_naive_parallel ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 206 of file gemm_kernels.c.

211{
213 gemm_naive_serial_float(A, B, bias, C, M, N, K);
214 return;
215 }
216#pragma omp parallel for
217 for (int i = 0; i < M; i++) {
218 for (int j = 0; j < N; j++) {
219 float sum = 0.0f;
220 for (int k = 0; k < K; k++) {
221 sum += A[i * K + k] * B[j * K + k];
222 }
223 float bias_val = bias ? bias[j] : 0.0f;
224 C[i * N + j] = sum + bias_val;
225 }
226 }
227}

References C, ck_strict_parity_enabled(), and gemm_naive_serial_float().

Referenced by ck_attention_project_head_major_ref(), ck_mlp_swiglu_forward_ref(), ck_qkv_project_head_major_ref(), gemm_avx512_parallel(), and gemm_fine_grained_parallel().

◆ gemm_naive_serial_double()

static void gemm_naive_serial_double ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)
static

Definition at line 112 of file gemm_kernels.c.

117{
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];
123 }
124 C[i * N + j] = (float)sum;
125 }
126 }
127}

References C.

◆ gemm_naive_serial_float()

static void gemm_naive_serial_float ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)
static

Definition at line 129 of file gemm_kernels.c.

134{
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];
140 }
141 C[i * N + j] = sum;
142 }
143 }
144}

References C.

Referenced by gemm_avx512_parallel(), gemm_blocked_serial(), gemm_fine_grained_parallel(), and gemm_naive_parallel().

◆ gemm_nn_avx512()

void gemm_nn_avx512 ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 420 of file gemm_kernels.c.

425{
427 gemm_nn_serial_double(A, B, bias, C, M, N, K);
428 return;
429 }
430#if defined(__AVX512F__)
431 // For gemm_nn, we can't vectorize over K easily since B[k,j] has stride N.
432 // Instead, vectorize over N (output columns) when N >= 16.
433#pragma omp parallel for
434 for (int i = 0; i < M; i++) {
435 int j = 0;
436 // Process 16 output columns at a time
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);
443 }
444 _mm512_storeu_ps(&C[i * N + j], sum_vec);
445 }
446 // Handle remaining columns
447 for (; j < N; j++) {
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];
451 }
452 C[i * N + j] = sum;
453 }
454 }
455#elif defined(__AVX__)
456 // AVX1: vectorize over N (8 columns at a time)
457#pragma omp parallel for
458 for (int i = 0; i < M; i++) {
459 int j = 0;
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);
467 }
468 _mm256_storeu_ps(&C[i * N + j], sum_vec);
469 }
470 for (; j < N; j++) {
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];
474 }
475 C[i * N + j] = sum;
476 }
477 }
478#else
479 gemm_nn_parallel(A, B, bias, C, M, N, K);
480#endif
481}
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_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)

References C, ck_strict_parity_enabled(), gemm_nn_parallel(), and gemm_nn_serial_double().

Referenced by gemm_nn_simd().

◆ gemm_nn_avx512_probe()

void gemm_nn_avx512_probe ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 493 of file gemm_kernels.c.

498{
500 gemm_nn_serial_double(A, B, bias, C, M, N, K);
501 return;
502 }
503#if defined(__AVX512F__)
504 // For gemm_nn, we can't vectorize over K easily since B[k,j] has stride N.
505 // Instead, vectorize over N (output columns) when N >= 16.
506#pragma omp parallel for
507 for (int i = 0; i < M; i++) {
508 int j = 0;
509 // Process 16 output columns at a time
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);
516 }
517 _mm512_storeu_ps(&C[i * N + j], sum_vec);
518 }
519 // Handle remaining columns
520 for (; j < N; j++) {
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];
524 }
525 C[i * N + j] = sum;
526 }
527 }
528#elif defined(__AVX__)
529 // AVX1: vectorize over N (8 columns at a time)
530#pragma omp parallel for
531 for (int i = 0; i < M; i++) {
532 int j = 0;
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);
540 }
541 _mm256_storeu_ps(&C[i * N + j], sum_vec);
542 }
543 for (; j < N; j++) {
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];
547 }
548 C[i * N + j] = sum;
549 }
550 }
551#else
552 gemm_nn_parallel(A, B, bias, C, M, N, K);
553#endif
554}

References C, ck_strict_parity_enabled(), gemm_nn_parallel(), and gemm_nn_serial_double().

Referenced by gemm_nn_simd().

◆ gemm_nn_blocked()

void gemm_nn_blocked ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 590 of file gemm_kernels.c.

595{
597 gemm_nn_serial_double(A, B, bias, C, M, N, K);
598 return;
599 }
600#if defined(__AVX512F__)
601 const int block_size = 64;
602#elif defined(__AVX__)
603 const int block_size = 32;
604#else
605 const int block_size = 32;
606#endif
607 // Initialize C with bias (parallelized)
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;
612 }
613 }
614 // Blocked multiply-accumulate (parallelized over M blocks)
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);
622
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);
628 int j;
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);
634 }
635 for (; j < j_end; j++) {
636 C[i * N + j] += a_val * B[k * N + j];
637 }
638#elif defined(__AVX__)
639 __m256 a_broadcast = _mm256_set1_ps(a_val);
640 int j;
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);
647 }
648 for (; j < j_end; j++) {
649 C[i * N + j] += a_val * B[k * N + j];
650 }
651#else
652 for (int j = jj; j < j_end; j++) {
653 C[i * N + j] += a_val * B[k * N + j];
654 }
655#endif
656 }
657 }
658 }
659 }
660 }
661}

References C, ck_min(), ck_strict_parity_enabled(), and gemm_nn_serial_double().

◆ gemm_nn_parallel()

void gemm_nn_parallel ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 398 of file gemm_kernels.c.

403{
405 gemm_nn_serial_double(A, B, bias, C, M, N, K);
406 return;
407 }
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];
414 }
415 C[i * N + j] = sum;
416 }
417 }
418}

References C, ck_strict_parity_enabled(), and gemm_nn_serial_double().

Referenced by gemm_nn_avx512(), and gemm_nn_avx512_probe().

◆ gemm_nn_serial_double()

static void gemm_nn_serial_double ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)
static

Definition at line 381 of file gemm_kernels.c.

386{
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];
392 }
393 C[i * N + j] = (float)sum;
394 }
395 }
396}

References C.

Referenced by gemm_nn_avx512(), gemm_nn_avx512_probe(), gemm_nn_blocked(), and gemm_nn_parallel().

◆ gemm_nn_simd()

void gemm_nn_simd ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 577 of file gemm_kernels.c.

582{
584 gemm_nn_avx512_probe(A, B, bias, C, M, N, K);
585 return;
586 }
587 gemm_nn_avx512(A, B, bias, C, M, N, K);
588}
static int ck_gemm_nn_impl_probe_enabled(void)
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_avx512(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)

References C, ck_gemm_nn_impl_probe_enabled(), gemm_nn_avx512(), and gemm_nn_avx512_probe().

Referenced by fc1_backward_kernel(), fc2_backward_kernel(), gemm_backward_f32_train_parallel_dispatch(), and gemm_backward_f32_train_parallel_dispatch_v2().

◆ gemm_nt_f32_llama_production()

void gemm_nt_f32_llama_production ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 1062 of file gemm_kernels.c.

1067{
1068 if (!A || !B || !C || M <= 0 || N <= 0 || K <= 0) {
1069 return;
1070 }
1071
1072#pragma omp parallel for schedule(static) if ((size_t)M * (size_t)N >= 96)
1073 for (int index = 0; index < M * N; ++index) {
1074 ck_gemm_nt_f32_llama_production_output(A, B, bias, C, M, N, K, index);
1075 }
1076}
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)

References C, and ck_gemm_nt_f32_llama_production_output().

Referenced by ck_moe_shared_q4k_gated_workspace(), and moe_swiglu_shared_forward_q8_0_gated_workspace().

◆ gemm_nt_f32_llama_production_output_range()

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 
)

Definition at line 1049 of file gemm_kernels.c.

1052{
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) {
1058 ck_gemm_nt_f32_llama_production_output(A, B, bias, C, M, N, K, index);
1059 }
1060}

References C, and ck_gemm_nt_f32_llama_production_output().

◆ gemm_nt_fp32_exact_parallel_dispatch()

void gemm_nt_fp32_exact_parallel_dispatch ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 179 of file gemm_kernels.c.

184{
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,
189 .bias_before_reduction = ck_strict_parity_enabled(),
190 };
191 ck_threadpool_t *pool = ck_threadpool_global();
192 if (!pool || ck_threadpool_n_threads(pool) <= 1 || M < 2 ||
193 (size_t)M * (size_t)N <= 4096) {
194 ck_gemm_nt_fp32_exact_rows(0, M, &args);
195 return;
196 }
197 int active = ck_threadpool_n_threads(pool);
198 if (active > M) active = M;
199 int grain = M / (active * 4);
200 if (grain < 1) grain = 1;
202 pool, active, 0, M, grain, ck_gemm_nt_fp32_exact_rows, &args);
203}
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)
static void ck_gemm_nt_fp32_exact_rows(int begin, int end, void *opaque)

References C, ck_gemm_nt_fp32_exact_rows(), ck_strict_parity_enabled(), ck_threadpool_global(), ck_threadpool_n_threads(), and ck_threadpool_parallel_for_n().

◆ gemm_nt_matvec_parallel()

static void gemm_nt_matvec_parallel ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  N,
int  K 
)
static

Definition at line 66 of file gemm_kernels.c.

72{
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;
77
78#if defined(__AVX512F__)
79 __m512 acc = _mm512_setzero_ps();
80 int k = 0;
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);
85 }
86 sum += _mm512_reduce_add_ps(acc);
87 for (; k < K; ++k) {
88 sum += A[k] * b_row[k];
89 }
90#elif defined(__AVX__)
91 __m256 acc = _mm256_setzero_ps();
92 int k = 0;
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));
97 }
98 sum += hsum256_ps(acc);
99 for (; k < K; ++k) {
100 sum += A[k] * b_row[k];
101 }
102#else
103 for (int k = 0; k < K; ++k) {
104 sum += A[k] * b_row[k];
105 }
106#endif
107
108 C[j] = sum;
109 }
110}

References C.

Referenced by gemm_blocked_serial().

◆ gemm_tn_avx512()

void gemm_tn_avx512 ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 709 of file gemm_kernels.c.

714{
716 gemm_tn_serial_double(A, B, bias, C, M, N, K);
717 return;
718 }
719#if defined(__AVX512F__)
720 // Vectorize over N (output columns)
721#pragma omp parallel for
722 for (int i = 0; i < M; i++) {
723 int j = 0;
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);
730 }
731 _mm512_storeu_ps(&C[i * N + j], sum_vec);
732 }
733 for (; j < N; j++) {
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];
737 }
738 C[i * N + j] = sum;
739 }
740 }
741#elif defined(__AVX__)
742 // AVX1: vectorize over N (8 columns at a time)
743#pragma omp parallel for
744 for (int i = 0; i < M; i++) {
745 int j = 0;
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);
753 }
754 _mm256_storeu_ps(&C[i * N + j], sum_vec);
755 }
756 for (; j < N; j++) {
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];
760 }
761 C[i * N + j] = sum;
762 }
763 }
764#else
765 gemm_tn_parallel(A, B, bias, C, M, N, K);
766#endif
767}
void gemm_tn_parallel(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)

References C, ck_strict_parity_enabled(), gemm_tn_parallel(), and gemm_tn_serial_double().

◆ gemm_tn_blocked()

void gemm_tn_blocked ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 769 of file gemm_kernels.c.

774{
776 gemm_tn_serial_double(A, B, bias, C, M, N, K);
777 return;
778 }
779#if defined(__AVX512F__)
780 const int block_size = 64;
781#elif defined(__AVX__)
782 const int block_size = 32;
783#else
784 const int block_size = 32;
785#endif
786 // Initialize C with bias (parallelized)
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;
791 }
792 }
793 // Blocked multiply-accumulate (parallelized over M blocks)
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);
801
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);
807 int j;
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);
813 }
814 for (; j < j_end; j++) {
815 C[i * N + j] += a_val * B[k * N + j];
816 }
817#elif defined(__AVX__)
818 __m256 a_broadcast = _mm256_set1_ps(a_val);
819 int j;
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);
826 }
827 for (; j < j_end; j++) {
828 C[i * N + j] += a_val * B[k * N + j];
829 }
830#else
831 for (int j = jj; j < j_end; j++) {
832 C[i * N + j] += a_val * B[k * N + j];
833 }
834#endif
835 }
836 }
837 }
838 }
839 }
840}

References C, ck_min(), ck_strict_parity_enabled(), and gemm_tn_serial_double().

◆ gemm_tn_parallel()

void gemm_tn_parallel ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)

Definition at line 687 of file gemm_kernels.c.

692{
694 gemm_tn_serial_double(A, B, bias, C, M, N, K);
695 return;
696 }
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];
703 }
704 C[i * N + j] = sum;
705 }
706 }
707}

References C, ck_strict_parity_enabled(), and gemm_tn_serial_double().

Referenced by fc1_backward_kernel(), fc2_backward_kernel(), and gemm_tn_avx512().

◆ gemm_tn_serial_double()

static void gemm_tn_serial_double ( const float *  A,
const float *  B,
const float *  bias,
float *  C,
int  M,
int  N,
int  K 
)
static

Definition at line 669 of file gemm_kernels.c.

674{
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++) {
679 // A.T[i,k] = A[k,i] = A[k*M + i]
680 sum += (double)A[k * M + i] * (double)B[k * N + j];
681 }
682 C[i * N + j] = (float)sum;
683 }
684 }
685}

References C.

Referenced by gemm_tn_avx512(), gemm_tn_blocked(), and gemm_tn_parallel().