17#ifndef EIGEN_BFLOAT16_H
18#define EIGEN_BFLOAT16_H
21#include "../../InternalHeaderCheck.h"
23#if defined(EIGEN_HAS_HIP_BF16)
30#pragma push_macro("EIGEN_CONSTEXPR")
32#define EIGEN_CONSTEXPR
35#define BF16_PACKET_FUNCTION(PACKET_F, PACKET_BF16, METHOD) \
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))); \
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)
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)
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)
82#if defined(EIGEN_HAS_HIP_BF16) && defined(EIGEN_GPU_COMPILE_PHASE)
83#define EIGEN_USE_HIP_BF16
92EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Eigen::bfloat16 bit_cast<Eigen::bfloat16, uint16_t>(
const uint16_t& src);
95EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC uint16_t bit_cast<uint16_t, Eigen::bfloat16>(
const Eigen::bfloat16& src);
97namespace bfloat16_impl {
99#if defined(EIGEN_USE_HIP_BF16)
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) {}
110struct __bfloat16_raw {
111#if defined(EIGEN_HAS_HIP_BF16) && !defined(EIGEN_GPU_COMPILE_PHASE)
112 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw() {}
114 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw() : value(0) {}
116 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR __bfloat16_raw(
unsigned short raw) : value(raw) {}
117 unsigned short value;
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);
128EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC __bfloat16_raw float_to_bfloat16_rtne<false>(
float ff);
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);
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) {}
141struct bfloat16 :
public bfloat16_impl::bfloat16_base {
142 using __bfloat16_raw = bfloat16_impl::__bfloat16_raw;
144 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16() {}
146 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16(
const __bfloat16_raw& h) : bfloat16_impl::bfloat16_base(h) {}
148 explicit EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR bfloat16(
bool b)
149 : bfloat16_impl::bfloat16_base(bfloat16_impl::raw_uint16_to_bfloat16(b ? 0x3f80 : 0)) {}
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))) {}
156 explicit EIGEN_DEVICE_FUNC bfloat16(
float f)
157 : bfloat16_impl::bfloat16_base(bfloat16_impl::float_to_bfloat16_rtne<false>(f)) {}
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()))) {}
165 EIGEN_DEVICE_FUNC
operator float()
const {
166 return bfloat16_impl::bfloat16_to_float(*
this);
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;
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;
204 static EIGEN_CONSTEXPR
const bool tinyness_before = std::numeric_limits<float>::tinyness_before;
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);
216 static EIGEN_CONSTEXPR Eigen::bfloat16 denorm_min() {
return Eigen::bfloat16_impl::raw_uint16_to_bfloat16(0x0001); }
220#if EIGEN_COMP_CXXVER < 17
222EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_specialized;
224EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_signed;
226EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_integer;
228EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_exact;
230EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::has_infinity;
232EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::has_quiet_NaN;
234EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::has_signaling_NaN;
235EIGEN_DIAGNOSTICS(push)
236EIGEN_DISABLE_DEPRECATED_WARNING
238EIGEN_CONSTEXPR
const std::float_denorm_style numeric_limits_bfloat16_impl<T>::has_denorm;
240EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::has_denorm_loss;
241EIGEN_DIAGNOSTICS(pop)
243EIGEN_CONSTEXPR
const std::float_round_style numeric_limits_bfloat16_impl<T>::round_style;
245EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_iec559;
247EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_bounded;
249EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::is_modulo;
251EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::digits;
253EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::digits10;
255EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::max_digits10;
257EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::radix;
259EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::min_exponent;
261EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::min_exponent10;
263EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::max_exponent;
265EIGEN_CONSTEXPR
const int numeric_limits_bfloat16_impl<T>::max_exponent10;
267EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::traps;
269EIGEN_CONSTEXPR
const bool numeric_limits_bfloat16_impl<T>::tinyness_before;
280class numeric_limits<Eigen::bfloat16> :
public Eigen::bfloat16_impl::numeric_limits_bfloat16_impl<> {};
282class numeric_limits<const Eigen::bfloat16> :
public numeric_limits<Eigen::bfloat16> {};
284class numeric_limits<volatile Eigen::bfloat16> :
public numeric_limits<Eigen::bfloat16> {};
286class numeric_limits<const volatile Eigen::bfloat16> :
public numeric_limits<Eigen::bfloat16> {};
291namespace bfloat16_impl {
296#if !defined(EIGEN_HAS_NATIVE_BF16) || (EIGEN_COMP_CLANG && !EIGEN_COMP_NVCC)
298#if EIGEN_COMP_CLANG && defined(EIGEN_GPUCC)
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__
305#define EIGEN_DEVICE_FUNC __host__ __device__
312EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator+(
const bfloat16& a,
const bfloat16& b) {
313 return bfloat16(
float(a) +
float(b));
315EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator+(
const bfloat16& a,
const int& b) {
316 return bfloat16(
float(a) +
static_cast<float>(b));
318EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator+(
const int& a,
const bfloat16& b) {
319 return bfloat16(
static_cast<float>(a) +
float(b));
321EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator*(
const bfloat16& a,
const bfloat16& b) {
322 return bfloat16(
float(a) *
float(b));
324EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator-(
const bfloat16& a,
const bfloat16& b) {
325 return bfloat16(
float(a) -
float(b));
327EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator/(
const bfloat16& a,
const bfloat16& b) {
328 return bfloat16(
float(a) /
float(b));
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);
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;
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;
343EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator+=(bfloat16& a,
const bfloat16& b) {
344 a = bfloat16(
float(a) +
float(b));
347EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator*=(bfloat16& a,
const bfloat16& b) {
348 a = bfloat16(
float(a) *
float(b));
351EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator-=(bfloat16& a,
const bfloat16& b) {
352 a = bfloat16(
float(a) -
float(b));
355EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16& operator/=(bfloat16& a,
const bfloat16& b) {
356 a = bfloat16(
float(a) /
float(b));
359EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator++(bfloat16& a) {
363EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator--(bfloat16& a) {
367EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator++(bfloat16& a,
int) {
368 bfloat16 original_value = a;
370 return original_value;
372EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator--(bfloat16& a,
int) {
373 bfloat16 original_value = a;
375 return original_value;
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));
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));
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));
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));
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));
411#if EIGEN_COMP_CLANG && defined(EIGEN_CUDACC)
412#pragma pop_macro("EIGEN_DEVICE_FUNC")
418EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 operator/(
const bfloat16& a, Index b) {
419 return bfloat16(
static_cast<float>(a) /
static_cast<float>(b));
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));
426 __bfloat16_raw output;
427 if (numext::isnan EIGEN_NOT_A_MACRO(v)) {
428 output.value = std::signbit(v) ? 0xFFC0 : 0x7FC0;
431 output.value =
static_cast<numext::uint16_t
>(numext::bit_cast<numext::uint32_t>(v) >> 16);
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)
442 return __bfloat16_raw(value);
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)
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));
462 __bfloat16_raw output;
464 if (numext::isnan EIGEN_NOT_A_MACRO(ff)) {
470 output.value = std::signbit(ff) ? 0xFFC0 : 0x7FC0;
621 output = float_to_bfloat16_rtne<true>(ff);
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));
636 numext::uint32_t input = numext::bit_cast<numext::uint32_t>(ff);
637 __bfloat16_raw output;
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);
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);
652 return numext::bit_cast<float>(
static_cast<numext::uint32_t
>(h.value) << 16);
658EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool(isinf)(
const bfloat16& a) {
659 return (raw_bfloat16_as_uint16(a) & 0x7fff) == 0x7f80;
661EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool(isnan)(
const bfloat16& a) {
662 return (raw_bfloat16_as_uint16(a) & 0x7fff) > 0x7f80;
664EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bool(isfinite)(
const bfloat16& a) {
665 return (raw_bfloat16_as_uint16(a) & 0x7fff) < 0x7f80;
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);
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))); }
676EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 ldexp(
const bfloat16& a,
int exponent) {
677 return bfloat16(numext::ldexp(
float(a), exponent));
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)));
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)));
690EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 atan2(
const bfloat16& a,
const bfloat16& b) {
691 return bfloat16(::atan2f(
float(a),
float(b)));
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));
708EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 floor(
const bfloat16& a) {
709 return exact_float_to_bfloat16(::floorf(
float(a)));
711EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 ceil(
const bfloat16& a) {
712 return exact_float_to_bfloat16(::ceilf(
float(a)));
714EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 rint(
const bfloat16& a) {
715 return exact_float_to_bfloat16(::rintf(
float(a)));
717EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 round(
const bfloat16& a) {
718 return exact_float_to_bfloat16(::roundf(
float(a)));
720EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 trunc(
const bfloat16& a) {
721 return exact_float_to_bfloat16(::truncf(
float(a)));
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)));
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;
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;
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));
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));
752EIGEN_DEVICE_FUNC
inline bfloat16 fma(
const bfloat16& a,
const bfloat16& b,
const bfloat16& c) {
754 return bfloat16(numext::fma(
static_cast<float>(a),
static_cast<float>(b),
static_cast<float>(c)));
758EIGEN_ALWAYS_INLINE std::ostream& operator<<(std::ostream& os,
const bfloat16& v) {
759 os << static_cast<float>(v);
769struct is_arithmetic<bfloat16> : std::true_type {};
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);
779 static EIGEN_DEVICE_FUNC
inline bfloat16 run() {
780 float result = Impl::run(MantissaBits);
781 return bfloat16_impl::exact_float_to_bfloat16(result);
788struct NumTraits<Eigen::bfloat16> : GenericNumTraits<Eigen::bfloat16> {
789 enum { IsSigned =
true, IsInteger =
false, IsComplex =
false, RequireInitialization =
false };
791 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR
static EIGEN_STRONG_INLINE Eigen::bfloat16 epsilon() {
792 return bfloat16_impl::raw_uint16_to_bfloat16(0x3c00);
794 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR
static EIGEN_STRONG_INLINE Eigen::bfloat16 dummy_precision() {
795 return bfloat16_impl::raw_uint16_to_bfloat16(0x3D4D);
797 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR
static EIGEN_STRONG_INLINE Eigen::bfloat16 highest() {
798 return bfloat16_impl::raw_uint16_to_bfloat16(0x7F7F);
800 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR
static EIGEN_STRONG_INLINE Eigen::bfloat16 lowest() {
801 return bfloat16_impl::raw_uint16_to_bfloat16(0xFF7F);
803 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR
static EIGEN_STRONG_INLINE Eigen::bfloat16 infinity() {
804 return bfloat16_impl::raw_uint16_to_bfloat16(0x7f80);
806 EIGEN_DEVICE_FUNC EIGEN_CONSTEXPR
static EIGEN_STRONG_INLINE Eigen::bfloat16 quiet_NaN() {
807 return bfloat16_impl::raw_uint16_to_bfloat16(0x7fc0);
813#if defined(EIGEN_HAS_HIP_BF16)
814#pragma pop_macro("EIGEN_CONSTEXPR")
821EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool(isnan)(
const Eigen::bfloat16& h) {
822 return (bfloat16_impl::isnan)(h);
826EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool(isinf)(
const Eigen::bfloat16& h) {
827 return (bfloat16_impl::isinf)(h);
831EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool(isfinite)(
const Eigen::bfloat16& h) {
832 return (bfloat16_impl::isfinite)(h);
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);
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);
845EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC bfloat16 nextafter(
const bfloat16& from,
const bfloat16& to) {
846 if (numext::isnan EIGEN_NOT_A_MACRO(from)) {
849 if (numext::isnan EIGEN_NOT_A_MACRO(to)) {
855 uint16_t from_bits = numext::bit_cast<uint16_t>(from);
856 bool from_sign = from_bits >> 15;
857 if ((from_bits & 0x7fff) == 0) {
860 from_bits = (to > from) ? uint16_t(0x0001) : uint16_t(0x8001);
861 }
else if ((to > from) != from_sign) {
867 return numext::bit_cast<bfloat16>(from_bits);
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));
881#if EIGEN_HAS_STD_HASH
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));
895#if defined(EIGEN_HIPCC)
897#if defined(EIGEN_HAS_HIP_BF16)
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)));
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)));
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)));
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)));
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)));
Holds information about the various numeric (i.e. scalar) types allowed by Eigen.
Definition NumTraits.h:233