Eigen  5.0.1
 
Loading...
Searching...
No Matches
BFloat16.h
1/* Copyright 2017 The TensorFlow Authors. All Rights Reserved.
2
3Licensed under the Apache License, Version 2.0 (the "License");
4you may not use this file except in compliance with the License.
5You may obtain a copy of the License at
6
7 http://www.apache.org/licenses/LICENSE-2.0
8
9Unless required by applicable law or agreed to in writing, software
10distributed under the License is distributed on an "AS IS" BASIS,
11WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12See the License for the specific language governing permissions and
13limitations under the License.
14==============================================================================*/
15// SPDX-License-Identifier: Apache-2.0
16
17#ifndef EIGEN_BFLOAT16_H
18#define EIGEN_BFLOAT16_H
19
20// IWYU pragma: private
21#include "../../InternalHeaderCheck.h"
22
23#if defined(EIGEN_HAS_HIP_BF16)
24// When compiling with GPU support, the "hip_bfloat16" base class as well as
25// some other routines are defined in the GPU compiler header files
26// (hip_bfloat16.h), and they are not tagged constexpr
27// As a consequence, we get compile failures when compiling Eigen with
28// GPU support. Hence the need to disable EIGEN_CONSTEXPR when building
29// Eigen with GPU support
30#pragma push_macro("EIGEN_CONSTEXPR")
31#undef EIGEN_CONSTEXPR
32#define EIGEN_CONSTEXPR
33#endif
34
35#define BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, METHOD) \
36 template <> \
37 EIGEN_DEFINE_FUNCTION_ALLOWING_MULTIPLE_DEFINITIONS EIGEN_UNUSED PACKET_BF16 METHOD<PACKET_BF16>( \
38 const PACKET_BF16& _x) { \
39 return F32ToBf16(METHOD<PACKET_F>(Bf16ToF32(_x))); \
40 }
41
42#define EIGEN_INSTANTIATE_GENERIC_MATH_FUNCS_BF16(PACKET_F, PACKET_BF16) \
43 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pcos) \
44 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, psin) \
45 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, psinh) \
46 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pcosh) \
47 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pasinh) \
48 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pacosh) \
49 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pexp) \
50 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pexp2) \
51 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pexpm1) \
52 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, plog) \
53 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, plog1p) \
54 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, plog2) \
55 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, plog10) \
56 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, preciprocal) \
57 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, prsqrt) \
58 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pcbrt) \
59 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, psqrt) \
60 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, ptanh)
61
62// BF16 wrappers for contrib/SpecialFunctions.
63#define EIGEN_INSTANTIATE_SPECIAL_FUNCS_BF16(PACKET_F, PACKET_BF16) \
64 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, perf) \
65 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pndtri)
66
67#define EIGEN_INSTANTIATE_BESSEL_FUNCS_BF16(PACKET_F, PACKET_BF16) \
68 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_i0) \
69 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_i0e) \
70 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_i1) \
71 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_i1e) \
72 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_j0) \
73 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_j1) \
74 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_k0) \
75 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_k0e) \
76 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_k1) \
77 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_k1e) \
78 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_y0) \
79 BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, pbessel_y1)
80
81// Only use HIP GPU bf16 in kernels
82#if defined(EIGEN_HAS_HIP_BF16) && defined(EIGEN_GPU_COMPILE_PHASE)
83#define EIGEN_USE_HIP_BF16
84#endif
85
86namespace Eigen {
87
88struct bfloat16;
89
90namespace numext {
91template <>
92EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Eigen::bfloat16 bit_cast<Eigen::bfloat16, uint16_t>(const uint16_t& src);
93
94template <>
95EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC uint16_t bit_cast<uint16_t, Eigen::bfloat16>(const Eigen::bfloat16& src);
96} // namespace numext
97namespace bfloat16_impl {
98
99#if defined(EIGEN_USE_HIP_BF16)
100
101struct __bfloat16_raw : public hip_bfloat16 {
102 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw() {}
103 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw(hip_bfloat16 hb) : hip_bfloat16(hb) {}
104 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw(unsigned short raw) : hip_bfloat16(raw) {}
105};
106
107#else
108
109// Make our own __bfloat16_raw definition.
110struct __bfloat16_raw {
111#if defined(EIGEN_HAS_HIP_BF16) && !defined(EIGEN_GPU_COMPILE_PHASE)
112 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw() {}
113#else
114 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw() : value(0) {}
115#endif
116 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw(unsigned short raw) : value(raw) {}
117 unsigned short value;
118};
119
120#endif // defined(EIGEN_USE_HIP_BF16)
121
122EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw raw_uint16_to_bfloat16(unsigned short value);
123template <bool AssumeArgumentIsNormalOrInfinityOrZero>
124EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw float_to_bfloat16_rtne(float ff);
125// Forward declarations of template specializations, to avoid Visual C++ 2019 errors, saying:
126// > error C2908: explicit specialization; 'float_to_bfloat16_rtne' has already been instantiated
127template <>
128EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw float_to_bfloat16_rtne<false>(float ff);
129template <>
130EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw float_to_bfloat16_rtne<true>(float ff);
131EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC float bfloat16_to_float(__bfloat16_raw h);
132
133struct bfloat16_base : public __bfloat16_raw {
134 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16_base() {}
135 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16_base(const __bfloat16_raw& h) : __bfloat16_raw(h) {}
136};
137
138} // namespace bfloat16_impl
139
140// Class definition.
141struct bfloat16 : public bfloat16_impl::bfloat16_base {
142 using __bfloat16_raw = bfloat16_impl::__bfloat16_raw;
143
144 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16() {}
145
146 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16(const __bfloat16_raw& h) : bfloat16_impl::bfloat16_base(h) {}
147
148 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16(bool b)
149 : bfloat16_impl::bfloat16_base(bfloat16_impl::raw_uint16_to_bfloat16(b ? 0x3f80 : 0)) {}
150
151 template <class T>
152 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16(T val)
153 : bfloat16_impl::bfloat16_base(
154 bfloat16_impl::float_to_bfloat16_rtne<std::is_integral<T>::value>(static_cast<float>(val))) {}
155
156 explicit EIGEN_DEVICE_FUNC bfloat16(float f)
157 : bfloat16_impl::bfloat16_base(bfloat16_impl::float_to_bfloat16_rtne<false>(f)) {}
158
159 // Following the convention of numpy, converting between complex and
160 // float will lead to loss of imag value.
161 template <typename RealScalar>
162 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16(const std::complex<RealScalar>& val)
163 : bfloat16_impl::bfloat16_base(bfloat16_impl::float_to_bfloat16_rtne<false>(static_cast<float>(val.real()))) {}
164
165 EIGEN_DEVICE_FUNC operator float() const { // NOLINT: Allow implicit conversion to float, because it is lossless.
166 return bfloat16_impl::bfloat16_to_float(*this);
167 }
168};
169
170// TODO(majnemer): Get rid of this once we can rely on C++17 inline variables to
171// solve the ODR issue.
172namespace bfloat16_impl {
173template <typename = void>
174struct numeric_limits_bfloat16_impl {
175 static EIGEN_CONSTEXPR const bool is_specialized = true;
176 static EIGEN_CONSTEXPR const bool is_signed = true;
177 static EIGEN_CONSTEXPR const bool is_integer = false;
178 static EIGEN_CONSTEXPR const bool is_exact = false;
179 static EIGEN_CONSTEXPR const bool has_infinity = true;
180 static EIGEN_CONSTEXPR const bool has_quiet_NaN = true;
181 static EIGEN_CONSTEXPR const bool has_signaling_NaN = true;
182 EIGEN_DIAGNOSTICS(push)
183 EIGEN_DISABLE_DEPRECATED_WARNING
184 static EIGEN_CONSTEXPR const std::float_denorm_style has_denorm = std::denorm_present;
185 static EIGEN_CONSTEXPR const bool has_denorm_loss = false;
186 EIGEN_DIAGNOSTICS(pop)
187 static EIGEN_CONSTEXPR const std::float_round_style round_style = std::numeric_limits<float>::round_style;
188 static EIGEN_CONSTEXPR const bool is_iec559 = true;
189 // The C++ standard defines this as "true if the set of values representable
190 // by the type is finite." BFloat16 has finite precision.
191 static EIGEN_CONSTEXPR const bool is_bounded = true;
192 static EIGEN_CONSTEXPR const bool is_modulo = false;
193 static EIGEN_CONSTEXPR const int digits = 8;
194 static EIGEN_CONSTEXPR const int digits10 = 2;
195 static EIGEN_CONSTEXPR const int max_digits10 = 4;
196 static EIGEN_CONSTEXPR const int radix = std::numeric_limits<float>::radix;
197 static EIGEN_CONSTEXPR const int min_exponent = std::numeric_limits<float>::min_exponent;
198 static EIGEN_CONSTEXPR const int min_exponent10 = std::numeric_limits<float>::min_exponent10;
199 static EIGEN_CONSTEXPR const int max_exponent = std::numeric_limits<float>::max_exponent;
200 static EIGEN_CONSTEXPR const int max_exponent10 = std::numeric_limits<float>::max_exponent10;
201 static EIGEN_CONSTEXPR const bool traps = std::numeric_limits<float>::traps;
202 // IEEE754: "The implementer shall choose how tininess is detected, but shall
203 // detect tininess in the same way for all operations in radix two"
204 static EIGEN_CONSTEXPR const bool tinyness_before = std::numeric_limits<float>::tinyness_before;
205
206 static EIGEN_CONSTEXPR Eigen::bfloat16(min)() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x0080); }
207 static EIGEN_CONSTEXPR Eigen::bfloat16 lowest() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0xff7f); }
208 static EIGEN_CONSTEXPR Eigen::bfloat16(max)() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x7f7f); }
209 static EIGEN_CONSTEXPR Eigen::bfloat16 epsilon() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x3c00); }
210 static EIGEN_CONSTEXPR Eigen::bfloat16 round_error() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x3f00); }
211 static EIGEN_CONSTEXPR Eigen::bfloat16 infinity() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x7f80); }
212 static EIGEN_CONSTEXPR Eigen::bfloat16 quiet_NaN() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x7fc0); }
213 static EIGEN_CONSTEXPR Eigen::bfloat16 signaling_NaN() {
214 return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x7fa0);
215 }
216 static EIGEN_CONSTEXPR Eigen::bfloat16 denorm_min() { return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x0001); }
217};
218
219// Redundant out-of-class definitions are required pre-C++17 but deprecated since.
220#if EIGEN_COMP_CXXVER < 17
221template <typename T>
222EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_specialized;
223template <typename T>
224EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_signed;
225template <typename T>
226EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_integer;
227template <typename T>
228EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_exact;
229template <typename T>
230EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::has_infinity;
231template <typename T>
232EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::has_quiet_NaN;
233template <typename T>
234EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::has_signaling_NaN;
235EIGEN_DIAGNOSTICS(push)
236EIGEN_DISABLE_DEPRECATED_WARNING
237template <typename T>
238EIGEN_CONSTEXPR const std::float_denorm_style numeric_limits_bfloat16_impl<T>::has_denorm;
239template <typename T>
240EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::has_denorm_loss;
241EIGEN_DIAGNOSTICS(pop)
242template <typename T>
243EIGEN_CONSTEXPR const std::float_round_style numeric_limits_bfloat16_impl<T>::round_style;
244template <typename T>
245EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_iec559;
246template <typename T>
247EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_bounded;
248template <typename T>
249EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::is_modulo;
250template <typename T>
251EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::digits;
252template <typename T>
253EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::digits10;
254template <typename T>
255EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::max_digits10;
256template <typename T>
257EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::radix;
258template <typename T>
259EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::min_exponent;
260template <typename T>
261EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::min_exponent10;
262template <typename T>
263EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::max_exponent;
264template <typename T>
265EIGEN_CONSTEXPR const int numeric_limits_bfloat16_impl<T>::max_exponent10;
266template <typename T>
267EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::traps;
268template <typename T>
269EIGEN_CONSTEXPR const bool numeric_limits_bfloat16_impl<T>::tinyness_before;
270#endif
271} // end namespace bfloat16_impl
272} // end namespace Eigen
273
274namespace std {
275// If std::numeric_limits<T> is specialized, should also specialize
276// std::numeric_limits<const T>, std::numeric_limits<volatile T>, and
277// std::numeric_limits<const volatile T>
278// https://stackoverflow.com/a/16519653/
279template <>
280class numeric_limits<Eigen::bfloat16> : public Eigen::bfloat16_impl::numeric_limits_bfloat16_impl<> {};
281template <>
282class numeric_limits<const Eigen::bfloat16> : public numeric_limits<Eigen::bfloat16> {};
283template <>
284class numeric_limits<volatile Eigen::bfloat16> : public numeric_limits<Eigen::bfloat16> {};
285template <>
286class numeric_limits<const volatile Eigen::bfloat16> : public numeric_limits<Eigen::bfloat16> {};
287} // end namespace std
288
289namespace Eigen {
290
291namespace bfloat16_impl {
292
293// We need to distinguish ‘clang as the CUDA compiler’ from ‘clang as the host compiler,
294// invoked by NVCC’ (e.g. on MacOS). The former needs to see both host and device implementation
295// of the functions, while the latter can only deal with one of them.
296#if !defined(EIGEN_HAS_NATIVE_BF16) || (EIGEN_COMP_CLANG && !EIGEN_COMP_NVCC) // Emulate support for bfloat16 floats
297
298#if EIGEN_COMP_CLANG && defined(EIGEN_GPUCC)
299// We need to provide emulated *host-side* BF16 operators for clang.
300#pragma push_macro("EIGEN_DEVICE_FUNC")
301#undef EIGEN_DEVICE_FUNC
302#if (defined(EIGEN_HAS_GPU_BF16) && defined(EIGEN_HAS_NATIVE_BF16))
303#define EIGEN_DEVICE_FUNC __host__
304#else // both host and device need emulated ops.
305#define EIGEN_DEVICE_FUNC __host__ __device__
306#endif
307#endif
308
309// Definitions for CPUs, mostly working through conversion
310// to/from fp32.
311
312EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator+(const bfloat16& a, const bfloat16& b) {
313 return bfloat16(float(a) + float(b));
314}
315EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator+(const bfloat16& a, const int& b) {
316 return bfloat16(float(a) + static_cast<float>(b));
317}
318EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator+(const int& a, const bfloat16& b) {
319 return bfloat16(static_cast<float>(a) + float(b));
320}
321EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator*(const bfloat16& a, const bfloat16& b) {
322 return bfloat16(float(a) * float(b));
323}
324EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator-(const bfloat16& a, const bfloat16& b) {
325 return bfloat16(float(a) - float(b));
326}
327EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator/(const bfloat16& a, const bfloat16& b) {
328 return bfloat16(float(a) / float(b));
329}
330EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator-(const bfloat16& a) {
331 numext::uint16_t x = numext::bit_cast<uint16_t>(a) ^ 0x8000;
332 return numext::bit_cast<bfloat16>(x);
333}
334EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC int16_t bfloat16_map_to_signed(numext::uint16_t bits) {
335 constexpr numext::uint16_t kAbsMask = 0x7fff;
336 return (bits >> 15) ? -(bits & kAbsMask) : bits;
337}
338EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool bfloat16_is_ordered(numext::uint16_t a, numext::uint16_t b) {
339 constexpr numext::uint16_t kAbsMask = 0x7fff;
340 constexpr numext::uint16_t kInf = 0x7f80;
341 return numext::maxi(a & kAbsMask, b & kAbsMask) <= kInf;
342}
343EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator+=(bfloat16& a, const bfloat16& b) {
344 a = bfloat16(float(a) + float(b));
345 return a;
346}
347EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator*=(bfloat16& a, const bfloat16& b) {
348 a = bfloat16(float(a) * float(b));
349 return a;
350}
351EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator-=(bfloat16& a, const bfloat16& b) {
352 a = bfloat16(float(a) - float(b));
353 return a;
354}
355EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator/=(bfloat16& a, const bfloat16& b) {
356 a = bfloat16(float(a) / float(b));
357 return a;
358}
359EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator++(bfloat16& a) {
360 a += bfloat16(1);
361 return a;
362}
363EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator--(bfloat16& a) {
364 a -= bfloat16(1);
365 return a;
366}
367EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator++(bfloat16& a, int) {
368 bfloat16 original_value = a;
369 ++a;
370 return original_value;
371}
372EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator--(bfloat16& a, int) {
373 bfloat16 original_value = a;
374 --a;
375 return original_value;
376}
377// Evaluate both predicates to keep comparison loops branch-free. Integer operands avoid Clang's
378// -Wbitwise-instead-of-logical warning.
379EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool operator==(const bfloat16& a, const bfloat16& b) {
380 const numext::uint16_t a_bits = numext::bit_cast<numext::uint16_t>(a);
381 const numext::uint16_t b_bits = numext::bit_cast<numext::uint16_t>(b);
382 return static_cast<unsigned int>(bfloat16_map_to_signed(a_bits) == bfloat16_map_to_signed(b_bits)) &
383 static_cast<unsigned int>(bfloat16_is_ordered(a_bits, b_bits));
384}
385EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool operator!=(const bfloat16& a, const bfloat16& b) { return !(a == b); }
386EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool operator<(const bfloat16& a, const bfloat16& b) {
387 const numext::uint16_t a_bits = numext::bit_cast<numext::uint16_t>(a);
388 const numext::uint16_t b_bits = numext::bit_cast<numext::uint16_t>(b);
389 return static_cast<unsigned int>(bfloat16_map_to_signed(a_bits) < bfloat16_map_to_signed(b_bits)) &
390 static_cast<unsigned int>(bfloat16_is_ordered(a_bits, b_bits));
391}
392EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool operator<=(const bfloat16& a, const bfloat16& b) {
393 const numext::uint16_t a_bits = numext::bit_cast<numext::uint16_t>(a);
394 const numext::uint16_t b_bits = numext::bit_cast<numext::uint16_t>(b);
395 return static_cast<unsigned int>(bfloat16_map_to_signed(a_bits) <= bfloat16_map_to_signed(b_bits)) &
396 static_cast<unsigned int>(bfloat16_is_ordered(a_bits, b_bits));
397}
398EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool operator>(const bfloat16& a, const bfloat16& b) {
399 const numext::uint16_t a_bits = numext::bit_cast<numext::uint16_t>(a);
400 const numext::uint16_t b_bits = numext::bit_cast<numext::uint16_t>(b);
401 return static_cast<unsigned int>(bfloat16_map_to_signed(a_bits) > bfloat16_map_to_signed(b_bits)) &
402 static_cast<unsigned int>(bfloat16_is_ordered(a_bits, b_bits));
403}
404EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool operator>=(const bfloat16& a, const bfloat16& b) {
405 const numext::uint16_t a_bits = numext::bit_cast<numext::uint16_t>(a);
406 const numext::uint16_t b_bits = numext::bit_cast<numext::uint16_t>(b);
407 return static_cast<unsigned int>(bfloat16_map_to_signed(a_bits) >= bfloat16_map_to_signed(b_bits)) &
408 static_cast<unsigned int>(bfloat16_is_ordered(a_bits, b_bits));
409}
410
411#if EIGEN_COMP_CLANG && defined(EIGEN_CUDACC)
412#pragma pop_macro("EIGEN_DEVICE_FUNC")
413#endif
414#endif // Emulate support for bfloat16 floats
415
416// Division by an index. Do it in full float precision to avoid accuracy
417// issues in converting the denominator to bfloat16.
418EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator/(const bfloat16& a, Index b) {
419 return bfloat16(static_cast<float>(a) / static_cast<float>(b));
420}
421
422EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw truncate_to_bfloat16(const float v) {
423#if defined(EIGEN_USE_HIP_BF16)
424 return __bfloat16_raw(__bfloat16_raw::round_to_bfloat16(v, __bfloat16_raw::truncate));
425#else
426 __bfloat16_raw output;
427 if (numext::isnan EIGEN_NOT_A_MACRO(v)) {
428 output.value = std::signbit(v) ? 0xFFC0 : 0x7FC0;
429 return output;
430 }
431 output.value = static_cast<numext::uint16_t>(numext::bit_cast<numext::uint32_t>(v) >> 16);
432 return output;
433#endif
434}
435
436EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw raw_uint16_to_bfloat16(numext::uint16_t value) {
437#if defined(EIGEN_USE_HIP_BF16)
438 __bfloat16_raw bf;
439 bf.data = value;
440 return bf;
441#else
442 return __bfloat16_raw(value);
443#endif
444}
445
446EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR numext::uint16_t raw_bfloat16_as_uint16(
447 const __bfloat16_raw& bf) {
448#if defined(EIGEN_USE_HIP_BF16)
449 return bf.data;
450#else
451 return bf.value;
452#endif
453}
454
455// float_to_bfloat16_rtne template specialization that does not make any
456// assumption about the value of its function argument (ff).
457template <>
458EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw float_to_bfloat16_rtne<false>(float ff) {
459#if defined(EIGEN_USE_HIP_BF16)
460 return __bfloat16_raw(__bfloat16_raw::round_to_bfloat16(ff));
461#else
462 __bfloat16_raw output;
463
464 if (numext::isnan EIGEN_NOT_A_MACRO(ff)) {
465 // If the value is a NaN, squash it to a qNaN with msb of fraction set,
466 // this makes sure after truncation we don't end up with an inf.
467 //
468 // qNaN magic: All exponent bits set + most significant bit of fraction
469 // set.
470 output.value = std::signbit(ff) ? 0xFFC0 : 0x7FC0;
471 } else {
472 // Fast rounding algorithm that rounds a half value to nearest even. This
473 // reduces expected error when we convert a large number of floats. Here
474 // is how it works:
475 //
476 // Definitions:
477 // To convert a float 32 to bfloat16, a float 32 can be viewed as 32 bits
478 // with the following tags:
479 //
480 // Sign | Exp (8 bits) | Frac (23 bits)
481 // S EEEEEEEE FFFFFFLRTTTTTTTTTTTTTTT
482 //
483 // S: Sign bit.
484 // E: Exponent bits.
485 // F: First 6 bits of fraction.
486 // L: Least significant bit of resulting bfloat16 if we truncate away the
487 // rest of the float32. This is also the 7th bit of fraction
488 // R: Rounding bit, 8th bit of fraction.
489 // T: Sticky bits, rest of fraction, 15 bits.
490 //
491 // To round half to nearest even, there are 3 cases where we want to round
492 // down (simply truncate the result of the bits away, which consists of
493 // rounding bit and sticky bits) and two cases where we want to round up
494 // (truncate then add one to the result).
495 //
496 // The fast converting algorithm simply adds lsb (L) to 0x7fff (15 bits of
497 // 1s) as the rounding bias, adds the rounding bias to the input, then
498 // truncates the last 16 bits away.
499 //
500 // To understand how it works, we can analyze this algorithm case by case:
501 //
502 // 1. L = 0, R = 0:
503 // Expect: round down, this is less than half value.
504 //
505 // Algorithm:
506 // - Rounding bias: 0x7fff + 0 = 0x7fff
507 // - Adding rounding bias to input may create any carry, depending on
508 // whether there is any value set to 1 in T bits.
509 // - R may be set to 1 if there is a carry.
510 // - L remains 0.
511 // - Note that this case also handles Inf and -Inf, where all fraction
512 // bits, including L, R and Ts are all 0. The output remains Inf after
513 // this algorithm.
514 //
515 // 2. L = 1, R = 0:
516 // Expect: round down, this is less than half value.
517 //
518 // Algorithm:
519 // - Rounding bias: 0x7fff + 1 = 0x8000
520 // - Adding rounding bias to input doesn't change sticky bits but
521 // adds 1 to rounding bit.
522 // - L remains 1.
523 //
524 // 3. L = 0, R = 1, all of T are 0:
525 // Expect: round down, this is exactly at half, the result is already
526 // even (L=0).
527 //
528 // Algorithm:
529 // - Rounding bias: 0x7fff + 0 = 0x7fff
530 // - Adding rounding bias to input sets all sticky bits to 1, but
531 // doesn't create a carry.
532 // - R remains 1.
533 // - L remains 0.
534 //
535 // 4. L = 1, R = 1:
536 // Expect: round up, this is exactly at half, the result needs to be
537 // round to the next even number.
538 //
539 // Algorithm:
540 // - Rounding bias: 0x7fff + 1 = 0x8000
541 // - Adding rounding bias to input doesn't change sticky bits, but
542 // creates a carry from rounding bit.
543 // - The carry sets L to 0, creates another carry bit and propagate
544 // forward to F bits.
545 // - If all the F bits are 1, a carry then propagates to the exponent
546 // bits, which then creates the minimum value with the next exponent
547 // value. Note that we won't have the case where exponents are all 1,
548 // since that's either a NaN (handled in the other if condition) or inf
549 // (handled in case 1).
550 //
551 // 5. L = 0, R = 1, any of T is 1:
552 // Expect: round up, this is greater than half.
553 //
554 // Algorithm:
555 // - Rounding bias: 0x7fff + 0 = 0x7fff
556 // - Adding rounding bias to input creates a carry from sticky bits,
557 // sets rounding bit to 0, then create another carry.
558 // - The second carry sets L to 1.
559 //
560 // Examples:
561 //
562 // Exact half value that is already even:
563 // Input:
564 // Sign | Exp (8 bit) | Frac (first 7 bit) | Frac (last 16 bit)
565 // S E E E E E E E E F F F F F F L RTTTTTTTTTTTTTTT
566 // 0 0 0 0 0 0 0 0 0 0 0 0 0 0 1 0 1000000000000000
567 //
568 // This falls into case 3. We truncate the rest of 16 bits and no
569 // carry is created into F and L:
570 //
571 // Output:
572 // Sign | Exp (8 bit) | Frac (first 7 bit)
573 // S E E E E E E E E F F F F F F L
574 // 0 0 0 0 0 0 0 0 0 0 0 0 0 0 1 0
575 //
576 // Exact half value, round to next even number:
577 // Input:
578 // Sign | Exp (8 bit) | Frac (first 7 bit) | Frac (last 16 bit)
579 // S E E E E E E E E F F F F F F L RTTTTTTTTTTTTTTT
580 // 0 0 0 0 0 0 0 0 0 0 0 0 0 0 0 1 1000000000000000
581 //
582 // This falls into case 4. We create a carry from R and T,
583 // which then propagates into L and F:
584 //
585 // Output:
586 // Sign | Exp (8 bit) | Frac (first 7 bit)
587 // S E E E E E E E E F F F F F F L
588 // 0 0 0 0 0 0 0 0 0 0 0 0 0 0 1 0
589 //
590 //
591 // Max denormal value round to min normal value:
592 // Input:
593 // Sign | Exp (8 bit) | Frac (first 7 bit) | Frac (last 16 bit)
594 // S E E E E E E E E F F F F F F L RTTTTTTTTTTTTTTT
595 // 0 0 0 0 0 0 0 0 0 1 1 1 1 1 1 1 1111111111111111
596 //
597 // This falls into case 4. We create a carry from R and T,
598 // propagate into L and F, which then propagates into exponent
599 // bits:
600 //
601 // Output:
602 // Sign | Exp (8 bit) | Frac (first 7 bit)
603 // S E E E E E E E E F F F F F F L
604 // 0 0 0 0 0 0 0 0 1 0 0 0 0 0 0 0
605 //
606 // Max normal value round to Inf:
607 // Input:
608 // Sign | Exp (8 bit) | Frac (first 7 bit) | Frac (last 16 bit)
609 // S E E E E E E E E F F F F F F L RTTTTTTTTTTTTTTT
610 // 0 1 1 1 1 1 1 1 0 1 1 1 1 1 1 1 1111111111111111
611 //
612 // This falls into case 4. We create a carry from R and T,
613 // propagate into L and F, which then propagates into exponent
614 // bits:
615 //
616 // Sign | Exp (8 bit) | Frac (first 7 bit)
617 // S E E E E E E E E F F F F F F L
618 // 0 1 1 1 1 1 1 1 1 0 0 0 0 0 0 0
619
620 // At this point, ff must be either a normal float, or +/-infinity.
621 output = float_to_bfloat16_rtne<true>(ff);
622 }
623 return output;
624#endif
625}
626
627// float_to_bfloat16_rtne template specialization that assumes that its function
628// argument (ff) is either a normal floating point number, or +/-infinity, or
629// zero. Used to improve the runtime performance of conversion from an integer
630// type to bfloat16.
631template <>
632EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw float_to_bfloat16_rtne<true>(float ff) {
633#if defined(EIGEN_USE_HIP_BF16)
634 return __bfloat16_raw(__bfloat16_raw::round_to_bfloat16(ff));
635#else
636 numext::uint32_t input = numext::bit_cast<numext::uint32_t>(ff);
637 __bfloat16_raw output;
638
639 // Least significant bit of resulting bfloat.
640 numext::uint32_t lsb = (input >> 16) & 1;
641 numext::uint32_t rounding_bias = 0x7fff + lsb;
642 input += rounding_bias;
643 output.value = static_cast<numext::uint16_t>(input >> 16);
644 return output;
645#endif
646}
647
648EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC float bfloat16_to_float(__bfloat16_raw h) {
649#if defined(EIGEN_USE_HIP_BF16)
650 return static_cast<float>(h);
651#else
652 return numext::bit_cast<float>(static_cast<numext::uint32_t>(h.value) << 16);
653#endif
654}
655
656// --- standard functions ---
657
658EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool(isinf)(const bfloat16& a) {
659 return (raw_bfloat16_as_uint16(a) & 0x7fff) == 0x7f80;
660}
661EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool(isnan)(const bfloat16& a) {
662 return (raw_bfloat16_as_uint16(a) & 0x7fff) > 0x7f80;
663}
664EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool(isfinite)(const bfloat16& a) {
665 return (raw_bfloat16_as_uint16(a) & 0x7fff) < 0x7f80;
666}
667
668EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 abs(const bfloat16& a) {
669 numext::uint16_t x = numext::bit_cast<numext::uint16_t>(a) & 0x7FFF;
670 return numext::bit_cast<bfloat16>(x);
671}
672EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 exp(const bfloat16& a) { return bfloat16(::expf(float(a))); }
673EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 exp2(const bfloat16& a) { return bfloat16(::exp2f(float(a))); }
674EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 expm1(const bfloat16& a) { return bfloat16(numext::expm1(float(a))); }
675// float covers the bfloat16 range, so the single conversion back rounds correctly.
676EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 ldexp(const bfloat16& a, int exponent) {
677 return bfloat16(numext::ldexp(float(a), exponent));
678}
679EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 log(const bfloat16& a) { return bfloat16(::logf(float(a))); }
680EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 log1p(const bfloat16& a) { return bfloat16(numext::log1p(float(a))); }
681EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 log10(const bfloat16& a) { return bfloat16(::log10f(float(a))); }
682EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 log2(const bfloat16& a) {
683 return bfloat16(static_cast<float>(EIGEN_LOG2E) * ::logf(float(a)));
684}
685EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 sqrt(const bfloat16& a) { return bfloat16(::sqrtf(float(a))); }
686EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 cbrt(const bfloat16& a) { return bfloat16(::cbrtf(float(a))); }
687EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 pow(const bfloat16& a, const bfloat16& b) {
688 return bfloat16(::powf(float(a), float(b)));
689}
690EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 atan2(const bfloat16& a, const bfloat16& b) {
691 return bfloat16(::atan2f(float(a), float(b)));
692}
693EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 sin(const bfloat16& a) { return bfloat16(::sinf(float(a))); }
694EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 cos(const bfloat16& a) { return bfloat16(::cosf(float(a))); }
695EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 tan(const bfloat16& a) { return bfloat16(::tanf(float(a))); }
696EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 asin(const bfloat16& a) { return bfloat16(::asinf(float(a))); }
697EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 acos(const bfloat16& a) { return bfloat16(::acosf(float(a))); }
698EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 atan(const bfloat16& a) { return bfloat16(::atanf(float(a))); }
699EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 sinh(const bfloat16& a) { return bfloat16(::sinhf(float(a))); }
700EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 cosh(const bfloat16& a) { return bfloat16(::coshf(float(a))); }
701EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 tanh(const bfloat16& a) { return bfloat16(::tanhf(float(a))); }
702EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 asinh(const bfloat16& a) { return bfloat16(::asinhf(float(a))); }
703EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 acosh(const bfloat16& a) { return bfloat16(::acoshf(float(a))); }
704EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 atanh(const bfloat16& a) { return bfloat16(::atanhf(float(a))); }
705EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 exact_float_to_bfloat16(float f) {
706 return raw_uint16_to_bfloat16(static_cast<numext::uint16_t>(numext::bit_cast<numext::uint32_t>(f) >> 16));
707}
708EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 floor(const bfloat16& a) {
709 return exact_float_to_bfloat16(::floorf(float(a)));
710}
711EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 ceil(const bfloat16& a) {
712 return exact_float_to_bfloat16(::ceilf(float(a)));
713}
714EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 rint(const bfloat16& a) {
715 return exact_float_to_bfloat16(::rintf(float(a)));
716}
717EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 round(const bfloat16& a) {
718 return exact_float_to_bfloat16(::roundf(float(a)));
719}
720EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 trunc(const bfloat16& a) {
721 return exact_float_to_bfloat16(::truncf(float(a)));
722}
723// fmod is exact: a - n*b is either a itself or a multiple of the ulp of b that is smaller than |b|.
724EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 fmod(const bfloat16& a, const bfloat16& b) {
725 return exact_float_to_bfloat16(::fmodf(float(a), float(b)));
726}
727
728EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16(min)(const bfloat16& a, const bfloat16& b) {
729 const float f1 = static_cast<float>(a);
730 const float f2 = static_cast<float>(b);
731 return f2 < f1 ? b : a;
732}
733
734EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16(max)(const bfloat16& a, const bfloat16& b) {
735 const float f1 = static_cast<float>(a);
736 const float f2 = static_cast<float>(b);
737 return f1 < f2 ? b : a;
738}
739
740EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 fmin(const bfloat16& a, const bfloat16& b) {
741 const float f1 = static_cast<float>(a);
742 const float f2 = static_cast<float>(b);
743 return exact_float_to_bfloat16(::fminf(f1, f2));
744}
745
746EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 fmax(const bfloat16& a, const bfloat16& b) {
747 const float f1 = static_cast<float>(a);
748 const float f2 = static_cast<float>(b);
749 return exact_float_to_bfloat16(::fmaxf(f1, f2));
750}
751
752EIGEN_DEVICE_FUNC inline bfloat16 fma(const bfloat16& a, const bfloat16& b, const bfloat16& c) {
753 // Emulate FMA via float.
754 return bfloat16(numext::fma(static_cast<float>(a), static_cast<float>(b), static_cast<float>(c)));
755}
756
757#ifndef EIGEN_NO_IO
758EIGEN_ALWAYS_INLINE std::ostream& operator<<(std::ostream& os, const bfloat16& v) {
759 os << static_cast<float>(v);
760 return os;
761}
762#endif
763
764} // namespace bfloat16_impl
765
766namespace internal {
767
768template <>
769struct is_arithmetic<bfloat16> : std::true_type {};
770
771template <>
772struct random_impl<bfloat16> {
773 enum : int { MantissaBits = 7 };
774 using Impl = random_impl<float>;
775 static EIGEN_DEVICE_FUNC inline bfloat16 run(const bfloat16& x, const bfloat16& y) {
776 float result = Impl::run(x, y, MantissaBits);
777 return bfloat16(result);
778 }
779 static EIGEN_DEVICE_FUNC inline bfloat16 run() {
780 float result = Impl::run(MantissaBits);
781 return bfloat16_impl::exact_float_to_bfloat16(result);
782 }
783};
784
785} // namespace internal
786
787template <>
788struct NumTraits<Eigen::bfloat16> : GenericNumTraits<Eigen::bfloat16> {
789 enum { IsSigned = true, IsInteger = false, IsComplex = false, RequireInitialization = false };
790
791 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR static EIGEN_STRONG_INLINE Eigen::bfloat16 epsilon() {
792 return bfloat16_impl::raw_uint16_to_bfloat16(0x3c00);
793 }
794 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR static EIGEN_STRONG_INLINE Eigen::bfloat16 dummy_precision() {
795 return bfloat16_impl::raw_uint16_to_bfloat16(0x3D4D); // bfloat16(5e-2f);
796 }
797 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR static EIGEN_STRONG_INLINE Eigen::bfloat16 highest() {
798 return bfloat16_impl::raw_uint16_to_bfloat16(0x7F7F);
799 }
800 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR static EIGEN_STRONG_INLINE Eigen::bfloat16 lowest() {
801 return bfloat16_impl::raw_uint16_to_bfloat16(0xFF7F);
802 }
803 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR static EIGEN_STRONG_INLINE Eigen::bfloat16 infinity() {
804 return bfloat16_impl::raw_uint16_to_bfloat16(0x7f80);
805 }
806 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR static EIGEN_STRONG_INLINE Eigen::bfloat16 quiet_NaN() {
807 return bfloat16_impl::raw_uint16_to_bfloat16(0x7fc0);
808 }
809};
810
811} // namespace Eigen
812
813#if defined(EIGEN_HAS_HIP_BF16)
814#pragma pop_macro("EIGEN_CONSTEXPR")
815#endif
816
817namespace Eigen {
818namespace numext {
819
820template <>
821EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool(isnan)(const Eigen::bfloat16& h) {
822 return (bfloat16_impl::isnan)(h);
823}
824
825template <>
826EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool(isinf)(const Eigen::bfloat16& h) {
827 return (bfloat16_impl::isinf)(h);
828}
829
830template <>
831EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool(isfinite)(const Eigen::bfloat16& h) {
832 return (bfloat16_impl::isfinite)(h);
833}
834
835template <>
836EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Eigen::bfloat16 bit_cast<Eigen::bfloat16, uint16_t>(const uint16_t& src) {
837 return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(src);
838}
839
840template <>
841EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC uint16_t bit_cast<uint16_t, Eigen::bfloat16>(const Eigen::bfloat16& src) {
842 return Eigen::bfloat16_impl::raw_bfloat16_as_uint16(src);
843}
844
845EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 nextafter(const bfloat16& from, const bfloat16& to) {
846 if (numext::isnan EIGEN_NOT_A_MACRO(from)) {
847 return from;
848 }
849 if (numext::isnan EIGEN_NOT_A_MACRO(to)) {
850 return to;
851 }
852 if (from == to) {
853 return to;
854 }
855 uint16_t from_bits = numext::bit_cast<uint16_t>(from);
856 bool from_sign = from_bits >> 15;
857 if ((from_bits & 0x7fff) == 0) {
858 // From ±0 toward a nonzero value: the neighbor is the smallest subnormal
859 // carrying the sign of the direction (IEEE-754 nextUp/nextDown of zero).
860 from_bits = (to > from) ? uint16_t(0x0001) : uint16_t(0x8001);
861 } else if ((to > from) != from_sign) {
862 // Toward the infinity with the same sign as from: increase the magnitude.
863 ++from_bits;
864 } else {
865 --from_bits;
866 }
867 return numext::bit_cast<bfloat16>(from_bits);
868}
869
870// Specialize multiply-add to match packet operations and reduce conversions to/from float.
871template <>
872EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Eigen::bfloat16 madd<Eigen::bfloat16>(const Eigen::bfloat16& x,
873 const Eigen::bfloat16& y,
874 const Eigen::bfloat16& z) {
875 return Eigen::bfloat16(static_cast<float>(x) * static_cast<float>(y) + static_cast<float>(z));
876}
877
878} // namespace numext
879} // namespace Eigen
880
881#if EIGEN_HAS_STD_HASH
882namespace std {
883template <>
884struct hash<Eigen::bfloat16> {
885 EIGEN_STRONG_INLINE std::size_t operator()(const Eigen::bfloat16& a) const {
886 return static_cast<std::size_t>(Eigen::numext::bit_cast<Eigen::numext::uint16_t>(a));
887 }
888};
889} // namespace std
890#endif
891
892// Warp shuffle overloads for Eigen::bfloat16.
893// HIP uses non-sync __shfl variants; CUDA has native __nv_bfloat16 support in __shfl_sync.
894// Note that the following are __device__ - only functions.
895#if defined(EIGEN_HIPCC)
896
897#if defined(EIGEN_HAS_HIP_BF16)
898
899__device__ EIGEN_STRONG_INLINE Eigen::bfloat16 __shfl(Eigen::bfloat16 var, int srcLane, int width = warpSize) {
900 const int ivar = static_cast<int>(Eigen::numext::bit_cast<Eigen::numext::uint16_t>(var));
901 return Eigen::numext::bit_cast<Eigen::bfloat16>(static_cast<Eigen::numext::uint16_t>(__shfl(ivar, srcLane, width)));
902}
903
904__device__ EIGEN_STRONG_INLINE Eigen::bfloat16 __shfl_up(Eigen::bfloat16 var, unsigned int delta,
905 int width = warpSize) {
906 const int ivar = static_cast<int>(Eigen::numext::bit_cast<Eigen::numext::uint16_t>(var));
907 return Eigen::numext::bit_cast<Eigen::bfloat16>(static_cast<Eigen::numext::uint16_t>(__shfl_up(ivar, delta, width)));
908}
909
910__device__ EIGEN_STRONG_INLINE Eigen::bfloat16 __shfl_down(Eigen::bfloat16 var, unsigned int delta,
911 int width = warpSize) {
912 const int ivar = static_cast<int>(Eigen::numext::bit_cast<Eigen::numext::uint16_t>(var));
913 return Eigen::numext::bit_cast<Eigen::bfloat16>(
914 static_cast<Eigen::numext::uint16_t>(__shfl_down(ivar, delta, width)));
915}
916
917__device__ EIGEN_STRONG_INLINE Eigen::bfloat16 __shfl_xor(Eigen::bfloat16 var, int laneMask, int width = warpSize) {
918 const int ivar = static_cast<int>(Eigen::numext::bit_cast<Eigen::numext::uint16_t>(var));
919 return Eigen::numext::bit_cast<Eigen::bfloat16>(
920 static_cast<Eigen::numext::uint16_t>(__shfl_xor(ivar, laneMask, width)));
921}
922
923#endif // HIP
924
925#endif // __shfl*
926
927#if defined(EIGEN_HIPCC)
928EIGEN_STRONG_INLINE __device__ Eigen::bfloat16 __ldg(const Eigen::bfloat16* ptr) {
929 return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(
930 __ldg(Eigen::numext::bit_cast<const Eigen::numext::uint16_t*>(ptr)));
931}
932#endif // __ldg
933
934#endif // EIGEN_BFLOAT16_H
Holds information about the various numeric (i.e. scalar) types allowed by Eigen.
Definition NumTraits.h:233