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

LayerNorm forward/backward kernels with SIMD (SSE/AVX/AVX512) More...

#include "ckernel_engine.h"
#include "bf16_utils.h"
#include <math.h>
#include <stdlib.h>

Go to the source code of this file.

Functions

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_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_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)
 
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_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_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)
 
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)
 
static void zero_layernorm_padding (float *out_ptr, int d_model, int aligned_embed_dim)
 

Detailed Description

LayerNorm forward/backward kernels with SIMD (SSE/AVX/AVX512)

CK-ENGINE KERNEL RULES:

  1. NO malloc/free - memory via bump allocator, pointers passed in
  2. NO OpenMP - parallelization at orchestrator/codegen layer
  3. API must define: inputs, outputs, workspace, and memory layouts
  4. Pure computation - deterministic, no side effects

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

LayerNorm: y = gamma * (x - mean) / sqrt(var + eps) + beta

Definition in file layernorm_kernels.c.

Function Documentation

◆ layernorm_backward_kernel()

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 
)

Definition at line 1007 of file layernorm_kernels.c.

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}

Referenced by layernorm_backward_kernel_bf16().

◆ layernorm_forward_ggml_exact()

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 
)
static

Definition at line 46 of file layernorm_kernels.c.

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}
static void zero_layernorm_padding(float *out_ptr, int d_model, int aligned_embed_dim)

References zero_layernorm_padding().

Referenced by layernorm_forward_rolled_slice(), layernorm_forward_unrolled_slice(), layernorm_naive_serial_bf16_storage(), and layernorm_naive_serial_matched_precision().

◆ layernorm_forward_rolled_slice()

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 
)

Definition at line 398 of file layernorm_kernels.c.

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}
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(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)

References ck_strict_parity_enabled(), layernorm_forward_ggml_exact(), and layernorm_naive_serial().

Referenced by layernorm_forward_rolled_slice_bf16().

◆ layernorm_forward_unrolled_slice()

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 
)

Definition at line 730 of file layernorm_kernels.c.

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}
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)

References ck_strict_parity_enabled(), layernorm_forward_ggml_exact(), and layernorm_forward_unrolled_slice_scalar().

Referenced by layernorm_forward_unrolled_slice_bf16().

◆ layernorm_forward_unrolled_slice_scalar()

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 
)
static

Definition at line 714 of file layernorm_kernels.c.

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}
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)

References layernorm_naive_serial_matched_precision().

Referenced by layernorm_forward_unrolled_slice().

◆ layernorm_naive_serial()

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 
)

Definition at line 175 of file layernorm_kernels.c.

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}

Referenced by layernorm_forward_rolled_slice().

◆ layernorm_naive_serial_bf16_storage()

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 
)

Definition at line 786 of file layernorm_kernels.c.

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}
static uint16_t float_to_bf16(float f)
Definition bf16_utils.h:90
static float bf16_to_float(uint16_t v)
Definition bf16_utils.h:38

References bf16_to_float(), float_to_bf16(), and layernorm_forward_ggml_exact().

◆ layernorm_naive_serial_matched_precision()

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 
)

Definition at line 764 of file layernorm_kernels.c.

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}

References layernorm_forward_ggml_exact().

Referenced by layernorm_forward_unrolled_slice_scalar().

◆ layernorm_pytorch_welford_bf16_storage()

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 
)

Definition at line 952 of file layernorm_kernels.c.

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}

References bf16_to_float(), and float_to_bf16().

◆ zero_layernorm_padding()

static void zero_layernorm_padding ( float *  out_ptr,
int  d_model,
int  aligned_embed_dim 
)
inlinestatic

Definition at line 26 of file layernorm_kernels.c.

29{
30 for (int idx = d_model; idx < aligned_embed_dim; ++idx) {
31 out_ptr[idx] = 0.0f;
32 }
33}

Referenced by layernorm_forward_ggml_exact().