Eigen  5.0.1
 
Loading...
Searching...
No Matches
GenericPacketMathFrexpLdexp.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2009-2019 Gael Guennebaud <gael.guennebaud@inria.fr>
5// Copyright (C) 2018-2025 Rasmus Munk Larsen <rmlarsen@gmail.com>
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_ARCH_GENERIC_PACKET_MATH_FREXP_LDEXP_H
13#define EIGEN_ARCH_GENERIC_PACKET_MATH_FREXP_LDEXP_H
14
15// IWYU pragma: private
16#include "../../InternalHeaderCheck.h"
17
18namespace Eigen {
19namespace internal {
20
21// Creates a Scalar integer type with same bit-width.
22template <typename T>
23struct make_integer;
24template <>
25struct make_integer<float> {
26 using type = numext::int32_t;
27};
28template <>
29struct make_integer<double> {
30 using type = numext::int64_t;
31};
32template <>
33struct make_integer<half> {
34 using type = numext::int16_t;
35};
36template <>
37struct make_integer<bfloat16> {
38 using type = numext::int16_t;
39};
40
41template <typename Packet>
42EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Packet pfrexp_generic_get_biased_exponent(const Packet& a) {
43 using Scalar = typename unpacket_traits<Packet>::type;
44 using PacketI = typename unpacket_traits<Packet>::integer_packet;
45 static constexpr int mantissa_bits = numext::numeric_limits<Scalar>::digits - 1;
46 return pcast<PacketI, Packet>(plogical_shift_right<mantissa_bits>(preinterpret<PacketI>(pabs(a))));
47}
48
49// Safely applies frexp, correctly handles denormals.
50// Assumes IEEE floating point format.
51template <typename Packet>
52EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Packet pfrexp_generic(const Packet& a, Packet& exponent) {
53 using Scalar = typename unpacket_traits<Packet>::type;
54 using ScalarUI = std::make_unsigned_t<typename make_integer<Scalar>::type>;
55 static constexpr int TotalBits = sizeof(Scalar) * CHAR_BIT, MantissaBits = numext::numeric_limits<Scalar>::digits - 1,
56 ExponentBits = TotalBits - MantissaBits - 1;
57
58 constexpr ScalarUI scalar_sign_mantissa_mask =
59 ~(((ScalarUI(1) << ExponentBits) - ScalarUI(1)) << MantissaBits); // ~0x7f800000
60 const Packet sign_mantissa_mask = pset1frombits<Packet>(static_cast<ScalarUI>(scalar_sign_mantissa_mask));
61 const Packet half = pset1<Packet>(Scalar(0.5));
62 const Packet zero = pzero(a);
63 const Packet normal_min = pset1<Packet>((numext::numeric_limits<Scalar>::min)()); // Minimum normal value, 2^-126
64
65 // To handle denormals, normalize by multiplying by 2^(int(MantissaBits)+1).
66 const Packet is_denormal = pcmp_lt(pabs(a), normal_min);
67 constexpr ScalarUI scalar_normalization_offset = ScalarUI(MantissaBits + 1); // 24
68 // The following cannot be constexpr because bfloat16(uint16_t) is not constexpr.
69 const Scalar scalar_normalization_factor = Scalar(ScalarUI(1) << int(scalar_normalization_offset)); // 2^24
70 const Packet normalization_factor = pset1<Packet>(scalar_normalization_factor);
71 const Packet normalized_a = pselect(is_denormal, pmul(a, normalization_factor), a);
72
73 // Determine exponent offset: -126 if normal, -126-24 if denormal
74 const Scalar scalar_exponent_offset = -Scalar((ScalarUI(1) << (ExponentBits - 1)) - ScalarUI(2)); // -126
75 Packet exponent_offset = pset1<Packet>(scalar_exponent_offset);
76 const Packet normalization_offset = pset1<Packet>(-Scalar(scalar_normalization_offset)); // -24
77 exponent_offset = pselect(is_denormal, padd(exponent_offset, normalization_offset), exponent_offset);
78
79 // Determine exponent and mantissa from normalized_a.
80 exponent = pfrexp_generic_get_biased_exponent(normalized_a);
81 // Zero, Inf and NaN return 'a' unmodified, exponent is zero
82 // (technically the exponent is unspecified for inf/NaN, but GCC/Clang set it to zero)
83 const Scalar scalar_non_finite_exponent = Scalar((ScalarUI(1) << ExponentBits) - ScalarUI(1)); // 255
84 const Packet non_finite_exponent = pset1<Packet>(scalar_non_finite_exponent);
85 const Packet is_zero_or_not_finite = por(pcmp_eq(a, zero), pcmp_eq(exponent, non_finite_exponent));
86 const Packet m = pselect(is_zero_or_not_finite, a, por(pand(normalized_a, sign_mantissa_mask), half));
87 exponent = pselect(is_zero_or_not_finite, zero, padd(exponent, exponent_offset));
88 return m;
89}
90
91// ((a * c1) * c1) * c2, with each partial product kept: -ffast-math would otherwise reassociate the factors, whose
92// products overflow or underflow where a * 2^e does not (see pldexp_generic).
93template <typename Packet>
94EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Packet pldexp_apply_factors(const Packet& a, const Packet& c1, const Packet& c2) {
95 Packet out = pmul(a, c1);
96 EIGEN_OPTIMIZATION_BARRIER(out)
97 out = pmul(out, c1);
98 EIGEN_OPTIMIZATION_BARRIER(out)
99 return pmul(out, c2);
100}
101
102// Safely applies ldexp, correctly handles overflows, underflows and denormals.
103// Assumes IEEE floating point format.
104template <typename Packet>
105EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Packet pldexp_generic(const Packet& a, const Packet& exponent) {
106 // We want to return a * 2^exponent, allowing for all possible integer
107 // exponents without overflowing or underflowing in intermediate
108 // computations.
109 //
110 // Since 'a' and the output can be denormal, the maximum range of 'exponent'
111 // to consider for a float is:
112 // -255-23 -> 255+23
113 // Below -278 any finite float 'a' will become zero, and above +278 any
114 // finite float will become inf, including when 'a' is the smallest possible
115 // denormal.
116 //
117 // Unfortunately, 2^(278) cannot be represented using either one or two
118 // finite normal floats, so we must split the scale factor into three parts.
119 // A product rounded in two steps can be off by one unit when an
120 // intermediate is subnormal, so only the last multiplication may round:
121 //
122 // Set e = min(max(exponent, -278), 278);
123 // b = min(max(exponent, -126), 126);
124 // t = (e - b) & ~1 (even, |t| <= 152)
125 // c1 = 2^(t/2)
126 // c2 = 2^(e - t) (|e - t| <= 127)
127 // out = ((a * c1) * c1) * c2 (= a * 2^e)
128 //
129 // For |e| <= 126, t = 0 and the result is the one product a * 2^e. For
130 // e > 126 the first two products are exact until they overflow, and then
131 // the result overflows. For e < -126, t >= e + 125, so an inexact partial
132 // product, below 2^-126, means the result is below 2^-251 and rounds to
133 // zero either way.
134 //
135 // Every partial product must contain 'a'. Reassociating scale factors can
136 // overflow (for example c1*c1 at e=256 for float), making pldexp(0, 256)
137 // NaN and finite denormal results infinite.
138 using PacketI = typename unpacket_traits<Packet>::integer_packet;
139 using Scalar = typename unpacket_traits<Packet>::type;
140 using ScalarI = typename unpacket_traits<PacketI>::type;
141 static constexpr int TotalBits = sizeof(Scalar) * CHAR_BIT, MantissaBits = numext::numeric_limits<Scalar>::digits - 1,
142 ExponentBits = TotalBits - MantissaBits - 1;
143
144 constexpr ScalarI bias_value = (ScalarI(1) << (ExponentBits - 1)) - ScalarI(1); // 127
145 constexpr ScalarI max_exp_value = (ScalarI(1) << ExponentBits) + ScalarI(MantissaBits - 1); // 278
146 constexpr ScalarI last_max_value = bias_value - ScalarI(1); // 126
147 const Packet max_exponent = pset1<Packet>(Scalar(max_exp_value));
148 const Packet neg_max_exponent = pset1<Packet>(Scalar(-max_exp_value));
149 const Packet last_max = pset1<Packet>(Scalar(last_max_value));
150 const Packet neg_last_max = pset1<Packet>(Scalar(-last_max_value));
151 const PacketI bias = pset1<PacketI>(bias_value);
152 const PacketI e = pcast<Packet, PacketI>(pmin(pmax(exponent, neg_max_exponent), max_exponent));
153 const PacketI b = pcast<Packet, PacketI>(pmin(pmax(exponent, neg_last_max), last_max));
154 const PacketI t = pandnot(psub(e, b), pset1<PacketI>(ScalarI(1))); // even
155 const Packet c1 = preinterpret<Packet>(plogical_shift_left<MantissaBits - 1>(padd(t, padd(bias, bias)))); // 2^(t/2)
156 const Packet c2 = preinterpret<Packet>(plogical_shift_left<MantissaBits>(padd(psub(e, t), bias))); // 2^(e-t)
157 return pldexp_apply_factors(a, c1, c2); // a * 2^e
158}
159
160// Explicitly multiplies
161// a * (2^e)
162// clamping e to the range
163// [NumTraits<Scalar>::min_exponent()-2, NumTraits<Scalar>::max_exponent()]
164//
165// This is approx 7x faster than pldexp_impl, but will prematurely over/underflow
166// if 2^e doesn't fit into a normal floating-point Scalar.
167//
168// Assumes IEEE floating point format
169template <typename Packet>
170EIGEN_STRONG_INLINE EIGEN_DEVICE_FUNC Packet pldexp_fast(const Packet& a, const Packet& exponent) {
171 using PacketI = typename unpacket_traits<Packet>::integer_packet;
172 using Scalar = typename unpacket_traits<Packet>::type;
173 using ScalarI = typename unpacket_traits<PacketI>::type;
174 static constexpr int TotalBits = sizeof(Scalar) * CHAR_BIT, MantissaBits = numext::numeric_limits<Scalar>::digits - 1,
175 ExponentBits = TotalBits - MantissaBits - 1;
176
177 const Packet bias = pset1<Packet>(Scalar((ScalarI(1) << (ExponentBits - 1)) - ScalarI(1))); // 127
178 const Packet limit = pset1<Packet>(Scalar((ScalarI(1) << ExponentBits) - ScalarI(1))); // 255
179 // restrict biased exponent between 0 and 255 for float.
180 const PacketI e = pcast<Packet, PacketI>(pmin(pmax(padd(exponent, bias), pzero(limit)), limit)); // exponent + 127
181 // return a * (2^e)
182 return pmul(a, preinterpret<Packet>(plogical_shift_left<MantissaBits>(e)));
183}
184
185} // end namespace internal
186} // end namespace Eigen
187
188#endif // EIGEN_ARCH_GENERIC_PACKET_MATH_FREXP_LDEXP_H