30#if defined(__AVX512F__) || defined(__AVX2__) || defined(__AVX__)
43 if (!a || !b || !y || n == 0) {
49#if defined(__AVX512F__)
51 for (; i + 16 <= n; i += 16) {
52 __m512 av = bf16_loadu_cvt_fp32(&a[i]);
53 __m512 bv = bf16_loadu_cvt_fp32(&b[i]);
54 __m512 yv = _mm512_add_ps(av, bv);
55 fp32_cvt_storeu_bf16(&y[i], yv);
71 int aligned_embed_dim)
73 const size_t count = (size_t)tokens * (
size_t)aligned_embed_dim;
74 for (
size_t i = 0; i < count; ++i) {
92 if (!a || !b || !y || n == 0) {
98#if defined(__AVX512F__)
99 __m512 alpha_v = _mm512_set1_ps(alpha);
100 for (; i + 16 <= n; i += 16) {
101 __m512 av = bf16_loadu_cvt_fp32(&a[i]);
102 __m512 bv = bf16_loadu_cvt_fp32(&b[i]);
103 __m512 yv = _mm512_fmadd_ps(bv, alpha_v, av);
104 fp32_cvt_storeu_bf16(&y[i], yv);
123 if (!a || !b || n == 0) {
129#if defined(__AVX512F__)
130 for (; i + 16 <= n; i += 16) {
131 __m512 av = bf16_loadu_cvt_fp32(&a[i]);
132 __m512 bv = bf16_loadu_cvt_fp32(&b[i]);
133 __m512 yv = _mm512_add_ps(av, bv);
134 fp32_cvt_storeu_bf16(&a[i], yv);
154 if (!a || !b || n == 0) {
160#if defined(__AVX512F__)
161 __m512 alpha_v = _mm512_set1_ps(alpha);
162 for (; i + 16 <= n; i += 16) {
163 __m512 av = bf16_loadu_cvt_fp32(&a[i]);
164 __m512 bv = bf16_loadu_cvt_fp32(&b[i]);
165 __m512 yv = _mm512_fmadd_ps(bv, alpha_v, av);
166 fp32_cvt_storeu_bf16(&a[i], yv);
192 if (!d_y || n == 0) {
199 if (d_a && d_a != d_y) {
200#if defined(__AVX512F__)
201 for (; i + 32 <= n; i += 32) {
202 __m512i v0 = _mm512_loadu_si512((
const __m512i*)&d_y[i]);
203 __m512i v1 = _mm512_loadu_si512((
const __m512i*)&d_y[i + 32]);
204 _mm512_storeu_si512((__m512i*)&d_a[i], v0);
205 _mm512_storeu_si512((__m512i*)&d_a[i + 32], v1);
215 if (d_b && d_b != d_y) {
216#if defined(__AVX512F__)
217 for (; i + 32 <= n; i += 32) {
218 __m512i v0 = _mm512_loadu_si512((
const __m512i*)&d_y[i]);
219 __m512i v1 = _mm512_loadu_si512((
const __m512i*)&d_y[i + 32]);
220 _mm512_storeu_si512((__m512i*)&d_b[i], v0);
221 _mm512_storeu_si512((__m512i*)&d_b[i + 32], v1);
242 if (!a || !b || !y || tokens <= 0 || dim <= 0) {
246 for (
int t = 0; t < tokens; ++t) {
247 const uint16_t *a_row = a + (size_t)t * aligned_dim;
248 const uint16_t *b_row = b + (size_t)t * aligned_dim;
249 uint16_t *y_row = y + (size_t)t * aligned_dim;
253#if defined(__AVX512F__)
254 for (; d + 16 <= dim; d += 16) {
255 __m512 av = bf16_loadu_cvt_fp32(&a_row[d]);
256 __m512 bv = bf16_loadu_cvt_fp32(&b_row[d]);
257 __m512 yv = _mm512_add_ps(av, bv);
258 fp32_cvt_storeu_bf16(&y_row[d], yv);
262 for (; d < dim; ++d) {
289 if (!a || !b || !y || n == 0) {
295#if defined(__AVX512F__)
296 for (; i + 16 <= n; i += 16) {
297 __m512 av = _mm512_loadu_ps(&a[i]);
298 __m512 bv = _mm512_loadu_ps(&b[i]);
299 __m512 yv = _mm512_add_ps(av, bv);
300 _mm512_storeu_ps(&y[i], yv);
305 for (; i + 8 <= n; i += 8) {
306 __m256 av = _mm256_loadu_ps(&a[i]);
307 __m256 bv = _mm256_loadu_ps(&b[i]);
308 __m256 yv = _mm256_add_ps(av, bv);
309 _mm256_storeu_ps(&y[i], yv);
322 if (!a || !b || n == 0) {
328#if defined(__AVX512F__)
329 for (; i + 16 <= n; i += 16) {
330 __m512 av = _mm512_loadu_ps(&a[i]);
331 __m512 bv = _mm512_loadu_ps(&b[i]);
332 __m512 yv = _mm512_add_ps(av, bv);
333 _mm512_storeu_ps(&a[i], yv);
338 for (; i + 8 <= n; i += 8) {
339 __m256 av = _mm256_loadu_ps(&a[i]);
340 __m256 bv = _mm256_loadu_ps(&b[i]);
341 __m256 yv = _mm256_add_ps(av, bv);
342 _mm256_storeu_ps(&a[i], yv);
void ck_residual_add_token_major_bf16_storage(const float *a, const float *b, float *out, int tokens, int aligned_embed_dim)
void add_scaled_forward_bf16(const uint16_t *a, const uint16_t *b, uint16_t *y, float alpha, size_t n)
void add_inplace_bf16(uint16_t *a, const uint16_t *b, size_t n)
void add_forward_f32(const float *a, const float *b, float *y, size_t n)
void add_inplace_f32(float *a, const float *b, size_t n)
void add_forward_bf16(const uint16_t *a, const uint16_t *b, uint16_t *y, size_t n)
void add_forward_2d_bf16(const uint16_t *a, const uint16_t *b, uint16_t *y, int tokens, int dim, int aligned_dim)
void add_backward_bf16(const uint16_t *d_y, uint16_t *d_a, uint16_t *d_b, size_t n)
void add_scaled_inplace_bf16(uint16_t *a, const uint16_t *b, float alpha, size_t n)
static uint16_t float_to_bf16(float f)
static float bf16_to_float(uint16_t v)