| 237 | } |
| 238 | |
| 239 | inline __m256i cvt_fp32_to_bf16(const __m512 src) { |
| 240 | #if defined(__AVX512BF16__) |
| 241 | if (has_avx512bf16()) { |
| 242 | return reinterpret_cast<__m256i>(_mm512_cvtneps_pbh(src)); |
| 243 | } |
| 244 | #endif |
| 245 | __m512i value = _mm512_castps_si512(src); |
| 246 | __m512i nan = _mm512_set1_epi32(0xffff); |
| 247 | auto mask_value = _mm512_cmp_ps_mask(src, src, _CMP_ORD_Q); |
| 248 | __m512i ones = _mm512_set1_epi32(0x1); |
| 249 | __m512i vec_bias = _mm512_set1_epi32(0x7fff); |
| 250 | // uint32_t lsb = (input >> 16) & 1; |
| 251 | auto t_value = _mm512_and_si512(_mm512_srli_epi32(value, 16), ones); |
| 252 | // uint32_t rounding_bias = 0x7fff + lsb; |
| 253 | t_value = _mm512_add_epi32(t_value, vec_bias); |
| 254 | // input += rounding_bias; |
| 255 | t_value = _mm512_add_epi32(t_value, value); |
| 256 | // input = input >> 16; |
| 257 | t_value = _mm512_srli_epi32(t_value, 16); |
| 258 | // Check NaN before converting back to bf16 |
| 259 | t_value = _mm512_mask_blend_epi32(mask_value, nan, t_value); |
| 260 | return _mm512_cvtusepi32_epi16(t_value); |
| 261 | } |
| 262 | |
| 263 | static inline __m512 set_nf4_lut() { |
| 264 | return _mm512_set_ps( |
no test coverage detected