122{
123 int T = tokens;
124 int D = d_model;
125 int aligned = aligned_embed_dim;
126
127 if (!d_output || !input || !gamma || !rstd_cache || !d_input || !d_gamma) {
128 return;
129 }
130
131
132#if defined(__AVX512F__)
133 {
134 int d = 0;
135 for (; d + 16 <= D; d += 16) {
136 _mm512_storeu_ps(&d_gamma[d], _mm512_setzero_ps());
137 }
138 for (; d < D; ++d) {
139 d_gamma[d] = 0.0f;
140 }
141 }
142#else
143 for (int d = 0; d < D; ++d) {
144 d_gamma[d] = 0.0f;
145 }
146#endif
147
148 for (int t = 0; t < T; ++t) {
149 const uint16_t *x_bf16 = input + (size_t)t * aligned;
150 const uint16_t *dY_bf16 = d_output + (size_t)t * aligned;
151 uint16_t *dX_bf16 = d_input + (size_t)t * aligned;
152 float rstd = rstd_cache[t];
153
154#if defined(__AVX512F__)
155
156 __m512 rstd_vec = _mm512_set1_ps(rstd);
157 __m512 sum_vec = _mm512_setzero_ps();
158 int d = 0;
159
160 for (; d + 16 <= D; d += 16) {
161 __m512 xv = bf16_loadu_cvt_fp32(&x_bf16[d]);
162 __m512 dyv = bf16_loadu_cvt_fp32(&dY_bf16[d]);
163 __m512 gv = _mm512_loadu_ps(&gamma[d]);
164 __m512 x_hat = _mm512_mul_ps(xv, rstd_vec);
165
166 __m512 prod = _mm512_mul_ps(dyv, gv);
167 sum_vec = _mm512_fmadd_ps(prod, x_hat, sum_vec);
168 }
169 float sum_dY_g_xhat = _mm512_reduce_add_ps(sum_vec);
170
171
172 for (; d < D; ++d) {
174 float x_hat = x * rstd;
176 sum_dY_g_xhat += dy * gamma[d] * x_hat;
177 }
178 float m = sum_dY_g_xhat / (float)D;
179
180
181 __m512 m_vec = _mm512_set1_ps(m);
182 d = 0;
183 for (; d + 16 <= D; d += 16) {
184 __m512 xv = bf16_loadu_cvt_fp32(&x_bf16[d]);
185 __m512 dyv = bf16_loadu_cvt_fp32(&dY_bf16[d]);
186 __m512 gv = _mm512_loadu_ps(&gamma[d]);
187 __m512 dgv = _mm512_loadu_ps(&d_gamma[d]);
188
189 __m512 x_hat = _mm512_mul_ps(xv, rstd_vec);
190
191
192 __m512 dy_g = _mm512_mul_ps(dyv, gv);
193 __m512 xhat_m = _mm512_mul_ps(x_hat, m_vec);
194 __m512 diff = _mm512_sub_ps(dy_g, xhat_m);
195 __m512 dxv = _mm512_mul_ps(rstd_vec, diff);
196 fp32_cvt_storeu_bf16(&dX_bf16[d], dxv);
197
198
199 dgv = _mm512_fmadd_ps(dyv, x_hat, dgv);
200 _mm512_storeu_ps(&d_gamma[d], dgv);
201 }
202
203 for (; d < D; ++d) {
205 float x_hat = x * rstd;
207 float dx = rstd * (dy * gamma[d] - x_hat * m);
209 d_gamma[d] += dy * x_hat;
210 }
211
212#else
213
214 double sum_dY_g_xhat = 0.0;
215 for (int d = 0; d < D; ++d) {
217 float x_hat = x * rstd;
219 sum_dY_g_xhat += (double)dy * (double)gamma[d] * (double)x_hat;
220 }
221 float m = (float)(sum_dY_g_xhat / (double)D);
222
223 for (int d = 0; d < D; ++d) {
225 float x_hat = x * rstd;
227 float dx = rstd * (dy * gamma[d] - x_hat * m);
229 d_gamma[d] += dy * x_hat;
230 }
231#endif
232
233
234 for (int d = D; d < aligned; ++d) {
235 dX_bf16[d] = 0;
236 }
237 }
238}
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)