24#if defined(__AVX512F__)
36#pragma GCC diagnostic push
37#pragma GCC diagnostic ignored "-Wmaybe-uninitialized"
41 const float *bias,
float *
C,
47 const float c = 0.7978845608f;
48 const float k = 0.044715f;
50 return 0.5f * x * (1.0f + tanhf(c * (x + k * x3)));
55 const float c = 0.7978845608f;
56 const float k = 0.044715f;
57 const float x2 = x * x;
58 const float x3 = x2 * x;
59 const float g = c * (x + k * x3);
60 const float tanh_g = tanhf(g);
61 const float sech2_g = 1.0f - tanh_g * tanh_g;
62 const float g_prime = c * (1.0f + 3.0f * k * x2);
63 return 0.5f * (1.0f + tanh_g) + 0.5f * x * sech2_g * g_prime;
66#if defined(__AVX512F__)
68static inline __m512 gelu_avx512(__m512 x)
70 const __m512 c = _mm512_set1_ps(0.7978845608f);
71 const __m512 k = _mm512_set1_ps(0.044715f);
72 const __m512 half = _mm512_set1_ps(0.5f);
73 const __m512 one = _mm512_set1_ps(1.0f);
75 __m512 x2 = _mm512_mul_ps(x, x);
76 __m512 x3 = _mm512_mul_ps(x2, x);
78 __m512 inner = _mm512_fmadd_ps(k, x3, x);
79 inner = _mm512_mul_ps(c, inner);
81 __m512 inner2 = _mm512_mul_ps(inner, inner);
82 __m512 num = _mm512_add_ps(_mm512_set1_ps(27.0f), inner2);
83 __m512 den = _mm512_fmadd_ps(_mm512_set1_ps(9.0f), inner2, _mm512_set1_ps(27.0f));
84 __m512 tanh_approx = _mm512_mul_ps(inner, _mm512_div_ps(num, den));
86 tanh_approx = _mm512_min_ps(tanh_approx, one);
87 tanh_approx = _mm512_max_ps(tanh_approx, _mm512_set1_ps(-1.0f));
89 __m512 result = _mm512_add_ps(one, tanh_approx);
90 result = _mm512_mul_ps(half, _mm512_mul_ps(x, result));
105 const uint16_t *W_fc1,
106 const uint16_t *b_fc1,
107 const uint16_t *W_fc2,
108 const uint16_t *b_fc2,
114 float *scratch_bias1_f,
115 float *scratch_bias2_f,
116 uint16_t *scratch_fc1_bf16)
118 if (!input || !W_fc1 || !b_fc1 || !W_fc2 || !b_fc2 || !fc1_output || !output)
return;
119 if (!scratch_bias1_f || !scratch_bias2_f || !scratch_fc1_bf16)
return;
122 const int D = aligned_dim;
123 const int fourD = 4 * D;
126 for (
int i = 0; i < fourD; ++i) {
129 for (
int i = 0; i < D; ++i) {
137#if defined(__AVX512F__)
138 #pragma omp parallel for
139 for (
int t = 0; t < T; ++t) {
140 float *row = fc1_output + (size_t)t * fourD;
142 for (; j <= fourD - 16; j += 16) {
143 __m512 x = _mm512_loadu_ps(row + j);
144 _mm512_storeu_ps(row + j, gelu_avx512(x));
146 for (; j < fourD; ++j) {
151 for (
int t = 0; t < T; ++t) {
152 for (
int j = 0; j < fourD; ++j) {
153 fc1_output[t * fourD + j] =
gelu_scalar(fc1_output[t * fourD + j]);
159#if defined(__AVX512F__)
160 #pragma omp parallel for
161 for (
int t = 0; t < T; ++t) {
162 float *src = fc1_output + (size_t)t * fourD;
163 uint16_t *dst = scratch_fc1_bf16 + (size_t)t * fourD;
165 for (; j <= fourD - 16; j += 16) {
166 __m512 fp32 = _mm512_loadu_ps(src + j);
167 __m512i as_int = _mm512_castps_si512(fp32);
168 __m512i lsb = _mm512_srli_epi32(as_int, 16);
169 lsb = _mm512_and_si512(lsb, _mm512_set1_epi32(1));
170 __m512i rounding = _mm512_add_epi32(_mm512_set1_epi32(0x7FFF), lsb);
171 __m512i rounded = _mm512_add_epi32(as_int, rounding);
172 __m512i shifted = _mm512_srli_epi32(rounded, 16);
173 __m256i bf16 = _mm512_cvtepi32_epi16(shifted);
174 _mm256_storeu_si256((__m256i *)(dst + j), bf16);
176 for (; j < fourD; ++j) {
181 for (
size_t i = 0; i < (size_t)T * fourD; ++i) {
187 gemm_bf16_fp32out(scratch_fc1_bf16, W_fc2, scratch_bias2_f, output, T, D, fourD);
200 const uint16_t *W_fc1,
201 const uint16_t *b_fc1,
202 const uint16_t *W_fc2,
203 const uint16_t *b_fc2,
209 float *scratch_input_f,
210 float *scratch_bias1_f,
211 float *scratch_bias2_f,
212 uint16_t *scratch_fc1_bf16)
214 if (!input || !W_fc1 || !b_fc1 || !W_fc2 || !b_fc2 || !fc1_output || !output)
return;
215 if (!scratch_input_f || !scratch_bias1_f || !scratch_bias2_f || !scratch_fc1_bf16)
return;
218 const int D = aligned_dim;
219 const int fourD = 4 * D;
230#if defined(__AVX512F__)
231 #pragma omp parallel for
232 for (
int t = 0; t < T; ++t) {
233 float *row = fc1_output + (size_t)t * fourD;
235 for (; j <= fourD - 16; j += 16) {
236 __m512 x = _mm512_loadu_ps(row + j);
237 _mm512_storeu_ps(row + j, gelu_avx512(x));
239 for (; j < fourD; ++j) {
244 for (
size_t i = 0; i < (size_t)T * fourD; ++i) {
251 gemm_bf16_fp32out(scratch_fc1_bf16, W_fc2, scratch_bias2_f, output, T, D, fourD);
269 const uint16_t *W_fc1,
270 const uint16_t *b_fc1,
271 const uint16_t *W_fc2,
272 const uint16_t *d_output,
281 float *scratch_fc1_pre,
282 uint16_t *scratch_fc1_act_bf16,
283 float *scratch_d_fc1)
285 if (!input || !W_fc1 || !b_fc1 || !W_fc2 || !d_output)
return;
286 if (!scratch_fc1_pre || !scratch_fc1_act_bf16 || !scratch_d_fc1)
return;
287 if (T <= 0 || aligned_dim <= 0)
return;
290 const int D = aligned_dim;
291 const int fourD = 4 * D;
294 for (
int t = 0; t < T; ++t) {
295 for (
int j = 0; j < fourD; ++j) {
297 for (
int i = 0; i < D; ++i) {
298 const float x =
bf16_to_float(input[(
size_t)t * (
size_t)D + (
size_t)i]);
299 const float w =
bf16_to_float(W_fc1[(
size_t)j * (
size_t)D + (
size_t)i]);
302 scratch_fc1_pre[(size_t)t * (
size_t)fourD + (size_t)j] = sum;
308 for (
int t = 0; t < T; ++t) {
309 for (
int i = 0; i < D; ++i) {
310 d_input[(size_t)t * (
size_t)D + (size_t)i] = 0.0f;
315 for (
size_t i = 0; i < (size_t)fourD * (
size_t)D; ++i) d_W_fc1[i] = 0.0f;
318 for (
int j = 0; j < fourD; ++j) d_b_fc1[j] = 0.0f;
321 for (
size_t i = 0; i < (size_t)D * (
size_t)fourD; ++i) d_W_fc2[i] = 0.0f;
324 for (
int o = 0; o < D; ++o) d_b_fc2[o] = 0.0f;
328 for (
int t = 0; t < T; ++t) {
329 for (
int j = 0; j < fourD; ++j) {
331 const float hq =
bf16_to_float(scratch_fc1_act_bf16[(
size_t)t * (
size_t)fourD + (
size_t)j]);
332 for (
int o = 0; o < D; ++o) {
333 const float dy =
bf16_to_float(d_output[(
size_t)t * (
size_t)D + (
size_t)o]);
334 const float w2 =
bf16_to_float(W_fc2[(
size_t)o * (
size_t)fourD + (
size_t)j]);
337 d_W_fc2[(size_t)o * (
size_t)fourD + (size_t)j] += dy * hq;
340 const float z = scratch_fc1_pre[(size_t)t * (
size_t)fourD + (size_t)j];
344 for (
int o = 0; o < D; ++o) {
345 d_b_fc2[o] +=
bf16_to_float(d_output[(
size_t)t * (
size_t)D + (
size_t)o]);
351 for (
int t = 0; t < T; ++t) {
352 for (
int j = 0; j < fourD; ++j) {
353 const float dz = scratch_d_fc1[(size_t)t * (
size_t)fourD + (size_t)j];
354 if (d_b_fc1) d_b_fc1[j] += dz;
355 for (
int i = 0; i < D; ++i) {
356 const float x =
bf16_to_float(input[(
size_t)t * (
size_t)D + (
size_t)i]);
357 const float w1 =
bf16_to_float(W_fc1[(
size_t)j * (
size_t)D + (
size_t)i]);
358 if (d_W_fc1) d_W_fc1[(size_t)j * (
size_t)D + (size_t)i] += dz * x;
359 if (d_input) d_input[(size_t)t * (
size_t)D + (size_t)i] += dz * w1;
365#pragma GCC diagnostic pop
static void float_tensor_to_bf16(const float *src, uint16_t *dst, size_t count)
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)
static void bf16_tensor_to_float(const uint16_t *src, float *dst, size_t count)
static float gelu_scalar(float x)
void mlp_token_parallel_bf16(const uint16_t *input, const uint16_t *W_fc1, const uint16_t *b_fc1, const uint16_t *W_fc2, const uint16_t *b_fc2, float *fc1_output, float *output, int T, int aligned_dim, int num_threads, float *scratch_bias1_f, float *scratch_bias2_f, uint16_t *scratch_fc1_bf16)
void mlp_token_parallel_bf16_fp32act(const uint16_t *input, const uint16_t *W_fc1, const uint16_t *b_fc1, const uint16_t *W_fc2, const uint16_t *b_fc2, float *fc1_output, float *output, int T, int aligned_dim, int num_threads, float *scratch_input_f, float *scratch_bias1_f, float *scratch_bias2_f, uint16_t *scratch_fc1_bf16)
void gemm_bf16_fp32out(const uint16_t *A, const uint16_t *B, const float *bias, float *C, int M, int N, int K)
void mlp_token_parallel_bf16_backward_mixed(const uint16_t *input, const uint16_t *W_fc1, const uint16_t *b_fc1, const uint16_t *W_fc2, const uint16_t *d_output, float *d_input, float *d_W_fc1, float *d_b_fc1, float *d_W_fc2, float *d_b_fc2, int T, int aligned_dim, int num_threads, float *scratch_fc1_pre, uint16_t *scratch_fc1_act_bf16, float *scratch_d_fc1)
static float gelu_derivative_scalar(float x)