27#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__)
127 const float *fc2_input,
144 T, aligned_in, aligned_out);
153 aligned_out, aligned_in, T);
156 for (
int out_idx = 0; out_idx < aligned_out; ++out_idx) {
157 float bias_grad = 0.0f;
158 for (
int t = 0; t < T; ++t) {
159 bias_grad += d_output[(size_t)t * aligned_out + out_idx];
161 d_b_fc2[out_idx] += bias_grad;
175 const float *fc1_input,
192 T, aligned_in, aligned_out);
199 aligned_out, aligned_in, T);
202 for (
int out_idx = 0; out_idx < aligned_out; ++out_idx) {
203 float bias_grad = 0.0f;
204 for (
int t = 0; t < T; ++t) {
205 bias_grad += d_output[(size_t)t * aligned_out + out_idx];
207 d_b_fc1[out_idx] += bias_grad;
void gelu_exact_inplace(float *data, size_t n)
void gemm_tn_parallel(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gemm_nn_simd(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void gelu_fast_inplace(float *data, size_t n)
void gemm_blocked_serial(const float *A, const float *B, const float *bias, float *C, int M, int N, int K)
void mlp_token_parallel(const float *input, const float *W_fc1, const float *b_fc1, const float *W_fc2, const float *b_fc2, float *fc1_output, float *output, int T, int aligned_dim, int num_threads)
void fc1_backward_kernel(const float *d_output, const float *fc1_input, const float *W_fc1, float *d_input, float *d_W_fc1, float *d_b_fc1, int T, int aligned_in, int aligned_out, int num_threads)
void mlp_token_parallel_exact(const float *input, const float *W_fc1, const float *b_fc1, const float *W_fc2, const float *b_fc2, float *fc1_output, float *output, int T, int aligned_dim, int num_threads)
void fc2_backward_kernel(const float *d_output, const float *fc2_input, const float *W_fc2, float *d_input, float *d_W_fc2, float *d_b_fc2, int T, int aligned_in, int aligned_out, int num_threads)