← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
gemm_kernels.c
Go to the documentation of this file.
1/**
2 * @file gemm_kernels.c
3 * @brief General matrix multiply (GEMM) kernels with SIMD (SSE/AVX/AVX512)
4 *
5 * CK-ENGINE KERNEL RULES:
6 * =======================
7 * 1. NO malloc/free - memory via bump allocator, pointers passed in
8 * 2. NO OpenMP - parallelization at orchestrator/codegen layer
9 * 3. API must define: inputs, outputs, workspace, and memory layouts
10 * 4. Pure computation - deterministic, no side effects
11 *
12 * After changes: make test && make llamacpp-parity-full
13 *
14 * LEGACY EXCEPTION: This file contains OpenMP for backward compatibility.
15 * New kernels should NOT use OpenMP internally.
16 *
17 * GEMM: C = alpha * A @ B + beta * C (with optional bias)
18 */
19
20#include "ckernel_engine.h"
21#include "ck_threadpool.h"
22#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
23#include <immintrin.h>
24#endif
25#ifdef _OPENMP
26#include <omp.h>
27#endif
28#include <stdlib.h>
29#include <string.h>
30
31static inline int ck_min(int a, int b) { return a < b ? a : b; }
32
33static inline void ck_gemm_add_bias(float *C, const float *bias, int M, int N)
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}
46
47// AVX1 horizontal sum helper (no _mm256_reduce_add_ps in AVX1)
48#if defined(__AVX__) && !defined(__AVX512F__)
49static inline float hsum256_ps(__m256 v) {
50 // Sum upper and lower 128-bit lanes
51 __m128 lo = _mm256_castps256_ps128(v);
52 __m128 hi = _mm256_extractf128_ps(v, 1);
53 __m128 sum128 = _mm_add_ps(lo, hi);
54 // Horizontal add within 128-bit
55 __m128 shuf = _mm_movehdup_ps(sum128); // [1,1,3,3]
56 __m128 sums = _mm_add_ps(sum128, shuf); // [0+1,1+1,2+3,3+3]
57 shuf = _mm_movehl_ps(shuf, sums); // [2+3,3+3,...]
58 sums = _mm_add_ss(sums, shuf); // [0+1+2+3,...]
59 return _mm_cvtss_f32(sums);
60}
61#endif
62
63// Fast path for M=1: parallelize across output channels (j).
64// This is the common decode-time shape (matrix-vector) and is otherwise single-threaded
65// in the blocked GEMM code because M=1 provides no parallelism on the row dimension.
66static void gemm_nt_matvec_parallel(const float *A, // [K]
67 const float *B, // [N x K] (row-major, transposed layout)
68 const float *bias, // [N] or NULL
69 float *C, // [N]
70 int N,
71 int K)
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}
111
112static void gemm_naive_serial_double(const float *A,
113 const float *B,
114 const float *bias,
115 float *C,
116 int M, int N, int K)
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}
128
129static void gemm_naive_serial_float(const float *A,
130 const float *B,
131 const float *bias,
132 float *C,
133 int M, int N, int K)
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}
145
146typedef struct {
147 const float *A;
148 const float *B;
149 const float *bias;
150 float *C;
151 int M;
152 int N;
153 int K;
154 int bias_before_reduction;
155} ck_gemm_nt_fp32_exact_args_t;
156
157static void ck_gemm_nt_fp32_exact_rows(int begin, int end, void *opaque)
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}
178
180 const float *B,
181 const float *bias,
182 float *C,
183 int M, int N, int K)
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}
204
205// Naive parallel GEMM (reference baseline) – copied from C-Transformer.
206void gemm_naive_parallel(const float *A,
207 const float *B,
208 const float *bias,
209 float *C,
210 int M, int N, int K)
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}
228
229// AVX-512 optimized GEMM with AVX1 fallback
230void gemm_avx512_parallel(const float *A,
231 const float *B,
232 const float *bias,
233 float *C,
234 int M, int N, int K)
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}
284
285// Cache-blocked GEMM with fine-grained parallelism and AVX1 fallback
286void gemm_fine_grained_parallel(const float *A,
287 const float *B,
288 const float *bias,
289 float *C,
290 int M, int N, int K)
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}
374
375// =============================================================================
376// GEMM_NN: C[M,N] = A[M,K] @ B[K,N] + bias[N]
377// B is stored row-major as [K,N] (no transpose)
378// Used for backward d_input computation: d_input = d_output @ W
379// =============================================================================
380
381static void gemm_nn_serial_double(const float *A,
382 const float *B,
383 const float *bias,
384 float *C,
385 int M, int N, int K)
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}
397
398void gemm_nn_parallel(const float *A,
399 const float *B,
400 const float *bias,
401 float *C,
402 int M, int N, int K)
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}
419
420void gemm_nn_avx512(const float *A,
421 const float *B,
422 const float *bias,
423 float *C,
424 int M, int N, int K)
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}
482
483/*
484 * Profiling duplicate of gemm_nn_avx512.
485 * Keep numerics and shape contract identical; this exists only so VTune/Advisor
486 * can attribute runtime to a distinct symbol during A/B experiments.
487 */
488#if defined(__GNUC__) && !defined(__INTEL_LLVM_COMPILER)
489__attribute__((noinline, noipa))
490#elif defined(__GNUC__)
491__attribute__((noinline))
492#endif
493void gemm_nn_avx512_probe(const float *A,
494 const float *B,
495 const float *bias,
496 float *C,
497 int M, int N, int K)
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}
555
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}
572
573/* Keep legacy symbol name for ABI stability.
574 * Actual ISA path is selected at compile time in gemm_nn_avx512().
575 * This wrapper avoids ISA-specific naming at call sites and in new code.
576 */
577void gemm_nn_simd(const float *A,
578 const float *B,
579 const float *bias,
580 float *C,
581 int M, int N, int K)
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}
589
590void gemm_nn_blocked(const float *A,
591 const float *B,
592 const float *bias,
593 float *C,
594 int M, int N, int K)
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}
662
663// =============================================================================
664// GEMM_TN: C[M,N] = A[K,M].T @ B[K,N] + bias[N]
665// A is stored row-major as [K,M], B is stored row-major as [K,N]
666// Used for backward d_W computation: d_W = d_output.T @ input
667// =============================================================================
668
669static void gemm_tn_serial_double(const float *A,
670 const float *B,
671 const float *bias,
672 float *C,
673 int M, int N, int K)
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}
686
687void gemm_tn_parallel(const float *A,
688 const float *B,
689 const float *bias,
690 float *C,
691 int M, int N, int K)
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}
708
709void gemm_tn_avx512(const float *A,
710 const float *B,
711 const float *bias,
712 float *C,
713 int M, int N, int K)
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}
768
769void gemm_tn_blocked(const float *A,
770 const float *B,
771 const float *bias,
772 float *C,
773 int M, int N, int K)
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}
841
842// =============================================================================
843// Original GEMM_NT: C[M,N] = A[M,K] @ B[N,K].T + bias[N]
844// B is stored row-major as [N,K] (transposed in the multiply)
845// =============================================================================
846
847// Serial cache-blocked GEMM with SIMD (AVX/AVX512).
848// Note: B is stored as [N x K] (transposed layout).
849void gemm_blocked_serial(const float *A,
850 const float *B,
851 const float *bias,
852 float *C,
853 int M, int N, int K)
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}
941
942/*
943 * FP32 provider for llama.cpp's production CPU graph.
944 *
945 * Multi-token F32 x F32 matmuls are accepted by llamafile_sgemm and accumulate
946 * each dot product in one native-width vector. A one-token matvec is rejected
947 * by that path and falls back to ggml_vec_dot_f32, which uses four independent
948 * vector accumulators before a fixed pairwise merge. The distinction is
949 * observable for Qwen3.5/3.6's narrow recurrent alpha/beta projections.
950 *
951 * Outputs are independent, so OpenMP changes ownership only, never arithmetic.
952 */
953#if defined(__AVX__) && !defined(__AVX512F__)
954static inline float ck_hsum256_llamafile(__m256 value)
955{
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);
962}
963#endif
964
966 const float *A, const float *B, const float *bias, float *C,
967 int M, int N, int K, int index)
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}
1048
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)
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}
1061
1063 const float *B,
1064 const float *bias,
1065 float *C,
1066 int M, int N, int K)
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}
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)
#define C(color)
Definition show_config.c:39
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)
uint32_t end
Definition utf8.c:215