31#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
41 return 1.0f / (1.0f + expf(-x));
48#if defined(__AVX512F__)
50static inline __m512 exp512_fast(__m512 x) {
52 x = _mm512_max_ps(x, _mm512_set1_ps(-88.0f));
53 x = _mm512_min_ps(x, _mm512_set1_ps(88.0f));
56 const __m512 log2e = _mm512_set1_ps(1.4426950408889634f);
57 __m512 z = _mm512_mul_ps(x, log2e);
60 __m512 zf = _mm512_roundscale_ps(z, _MM_FROUND_TO_NEAREST_INT);
61 __m512 f = _mm512_sub_ps(z, zf);
64 const __m512 c0 = _mm512_set1_ps(1.0f);
65 const __m512 c1 = _mm512_set1_ps(0.6931471805599453f);
66 const __m512 c2 = _mm512_set1_ps(0.2402265069591007f);
67 const __m512 c3 = _mm512_set1_ps(0.05550410866482158f);
68 const __m512 c4 = _mm512_set1_ps(0.009618129107628478f);
70 __m512 poly = _mm512_fmadd_ps(f, c4, c3);
71 poly = _mm512_fmadd_ps(f, poly, c2);
72 poly = _mm512_fmadd_ps(f, poly, c1);
73 poly = _mm512_fmadd_ps(f, poly, c0);
76 __m512i zi = _mm512_cvtps_epi32(zf);
77 zi = _mm512_add_epi32(zi, _mm512_set1_epi32(127));
78 zi = _mm512_slli_epi32(zi, 23);
79 __m512 scale = _mm512_castsi512_ps(zi);
81 return _mm512_mul_ps(poly, scale);
85static inline __m512 sigmoid512_fast(__m512 x) {
86 __m512 neg_x = _mm512_sub_ps(_mm512_setzero_ps(), x);
87 __m512 exp_neg = exp512_fast(neg_x);
88 __m512 one = _mm512_set1_ps(1.0f);
89 return _mm512_div_ps(one, _mm512_add_ps(one, exp_neg));
95static inline __m256 exp256_fast(__m256 x) {
97 x = _mm256_max_ps(x, _mm256_set1_ps(-88.0f));
98 x = _mm256_min_ps(x, _mm256_set1_ps(88.0f));
101 const __m256 log2e = _mm256_set1_ps(1.4426950408889634f);
102 __m256 z = _mm256_mul_ps(x, log2e);
105 __m256 zf = _mm256_round_ps(z, _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC);
106 __m256 f = _mm256_sub_ps(z, zf);
109 const __m256 c0 = _mm256_set1_ps(1.0f);
110 const __m256 c1 = _mm256_set1_ps(0.6931471805599453f);
111 const __m256 c2 = _mm256_set1_ps(0.2402265069591007f);
112 const __m256 c3 = _mm256_set1_ps(0.05550410866482158f);
113 const __m256 c4 = _mm256_set1_ps(0.009618129107628478f);
115 __m256 poly = _mm256_fmadd_ps(f, c4, c3);
116 poly = _mm256_fmadd_ps(f, poly, c2);
117 poly = _mm256_fmadd_ps(f, poly, c1);
118 poly = _mm256_fmadd_ps(f, poly, c0);
121 __m256i zi = _mm256_cvtps_epi32(zf);
122 zi = _mm256_add_epi32(zi, _mm256_set1_epi32(127));
123 zi = _mm256_slli_epi32(zi, 23);
124 __m256 scale = _mm256_castsi256_ps(zi);
126 return _mm256_mul_ps(poly, scale);
130static inline __m256 sigmoid256_fast(__m256 x) {
131 __m256 neg_x = _mm256_sub_ps(_mm256_setzero_ps(), x);
132 __m256 exp_neg = exp256_fast(neg_x);
133 __m256 one = _mm256_set1_ps(1.0f);
134 return _mm256_div_ps(one, _mm256_add_ps(one, exp_neg));
138static inline __m256 ck_ggml_expf_avx2(__m256 x) {
139 const __m256 r = _mm256_set1_ps(0x1.8p23f);
140 const __m256 z = _mm256_fmadd_ps(x, _mm256_set1_ps(0x1.715476p+0f), r);
141 const __m256 n = _mm256_sub_ps(z, r);
142 const __m256 b = _mm256_fnmadd_ps(
144 _mm256_set1_ps(0x1.7f7d1cp-20f),
145 _mm256_fnmadd_ps(n, _mm256_set1_ps(0x1.62e4p-1f), x));
146 const __m256i e = _mm256_slli_epi32(_mm256_castps_si256(z), 23);
147 const __m256 k = _mm256_castsi256_ps(
148 _mm256_add_epi32(e, _mm256_castps_si256(_mm256_set1_ps(1))));
149 const __m256i c = _mm256_castps_si256(
150 _mm256_cmp_ps(_mm256_andnot_ps(_mm256_set1_ps(-0.0f), n),
151 _mm256_set1_ps(126), _CMP_GT_OQ));
152 const __m256 u = _mm256_mul_ps(b, b);
153 const __m256 j = _mm256_fmadd_ps(
155 _mm256_fmadd_ps(_mm256_set1_ps(0x1.0e4020p-7f), b,
156 _mm256_set1_ps(0x1.573e2ep-5f)),
158 _mm256_fmadd_ps(_mm256_set1_ps(0x1.555e66p-3f), b,
159 _mm256_set1_ps(0x1.fffdb6p-2f))),
161 _mm256_mul_ps(_mm256_set1_ps(0x1.ffffecp-1f), b));
162 if (!_mm256_movemask_ps(_mm256_castsi256_ps(c))) {
163 return _mm256_fmadd_ps(j, k, k);
165 const __m256i g = _mm256_and_si256(
166 _mm256_castps_si256(_mm256_cmp_ps(n, _mm256_setzero_ps(), _CMP_LE_OQ)),
167 _mm256_set1_epi32((
int)0x82000000u));
168 const __m256 s1 = _mm256_castsi256_ps(
169 _mm256_add_epi32(g, _mm256_set1_epi32(0x7f000000u)));
170 const __m256 s2 = _mm256_castsi256_ps(_mm256_sub_epi32(e, g));
171 const __m256i d = _mm256_castps_si256(
172 _mm256_cmp_ps(_mm256_andnot_ps(_mm256_set1_ps(-0.0f), n),
173 _mm256_set1_ps(192), _CMP_GT_OQ));
175 _mm256_and_ps(_mm256_castsi256_ps(d), _mm256_mul_ps(s1, s1)),
177 _mm256_castsi256_ps(d),
179 _mm256_and_ps(_mm256_castsi256_ps(c),
180 _mm256_mul_ps(_mm256_fmadd_ps(s2, j, s2), s1)),
181 _mm256_andnot_ps(_mm256_castsi256_ps(c),
182 _mm256_fmadd_ps(k, j, k)))));
186#if defined(__AVX512F__) && defined(__AVX512DQ__)
189static inline __m512 ck_ggml_expf_avx512(__m512 x) {
190 const __m512 r = _mm512_set1_ps(0x1.8p23f);
191 const __m512 z = _mm512_fmadd_ps(x, _mm512_set1_ps(0x1.715476p+0f), r);
192 const __m512 n = _mm512_sub_ps(z, r);
193 const __m512 b = _mm512_fnmadd_ps(
194 n, _mm512_set1_ps(0x1.7f7d1cp-20f),
195 _mm512_fnmadd_ps(n, _mm512_set1_ps(0x1.62e4p-1f), x));
196 const __mmask16 d = _mm512_cmp_ps_mask(
197 _mm512_abs_ps(n), _mm512_set1_ps(192.0f), _CMP_GT_OQ);
198 const __m512 u = _mm512_mul_ps(b, b);
199 const __m512 j = _mm512_fmadd_ps(
202 _mm512_set1_ps(0x1.0e4020p-7f), b,
203 _mm512_set1_ps(0x1.573e2ep-5f)),
206 _mm512_set1_ps(0x1.555e66p-3f), b,
207 _mm512_set1_ps(0x1.fffdb6p-2f))),
210 _mm512_set1_ps(0x1.ffffecp-1f), b,
211 _mm512_set1_ps(1.0f)));
212 const __m512 res = _mm512_scalef_ps(j, n);
213 if (_mm512_kortestz(d, d)) {
216 const __m512 zero = _mm512_setzero_ps();
217 const __m512 alt = _mm512_mask_blend_ps(
218 _mm512_cmp_ps_mask(n, zero, _CMP_LE_OQ),
219 _mm512_set1_ps(INFINITY),
221 return _mm512_mask_blend_ps(d, res, alt);
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)) {
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;
259#if defined(__AVX512F__)
261 for (; d + 16 <= D; d += 16) {
262 __m512 a = _mm512_loadu_ps(&row[d]);
263 __m512 b = _mm512_loadu_ps(&row[D + d]);
265 __m512 s = sigmoid512_fast(a);
266 __m512
silu = _mm512_mul_ps(a, s);
267 __m512 y = _mm512_mul_ps(
silu, b);
269 _mm512_storeu_ps(&out_row[d], y);
271#elif defined(__AVX2__)
273 for (; d + 8 <= D; d += 8) {
274 __m256 a = _mm256_loadu_ps(&row[d]);
275 __m256 b = _mm256_loadu_ps(&row[D + d]);
277 __m256 s = sigmoid256_fast(a);
278 __m256
silu = _mm256_mul_ps(a, s);
279 __m256 y = _mm256_mul_ps(
silu, b);
281 _mm256_storeu_ps(&out_row[d], y);
283#elif defined(__AVX__)
288 for (; d + 8 <= D; d += 8) {
289 __m256 a = _mm256_loadu_ps(&row[d]);
290 __m256 b = _mm256_loadu_ps(&row[D + d]);
293 _mm256_store_ps(a_arr, a);
294 for (
int j = 0; j < 8; ++j) {
297 __m256 s = _mm256_load_ps(s_arr);
299 __m256
silu = _mm256_mul_ps(a, s);
300 __m256 y = _mm256_mul_ps(
silu, b);
302 _mm256_storeu_ps(&out_row[d], y);
309 float b = row[D + d];
314 out_row[d] =
silu * b;
324 if (!input || !output_q8 || tokens <= 0 || dim <= 0) {
327 if ((dim %
QK_K) != 0) {
331 const char *fast_env = getenv(
"CK_SWIGLU_FAST");
332 const char *exact_env = getenv(
"CK_SWIGLU_EXACT");
334 (fast_env && atoi(fast_env) != 0) &&
335 !(exact_env && atoi(exact_env) != 0);
337 const int blocks_per_row = dim /
QK_K;
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;
345 for (
int block = 0; block < blocks_per_row; ++block) {
346 const int base = block *
QK_K;
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);
363 for (; d <
QK_K; ++d) {
364 const float a = row[base + d];
365 const float b = row[dim + base + d];
367 tmp[d] = (a * s) * b;
387 const float *d_output,
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);
406#if defined(__AVX512F__)
408 __m512 one = _mm512_set1_ps(1.0f);
409 for (; d + 16 <= D; d += 16) {
410 __m512 a = _mm512_loadu_ps(&row[d]);
411 __m512 b = _mm512_loadu_ps(&row[D + d]);
412 __m512 dy = _mm512_loadu_ps(&dy_row[d]);
414 __m512 s = sigmoid512_fast(a);
415 __m512
silu = _mm512_mul_ps(a, s);
416 __m512 one_minus_s = _mm512_sub_ps(one, s);
417 __m512 inner = _mm512_fmadd_ps(a, one_minus_s, one);
418 __m512 silu_prime = _mm512_mul_ps(s, inner);
421 __m512 dA = _mm512_mul_ps(dy, _mm512_mul_ps(b, silu_prime));
423 __m512 dB = _mm512_mul_ps(dy,
silu);
425 _mm512_storeu_ps(&dx_row[d], dA);
426 _mm512_storeu_ps(&dx_row[D + d], dB);
428#elif defined(__AVX2__)
430 __m256 one = _mm256_set1_ps(1.0f);
431 for (; d + 8 <= D; d += 8) {
432 __m256 a = _mm256_loadu_ps(&row[d]);
433 __m256 b = _mm256_loadu_ps(&row[D + d]);
434 __m256 dy = _mm256_loadu_ps(&dy_row[d]);
436 __m256 s = sigmoid256_fast(a);
437 __m256
silu = _mm256_mul_ps(a, s);
438 __m256 one_minus_s = _mm256_sub_ps(one, s);
439 __m256 inner = _mm256_fmadd_ps(a, one_minus_s, one);
440 __m256 silu_prime = _mm256_mul_ps(s, inner);
443 __m256 dA = _mm256_mul_ps(dy, _mm256_mul_ps(b, silu_prime));
445 __m256 dB = _mm256_mul_ps(dy,
silu);
447 _mm256_storeu_ps(&dx_row[d], dA);
448 _mm256_storeu_ps(&dx_row[D + d], dB);
450#elif defined(__AVX__)
452 __m256 one = _mm256_set1_ps(1.0f);
456 for (; d + 8 <= D; d += 8) {
457 __m256 a = _mm256_loadu_ps(&row[d]);
458 __m256 b = _mm256_loadu_ps(&row[D + d]);
459 __m256 dy = _mm256_loadu_ps(&dy_row[d]);
462 _mm256_store_ps(a_arr, a);
463 for (
int j = 0; j < 8; ++j) {
466 __m256 s = _mm256_load_ps(s_arr);
468 __m256
silu = _mm256_mul_ps(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);
472 __m256 silu_prime = _mm256_mul_ps(s, inner);
475 __m256 dA = _mm256_mul_ps(dy, _mm256_mul_ps(b, silu_prime));
477 __m256 dB = _mm256_mul_ps(dy,
silu);
479 _mm256_storeu_ps(&dx_row[d], dA);
480 _mm256_storeu_ps(&dx_row[D + d], dB);
487 float b = row[D + d];
488 float dy = dy_row[d];
492 float silu_prime = s * (1.0f + a * (1.0f - s));
494 float dA = dy * b * silu_prime;
495 float dB = dy *
silu;
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;
527 for (
int d = 0; d < D; ++d) {
529 float b = row[D + d];
533 out_row[d] =
silu * b;
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;
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));
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));
569 for (; d < dim; ++d) {
570 const float gate = row[d];
571 out_row[d] = (gate / (1.0f + expf(-gate))) * row[dim + d];
582 if (!gate || !up || !output || tokens <= 0 || dim <= 0) {
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;
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));
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));
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];
619#if defined(__AVX512F__)
620typedef __m512 (*ck_sleef_expf16_fn)(__m512);
622static ck_sleef_expf16_fn ck_pytorch_swiglu_expf16 = NULL;
623static void *ck_pytorch_swiglu_sleef_handle = NULL;
624static pthread_once_t ck_pytorch_swiglu_once = PTHREAD_ONCE_INIT;
626static void ck_bind_pytorch_swiglu_sleef(
void)
628 const char *library = getenv(
"CK_SLEEF_LIBRARY");
629 if (library && *library) {
630 ck_pytorch_swiglu_sleef_handle = dlopen(library, RTLD_NOW | RTLD_LOCAL);
631 if (ck_pytorch_swiglu_sleef_handle) {
632 ck_pytorch_swiglu_expf16 = (ck_sleef_expf16_fn)dlsym(
633 ck_pytorch_swiglu_sleef_handle,
"Sleef_expf16_u10");
636 ck_pytorch_swiglu_expf16 =
637 (ck_sleef_expf16_fn)dlsym(
RTLD_DEFAULT,
"Sleef_expf16_u10");
663 if (!input || !output || tokens < 0 || dim < 0) {
664 fprintf(stderr,
"HARD KERNEL CONTRACT FAULT: invalid PyTorch BF16 SwiGLU arguments\n");
668#if defined(__AVX512F__)
669 pthread_once(&ck_pytorch_swiglu_once, ck_bind_pytorch_swiglu_sleef);
670 if (!ck_pytorch_swiglu_expf16) {
672 "HARD KERNEL CONTRACT FAULT: PyTorch BF16 SwiGLU requires "
673 "SLEEF Sleef_expf16_u10; set CK_SLEEF_LIBRARY\n");
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;
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);
691 _mm512_store_ps(silu_lanes,
silu);
692 for (
int lane = 0; lane < 16; ++lane) {
700 for (; d < dim; ++d) {
703 const float silu = gate_bf16 / (1.0f + expf(-gate_bf16));
720 const float *d_output,
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);
733 for (
int d = 0; d < D; ++d) {
735 float b = row[D + d];
736 float dy = dy_row[d];
740 float silu_prime = s * (1.0f + a * (1.0f - s));
742 float dA = dy * b * silu_prime;
743 float dB = dy *
silu;
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
void quantize_row_q8_k(const float *x, void *y, int k)
float sigmoid_scalar(float x)
int ck_strict_parity_enabled(void)
Quantization block structures for weight-only quantization.
void swiglu_forward_exact(const float *input, float *output, int tokens, int dim)
void swiglu_forward_pytorch_bf16_storage(const float *input, float *output, int tokens, int dim)
void swiglu_forward(const float *input, float *output, int tokens, int dim)
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_q8_k(const float *input, void *output_q8, int tokens, int dim)
static float sigmoid_scalar_parity(float x)
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)
__attribute__((visibility("default"))) CKTokenizer *ck_tokenizer_create(CKTokenizerType type)
static void silu(float *x, int n)