← Back to C-Kernel-Engine Docs Doxygen Source Documentation
 
Loading...
Searching...
No Matches
layernorm_kernels.c
Go to the documentation of this file.
1/**
2 * @file layernorm_kernels.c
3 * @brief LayerNorm forward/backward 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 * LayerNorm: y = gamma * (x - mean) / sqrt(var + eps) + beta
15 */
16
17#include "ckernel_engine.h"
18#include "bf16_utils.h"
19
20#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE2__)
21#include <immintrin.h>
22#endif
23#include <math.h>
24#include <stdlib.h>
25
26static inline void zero_layernorm_padding(float *out_ptr,
27 int d_model,
28 int aligned_embed_dim)
29{
30 for (int idx = d_model; idx < aligned_embed_dim; ++idx) {
31 out_ptr[idx] = 0.0f;
32 }
33}
34
35void layernorm_naive_serial_matched_precision(const float *input,
36 const float *gamma,
37 const float *beta,
38 float *output,
39 float *mean_cache,
40 float *rstd_cache,
41 int tokens, int d_model, float eps);
42
43#if defined(__clang__) || defined(__INTEL_LLVM_COMPILER)
44#pragma float_control(precise, on, push)
45#endif
46static void layernorm_forward_ggml_exact(const float *input,
47 const float *gamma,
48 const float *beta,
49 float *output,
50 float *mean_cache,
51 float *rstd_cache,
52 int tokens,
53 int d_model,
54 int input_stride,
55 int output_stride,
56 int aligned_embed_dim,
57 float eps)
58{
59 for (int t = 0; t < tokens; ++t) {
60 const float *x = input + (size_t)t * (size_t)input_stride;
61 float *y = output + (size_t)t * (size_t)output_stride;
62
63 double sum_acc = 0.0;
64#if defined(__clang__)
65#pragma clang loop vectorize(disable)
66#pragma clang loop interleave(disable)
67#endif
68 for (int i = 0; i < d_model; ++i) {
69 sum_acc += (double)x[i];
70 }
71 const float sum = (float)sum_acc;
72 const float mean = sum / (float)d_model;
73
74 double var_acc = 0.0;
75 int i = 0;
76#if defined(__AVX2__) && defined(__FMA__)
77 /* This provider promises ggml's ordered AVX2 reduction. Keep that
78 * arithmetic identity on AVX-512 hosts instead of widening it based
79 * on the compiler target. */
80 for (; i + 7 < d_model; i += 8) {
81 __m256 val = _mm256_sub_ps(_mm256_loadu_ps(x + i),
82 _mm256_set1_ps(mean));
83 _mm256_storeu_ps(y + i, val);
84 val = _mm256_mul_ps(val, val);
85 __m128 val2 = _mm_add_ps(_mm256_extractf128_ps(val, 1),
86 _mm256_castps256_ps128(val));
87 val2 = _mm_add_ps(val2, _mm_movehl_ps(val2, val2));
88 val2 = _mm_add_ss(val2, _mm_movehdup_ps(val2));
89 var_acc += (double)_mm_cvtss_f32(val2);
90 }
91#elif defined(__SSE2__)
92 for (; i + 3 < d_model; i += 4) {
93 __m128 val = _mm_sub_ps(_mm_loadu_ps(x + i),
94 _mm_set1_ps(mean));
95 _mm_storeu_ps(y + i, val);
96 val = _mm_mul_ps(val, val);
97#if defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)
98 val = _mm_add_ps(val, _mm_movehl_ps(val, val));
99 val = _mm_add_ss(val, _mm_movehdup_ps(val));
100#else
101 __m128 tmp = _mm_shuffle_ps(val, val, _MM_SHUFFLE(2, 3, 0, 1));
102 val = _mm_add_ps(val, tmp);
103 tmp = _mm_movehl_ps(tmp, val);
104 val = _mm_add_ss(val, tmp);
105#endif
106 var_acc += (double)_mm_cvtss_f32(val);
107 }
108#endif
109#if defined(__clang__)
110#pragma clang loop vectorize(disable)
111#pragma clang loop interleave(disable)
112#endif
113 for (; i < d_model; ++i) {
114 const float centered = x[i] - mean;
115 y[i] = centered;
116 var_acc += (double)(centered * centered);
117 }
118 const float variance = (float)(var_acc / (double)d_model);
119 const float scale = 1.0f / sqrtf(variance + eps);
120
121 if (mean_cache) {
122 mean_cache[t] = mean;
123 }
124 if (rstd_cache) {
125 rstd_cache[t] = scale;
126 }
127
128#if defined(__clang__)
129#pragma clang loop vectorize(disable)
130#pragma clang loop interleave(disable)
131#endif
132 for (int i = 0; i < d_model; ++i) {
133 y[i] *= scale;
134 }
135 if (gamma) {
136#if defined(__clang__)
137#pragma clang loop vectorize(disable)
138#pragma clang loop interleave(disable)
139#endif
140 for (int i = 0; i < d_model; ++i) {
141 y[i] *= gamma[i];
142 }
143 }
144 if (beta) {
145#if defined(__clang__)
146#pragma clang loop vectorize(disable)
147#pragma clang loop interleave(disable)
148#endif
149 for (int i = 0; i < d_model; ++i) {
150 y[i] += beta[i];
151 }
152 }
153
154 if (aligned_embed_dim > d_model) {
155 zero_layernorm_padding(y, d_model, aligned_embed_dim);
156 }
157 }
158}
159#if defined(__clang__) || defined(__INTEL_LLVM_COMPILER)
160#pragma float_control(pop)
161#endif
162
163#if defined(__AVX2__) || defined(__AVX__)
164static inline float hsum256_ps(__m256 v)
165{
166 __m128 low = _mm256_castps256_ps128(v);
167 __m128 high = _mm256_extractf128_ps(v, 1);
168 __m128 sum = _mm_add_ps(low, high);
169 sum = _mm_hadd_ps(sum, sum);
170 sum = _mm_hadd_ps(sum, sum);
171 return _mm_cvtss_f32(sum);
172}
173#endif
174// Naive serial LayerNorm implementation (forward only), copied from C-Transformer.
175void layernorm_naive_serial(const float *input,
176 const float *gamma,
177 const float *beta,
178 float *output,
179 float *mean_cache,
180 float *rstd_cache,
181 int tokens, int d_model, int aligned_embed_dim,
182 float eps)
183{
184 for (int t = 0; t < tokens; ++t) {
185 const float *in_ptr = input + t * aligned_embed_dim;
186 float *out_ptr = output + t * aligned_embed_dim;
187
188 float sum_val = 0.0f;
189 for (int i = 0; i < d_model; ++i) {
190 sum_val += in_ptr[i];
191 }
192 float mean = sum_val / (float)d_model;
193
194 float sum_sq_diff = 0.0f;
195 for (int i = 0; i < d_model; ++i) {
196 float diff = in_ptr[i] - mean;
197 sum_sq_diff += diff * diff;
198 }
199 float variance = sum_sq_diff / (float)d_model + eps;
200
201 double var_double = (double)variance;
202 float inv_std = (float)(1.0 / sqrt(var_double));
203
204 for (int i = 0; i < d_model; ++i) {
205 float normalized_val = (in_ptr[i] - mean) * inv_std;
206 out_ptr[i] = normalized_val * gamma[i] + beta[i];
207 }
208
209 if (mean_cache) {
210 mean_cache[t] = mean;
211 }
212 if (rstd_cache) {
213 rstd_cache[t] = inv_std;
214 }
215 /* Keep aligned padding quiet so future GEMMs see deterministic memory. */
216 if (aligned_embed_dim > d_model) {
217 /* Keep padded lanes zeroed so subsequent GEMMs never read stale data. */
218 for (int i = d_model; i < aligned_embed_dim; ++i) {
219 out_ptr[i] = 0.0f;
220 }
221 }
222 }
223}
224
225#if defined(__AVX512F__)
226// AVX-512 rolled slice kernel, copied from C-Transformer (model-agnostic).
227static void layernorm_forward_rolled_slice_avx512(const float *__restrict input_slice_base,
228 const float *__restrict gamma,
229 const float *__restrict beta,
230 float *__restrict output_slice_base,
231 float *__restrict mean_cache_slice,
232 float *__restrict rstd_cache_slice,
233 int num_tokens_in_slice,
234 int d_model,
235 int aligned_embed_dim,
236 float eps)
237{
238 for (int t = 0; t < num_tokens_in_slice; ++t) {
239 const float *in_ptr_token = input_slice_base + t * aligned_embed_dim;
240 float *out_ptr_token = output_slice_base + t * aligned_embed_dim;
241
242 __m512 acc_sum_vec = _mm512_setzero_ps();
243 int j = 0;
244 for (; j <= d_model - 16; j += 16) {
245 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
246 __m512 v = _mm512_load_ps(in_ptr_token + j);
247 acc_sum_vec = _mm512_add_ps(acc_sum_vec, v);
248 }
249 float mean = _mm512_reduce_add_ps(acc_sum_vec);
250 for (; j < d_model; ++j) {
251 mean += in_ptr_token[j];
252 }
253 mean /= (float)d_model;
254 __m512 mean_vec = _mm512_set1_ps(mean);
255
256 __m512 acc_var_vec = _mm512_setzero_ps();
257 j = 0;
258 for (; j <= d_model - 16; j += 16) {
259 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
260 __m512 v = _mm512_load_ps(in_ptr_token + j);
261 __m512 diff = _mm512_sub_ps(v, mean_vec);
262 acc_var_vec = _mm512_fmadd_ps(diff, diff, acc_var_vec);
263 }
264 float var = _mm512_reduce_add_ps(acc_var_vec);
265 for (; j < d_model; ++j) {
266 float diff = in_ptr_token[j] - mean;
267 var += diff * diff;
268 }
269 var = var / (float)d_model + eps;
270 double var_double = (double)var;
271 float inv_std = (float)(1.0 / sqrt(var_double));
272 __m512 inv_std_vec = _mm512_set1_ps(inv_std);
273
274 if (mean_cache_slice) {
275 mean_cache_slice[t] = mean;
276 }
277 if (rstd_cache_slice) {
278 rstd_cache_slice[t] = inv_std;
279 }
280
281 j = 0;
282 for (; j <= d_model - 16; j += 16) {
283 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
284 _mm_prefetch((const char *)(gamma + j + 128), _MM_HINT_T0);
285 _mm_prefetch((const char *)(beta + j + 128), _MM_HINT_T0);
286
287 __m512 v = _mm512_load_ps(in_ptr_token + j);
288 __m512 g = _mm512_load_ps(gamma + j);
289 __m512 b = _mm512_load_ps(beta + j);
290
291 __m512 n = _mm512_mul_ps(_mm512_sub_ps(v, mean_vec), inv_std_vec);
292 __m512 o = _mm512_fmadd_ps(n, g, b);
293
294 _mm512_store_ps(out_ptr_token + j, o);
295 }
296 for (; j < d_model; ++j) {
297 float normed = (in_ptr_token[j] - mean) * inv_std;
298 out_ptr_token[j] = normed * gamma[j] + beta[j];
299 }
300
301 if (aligned_embed_dim > d_model) {
302 /* Keep the padded lanes zeroed so later GEMMs see deterministic memory. */
303 zero_layernorm_padding(out_ptr_token, d_model, aligned_embed_dim);
304 }
305 }
306}
307#elif defined(__AVX2__) || defined(__AVX__)
308// AVX/AVX2 rolled slice kernel (8-float vectors).
309static void layernorm_forward_rolled_slice_avx256(const float *__restrict input_slice_base,
310 const float *__restrict gamma,
311 const float *__restrict beta,
312 float *__restrict output_slice_base,
313 float *__restrict mean_cache_slice,
314 float *__restrict rstd_cache_slice,
315 int num_tokens_in_slice,
316 int d_model,
317 int aligned_embed_dim,
318 float eps)
319{
320 for (int t = 0; t < num_tokens_in_slice; ++t) {
321 const float *in_ptr_token = input_slice_base + t * aligned_embed_dim;
322 float *out_ptr_token = output_slice_base + t * aligned_embed_dim;
323
324 __m256 acc_sum_vec = _mm256_setzero_ps();
325 int j = 0;
326 for (; j <= d_model - 8; j += 8) {
327 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
328 __m256 v = _mm256_load_ps(in_ptr_token + j);
329 acc_sum_vec = _mm256_add_ps(acc_sum_vec, v);
330 }
331 float mean = hsum256_ps(acc_sum_vec);
332 for (; j < d_model; ++j) {
333 mean += in_ptr_token[j];
334 }
335 mean /= (float)d_model;
336 __m256 mean_vec = _mm256_set1_ps(mean);
337
338 __m256 acc_var_vec = _mm256_setzero_ps();
339 j = 0;
340 for (; j <= d_model - 8; j += 8) {
341 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
342 __m256 v = _mm256_load_ps(in_ptr_token + j);
343 __m256 diff = _mm256_sub_ps(v, mean_vec);
344#if defined(__FMA__)
345 acc_var_vec = _mm256_fmadd_ps(diff, diff, acc_var_vec);
346#else
347 acc_var_vec = _mm256_add_ps(acc_var_vec, _mm256_mul_ps(diff, diff));
348#endif
349 }
350 float var = hsum256_ps(acc_var_vec);
351 for (; j < d_model; ++j) {
352 float diff = in_ptr_token[j] - mean;
353 var += diff * diff;
354 }
355 var = var / (float)d_model + eps;
356 double var_double = (double)var;
357 float inv_std = (float)(1.0 / sqrt(var_double));
358 __m256 inv_std_vec = _mm256_set1_ps(inv_std);
359
360 if (mean_cache_slice) {
361 mean_cache_slice[t] = mean;
362 }
363 if (rstd_cache_slice) {
364 rstd_cache_slice[t] = inv_std;
365 }
366
367 j = 0;
368 for (; j <= d_model - 8; j += 8) {
369 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
370 _mm_prefetch((const char *)(gamma + j + 128), _MM_HINT_T0);
371 _mm_prefetch((const char *)(beta + j + 128), _MM_HINT_T0);
372
373 __m256 v = _mm256_load_ps(in_ptr_token + j);
374 __m256 g = _mm256_load_ps(gamma + j);
375 __m256 b = _mm256_load_ps(beta + j);
376
377 __m256 n = _mm256_mul_ps(_mm256_sub_ps(v, mean_vec), inv_std_vec);
378#if defined(__FMA__)
379 __m256 o = _mm256_fmadd_ps(n, g, b);
380#else
381 __m256 o = _mm256_add_ps(_mm256_mul_ps(n, g), b);
382#endif
383
384 _mm256_store_ps(out_ptr_token + j, o);
385 }
386 for (; j < d_model; ++j) {
387 float normed = (in_ptr_token[j] - mean) * inv_std;
388 out_ptr_token[j] = normed * gamma[j] + beta[j];
389 }
390
391 if (aligned_embed_dim > d_model) {
392 zero_layernorm_padding(out_ptr_token, d_model, aligned_embed_dim);
393 }
394 }
395}
396#endif
397
398void layernorm_forward_rolled_slice(const float *__restrict input_slice_base,
399 const float *__restrict gamma,
400 const float *__restrict beta,
401 float *__restrict output_slice_base,
402 float *__restrict mean_cache_slice,
403 float *__restrict rstd_cache_slice,
404 int num_tokens_in_slice,
405 int d_model,
406 int aligned_embed_dim,
407 float eps)
408{
410 layernorm_forward_ggml_exact(input_slice_base, gamma, beta,
411 output_slice_base, mean_cache_slice, rstd_cache_slice,
412 num_tokens_in_slice, d_model,
413 aligned_embed_dim, aligned_embed_dim, aligned_embed_dim, eps);
414 return;
415 }
416
417#if defined(__AVX512F__)
418 layernorm_forward_rolled_slice_avx512(input_slice_base, gamma, beta,
419 output_slice_base, mean_cache_slice, rstd_cache_slice,
420 num_tokens_in_slice, d_model, aligned_embed_dim, eps);
421#elif defined(__AVX2__) || defined(__AVX__)
422 layernorm_forward_rolled_slice_avx256(input_slice_base, gamma, beta,
423 output_slice_base, mean_cache_slice, rstd_cache_slice,
424 num_tokens_in_slice, d_model, aligned_embed_dim, eps);
425#else
426 layernorm_naive_serial(input_slice_base, gamma, beta,
427 output_slice_base, mean_cache_slice, rstd_cache_slice,
428 num_tokens_in_slice, d_model, aligned_embed_dim, eps);
429#endif
430}
431
432#if defined(__AVX512F__)
433// AVX-512 unrolled slice kernel, copied from C-Transformer (model-agnostic).
434static void layernorm_forward_unrolled_slice_avx512(const float *__restrict input_slice_base,
435 const float *__restrict gamma,
436 const float *__restrict beta,
437 float *__restrict output_slice_base,
438 float *__restrict mean_cache_slice,
439 float *__restrict rstd_cache_slice,
440 int num_tokens_in_slice,
441 int d_model,
442 float eps)
443{
444 for (int t = 0; t < num_tokens_in_slice; ++t) {
445 const float *in_ptr_token = input_slice_base + t * d_model;
446 float *out_ptr_token = output_slice_base + t * d_model;
447
448 __m512 acc0 = _mm512_setzero_ps();
449 __m512 acc1 = _mm512_setzero_ps();
450 __m512 acc2 = _mm512_setzero_ps();
451 __m512 acc3 = _mm512_setzero_ps();
452
453 int j = 0;
454 int unroll_factor_floats = 64;
455
456 for (; j <= d_model - unroll_factor_floats; j += unroll_factor_floats) {
457 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
458
459 __m512 v0 = _mm512_load_ps(in_ptr_token + j);
460 __m512 v1 = _mm512_load_ps(in_ptr_token + j + 16);
461 __m512 v2 = _mm512_load_ps(in_ptr_token + j + 32);
462 __m512 v3 = _mm512_load_ps(in_ptr_token + j + 48);
463
464 acc0 = _mm512_add_ps(acc0, v0);
465 acc1 = _mm512_add_ps(acc1, v1);
466 acc2 = _mm512_add_ps(acc2, v2);
467 acc3 = _mm512_add_ps(acc3, v3);
468 }
469 __m512 acc_sum = _mm512_add_ps(_mm512_add_ps(acc0, acc1),
470 _mm512_add_ps(acc2, acc3));
471 float mean = _mm512_reduce_add_ps(acc_sum);
472
473 for (; j < d_model; ++j) {
474 mean += in_ptr_token[j];
475 }
476 mean /= (float)d_model;
477 __m512 mean_vec = _mm512_set1_ps(mean);
478
479 acc0 = _mm512_setzero_ps();
480 acc1 = _mm512_setzero_ps();
481 acc2 = _mm512_setzero_ps();
482 acc3 = _mm512_setzero_ps();
483
484 j = 0;
485 for (; j <= d_model - unroll_factor_floats; j += unroll_factor_floats) {
486 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
487
488 __m512 v0 = _mm512_load_ps(in_ptr_token + j);
489 __m512 v1 = _mm512_load_ps(in_ptr_token + j + 16);
490 __m512 v2 = _mm512_load_ps(in_ptr_token + j + 32);
491 __m512 v3 = _mm512_load_ps(in_ptr_token + j + 48);
492
493 __m512 d0 = _mm512_sub_ps(v0, mean_vec);
494 __m512 d1 = _mm512_sub_ps(v1, mean_vec);
495 __m512 d2 = _mm512_sub_ps(v2, mean_vec);
496 __m512 d3 = _mm512_sub_ps(v3, mean_vec);
497
498 acc0 = _mm512_fmadd_ps(d0, d0, acc0);
499 acc1 = _mm512_fmadd_ps(d1, d1, acc1);
500 acc2 = _mm512_fmadd_ps(d2, d2, acc2);
501 acc3 = _mm512_fmadd_ps(d3, d3, acc3);
502 }
503 acc_sum = _mm512_add_ps(_mm512_add_ps(acc0, acc1),
504 _mm512_add_ps(acc2, acc3));
505 float var = _mm512_reduce_add_ps(acc_sum);
506
507 for (; j < d_model; ++j) {
508 float diff = in_ptr_token[j] - mean;
509 var += diff * diff;
510 }
511 var = var / (float)d_model + eps;
512 double var_double = (double)var;
513 float inv_std = (float)(1.0 / sqrt(var_double));
514 __m512 inv_std_vec = _mm512_set1_ps(inv_std);
515
516 if (mean_cache_slice) {
517 mean_cache_slice[t] = mean;
518 }
519 if (rstd_cache_slice) {
520 rstd_cache_slice[t] = inv_std;
521 }
522
523 j = 0;
524 for (; j <= d_model - unroll_factor_floats; j += unroll_factor_floats) {
525 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
526 _mm_prefetch((const char *)(gamma + j + 128), _MM_HINT_T0);
527 _mm_prefetch((const char *)(beta + j + 128), _MM_HINT_T0);
528
529 __m512 v0 = _mm512_load_ps(in_ptr_token + j);
530 __m512 v1 = _mm512_load_ps(in_ptr_token + j + 16);
531 __m512 v2 = _mm512_load_ps(in_ptr_token + j + 32);
532 __m512 v3 = _mm512_load_ps(in_ptr_token + j + 48);
533
534 __m512 g0 = _mm512_load_ps(gamma + j);
535 __m512 g1 = _mm512_load_ps(gamma + j + 16);
536 __m512 g2 = _mm512_load_ps(gamma + j + 32);
537 __m512 g3 = _mm512_load_ps(gamma + j + 48);
538
539 __m512 b0 = _mm512_load_ps(beta + j);
540 __m512 b1 = _mm512_load_ps(beta + j + 16);
541 __m512 b2 = _mm512_load_ps(beta + j + 32);
542 __m512 b3 = _mm512_load_ps(beta + j + 48);
543
544 __m512 n0 = _mm512_mul_ps(_mm512_sub_ps(v0, mean_vec), inv_std_vec);
545 __m512 n1 = _mm512_mul_ps(_mm512_sub_ps(v1, mean_vec), inv_std_vec);
546 __m512 n2 = _mm512_mul_ps(_mm512_sub_ps(v2, mean_vec), inv_std_vec);
547 __m512 n3 = _mm512_mul_ps(_mm512_sub_ps(v3, mean_vec), inv_std_vec);
548
549 __m512 o0 = _mm512_fmadd_ps(n0, g0, b0);
550 __m512 o1 = _mm512_fmadd_ps(n1, g1, b1);
551 __m512 o2 = _mm512_fmadd_ps(n2, g2, b2);
552 __m512 o3 = _mm512_fmadd_ps(n3, g3, b3);
553
554 _mm512_store_ps(out_ptr_token + j, o0);
555 _mm512_store_ps(out_ptr_token + j + 16, o1);
556 _mm512_store_ps(out_ptr_token + j + 32, o2);
557 _mm512_store_ps(out_ptr_token + j + 48, o3);
558 }
559 for (; j < d_model; ++j) {
560 float normed = (in_ptr_token[j] - mean) * inv_std;
561 out_ptr_token[j] = normed * gamma[j] + beta[j];
562 }
563 }
564}
565#elif defined(__AVX2__) || defined(__AVX__)
566// AVX/AVX2 unrolled slice kernel (8-float vectors).
567static void layernorm_forward_unrolled_slice_avx256(const float *__restrict input_slice_base,
568 const float *__restrict gamma,
569 const float *__restrict beta,
570 float *__restrict output_slice_base,
571 float *__restrict mean_cache_slice,
572 float *__restrict rstd_cache_slice,
573 int num_tokens_in_slice,
574 int d_model,
575 float eps)
576{
577 for (int t = 0; t < num_tokens_in_slice; ++t) {
578 const float *in_ptr_token = input_slice_base + t * d_model;
579 float *out_ptr_token = output_slice_base + t * d_model;
580
581 __m256 acc0 = _mm256_setzero_ps();
582 __m256 acc1 = _mm256_setzero_ps();
583 __m256 acc2 = _mm256_setzero_ps();
584 __m256 acc3 = _mm256_setzero_ps();
585
586 int j = 0;
587 int unroll_factor_floats = 32;
588
589 for (; j <= d_model - unroll_factor_floats; j += unroll_factor_floats) {
590 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
591
592 __m256 v0 = _mm256_load_ps(in_ptr_token + j);
593 __m256 v1 = _mm256_load_ps(in_ptr_token + j + 8);
594 __m256 v2 = _mm256_load_ps(in_ptr_token + j + 16);
595 __m256 v3 = _mm256_load_ps(in_ptr_token + j + 24);
596
597 acc0 = _mm256_add_ps(acc0, v0);
598 acc1 = _mm256_add_ps(acc1, v1);
599 acc2 = _mm256_add_ps(acc2, v2);
600 acc3 = _mm256_add_ps(acc3, v3);
601 }
602 __m256 acc_sum = _mm256_add_ps(_mm256_add_ps(acc0, acc1),
603 _mm256_add_ps(acc2, acc3));
604 float mean = hsum256_ps(acc_sum);
605
606 for (; j < d_model; ++j) {
607 mean += in_ptr_token[j];
608 }
609 mean /= (float)d_model;
610 __m256 mean_vec = _mm256_set1_ps(mean);
611
612 acc0 = _mm256_setzero_ps();
613 acc1 = _mm256_setzero_ps();
614 acc2 = _mm256_setzero_ps();
615 acc3 = _mm256_setzero_ps();
616
617 j = 0;
618 for (; j <= d_model - unroll_factor_floats; j += unroll_factor_floats) {
619 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
620
621 __m256 v0 = _mm256_load_ps(in_ptr_token + j);
622 __m256 v1 = _mm256_load_ps(in_ptr_token + j + 8);
623 __m256 v2 = _mm256_load_ps(in_ptr_token + j + 16);
624 __m256 v3 = _mm256_load_ps(in_ptr_token + j + 24);
625
626 __m256 d0 = _mm256_sub_ps(v0, mean_vec);
627 __m256 d1 = _mm256_sub_ps(v1, mean_vec);
628 __m256 d2 = _mm256_sub_ps(v2, mean_vec);
629 __m256 d3 = _mm256_sub_ps(v3, mean_vec);
630
631#if defined(__FMA__)
632 acc0 = _mm256_fmadd_ps(d0, d0, acc0);
633 acc1 = _mm256_fmadd_ps(d1, d1, acc1);
634 acc2 = _mm256_fmadd_ps(d2, d2, acc2);
635 acc3 = _mm256_fmadd_ps(d3, d3, acc3);
636#else
637 acc0 = _mm256_add_ps(acc0, _mm256_mul_ps(d0, d0));
638 acc1 = _mm256_add_ps(acc1, _mm256_mul_ps(d1, d1));
639 acc2 = _mm256_add_ps(acc2, _mm256_mul_ps(d2, d2));
640 acc3 = _mm256_add_ps(acc3, _mm256_mul_ps(d3, d3));
641#endif
642 }
643 acc_sum = _mm256_add_ps(_mm256_add_ps(acc0, acc1),
644 _mm256_add_ps(acc2, acc3));
645 float var = hsum256_ps(acc_sum);
646
647 for (; j < d_model; ++j) {
648 float diff = in_ptr_token[j] - mean;
649 var += diff * diff;
650 }
651 var = var / (float)d_model + eps;
652 double var_double = (double)var;
653 float inv_std = (float)(1.0 / sqrt(var_double));
654 __m256 inv_std_vec = _mm256_set1_ps(inv_std);
655
656 if (mean_cache_slice) {
657 mean_cache_slice[t] = mean;
658 }
659 if (rstd_cache_slice) {
660 rstd_cache_slice[t] = inv_std;
661 }
662
663 j = 0;
664 for (; j <= d_model - unroll_factor_floats; j += unroll_factor_floats) {
665 _mm_prefetch((const char *)(in_ptr_token + j + 128), _MM_HINT_T0);
666 _mm_prefetch((const char *)(gamma + j + 128), _MM_HINT_T0);
667 _mm_prefetch((const char *)(beta + j + 128), _MM_HINT_T0);
668
669 __m256 v0 = _mm256_load_ps(in_ptr_token + j);
670 __m256 v1 = _mm256_load_ps(in_ptr_token + j + 8);
671 __m256 v2 = _mm256_load_ps(in_ptr_token + j + 16);
672 __m256 v3 = _mm256_load_ps(in_ptr_token + j + 24);
673
674 __m256 g0 = _mm256_load_ps(gamma + j);
675 __m256 g1 = _mm256_load_ps(gamma + j + 8);
676 __m256 g2 = _mm256_load_ps(gamma + j + 16);
677 __m256 g3 = _mm256_load_ps(gamma + j + 24);
678
679 __m256 b0 = _mm256_load_ps(beta + j);
680 __m256 b1 = _mm256_load_ps(beta + j + 8);
681 __m256 b2 = _mm256_load_ps(beta + j + 16);
682 __m256 b3 = _mm256_load_ps(beta + j + 24);
683
684 __m256 n0 = _mm256_mul_ps(_mm256_sub_ps(v0, mean_vec), inv_std_vec);
685 __m256 n1 = _mm256_mul_ps(_mm256_sub_ps(v1, mean_vec), inv_std_vec);
686 __m256 n2 = _mm256_mul_ps(_mm256_sub_ps(v2, mean_vec), inv_std_vec);
687 __m256 n3 = _mm256_mul_ps(_mm256_sub_ps(v3, mean_vec), inv_std_vec);
688
689#if defined(__FMA__)
690 __m256 o0 = _mm256_fmadd_ps(n0, g0, b0);
691 __m256 o1 = _mm256_fmadd_ps(n1, g1, b1);
692 __m256 o2 = _mm256_fmadd_ps(n2, g2, b2);
693 __m256 o3 = _mm256_fmadd_ps(n3, g3, b3);
694#else
695 __m256 o0 = _mm256_add_ps(_mm256_mul_ps(n0, g0), b0);
696 __m256 o1 = _mm256_add_ps(_mm256_mul_ps(n1, g1), b1);
697 __m256 o2 = _mm256_add_ps(_mm256_mul_ps(n2, g2), b2);
698 __m256 o3 = _mm256_add_ps(_mm256_mul_ps(n3, g3), b3);
699#endif
700
701 _mm256_store_ps(out_ptr_token + j, o0);
702 _mm256_store_ps(out_ptr_token + j + 8, o1);
703 _mm256_store_ps(out_ptr_token + j + 16, o2);
704 _mm256_store_ps(out_ptr_token + j + 24, o3);
705 }
706 for (; j < d_model; ++j) {
707 float normed = (in_ptr_token[j] - mean) * inv_std;
708 out_ptr_token[j] = normed * gamma[j] + beta[j];
709 }
710 }
711}
712#else
713// Scalar fallback when AVX-512 is unavailable.
714static void layernorm_forward_unrolled_slice_scalar(const float *__restrict input_slice_base,
715 const float *__restrict gamma,
716 const float *__restrict beta,
717 float *__restrict output_slice_base,
718 float *__restrict mean_cache_slice,
719 float *__restrict rstd_cache_slice,
720 int num_tokens_in_slice,
721 int d_model,
722 float eps)
723{
724 layernorm_naive_serial_matched_precision(input_slice_base, gamma, beta,
725 output_slice_base, mean_cache_slice, rstd_cache_slice,
726 num_tokens_in_slice, d_model, eps);
727}
728#endif
729
730void layernorm_forward_unrolled_slice(const float *__restrict input_slice_base,
731 const float *__restrict gamma,
732 const float *__restrict beta,
733 float *__restrict output_slice_base,
734 float *__restrict mean_cache_slice,
735 float *__restrict rstd_cache_slice,
736 int num_tokens_in_slice,
737 int d_model,
738 float eps)
739{
741 layernorm_forward_ggml_exact(input_slice_base, gamma, beta,
742 output_slice_base, mean_cache_slice, rstd_cache_slice,
743 num_tokens_in_slice, d_model,
744 d_model, d_model, d_model, eps);
745 return;
746 }
747
748#if defined(__AVX512F__)
749 layernorm_forward_unrolled_slice_avx512(input_slice_base, gamma, beta,
750 output_slice_base, mean_cache_slice, rstd_cache_slice,
751 num_tokens_in_slice, d_model, eps);
752#elif defined(__AVX2__) || defined(__AVX__)
753 layernorm_forward_unrolled_slice_avx256(input_slice_base, gamma, beta,
754 output_slice_base, mean_cache_slice, rstd_cache_slice,
755 num_tokens_in_slice, d_model, eps);
756#else
757 layernorm_forward_unrolled_slice_scalar(input_slice_base, gamma, beta,
758 output_slice_base, mean_cache_slice, rstd_cache_slice,
759 num_tokens_in_slice, d_model, eps);
760#endif
761}
762
763// Precision-matched naive LayerNorm used for benchmarking, copied from C-Transformer.
765 const float *gamma,
766 const float *beta,
767 float *output,
768 float *mean_cache,
769 float *rstd_cache,
770 int tokens, int d_model, float eps)
771{
772 layernorm_forward_ggml_exact(input, gamma, beta,
773 output, mean_cache, rstd_cache,
774 tokens, d_model,
775 d_model, d_model, d_model, eps);
776}
777
778/*
779 * Float-buffer LayerNorm with BF16 storage boundaries.
780 *
781 * The physical activation arena remains FP32 so this variant composes with the
782 * existing graph ABI. Inputs are first rounded as if loaded from BF16 storage,
783 * the established matched reduction executes in FP32, and outputs are rounded
784 * back to BF16 values while remaining represented as float.
785 */
787 const float *gamma,
788 const float *beta,
789 float *output,
790 float *mean_cache,
791 float *rstd_cache,
792 int tokens, int d_model, float eps)
793{
794 const size_t count = (size_t)tokens * (size_t)d_model;
795 for (size_t i = 0; i < count; ++i) {
796 output[i] = bf16_to_float(float_to_bf16(input[i]));
797 }
798 layernorm_forward_ggml_exact(output, gamma, beta,
799 output, mean_cache, rstd_cache,
800 tokens, d_model,
801 d_model, d_model, d_model, eps);
802 for (size_t i = 0; i < count; ++i) {
803 output[i] = bf16_to_float(float_to_bf16(output[i]));
804 }
805}
806
807#if defined(__AVX2__) && defined(__FMA__)
808static inline void layernorm_pytorch_add_moments_vec(
809 int m0_add,
810 __m256 m1_add,
811 __m256 m2_add,
812 int *m0,
813 __m256 *m1,
814 __m256 *m2)
815{
816 const int n = *m0 + m0_add;
817 const float c = n == 0 ? 0.0f : (float)m0_add / (float)n;
818 const __m256 delta = _mm256_sub_ps(m1_add, *m1);
819 const __m256 c_vec = _mm256_set1_ps(c);
820 const __m256 m2_tmp = _mm256_add_ps(*m2, m2_add);
821 const __m256 c_delta = _mm256_mul_ps(c_vec, delta);
822 const __m256 m0_delta = _mm256_mul_ps(
823 delta, _mm256_set1_ps((float)*m0));
824 *m1 = _mm256_add_ps(*m1, c_delta);
825 *m2 = _mm256_fmadd_ps(m0_delta, c_delta, m2_tmp);
826 *m0 = n;
827}
828
829static inline void layernorm_pytorch_add_moments_scalar(
830 int m0_add,
831 float m1_add,
832 float m2_add,
833 int *m0,
834 float *m1,
835 float *m2)
836{
837 const int n = *m0 + m0_add;
838 const float c = n == 0 ? 0.0f : (float)m0_add / (float)n;
839 const float delta = m1_add - *m1;
840 *m1 = fmaf(c, delta, *m1);
841 const float delta_term = (delta * delta) * c;
842 *m2 += fmaf(delta_term, (float)*m0, m2_add);
843 *m0 = n;
844}
845
846/*
847 * Match the AVX2 ATen BF16 RowwiseMoments provider selected by PyTorch 2.8.
848 * It consumes 16 BF16 values per vector, converts its lower and upper halves
849 * to two 8-lane FP32 vectors, and combines 16-vector Welford chunks
850 * through a binary cascade. The activation arena contains BF16-rounded FP32
851 * values, so the same lane grouping can be reproduced without changing ABI.
852 */
853static inline void layernorm_pytorch_bf16_rowwise_moments_avx2(
854 const float *x,
855 int n_values,
856 float *mean,
857 float *variance)
858{
859 enum { BF16_VEC_SIZE = 16, ACC_VEC_SIZE = 8, CHUNK_SIZE = 16, MAX_DEPTH = 64 };
860 const int n_vec = n_values / BF16_VEC_SIZE;
861 const int chunks = (n_vec + CHUNK_SIZE - 1) / CHUNK_SIZE;
862 int depth = 0;
863 for (int v = chunks; v > 1; v = (v + 1) / 2) {
864 ++depth;
865 }
866 if (depth == 0) {
867 depth = 1;
868 }
869
870 int m0_stack[MAX_DEPTH] = {0};
871 __m256 m1_stack[MAX_DEPTH];
872 __m256 m2_stack[MAX_DEPTH];
873 const __m256 zero = _mm256_setzero_ps();
874 for (int i = 0; i < depth; ++i) {
875 m1_stack[i] = zero;
876 m2_stack[i] = zero;
877 }
878
879 for (int chunk = 0; chunk < chunks; ++chunk) {
880 const int remaining = n_vec - chunk * CHUNK_SIZE;
881 const int count = remaining < CHUNK_SIZE ? remaining : CHUNK_SIZE;
882 __m256 local_m1_lo = zero;
883 __m256 local_m1_hi = zero;
884 __m256 local_m2_lo = zero;
885 __m256 local_m2_hi = zero;
886 const float *chunk_x = x + (size_t)chunk * CHUNK_SIZE * BF16_VEC_SIZE;
887
888 for (int j = 0; j < count; ++j) {
889 const float c = 1.0f / (float)(j + 1);
890 const __m256 c_vec = _mm256_set1_ps(c);
891 const float *row = chunk_x + (size_t)j * BF16_VEC_SIZE;
892 const __m256 x_lo = _mm256_loadu_ps(row);
893 const __m256 x_hi = _mm256_loadu_ps(row + ACC_VEC_SIZE);
894 const __m256 delta_lo = _mm256_sub_ps(x_lo, local_m1_lo);
895 const __m256 delta_hi = _mm256_sub_ps(x_hi, local_m1_hi);
896 local_m1_lo = _mm256_fmadd_ps(delta_lo, c_vec, local_m1_lo);
897 local_m1_hi = _mm256_fmadd_ps(delta_hi, c_vec, local_m1_hi);
898 local_m2_lo = _mm256_fmadd_ps(
899 delta_lo, _mm256_sub_ps(x_lo, local_m1_lo), local_m2_lo);
900 local_m2_hi = _mm256_fmadd_ps(
901 delta_hi, _mm256_sub_ps(x_hi, local_m1_hi), local_m2_hi);
902 }
903
904 layernorm_pytorch_add_moments_vec(
905 count, local_m1_lo, local_m2_lo,
906 &m0_stack[0], &m1_stack[0], &m2_stack[0]);
907 layernorm_pytorch_add_moments_vec(
908 count, local_m1_hi, local_m2_hi,
909 &m0_stack[0], &m1_stack[0], &m2_stack[0]);
910
911 int mask = chunk + 1;
912 for (int level = 1; level < depth && (mask & 1) == 0; ++level) {
913 layernorm_pytorch_add_moments_vec(
914 m0_stack[level - 1], m1_stack[level - 1], m2_stack[level - 1],
915 &m0_stack[level], &m1_stack[level], &m2_stack[level]);
916 m0_stack[level - 1] = 0;
917 m1_stack[level - 1] = zero;
918 m2_stack[level - 1] = zero;
919 mask >>= 1;
920 }
921 }
922 for (int level = 1; level < depth; ++level) {
923 layernorm_pytorch_add_moments_vec(
924 m0_stack[level], m1_stack[level], m2_stack[level],
925 &m0_stack[0], &m1_stack[0], &m2_stack[0]);
926 }
927
928 float m1_lanes[ACC_VEC_SIZE];
929 float m2_lanes[ACC_VEC_SIZE];
930 _mm256_storeu_ps(m1_lanes, m1_stack[0]);
931 _mm256_storeu_ps(m2_lanes, m2_stack[0]);
932 int scalar_count = 0;
933 float scalar_m1 = 0.0f;
934 float scalar_m2 = 0.0f;
935 for (int i = n_vec * BF16_VEC_SIZE; i < n_values; ++i) {
936 const float delta = x[i] - scalar_m1;
937 ++scalar_count;
938 scalar_m1 += delta / (float)scalar_count;
939 scalar_m2 += delta * (x[i] - scalar_m1);
940 }
941 const int lane_count = n_vec * BF16_VEC_SIZE / ACC_VEC_SIZE;
942 for (int lane = 0; lane < ACC_VEC_SIZE; ++lane) {
943 layernorm_pytorch_add_moments_scalar(
944 lane_count, m1_lanes[lane], m2_lanes[lane],
945 &scalar_count, &scalar_m1, &scalar_m2);
946 }
947 *mean = scalar_m1;
948 *variance = scalar_m2 / (float)n_values;
949}
950#endif
951
953 const float *gamma,
954 const float *beta,
955 float *output,
956 float *mean_cache,
957 float *rstd_cache,
958 int tokens,
959 int d_model,
960 float eps)
961{
962#if !defined(__AVX2__) || !defined(__FMA__)
963 (void)input; (void)gamma; (void)beta; (void)output;
964 (void)mean_cache; (void)rstd_cache; (void)tokens; (void)d_model; (void)eps;
965 abort();
966#else
967 for (int t = 0; t < tokens; ++t) {
968 const float *x = input + (size_t)t * (size_t)d_model;
969 float *y = output + (size_t)t * (size_t)d_model;
970 float mean;
971 float variance;
972 layernorm_pytorch_bf16_rowwise_moments_avx2(x, d_model, &mean, &variance);
973 const float rstd = 1.0f / sqrtf(variance + eps);
974 const float bias = -rstd * mean;
975 int i = 0;
976 for (; i + 7 < d_model; i += 8) {
977 const __m256 x_vec = _mm256_loadu_ps(x + i);
978 const __m256 gamma_vec = gamma ? _mm256_loadu_ps(gamma + i) : _mm256_set1_ps(1.0f);
979 const __m256 beta_vec = beta ? _mm256_loadu_ps(beta + i) : _mm256_setzero_ps();
980 const __m256 normalized = _mm256_fmadd_ps(
981 x_vec, _mm256_set1_ps(rstd), _mm256_set1_ps(bias));
982 const __m256 transformed = _mm256_fmadd_ps(normalized, gamma_vec, beta_vec);
983 float lanes[8];
984 _mm256_storeu_ps(lanes, transformed);
985 for (int lane = 0; lane < 8; ++lane) {
986 y[i + lane] = bf16_to_float(float_to_bf16(lanes[lane]));
987 }
988 }
989 for (; i < d_model; ++i) {
990 const float gamma_v = gamma ? gamma[i] : 1.0f;
991 const float beta_v = beta ? beta[i] : 0.0f;
992 const float value = fmaf(fmaf(x[i], rstd, bias), gamma_v, beta_v);
993 y[i] = bf16_to_float(float_to_bf16(value));
994 }
995 if (mean_cache) {
996 mean_cache[t] = mean;
997 }
998 if (rstd_cache) {
999 rstd_cache[t] = rstd;
1000 }
1001 }
1002#endif
1003}
1004
1005// LayerNorm backward kernel (model-agnostic), adapted from C-Transformer's
1006// backward_layernorm. Computes gradients w.r.t. input, gamma, and beta.
1007void layernorm_backward_kernel(const float *d_output, // [T×aligned_D]
1008 const float *input, // [T×aligned_D]
1009 const float *gamma, // [D]
1010 const float *mean, // [T]
1011 const float *rstd, // [T]
1012 float *d_input, // [T×aligned_D]
1013 float *d_gamma, // [D] (accumulated)
1014 float *d_beta, // [D] (accumulated)
1015 int tokens, int d_model, int aligned_embed_dim)
1016{
1017 int T = tokens;
1018 int D = d_model;
1019 int aligned_D = aligned_embed_dim;
1020
1021 // Per-token input gradients
1022 for (int t = 0; t < T; ++t) {
1023 float mean_t = mean[t];
1024 float rstd_t = rstd[t];
1025
1026 float d_y_gamma_sum = 0.0f;
1027 float d_y_gamma_xhat_sum = 0.0f;
1028
1029 // First pass: compute sums
1030 for (int d = 0; d < D; ++d) {
1031 float x = input[t * aligned_D + d];
1032 float x_hat = (x - mean_t) * rstd_t;
1033 float d_y = d_output[t * aligned_D + d];
1034 float d_y_gamma = d_y * gamma[d];
1035
1036 d_y_gamma_sum += d_y_gamma;
1037 d_y_gamma_xhat_sum += d_y_gamma * x_hat;
1038 }
1039
1040 // Second pass: compute input gradients
1041 float scale = rstd_t / (float)D;
1042 for (int d = 0; d < D; ++d) {
1043 float x = input[t * aligned_D + d];
1044 float x_hat = (x - mean_t) * rstd_t;
1045 float d_y = d_output[t * aligned_D + d];
1046
1047 d_input[t * aligned_D + d] =
1048 scale * ((float)D * d_y * gamma[d] - d_y_gamma_sum - x_hat * d_y_gamma_xhat_sum);
1049 }
1050
1051 // Zero padding for aligned dimension beyond D
1052 for (int d = D; d < aligned_D; ++d) {
1053 d_input[t * aligned_D + d] = 0.0f;
1054 }
1055 }
1056
1057 // Parameter gradients (gamma, beta)
1058 for (int d = 0; d < D; ++d) {
1059 float gamma_grad = 0.0f;
1060 float beta_grad = 0.0f;
1061
1062 for (int t = 0; t < T; ++t) {
1063 float x = input[t * aligned_D + d];
1064 float x_hat = (x - mean[t]) * rstd[t];
1065 float d_y = d_output[t * aligned_D + d];
1066
1067 gamma_grad += d_y * x_hat;
1068 beta_grad += d_y;
1069 }
1070
1071 d_gamma[d] += gamma_grad;
1072 d_beta[d] += beta_grad;
1073 }
1074}
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
int ck_strict_parity_enabled(void)
static void layernorm_forward_ggml_exact(const float *input, const float *gamma, const float *beta, float *output, float *mean_cache, float *rstd_cache, int tokens, int d_model, int input_stride, int output_stride, int aligned_embed_dim, float eps)
void layernorm_naive_serial_matched_precision(const float *input, const float *gamma, const float *beta, float *output, float *mean_cache, float *rstd_cache, int tokens, int d_model, float eps)
void layernorm_pytorch_welford_bf16_storage(const float *input, const float *gamma, const float *beta, float *output, float *mean_cache, float *rstd_cache, int tokens, int d_model, float eps)
void layernorm_naive_serial_bf16_storage(const float *input, const float *gamma, const float *beta, float *output, float *mean_cache, float *rstd_cache, int tokens, int d_model, float eps)
static void zero_layernorm_padding(float *out_ptr, int d_model, int aligned_embed_dim)
void layernorm_naive_serial(const float *input, const float *gamma, const float *beta, float *output, float *mean_cache, float *rstd_cache, int tokens, int d_model, int aligned_embed_dim, float eps)
void layernorm_backward_kernel(const float *d_output, const float *input, const float *gamma, const float *mean, const float *rstd, float *d_input, float *d_gamma, float *d_beta, int tokens, int d_model, int aligned_embed_dim)
static void layernorm_forward_unrolled_slice_scalar(const float *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, float *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, float eps)
void layernorm_forward_rolled_slice(const float *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, float *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, int aligned_embed_dim, float eps)
void layernorm_forward_unrolled_slice(const float *__restrict input_slice_base, const float *__restrict gamma, const float *__restrict beta, float *__restrict output_slice_base, float *__restrict mean_cache_slice, float *__restrict rstd_cache_slice, int num_tokens_in_slice, int d_model, float eps)
int32_t int32_t int32_t int32_t int32_t mask
Definition tokenizer.h:234