20#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__) || defined(__SSE2__)
28 int aligned_embed_dim)
30 for (
int idx = d_model; idx < aligned_embed_dim; ++idx) {
41 int tokens,
int d_model,
float eps);
43#if defined(__clang__) || defined(__INTEL_LLVM_COMPILER)
44#pragma float_control(precise, on, push)
56 int aligned_embed_dim,
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;
65#pragma clang loop vectorize(disable)
66#pragma clang loop interleave(disable)
68 for (
int i = 0; i < d_model; ++i) {
69 sum_acc += (double)x[i];
71 const float sum = (float)sum_acc;
72 const float mean = sum / (float)d_model;
76#if defined(__AVX2__) && defined(__FMA__)
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);
91#elif defined(__SSE2__)
92 for (; i + 3 < d_model; i += 4) {
93 __m128 val = _mm_sub_ps(_mm_loadu_ps(x + i),
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));
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);
106 var_acc += (double)_mm_cvtss_f32(val);
109#if defined(__clang__)
110#pragma clang loop vectorize(disable)
111#pragma clang loop interleave(disable)
113 for (; i < d_model; ++i) {
114 const float centered = x[i] - mean;
116 var_acc += (double)(centered * centered);
118 const float variance = (float)(var_acc / (
double)d_model);
119 const float scale = 1.0f / sqrtf(variance + eps);
122 mean_cache[t] = mean;
125 rstd_cache[t] = scale;
128#if defined(__clang__)
129#pragma clang loop vectorize(disable)
130#pragma clang loop interleave(disable)
132 for (
int i = 0; i < d_model; ++i) {
136#if defined(__clang__)
137#pragma clang loop vectorize(disable)
138#pragma clang loop interleave(disable)
140 for (
int i = 0; i < d_model; ++i) {
145#if defined(__clang__)
146#pragma clang loop vectorize(disable)
147#pragma clang loop interleave(disable)
149 for (
int i = 0; i < d_model; ++i) {
154 if (aligned_embed_dim > d_model) {
159#if defined(__clang__) || defined(__INTEL_LLVM_COMPILER)
160#pragma float_control(pop)
163#if defined(__AVX2__) || defined(__AVX__)
164static inline float hsum256_ps(__m256 v)
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);
181 int tokens,
int d_model,
int aligned_embed_dim,
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;
188 float sum_val = 0.0f;
189 for (
int i = 0; i < d_model; ++i) {
190 sum_val += in_ptr[i];
192 float mean = sum_val / (float)d_model;
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;
199 float variance = sum_sq_diff / (float)d_model + eps;
201 double var_double = (double)variance;
202 float inv_std = (float)(1.0 / sqrt(var_double));
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];
210 mean_cache[t] = mean;
213 rstd_cache[t] = inv_std;
216 if (aligned_embed_dim > d_model) {
218 for (
int i = d_model; i < aligned_embed_dim; ++i) {
225#if defined(__AVX512F__)
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,
235 int aligned_embed_dim,
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;
242 __m512 acc_sum_vec = _mm512_setzero_ps();
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);
249 float mean = _mm512_reduce_add_ps(acc_sum_vec);
250 for (; j < d_model; ++j) {
251 mean += in_ptr_token[j];
253 mean /= (float)d_model;
254 __m512 mean_vec = _mm512_set1_ps(mean);
256 __m512 acc_var_vec = _mm512_setzero_ps();
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);
264 float var = _mm512_reduce_add_ps(acc_var_vec);
265 for (; j < d_model; ++j) {
266 float diff = in_ptr_token[j] - mean;
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);
274 if (mean_cache_slice) {
275 mean_cache_slice[t] = mean;
277 if (rstd_cache_slice) {
278 rstd_cache_slice[t] = inv_std;
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);
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);
291 __m512 n = _mm512_mul_ps(_mm512_sub_ps(v, mean_vec), inv_std_vec);
292 __m512 o = _mm512_fmadd_ps(n, g, b);
294 _mm512_store_ps(out_ptr_token + j, o);
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];
301 if (aligned_embed_dim > d_model) {
307#elif defined(__AVX2__) || defined(__AVX__)
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,
317 int aligned_embed_dim,
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;
324 __m256 acc_sum_vec = _mm256_setzero_ps();
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);
331 float mean = hsum256_ps(acc_sum_vec);
332 for (; j < d_model; ++j) {
333 mean += in_ptr_token[j];
335 mean /= (float)d_model;
336 __m256 mean_vec = _mm256_set1_ps(mean);
338 __m256 acc_var_vec = _mm256_setzero_ps();
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);
345 acc_var_vec = _mm256_fmadd_ps(diff, diff, acc_var_vec);
347 acc_var_vec = _mm256_add_ps(acc_var_vec, _mm256_mul_ps(diff, diff));
350 float var = hsum256_ps(acc_var_vec);
351 for (; j < d_model; ++j) {
352 float diff = in_ptr_token[j] - mean;
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);
360 if (mean_cache_slice) {
361 mean_cache_slice[t] = mean;
363 if (rstd_cache_slice) {
364 rstd_cache_slice[t] = inv_std;
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);
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);
377 __m256 n = _mm256_mul_ps(_mm256_sub_ps(v, mean_vec), inv_std_vec);
379 __m256 o = _mm256_fmadd_ps(n, g, b);
381 __m256 o = _mm256_add_ps(_mm256_mul_ps(n, g), b);
384 _mm256_store_ps(out_ptr_token + j, o);
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];
391 if (aligned_embed_dim > d_model) {
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,
406 int aligned_embed_dim,
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);
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);
427 output_slice_base, mean_cache_slice, rstd_cache_slice,
428 num_tokens_in_slice, d_model, aligned_embed_dim, eps);
432#if defined(__AVX512F__)
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,
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;
448 __m512 acc0 = _mm512_setzero_ps();
449 __m512 acc1 = _mm512_setzero_ps();
450 __m512 acc2 = _mm512_setzero_ps();
451 __m512 acc3 = _mm512_setzero_ps();
454 int unroll_factor_floats = 64;
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);
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);
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);
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);
473 for (; j < d_model; ++j) {
474 mean += in_ptr_token[j];
476 mean /= (float)d_model;
477 __m512 mean_vec = _mm512_set1_ps(mean);
479 acc0 = _mm512_setzero_ps();
480 acc1 = _mm512_setzero_ps();
481 acc2 = _mm512_setzero_ps();
482 acc3 = _mm512_setzero_ps();
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);
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);
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);
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);
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);
507 for (; j < d_model; ++j) {
508 float diff = in_ptr_token[j] - mean;
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);
516 if (mean_cache_slice) {
517 mean_cache_slice[t] = mean;
519 if (rstd_cache_slice) {
520 rstd_cache_slice[t] = inv_std;
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);
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);
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);
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);
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);
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);
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);
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];
565#elif defined(__AVX2__) || defined(__AVX__)
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,
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;
581 __m256 acc0 = _mm256_setzero_ps();
582 __m256 acc1 = _mm256_setzero_ps();
583 __m256 acc2 = _mm256_setzero_ps();
584 __m256 acc3 = _mm256_setzero_ps();
587 int unroll_factor_floats = 32;
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);
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);
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);
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);
606 for (; j < d_model; ++j) {
607 mean += in_ptr_token[j];
609 mean /= (float)d_model;
610 __m256 mean_vec = _mm256_set1_ps(mean);
612 acc0 = _mm256_setzero_ps();
613 acc1 = _mm256_setzero_ps();
614 acc2 = _mm256_setzero_ps();
615 acc3 = _mm256_setzero_ps();
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);
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);
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);
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);
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));
643 acc_sum = _mm256_add_ps(_mm256_add_ps(acc0, acc1),
644 _mm256_add_ps(acc2, acc3));
645 float var = hsum256_ps(acc_sum);
647 for (; j < d_model; ++j) {
648 float diff = in_ptr_token[j] - mean;
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);
656 if (mean_cache_slice) {
657 mean_cache_slice[t] = mean;
659 if (rstd_cache_slice) {
660 rstd_cache_slice[t] = inv_std;
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);
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);
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);
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);
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);
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);
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);
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);
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];
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,
725 output_slice_base, mean_cache_slice, rstd_cache_slice,
726 num_tokens_in_slice, d_model, eps);
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,
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);
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);
758 output_slice_base, mean_cache_slice, rstd_cache_slice,
759 num_tokens_in_slice, d_model, eps);
770 int tokens,
int d_model,
float eps)
773 output, mean_cache, rstd_cache,
775 d_model, d_model, d_model, eps);
792 int tokens,
int d_model,
float eps)
794 const size_t count = (size_t)tokens * (
size_t)d_model;
795 for (
size_t i = 0; i < count; ++i) {
799 output, mean_cache, rstd_cache,
801 d_model, d_model, d_model, eps);
802 for (
size_t i = 0; i < count; ++i) {
807#if defined(__AVX2__) && defined(__FMA__)
808static inline void layernorm_pytorch_add_moments_vec(
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);
829static inline void layernorm_pytorch_add_moments_scalar(
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);
853static inline void layernorm_pytorch_bf16_rowwise_moments_avx2(
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;
863 for (
int v = chunks; v > 1; v = (v + 1) / 2) {
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) {
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;
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);
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]);
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;
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]);
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;
938 scalar_m1 += delta / (float)scalar_count;
939 scalar_m2 += delta * (x[i] - scalar_m1);
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);
948 *variance = scalar_m2 / (float)n_values;
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;
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;
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;
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);
984 _mm256_storeu_ps(lanes, transformed);
985 for (
int lane = 0; lane < 8; ++lane) {
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);
996 mean_cache[t] = mean;
999 rstd_cache[t] = rstd;
1015 int tokens,
int d_model,
int aligned_embed_dim)
1019 int aligned_D = aligned_embed_dim;
1022 for (
int t = 0; t < T; ++t) {
1023 float mean_t = mean[t];
1024 float rstd_t = rstd[t];
1026 float d_y_gamma_sum = 0.0f;
1027 float d_y_gamma_xhat_sum = 0.0f;
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];
1036 d_y_gamma_sum += d_y_gamma;
1037 d_y_gamma_xhat_sum += d_y_gamma * x_hat;
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];
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);
1052 for (
int d = D; d < aligned_D; ++d) {
1053 d_input[t * aligned_D + d] = 0.0f;
1058 for (
int d = 0; d < D; ++d) {
1059 float gamma_grad = 0.0f;
1060 float beta_grad = 0.0f;
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];
1067 gamma_grad += d_y * x_hat;
1071 d_gamma[d] += gamma_grad;
1072 d_beta[d] += beta_grad;
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
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