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

SwiGLU activation kernels with SIMD (SSE/AVX/AVX512) More...

#include "bf16_utils.h"
#include "ckernel_engine.h"
#include "ckernel_quant.h"
#include <dlfcn.h>
#include <math.h>
#include <pthread.h>
#include <stddef.h>
#include <stdio.h>
#include <stdlib.h>

Go to the source code of this file.

Macros

#define _GNU_SOURCE
 

Functions

static float sigmoid_scalar_parity (float x)
 
void swiglu_backward (const float *input, const float *d_output, float *d_input, int tokens, int dim)
 
void swiglu_backward_exact (const float *input, const float *d_output, float *d_input, int tokens, int dim)
 
void swiglu_forward (const float *input, float *output, int tokens, int dim)
 
void swiglu_forward_exact (const float *input, float *output, int tokens, int dim)
 
void swiglu_forward_ggml (const float *input, float *output, int tokens, int dim)
 
void swiglu_forward_ggml_split (const float *gate, const float *up, float *output, int tokens, int dim)
 
void swiglu_forward_pytorch_bf16_storage (const float *input, float *output, int tokens, int dim)
 
void swiglu_forward_q8_k (const float *input, void *output_q8, int tokens, int dim)
 

Detailed Description

SwiGLU activation 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

SwiGLU: y = silu(gate) * up = (gate * sigmoid(gate)) * up

Definition in file swiglu_kernels.c.

Macro Definition Documentation

◆ _GNU_SOURCE

#define _GNU_SOURCE

Definition at line 18 of file swiglu_kernels.c.

Function Documentation

◆ sigmoid_scalar_parity()

static float sigmoid_scalar_parity ( float  x)
inlinestatic

Definition at line 39 of file swiglu_kernels.c.

40{
41 return 1.0f / (1.0f + expf(-x));
42}

Referenced by swiglu_backward_exact(), swiglu_forward_exact(), and swiglu_forward_q8_k().

◆ swiglu_backward()

void swiglu_backward ( const float *  input,
const float *  d_output,
float *  d_input,
int  tokens,
int  dim 
)

SwiGLU backward pass

Test:

test_swiglu.py::TestSwiGLUBackward::test_backward_tokens

test_swiglu.py::TestSwiGLUBackward::test_backward_single

test_parity.py::test_swiglu_backward_parity

Computes dGate and dUp given dY. dGate = dy * b * silu'(a), dUp = dy * silu(a)

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

Definition at line 386 of file swiglu_kernels.c.

