← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
swiglu_kernels.c
Go to the documentation of this file.
1/**
2 * @file swiglu_kernels.c
3 * @brief SwiGLU activation 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 * SwiGLU: y = silu(gate) * up = (gate * sigmoid(gate)) * up
15 */
16
17#ifndef _GNU_SOURCE
18#define _GNU_SOURCE
19#endif
20
21#include "bf16_utils.h"
22#include "ckernel_engine.h"
23#include "ckernel_quant.h"
24#include <dlfcn.h>
25#include <math.h>
26#include <pthread.h>
27#include <stddef.h>
28#include <stdio.h>
29#include <stdlib.h>
30
31#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
32#include <immintrin.h>
33#endif
34
35/*
36 * PyTorch-parity sigmoid for strict/exact SwiGLU paths.
37 * Keep this in fp32 to match ATen CPU opmath more closely.
38 */
39static inline float sigmoid_scalar_parity(float x)
40{
41 return 1.0f / (1.0f + expf(-x));
42}
43
44/* ========================================================================== */
45/* Fast exp approximation for SIMD */
46/* ========================================================================== */
47
48#if defined(__AVX512F__)
49// AVX-512 fast exp approximation
50static inline __m512 exp512_fast(__m512 x) {
51 // Clamp to avoid overflow/underflow
52 x = _mm512_max_ps(x, _mm512_set1_ps(-88.0f));
53 x = _mm512_min_ps(x, _mm512_set1_ps(88.0f));
54
55 // exp(x) = 2^(x * log2(e))
56 const __m512 log2e = _mm512_set1_ps(1.4426950408889634f);
57 __m512 z = _mm512_mul_ps(x, log2e);
58
59 // Split into integer and fractional parts
60 __m512 zf = _mm512_roundscale_ps(z, _MM_FROUND_TO_NEAREST_INT);
61 __m512 f = _mm512_sub_ps(z, zf);
62
63 // Polynomial for 2^f, f in [-0.5, 0.5]
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);
69
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);
74
75 // Scale by 2^n
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);
80
81 return _mm512_mul_ps(poly, scale);
82}
83
84// AVX-512 sigmoid: 1 / (1 + exp(-x))
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));
90}
91#endif
92
93#if defined(__AVX2__)
94// AVX2 fast exp approximation (needs FMA and integer ops)
95static inline __m256 exp256_fast(__m256 x) {
96 // Clamp
97 x = _mm256_max_ps(x, _mm256_set1_ps(-88.0f));
98 x = _mm256_min_ps(x, _mm256_set1_ps(88.0f));
99
100 // exp(x) = 2^(x * log2(e))
101 const __m256 log2e = _mm256_set1_ps(1.4426950408889634f);
102 __m256 z = _mm256_mul_ps(x, log2e);
103
104 // Round to nearest integer
105 __m256 zf = _mm256_round_ps(z, _MM_FROUND_TO_NEAREST_INT | _MM_FROUND_NO_EXC);
106 __m256 f = _mm256_sub_ps(z, zf);
107
108 // Polynomial for 2^f
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);
114
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);
119
120 // Scale by 2^n
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);
125
126 return _mm256_mul_ps(poly, scale);
127}
128
129// AVX2 sigmoid
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));
135}
136
137/* GGML's parity path uses this specific vector exponential approximation. */
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(
143 n,
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(
154 _mm256_fmadd_ps(
155 _mm256_fmadd_ps(_mm256_set1_ps(0x1.0e4020p-7f), b,
156 _mm256_set1_ps(0x1.573e2ep-5f)),
157 u,
158 _mm256_fmadd_ps(_mm256_set1_ps(0x1.555e66p-3f), b,
159 _mm256_set1_ps(0x1.fffdb6p-2f))),
160 u,
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);
164 }
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));
174 return _mm256_or_ps(
175 _mm256_and_ps(_mm256_castsi256_ps(d), _mm256_mul_ps(s1, s1)),
176 _mm256_andnot_ps(
177 _mm256_castsi256_ps(d),
178 _mm256_or_ps(
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)))));
183}
184#endif
185
186#if defined(__AVX512F__) && defined(__AVX512DQ__)
187/* Keep the production parity provider aligned with llama.cpp's AVX-512
188 * exponential approximation and instruction grouping. */
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(
200 _mm512_fmadd_ps(
201 _mm512_fmadd_ps(
202 _mm512_set1_ps(0x1.0e4020p-7f), b,
203 _mm512_set1_ps(0x1.573e2ep-5f)),
204 u,
205 _mm512_fmadd_ps(
206 _mm512_set1_ps(0x1.555e66p-3f), b,
207 _mm512_set1_ps(0x1.fffdb6p-2f))),
208 u,
209 _mm512_fmadd_ps(
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)) {
214 return res;
215 }
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),
220 zero);
221 return _mm512_mask_blend_ps(d, res, alt);
222}
223#endif
224
225/**
226 * SwiGLU forward pass
227 * @test test_swiglu.py::TestSwiGLUForward::test_forward_tokens
228 * @test test_swiglu.py::TestSwiGLUForward::test_forward_single
229 * @test test_mlp.py::TestMLPForward::test_swiglu_mlp
230 * @test test_fused_swiglu_decode.py::TestFusedSwiGLUDecode::test_fused_swiglu_decode
231 * @test test_parity.py::test_swiglu_parity
232 *
233 * SwiGLU: y = silu(gate) * up where silu(x) = x * sigmoid(x)
234 *
235 * After changes: make test && make llamacpp-parity-full
236 */
237void swiglu_forward(const float *input,
238 float *output,
239 int tokens,
240 int dim)
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}
318
319void swiglu_forward_q8_k(const float *input,
320 void *output_q8,
321 int tokens,
322 int dim)
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}
374
375/**
376 * SwiGLU backward pass
377 * @test test_swiglu.py::TestSwiGLUBackward::test_backward_tokens
378 * @test test_swiglu.py::TestSwiGLUBackward::test_backward_single
379 * @test test_parity.py::test_swiglu_backward_parity
380 *
381 * Computes dGate and dUp given dY.
382 * dGate = dy * b * silu'(a), dUp = dy * silu(a)
383 *
384 * After changes: make test && make llamacpp-parity-full
385 */
386void swiglu_backward(const float *input,
387 const float *d_output,
388 float *d_input,
389 int tokens,
390 int dim)
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}
502// ============================================================================
503// Exact versions using standard library expf (slower but accurate)
504// ============================================================================
505
506/**
507 * SwiGLU forward pass (exact version using stdlib sigmoid)
508 * @test test_swiglu.py::TestSwiGLUForward::test_exact_vs_fast
509 * @test test_swiglu.py::TestSwiGLUForward::test_exact_single
510 *
511 * Uses standard library expf for numerical accuracy reference.
512 *
513 * After changes: make test
514 */
515void swiglu_forward_exact(const float *input,
516 float *output,
517 int tokens,
518 int dim)
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}
537
538void swiglu_forward_ggml(const float *input,
539 float *output,
540 int tokens,
541 int dim)
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}
575
576void swiglu_forward_ggml_split(const float *gate,
577 const float *up,
578 float *output,
579 int tokens,
580 int dim)
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}
618
619#if defined(__AVX512F__)
620typedef __m512 (*ck_sleef_expf16_fn)(__m512);
621
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;
625
626static void ck_bind_pytorch_swiglu_sleef(void)
627{
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");
634 }
635 } else {
636 ck_pytorch_swiglu_expf16 =
637 (ck_sleef_expf16_fn)dlsym(RTLD_DEFAULT, "Sleef_expf16_u10");
638 }
639}
640#endif
641
642/*
643 * Match the ATen BF16 Qwen MLP expression, not the GGML FP32 expression:
644 *
645 * silu_bf16 = aten::silu(gate_bf16)
646 * out_bf16 = silu_bf16 * up_bf16
647 *
648 * PyTorch 2.8's reduced-floating CPU SiLU converts BF16 lanes to FP32, uses
649 * its SLEEF vector exponential, evaluates x / (1 + exp(-x)), and rounds the
650 * SiLU result to BF16. The following TensorIterator multiplication converts
651 * both BF16 operands to FP32 and performs a second BF16 store. Keeping the
652 * SiLU intermediate in FP32 is close, but is a different storage contract and
653 * changes long teacher-forced generation after repeated MLP blocks.
654 *
655 * Source oracle: PyTorch aten/src/ATen/native/cpu/Activation.cpp,
656 * silu_kernel(), commit a1cb3cc05d46d198467bebbb6e8fba50a325d4e7.
657 */
659 float *output,
660 int tokens,
661 int dim)
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}
709
710/**
711 * SwiGLU backward pass (exact version using stdlib sigmoid)
712 * @test test_swiglu.py::TestSwiGLUBackward::test_exact_vs_fast
713 * @test test_swiglu.py::TestSwiGLUBackward::test_exact_single
714 *
715 * Uses standard library expf for numerical accuracy reference.
716 *
717 * After changes: make test
718 */
719void swiglu_backward_exact(const float *input,
720 const float *d_output,
721 float *d_input,
722 int tokens,
723 int dim)
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}
#define RTLD_DEFAULT
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
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.
#define QK_K
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)