391{
393 swiglu_backward_exact(input, d_output, d_input, tokens, dim);
394 return;
395 }
396
397 int T = tokens;
398 int D = dim;
399
400 for (int t = 0; t < T; ++t) {
401 const float *row = input + (size_t)t * (2 * D);
402 const float *dy_row = d_output + (size_t)t * D;
403 float *dx_row = d_input + (size_t)t * (2 * D);
404 int d = 0;
405
406#if defined(__AVX512F__)
407 // AVX-512: Process 16 floats at a time
408 __m512 one = _mm512_set1_ps(1.0f);
409 for (; d + 16 <= D; d += 16) {
410 __m512 a = _mm512_loadu_ps(&row[d]); // gate
411 __m512 b = _mm512_loadu_ps(&row[D + d]); // value
412 __m512 dy = _mm512_loadu_ps(&dy_row[d]);
413
414 __m512 s = sigmoid512_fast(a); // sigmoid(a)
415 __m512 silu = _mm512_mul_ps(a, s); // silu(a) = a * s
416 __m512 one_minus_s = _mm512_sub_ps(one, s);
417 __m512 inner = _mm512_fmadd_ps(a, one_minus_s, one); // 1 + a * (1 - s)
418 __m512 silu_prime = _mm512_mul_ps(s, inner); // s * (1 + a * (1 - s))
419
420 // dA = dy * b * silu_prime
421 __m512 dA = _mm512_mul_ps(dy, _mm512_mul_ps(b, silu_prime));
422 // dB = dy * silu
423 __m512 dB = _mm512_mul_ps(dy, silu);
424
425 _mm512_storeu_ps(&dx_row[d], dA);
426 _mm512_storeu_ps(&dx_row[D + d], dB);
427 }
428#elif defined(__AVX2__)
429 // AVX2: Process 8 floats at a time
430 __m256 one = _mm256_set1_ps(1.0f);
431 for (; d + 8 <= D; d += 8) {
432 __m256 a = _mm256_loadu_ps(&row[d]); // gate
433 __m256 b = _mm256_loadu_ps(&row[D + d]); // value
434 __m256 dy = _mm256_loadu_ps(&dy_row[d]);
435
436 __m256 s = sigmoid256_fast(a); // sigmoid(a)
437 __m256 silu = _mm256_mul_ps(a, s); // silu(a) = a * s
438 __m256 one_minus_s = _mm256_sub_ps(one, s);
439 __m256 inner = _mm256_fmadd_ps(a, one_minus_s, one); // 1 + a * (1 - s)
440 __m256 silu_prime = _mm256_mul_ps(s, inner); // s * (1 + a * (1 - s))
441
442 // dA = dy * b * silu_prime
443 __m256 dA = _mm256_mul_ps(dy, _mm256_mul_ps(b, silu_prime));
444 // dB = dy * silu
445 __m256 dB = _mm256_mul_ps(dy, silu);
446
447 _mm256_storeu_ps(&dx_row[d], dA);
448 _mm256_storeu_ps(&dx_row[D + d], dB);
449 }
450#elif defined(__AVX__)
451 // AVX1: Vectorize arithmetic, use scalar sigmoid
452 __m256 one = _mm256_set1_ps(1.0f);
453 float a_arr[8] __attribute__((aligned(32)));
454 float s_arr[8] __attribute__((aligned(32)));
455
456 for (; d + 8 <= D; d += 8) {
457 __m256 a = _mm256_loadu_ps(&row[d]); // gate
458 __m256 b = _mm256_loadu_ps(&row[D + d]); // value
459 __m256 dy = _mm256_loadu_ps(&dy_row[d]);
460
461 // Compute sigmoid scalarly
462 _mm256_store_ps(a_arr, a);
463 for (int j = 0; j < 8; ++j) {
464 s_arr[j] = sigmoid_scalar(a_arr[j]);
465 }
466 __m256 s = _mm256_load_ps(s_arr);
467
468 __m256 silu = _mm256_mul_ps(a, s); // silu(a) = a * s
469 __m256 one_minus_s = _mm256_sub_ps(one, s);
470 __m256 a_one_minus_s = _mm256_mul_ps(a, one_minus_s);
471 __m256 inner = _mm256_add_ps(one, a_one_minus_s); // 1 + a * (1 - s)
472 __m256 silu_prime = _mm256_mul_ps(s, inner); // s * (1 + a * (1 - s))
473
474 // dA = dy * b * silu_prime
475 __m256 dA = _mm256_mul_ps(dy, _mm256_mul_ps(b, silu_prime));
476 // dB = dy * silu
477 __m256 dB = _mm256_mul_ps(dy, silu);
478
479 _mm256_storeu_ps(&dx_row[d], dA);
480 _mm256_storeu_ps(&dx_row[D + d], dB);
481 }
482#endif
483
484 // Scalar fallback for remaining elements
485 for (; d < D; ++d) {
486 float a = row[d]; // gate
487 float b = row[D + d]; // value
488 float dy = dy_row[d];
489
490 float s = sigmoid_scalar(a); // sigmoid(a)
491 float silu = a * s; // silu(a)
492 float silu_prime = s * (1.0f + a * (1.0f - s)); // silu'(a), PyTorch form
493
494 float dA = dy * b * silu_prime;
495 float dB = dy * silu;
496
497 dx_row[d] = dA;
498 dx_row[D + d] = dB;
499 }
500 }
501}
float sigmoid_scalar(float x)
int ck_strict_parity_enabled(void)
void swiglu_backward_exact(const float *input, const float *d_output, float *d_input, int tokens, int dim)
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)
static void silu(float *x, int n)

References __attribute__(), ck_strict_parity_enabled(), sigmoid_scalar(), silu(), and swiglu_backward_exact().

Referenced by ck_layer_backward_rmsnorm_swiglu().

◆ swiglu_backward_exact()

void swiglu_backward_exact ( const float *  input,
const float *  d_output,
float *  d_input,
int  tokens,
int  dim 
)

SwiGLU backward pass (exact version using stdlib sigmoid)

Test:

test_swiglu.py::TestSwiGLUBackward::test_exact_vs_fast

test_swiglu.py::TestSwiGLUBackward::test_exact_single

Uses standard library expf for numerical accuracy reference.

After changes: make test

Definition at line 719 of file swiglu_kernels.c.

724{
725 int T = tokens;
726 int D = dim;
727
728 for (int t = 0; t < T; ++t) {
729 const float *row = input + (size_t)t * (2 * D);
730 const float *dy_row = d_output + (size_t)t * D;
731 float *dx_row = d_input + (size_t)t * (2 * D);
732
733 for (int d = 0; d < D; ++d) {
734 float a = row[d]; // gate
735 float b = row[D + d]; // value
736 float dy = dy_row[d];
737
738 float s = sigmoid_scalar_parity(a); // sigmoid(a)
739 float silu = a * s; // silu(a)
740 float silu_prime = s * (1.0f + a * (1.0f - s)); // silu'(a), PyTorch form
741
742 float dA = dy * b * silu_prime;
743 float dB = dy * silu;
744
745 dx_row[d] = dA;
746 dx_row[D + d] = dB;
747 }
748 }
749}
static float sigmoid_scalar_parity(float x)

References sigmoid_scalar_parity(), and silu().

Referenced by swiglu_backward().

◆ swiglu_forward()

void swiglu_forward ( const float *  input,
float *  output,
int  tokens,
int  dim 
)

SwiGLU forward pass

Test:

test_swiglu.py::TestSwiGLUForward::test_forward_tokens

test_swiglu.py::TestSwiGLUForward::test_forward_single

test_mlp.py::TestMLPForward::test_swiglu_mlp

test_fused_swiglu_decode.py::TestFusedSwiGLUDecode::test_fused_swiglu_decode

test_parity.py::test_swiglu_parity

SwiGLU: y = silu(gate) * up where silu(x) = x * sigmoid(x)

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

Definition at line 237 of file swiglu_kernels.c.

241{
242 const char *fast_env = getenv("CK_SWIGLU_FAST");
243 const char *exact_env = getenv("CK_SWIGLU_EXACT");
245 !(fast_env && atoi(fast_env) != 0) ||
246 (exact_env && atoi(exact_env) != 0)) {
247 swiglu_forward_exact(input, output, tokens, dim);
248 return;
249 }
250
251 int T = tokens;
252 int D = dim;
253
254 for (int t = 0; t < T; ++t) {
255 const float *row = input + (size_t)t * (2 * D);
256 float *out_row = output + (size_t)t * D;
257 int d = 0;
258
259#if defined(__AVX512F__)
260 // AVX-512: Process 16 floats at a time
261 for (; d + 16 <= D; d += 16) {
262 __m512 a = _mm512_loadu_ps(&row[d]); // gate
263 __m512 b = _mm512_loadu_ps(&row[D + d]); // value
264
265 __m512 s = sigmoid512_fast(a); // sigmoid(a)
266 __m512 silu = _mm512_mul_ps(a, s); // silu(a) = a * sigmoid(a)
267 __m512 y = _mm512_mul_ps(silu, b); // y = silu(a) * b
268
269 _mm512_storeu_ps(&out_row[d], y);
270 }
271#elif defined(__AVX2__)
272 // AVX2: Process 8 floats at a time
273 for (; d + 8 <= D; d += 8) {
274 __m256 a = _mm256_loadu_ps(&row[d]); // gate
275 __m256 b = _mm256_loadu_ps(&row[D + d]); // value
276
277 __m256 s = sigmoid256_fast(a); // sigmoid(a)
278 __m256 silu = _mm256_mul_ps(a, s); // silu(a) = a * sigmoid(a)
279 __m256 y = _mm256_mul_ps(silu, b); // y = silu(a) * b
280
281 _mm256_storeu_ps(&out_row[d], y);
282 }
283#elif defined(__AVX__)
284 // AVX1: Vectorize arithmetic, use scalar sigmoid
285 float a_arr[8] __attribute__((aligned(32)));
286 float s_arr[8] __attribute__((aligned(32)));
287
288 for (; d + 8 <= D; d += 8) {
289 __m256 a = _mm256_loadu_ps(&row[d]); // gate
290 __m256 b = _mm256_loadu_ps(&row[D + d]); // value
291
292 // Compute sigmoid scalarly
293 _mm256_store_ps(a_arr, a);
294 for (int j = 0; j < 8; ++j) {
295 s_arr[j] = sigmoid_scalar(a_arr[j]);
296 }
297 __m256 s = _mm256_load_ps(s_arr);
298
299 __m256 silu = _mm256_mul_ps(a, s); // silu(a) = a * sigmoid(a)
300 __m256 y = _mm256_mul_ps(silu, b); // y = silu(a) * b
301
302 _mm256_storeu_ps(&out_row[d], y);
303 }
304#endif
305
306 // Scalar fallback for remaining elements
307 for (; d < D; ++d) {
308 float a = row[d]; // gate
309 float b = row[D + d]; // value
310
311 float s = sigmoid_scalar(a); // sigmoid(a)
312 float silu = a * s; // silu(a) = a * sigmoid(a)
313
314 out_row[d] = silu * b;
315 }
316 }
317}
void swiglu_forward_exact(const float *input, float *output, int tokens, int dim)

References __attribute__(), ck_strict_parity_enabled(), sigmoid_scalar(), silu(), and swiglu_forward_exact().

Referenced by ck_mlp_swiglu_forward(), ck_mlp_swiglu_forward_q4_k(), ck_mlp_swiglu_forward_q4_k_q8_k(), ck_mlp_swiglu_forward_q4_k_q8_k_prefill(), ck_mlp_swiglu_forward_quant(), ck_mlp_swiglu_forward_ref(), ck_test_swiglu(), model_layer_0_decode(), model_layer_0_decode(), model_layer_0_decode(), model_layer_0_decode(), model_layer_0_prefill(), model_layer_0_prefill(), model_layer_0_prefill(), model_layer_0_prefill(), model_layer_10_decode(), model_layer_10_decode(), model_layer_10_decode(), model_layer_10_decode(), model_layer_10_prefill(), model_layer_10_prefill(), model_layer_10_prefill(), model_layer_10_prefill(), model_layer_11_decode(), model_layer_11_decode(), model_layer_11_decode(), model_layer_11_decode(), model_layer_11_prefill(), model_layer_11_prefill(), model_layer_11_prefill(), model_layer_11_prefill(), model_layer_12_decode(), model_layer_12_decode(), model_layer_12_decode(), model_layer_12_decode(), model_layer_12_prefill(), model_layer_12_prefill(), model_layer_12_prefill(), model_layer_12_prefill(), model_layer_13_decode(), model_layer_13_decode(), model_layer_13_decode(), model_layer_13_decode(), model_layer_13_prefill(), model_layer_13_prefill(), model_layer_13_prefill(), model_layer_13_prefill(), model_layer_14_decode(), model_layer_14_decode(), model_layer_14_decode(), model_layer_14_decode(), model_layer_14_prefill(), model_layer_14_prefill(), model_layer_14_prefill(), model_layer_14_prefill(), model_layer_15_decode(), model_layer_15_decode(), model_layer_15_decode(), model_layer_15_decode(), model_layer_15_prefill(), model_layer_15_prefill(), model_layer_15_prefill(), model_layer_15_prefill(), model_layer_16_decode(), model_layer_16_decode(), model_layer_16_decode(), model_layer_16_decode(), model_layer_16_prefill(), model_layer_16_prefill(), model_layer_16_prefill(), model_layer_16_prefill(), model_layer_17_decode(), model_layer_17_decode(), model_layer_17_decode(), model_layer_17_decode(), model_layer_17_prefill(), model_layer_17_prefill(), model_layer_17_prefill(), model_layer_17_prefill(), model_layer_18_decode(), model_layer_18_decode(), model_layer_18_decode(), model_layer_18_decode(), model_layer_18_prefill(), model_layer_18_prefill(), model_layer_18_prefill(), model_layer_18_prefill(), model_layer_19_decode(), model_layer_19_decode(), model_layer_19_decode(), model_layer_19_decode(), model_layer_19_prefill(), model_layer_19_prefill(), model_layer_19_prefill(), model_layer_19_prefill(), model_layer_1_decode(), model_layer_1_decode(), model_layer_1_decode(), model_layer_1_decode(), model_layer_1_prefill(), model_layer_1_prefill(), model_layer_1_prefill(), model_layer_1_prefill(), model_layer_20_decode(), model_layer_20_decode(), model_layer_20_decode(), model_layer_20_decode(), model_layer_20_prefill(), model_layer_20_prefill(), model_layer_20_prefill(), model_layer_20_prefill(), model_layer_21_decode(), model_layer_21_decode(), model_layer_21_decode(), model_layer_21_decode(), model_layer_21_prefill(), model_layer_21_prefill(), model_layer_21_prefill(), model_layer_21_prefill(), model_layer_22_decode(), model_layer_22_decode(), model_layer_22_decode(), model_layer_22_decode(), model_layer_22_prefill(), model_layer_22_prefill(), model_layer_22_prefill(), model_layer_22_prefill(), model_layer_23_decode(), model_layer_23_decode(), model_layer_23_decode(), model_layer_23_decode(), model_layer_23_prefill(), model_layer_23_prefill(), model_layer_23_prefill(), model_layer_23_prefill(), model_layer_2_decode(), model_layer_2_decode(), model_layer_2_decode(), model_layer_2_decode(), model_layer_2_prefill(), model_layer_2_prefill(), model_layer_2_prefill(), model_layer_2_prefill(), model_layer_3_decode(), model_layer_3_decode(), model_layer_3_decode(), model_layer_3_decode(), model_layer_3_prefill(), model_layer_3_prefill(), model_layer_3_prefill(), model_layer_3_prefill(), model_layer_4_decode(), model_layer_4_decode(), model_layer_4_decode(), model_layer_4_decode(), model_layer_4_prefill(), model_layer_4_prefill(), model_layer_4_prefill(), model_layer_4_prefill(), model_layer_5_decode(), model_layer_5_decode(), model_layer_5_decode(), model_layer_5_decode(), model_layer_5_prefill(), model_layer_5_prefill(), model_layer_5_prefill(), model_layer_5_prefill(), model_layer_6_decode(), model_layer_6_decode(), model_layer_6_decode(), model_layer_6_decode(), model_layer_6_prefill(), model_layer_6_prefill(), model_layer_6_prefill(), model_layer_6_prefill(), model_layer_7_decode(), model_layer_7_decode(), model_layer_7_decode(), model_layer_7_decode(), model_layer_7_prefill(), model_layer_7_prefill(), model_layer_7_prefill(), model_layer_7_prefill(), model_layer_8_decode(), model_layer_8_decode(), model_layer_8_decode(), model_layer_8_decode(), model_layer_8_prefill(), model_layer_8_prefill(), model_layer_8_prefill(), model_layer_8_prefill(), model_layer_9_decode(), model_layer_9_decode(), model_layer_9_decode(), model_layer_9_decode(), model_layer_9_prefill(), model_layer_9_prefill(), model_layer_9_prefill(), model_layer_9_prefill(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_decode(), qwen2_0_5b_decode_layer_0_prefill(), qwen2_0_5b_decode_layer_0_prefill(), qwen2_0_5b_decode_layer_0_prefill(), qwen2_0_5b_decode_layer_0_prefill(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_decode(), qwen2_0_5b_decode_layer_10_prefill(), qwen2_0_5b_decode_layer_10_prefill(), qwen2_0_5b_decode_layer_10_prefill(), qwen2_0_5b_decode_layer_10_prefill(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_decode(), qwen2_0_5b_decode_layer_11_prefill(), qwen2_0_5b_decode_layer_11_prefill(), qwen2_0_5b_decode_layer_11_prefill(), qwen2_0_5b_decode_layer_11_prefill(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_decode(), qwen2_0_5b_decode_layer_12_prefill(), qwen2_0_5b_decode_layer_12_prefill(), qwen2_0_5b_decode_layer_12_prefill(), qwen2_0_5b_decode_layer_12_prefill(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_decode(), qwen2_0_5b_decode_layer_13_prefill(), qwen2_0_5b_decode_layer_13_prefill(), qwen2_0_5b_decode_layer_13_prefill(), qwen2_0_5b_decode_layer_13_prefill(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_decode(), qwen2_0_5b_decode_layer_14_prefill(), qwen2_0_5b_decode_layer_14_prefill(), qwen2_0_5b_decode_layer_14_prefill(), qwen2_0_5b_decode_layer_14_prefill(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_decode(), qwen2_0_5b_decode_layer_15_prefill(), qwen2_0_5b_decode_layer_15_prefill(), qwen2_0_5b_decode_layer_15_prefill(), qwen2_0_5b_decode_layer_15_prefill(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_decode(), qwen2_0_5b_decode_layer_16_prefill(), qwen2_0_5b_decode_layer_16_prefill(), qwen2_0_5b_decode_layer_16_prefill(), qwen2_0_5b_decode_layer_16_prefill(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_decode(), qwen2_0_5b_decode_layer_17_prefill(), qwen2_0_5b_decode_layer_17_prefill(), qwen2_0_5b_decode_layer_17_prefill(), qwen2_0_5b_decode_layer_17_prefill(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_decode(), qwen2_0_5b_decode_layer_18_prefill(), qwen2_0_5b_decode_layer_18_prefill(), qwen2_0_5b_decode_layer_18_prefill(), qwen2_0_5b_decode_layer_18_prefill(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_decode(), qwen2_0_5b_decode_layer_19_prefill(), qwen2_0_5b_decode_layer_19_prefill(), qwen2_0_5b_decode_layer_19_prefill(), qwen2_0_5b_decode_layer_19_prefill(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_decode(), qwen2_0_5b_decode_layer_1_prefill(), qwen2_0_5b_decode_layer_1_prefill(), qwen2_0_5b_decode_layer_1_prefill(), qwen2_0_5b_decode_layer_1_prefill(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_decode(), qwen2_0_5b_decode_layer_20_prefill(), qwen2_0_5b_decode_layer_20_prefill(), qwen2_0_5b_decode_layer_20_prefill(), qwen2_0_5b_decode_layer_20_prefill(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_decode(), qwen2_0_5b_decode_layer_21_prefill(), qwen2_0_5b_decode_layer_21_prefill(), qwen2_0_5b_decode_layer_21_prefill(), qwen2_0_5b_decode_layer_21_prefill(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_decode(), qwen2_0_5b_decode_layer_22_prefill(), qwen2_0_5b_decode_layer_22_prefill(), qwen2_0_5b_decode_layer_22_prefill(), qwen2_0_5b_decode_layer_22_prefill(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_decode(), qwen2_0_5b_decode_layer_23_prefill(), qwen2_0_5b_decode_layer_23_prefill(), qwen2_0_5b_decode_layer_23_prefill(), qwen2_0_5b_decode_layer_23_prefill(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_decode(), qwen2_0_5b_decode_layer_2_prefill(), qwen2_0_5b_decode_layer_2_prefill(), qwen2_0_5b_decode_layer_2_prefill(), qwen2_0_5b_decode_layer_2_prefill(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_decode(), qwen2_0_5b_decode_layer_3_prefill(), qwen2_0_5b_decode_layer_3_prefill(), qwen2_0_5b_decode_layer_3_prefill(), qwen2_0_5b_decode_layer_3_prefill(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_decode(), qwen2_0_5b_decode_layer_4_prefill(), qwen2_0_5b_decode_layer_4_prefill(), qwen2_0_5b_decode_layer_4_prefill(), qwen2_0_5b_decode_layer_4_prefill(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_decode(), qwen2_0_5b_decode_layer_5_prefill(), qwen2_0_5b_decode_layer_5_prefill(), qwen2_0_5b_decode_layer_5_prefill(), qwen2_0_5b_decode_layer_5_prefill(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_decode(), qwen2_0_5b_decode_layer_6_prefill(), qwen2_0_5b_decode_layer_6_prefill(), qwen2_0_5b_decode_layer_6_prefill(), qwen2_0_5b_decode_layer_6_prefill(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_decode(), qwen2_0_5b_decode_layer_7_prefill(), qwen2_0_5b_decode_layer_7_prefill(), qwen2_0_5b_decode_layer_7_prefill(), qwen2_0_5b_decode_layer_7_prefill(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_decode(), qwen2_0_5b_decode_layer_8_prefill(), qwen2_0_5b_decode_layer_8_prefill(), qwen2_0_5b_decode_layer_8_prefill(), qwen2_0_5b_decode_layer_8_prefill(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_decode(), qwen2_0_5b_decode_layer_9_prefill(), qwen2_0_5b_decode_layer_9_prefill(), qwen2_0_5b_decode_layer_9_prefill(), and qwen2_0_5b_decode_layer_9_prefill().

◆ swiglu_forward_exact()

void swiglu_forward_exact ( const float *  input,
float *  output,
int  tokens,
int  dim 
)

SwiGLU forward pass (exact version using stdlib sigmoid)

Test:

test_swiglu.py::TestSwiGLUForward::test_exact_vs_fast

test_swiglu.py::TestSwiGLUForward::test_exact_single

Uses standard library expf for numerical accuracy reference.

After changes: make test

Definition at line 515 of file swiglu_kernels.c.

519{
520 int T = tokens;
521 int D = dim;
522
523 for (int t = 0; t < T; ++t) {
524 const float *row = input + (size_t)t * (2 * D);
525 float *out_row = output + (size_t)t * D;
526
527 for (int d = 0; d < D; ++d) {
528 float a = row[d]; // gate
529 float b = row[D + d]; // value
530
531 float s = sigmoid_scalar_parity(a); // sigmoid(a)
532 float silu = a * s; // silu(a)
533 out_row[d] = silu * b;
534 }
535 }
536}

References sigmoid_scalar_parity(), and silu().

Referenced by ck_mlp_swiglu_forward(), ck_mlp_swiglu_forward_ref(), and swiglu_forward().

◆ swiglu_forward_ggml()

void swiglu_forward_ggml ( const float *  input,
float *  output,
int  tokens,
int  dim 
)

Definition at line 538 of file swiglu_kernels.c.

542{
543 for (int t = 0; t < tokens; ++t) {
544 const float *row = input + (size_t)t * (2 * dim);
545 float *out_row = output + (size_t)t * dim;
546 int d = 0;
547
548#if defined(__AVX512F__) && defined(__AVX512DQ__)
549 for (; d + 16 <= dim; d += 16) {
550 const __m512 gate = _mm512_loadu_ps(row + d);
551 const __m512 up = _mm512_loadu_ps(row + dim + d);
552 const __m512 neg_gate = _mm512_sub_ps(_mm512_setzero_ps(), gate);
553 const __m512 denom = _mm512_add_ps(
554 _mm512_set1_ps(1.0f), ck_ggml_expf_avx512(neg_gate));
555 const __m512 silu = _mm512_div_ps(gate, denom);
556 _mm512_storeu_ps(out_row + d, _mm512_mul_ps(silu, up));
557 }
558#elif defined(__AVX2__) && defined(__FMA__)
559 for (; d + 8 <= dim; d += 8) {
560 const __m256 gate = _mm256_loadu_ps(row + d);
561 const __m256 up = _mm256_loadu_ps(row + dim + d);
562 const __m256 neg_gate = _mm256_sub_ps(_mm256_setzero_ps(), gate);
563 const __m256 denom = _mm256_add_ps(
564 _mm256_set1_ps(1.0f), ck_ggml_expf_avx2(neg_gate));
565 const __m256 silu = _mm256_div_ps(gate, denom);
566 _mm256_storeu_ps(out_row + d, _mm256_mul_ps(silu, up));
567 }
568#endif
569 for (; d < dim; ++d) {
570 const float gate = row[d];
571 out_row[d] = (gate / (1.0f + expf(-gate))) * row[dim + d];
572 }
573 }
574}

References silu().

Referenced by ck_moe_q4k_mixed_route_work(), ck_moe_q4k_q5k_bucket_work(), ck_moe_q4k_q5k_route_work(), moe_swiglu_expert_forward_q4k_q4k_workspace(), moe_swiglu_expert_forward_q4k_q5_0_workspace(), moe_swiglu_expert_forward_q4k_q5k_workspace(), moe_swiglu_expert_forward_q4k_q6k_workspace(), moe_swiglu_expert_forward_q4k_q8_0_workspace(), moe_swiglu_shared_forward_q4k_q4k_workspace(), moe_swiglu_shared_forward_q4k_q6k_workspace(), and moe_swiglu_shared_forward_q8_0_gated_workspace().

◆ swiglu_forward_ggml_split()

void swiglu_forward_ggml_split ( const float *  gate,
const float *  up,
float *  output,
int  tokens,
int  dim 
)

Definition at line 576 of file swiglu_kernels.c.

581{
582 if (!gate || !up || !output || tokens <= 0 || dim <= 0) {
583 return;
584 }
585 for (int t = 0; t < tokens; ++t) {
586 const float *gate_row = gate + (size_t)t * (size_t)dim;
587 const float *up_row = up + (size_t)t * (size_t)dim;
588 float *out_row = output + (size_t)t * (size_t)dim;
589 int d = 0;
590
591#if defined(__AVX512F__) && defined(__AVX512DQ__)
592 for (; d + 16 <= dim; d += 16) {
593 const __m512 gate_v = _mm512_loadu_ps(gate_row + d);
594 const __m512 up_v = _mm512_loadu_ps(up_row + d);
595 const __m512 neg_gate = _mm512_sub_ps(_mm512_setzero_ps(), gate_v);
596 const __m512 denom = _mm512_add_ps(
597 _mm512_set1_ps(1.0f), ck_ggml_expf_avx512(neg_gate));
598 const __m512 silu = _mm512_div_ps(gate_v, denom);
599 _mm512_storeu_ps(out_row + d, _mm512_mul_ps(silu, up_v));
600 }
601#elif defined(__AVX2__) && defined(__FMA__)
602 for (; d + 8 <= dim; d += 8) {
603 const __m256 gate_v = _mm256_loadu_ps(gate_row + d);
604 const __m256 up_v = _mm256_loadu_ps(up_row + d);
605 const __m256 neg_gate = _mm256_sub_ps(_mm256_setzero_ps(), gate_v);
606 const __m256 denom = _mm256_add_ps(
607 _mm256_set1_ps(1.0f), ck_ggml_expf_avx2(neg_gate));
608 const __m256 silu = _mm256_div_ps(gate_v, denom);
609 _mm256_storeu_ps(out_row + d, _mm256_mul_ps(silu, up_v));
610 }
611#endif
612 for (; d < dim; ++d) {
613 const float gate_v = gate_row[d];
614 out_row[d] = (gate_v / (1.0f + expf(-gate_v))) * up_row[d];
615 }
616 }
617}

References silu().

Referenced by ck_moe_q4k_q5k_bucket_work(), and ck_moe_shared_q4k_gated_workspace().

◆ swiglu_forward_pytorch_bf16_storage()

void swiglu_forward_pytorch_bf16_storage ( const float *  input,
float *  output,
int  tokens,
int  dim 
)

Definition at line 658 of file swiglu_kernels.c.

662{
663 if (!input || !output || tokens < 0 || dim < 0) {
664 fprintf(stderr, "HARD KERNEL CONTRACT FAULT: invalid PyTorch BF16 SwiGLU arguments\n");
665 abort();
666 }
667
668#if defined(__AVX512F__)
669 pthread_once(&ck_pytorch_swiglu_once, ck_bind_pytorch_swiglu_sleef);
670 if (!ck_pytorch_swiglu_expf16) {
671 fprintf(stderr,
672 "HARD KERNEL CONTRACT FAULT: PyTorch BF16 SwiGLU requires "
673 "SLEEF Sleef_expf16_u10; set CK_SLEEF_LIBRARY\n");
674 abort();
675 }
676#endif
677
678 for (int t = 0; t < tokens; ++t) {
679 const float *row = input + (size_t)t * (size_t)(2 * dim);
680 float *out_row = output + (size_t)t * (size_t)dim;
681 int d = 0;
682
683#if defined(__AVX512F__)
684 for (; d + 16 <= dim; d += 16) {
685 const __m512 gate = _mm512_loadu_ps(row + d);
686 const __m512 denominator = _mm512_add_ps(
687 _mm512_set1_ps(1.0f),
688 ck_pytorch_swiglu_expf16(_mm512_sub_ps(_mm512_setzero_ps(), gate)));
689 const __m512 silu = _mm512_div_ps(gate, denominator);
690 float silu_lanes[16] __attribute__((aligned(64)));
691 _mm512_store_ps(silu_lanes, silu);
692 for (int lane = 0; lane < 16; ++lane) {
693 const float silu_bf16 = bf16_to_float(float_to_bf16(silu_lanes[lane]));
694 const float up_bf16 = bf16_to_float(float_to_bf16(row[dim + d + lane]));
695 out_row[d + lane] = bf16_to_float(
696 float_to_bf16(silu_bf16 * up_bf16));
697 }
698 }
699#endif
700 for (; d < dim; ++d) {
701 const float gate_bf16 = bf16_to_float(float_to_bf16(row[d]));
702 const float up_bf16 = bf16_to_float(float_to_bf16(row[dim + d]));
703 const float silu = gate_bf16 / (1.0f + expf(-gate_bf16));
704 const float silu_bf16 = bf16_to_float(float_to_bf16(silu));
705 out_row[d] = bf16_to_float(float_to_bf16(silu_bf16 * up_bf16));
706 }
707 }
708}
static uint16_t float_to_bf16(float f)
Definition bf16_utils.h:90
static float bf16_to_float(uint16_t v)
Definition bf16_utils.h:38

References __attribute__(), bf16_to_float(), float_to_bf16(), and silu().

Referenced by moe_swiglu_packed_expert_forward_bf16().

◆ swiglu_forward_q8_k()

void swiglu_forward_q8_k ( const float *  input,
void *  output_q8,
int  tokens,
int  dim 
)

Definition at line 319 of file swiglu_kernels.c.

323{
324 if (!input || !output_q8 || tokens <= 0 || dim <= 0) {
325 return;
326 }
327 if ((dim % QK_K) != 0) {
328 return;
329 }
330
331 const char *fast_env = getenv("CK_SWIGLU_FAST");
332 const char *exact_env = getenv("CK_SWIGLU_EXACT");
333 const int use_fast = !ck_strict_parity_enabled() &&
334 (fast_env && atoi(fast_env) != 0) &&
335 !(exact_env && atoi(exact_env) != 0);
336
337 const int blocks_per_row = dim / QK_K;
338 block_q8_K *q8 = (block_q8_K *)output_q8;
339 float tmp[QK_K];
340
341 for (int t = 0; t < tokens; ++t) {
342 const float *row = input + (size_t)t * (size_t)(2 * dim);
343 block_q8_K *q8_row = q8 + (size_t)t * (size_t)blocks_per_row;
344
345 for (int block = 0; block < blocks_per_row; ++block) {
346 const int base = block * QK_K;
347 int d = 0;
348
349#if defined(__AVX2__)
350 if (use_fast) {
351 for (; d + 8 <= QK_K; d += 8) {
352 const __m256 a = _mm256_loadu_ps(row + base + d);
353 const __m256 b = _mm256_loadu_ps(row + dim + base + d);
354 const __m256 s = sigmoid256_fast(a);
355 const __m256 y = _mm256_mul_ps(_mm256_mul_ps(a, s), b);
356 _mm256_storeu_ps(tmp + d, y);
357 }
358 }
359#else
360 (void)use_fast;
361#endif
362
363 for (; d < QK_K; ++d) {
364 const float a = row[base + d];
365 const float b = row[dim + base + d];
366 const float s = use_fast ? sigmoid_scalar(a) : sigmoid_scalar_parity(a);
367 tmp[d] = (a * s) * b;
368 }
369
370 quantize_row_q8_k(tmp, (void *)&q8_row[block], QK_K);
371 }
372 }
373}
void quantize_row_q8_k(const float *x, void *y, int k)
#define QK_K

References ck_strict_parity_enabled(), QK_K, quantize_row_q8_k(), sigmoid_scalar(), and sigmoid_scalar_parity().