Eigen  5.0.1
 
Loading...
Searching...
No Matches
SafeScaling.h
1// SPDX-FileCopyrightText: The Eigen Authors
2// SPDX-License-Identifier: MPL-2.0
3
4#ifndef EIGEN_SAFE_SCALING_H
5#define EIGEN_SAFE_SCALING_H
6
7// IWYU pragma: private
8#include "./InternalHeaderCheck.h"
9
10namespace Eigen {
11namespace internal {
12
13// Select factors for a positive finite scale. Supported binary floating-point scalars use normal powers of two;
14// other scalar types use division for matrix scaling and can request a clamped reciprocal for stableNorm(). General
15// scaling rounds down so it does not discard a representable tail that division by the original value would preserve.
16// Stable reductions separately round up to keep scaled magnitudes at most one. At the upper exponent boundary, the
17// scale is clamped to keep its reciprocal normal; for the standard binary formats this leaves magnitudes below four.
18template <typename Scalar>
19struct supports_power_of_two_scaling
20 : bool_constant<(std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value) &&
21 std::numeric_limits<Scalar>::is_iec559 && std::numeric_limits<Scalar>::radix == 2 &&
22 (sizeof(Scalar) == sizeof(numext::uint32_t) || sizeof(Scalar) == sizeof(numext::uint64_t))> {};
23
24#if !defined(EIGEN_GPU_COMPILE_PHASE)
25template <>
26struct supports_power_of_two_scaling<long double> : bool_constant<std::numeric_limits<long double>::radix == 2> {};
27#endif
28
29template <>
30struct supports_power_of_two_scaling<half> : true_type {};
31
32template <>
33struct supports_power_of_two_scaling<bfloat16> : true_type {};
34
35template <typename Scalar, bool = supports_power_of_two_scaling<Scalar>::value>
36struct safe_scaling;
37
38template <typename Scalar>
39struct safe_scaling_factors {
40 Scalar scale = Scalar(1);
41 // Division-only arithmetic scaling leaves this unused reciprocal at one.
42 Scalar invScale = Scalar(1);
43};
44
45// value * 2^exponent through the integer significand of the value with bit pattern valueBits, so FTZ/DAZ hardware and
46// ARMv7 NEON can flush neither a subnormal input nor a subnormal result. Exact wherever the result is representable;
47// a subnormal result rounds to nearest, ties to even, as IEEE 754 multiplication does; below denorm_min / 2 the result
48// is a signed zero and above the largest finite value a signed infinity. Zeros, infinities and NaNs pass through. Keep
49// the integer-only ABI so compilers cannot replace the reconstruction with FTZ/DAZ-sensitive floating-point arithmetic.
50template <typename Scalar>
51EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE typename binary_floating_point_traits<Scalar>::Bits scale_binary_bits_by_exponent(
52 const typename binary_floating_point_traits<Scalar>::Bits valueBits, int exponent) {
53 using Binary = binary_floating_point_traits<Scalar>;
54 using Bits = typename Binary::Bits;
55 constexpr int kInfinityExponent = int(Binary::kExponentMask >> Binary::kFractionBits);
56
57 const Bits valueExponentBits = valueBits & Binary::kExponentMask;
58 const Bits fraction = valueBits & Binary::kFractionMask;
59 if (valueExponentBits == Binary::kExponentMask || (valueExponentBits == 0 && fraction == 0)) return valueBits;
60 const Bits sign = valueBits & Binary::kSignBit;
61
62 Bits significand = fraction;
63 int resultExponent = int(valueExponentBits >> Binary::kFractionBits);
64 if (resultExponent == 0) {
65 resultExponent = 1;
66 while (significand < Binary::kExponentUnit) {
67 significand <<= 1;
68 --resultExponent;
69 }
70 }
71 // Every finite value saturates beyond this magnitude of exponent, so the clamp changes no result.
72 constexpr int kSaturatingExponent = kInfinityExponent + Binary::kFractionBits + 2;
73 resultExponent += numext::mini(numext::maxi(exponent, -kSaturatingExponent), kSaturatingExponent);
74 if (resultExponent >= kInfinityExponent) return sign | Binary::kExponentMask;
75 if (resultExponent > 0)
76 return sign | (Bits(resultExponent) << Binary::kFractionBits) | (significand & Binary::kFractionMask);
77
78 // Subnormal result m * 2^(resultExponent - 1) denorm_min, m = significand | kExponentUnit < 2 * kExponentUnit: round
79 // to nearest, ties to even, as IEEE 754 multiplication does. A carry into kExponentUnit encodes the smallest normal;
80 // for shift > kFractionBits + 1 the result is below denorm_min / 2 and rounds to a signed zero.
81 const int shift = 1 - resultExponent;
82 if (shift > Binary::kFractionBits + 1) return sign;
83 const Bits mantissa = significand | Binary::kExponentUnit;
84 const Bits halfway = Bits(1) << (shift - 1);
85 const Bits remainder = mantissa & ((halfway << 1) - 1);
86 Bits rounded = mantissa >> shift;
87 if (remainder > halfway || (remainder == halfway && (rounded & Bits(1)) != 0)) ++rounded;
88 return sign | rounded;
89}
90
91// The same scaling by a positive normal power of two given as the bit pattern of its exponent field.
92template <typename Scalar>
93EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE typename binary_floating_point_traits<Scalar>::Bits
94scale_binary_bits_by_power_of_two(const typename binary_floating_point_traits<Scalar>::Bits valueBits,
95 const typename binary_floating_point_traits<Scalar>::Bits factorExponentBits) {
96 using Binary = binary_floating_point_traits<Scalar>;
97 return scale_binary_bits_by_exponent<Scalar>(
98 valueBits, int(factorExponentBits >> Binary::kFractionBits) - Binary::kExponentBias);
99}
100
101// value * 2^exponent for the binary floating-point scalars, with the rounding and saturation of
102// scale_binary_bits_by_exponent(): the ldexp that FTZ/DAZ hardware cannot flush, for any int exponent.
103template <typename Scalar>
104EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar scale_binary_by_exponent(const Scalar& value, int exponent) {
105 using Binary = binary_floating_point_traits<Scalar>;
106 return numext::bit_cast<Scalar>(scale_binary_bits_by_exponent<Scalar>(Binary::bits(value), exponent));
107}
108
109template <typename Scalar>
110EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::complex<Scalar> scale_binary_by_exponent(const std::complex<Scalar>& value,
111 int exponent) {
112 return std::complex<Scalar>(scale_binary_by_exponent(value.real(), exponent),
113 scale_binary_by_exponent(value.imag(), exponent));
114}
115
116// frexp's exponent e, 2^(e - 1) <= |value| < 2^e, read from the representation of a finite value so that FTZ/DAZ
117// hardware cannot flush a subnormal; 0 for a zero. Infinities and NaNs are not classified: they come back one above
118// the largest finite exponent. The classification takes the magnitude bits out of line, like
119// scale_binary_bits_by_exponent(): inlined, a compiler folds an integer test on the bit pattern of a float back into
120// a floating-point comparison (`magnitude == 0` into `value == 0`), which FTZ/DAZ hardware flushes exactly like the
121// arithmetic the bits were meant to bypass.
122template <typename Scalar>
123EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE int binary_frexp_exponent_of_magnitude(
124 const typename binary_floating_point_traits<Scalar>::Bits magnitude) {
125 using Binary = binary_floating_point_traits<Scalar>;
126 using Bits = typename Binary::Bits;
127 if (magnitude == 0) return 0;
128 const Bits exponentBits = magnitude & Binary::kExponentMask;
129 if (exponentBits != 0) return int(exponentBits >> Binary::kFractionBits) - Binary::kExponentBias + 1;
130 // A subnormal is its fraction times 2^(min_exponent - digits).
131 return log_2_impl<Bits>::run_floor(magnitude) + 1 + std::numeric_limits<Scalar>::min_exponent -
132 std::numeric_limits<Scalar>::digits;
133}
134
135template <typename Scalar>
136EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE int binary_frexp_exponent(const Scalar& value) {
137 return binary_frexp_exponent_of_magnitude<Scalar>(binary_floating_point_traits<Scalar>::magnitude(value));
138}
139
140// The larger of two magnitude bit patterns, out of line for the same reason.
141template <typename Scalar>
142EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE typename binary_floating_point_traits<Scalar>::Bits larger_magnitude_bits(
143 const typename binary_floating_point_traits<Scalar>::Bits a,
144 const typename binary_floating_point_traits<Scalar>::Bits b) {
145 return a < b ? b : a;
146}
147
148// |value| < the smallest normal from the magnitude bits, out of line for the same reason.
149template <typename Scalar>
150EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE bool is_below_normal_magnitude_bits(
151 const typename binary_floating_point_traits<Scalar>::Bits magnitude) {
152 return magnitude < binary_floating_point_traits<Scalar>::kExponentUnit;
153}
154
155// frexp through the representation: the significand in [0.5, 1), with the sign of value, and its exponent.
156template <typename Scalar>
157EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar binary_frexp(const Scalar& value, int& exponent) {
158 exponent = binary_frexp_exponent(value);
159 return scale_binary_by_exponent(value, -exponent);
160}
161
162// Multiplication by a positive normal power of two through integer significands, so FTZ/DAZ cannot flush a subnormal
163// input or result; a subnormal result rounds as IEEE 754 multiplication does, and a result above the largest finite
164// value saturates to infinity.
165template <typename Scalar>
166EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar scale_binary_by_power_of_two(const Scalar& value, const Scalar& factor) {
167 using Binary = binary_floating_point_traits<Scalar>;
168 using Bits = typename Binary::Bits;
169 const Bits factorExponentBits = Binary::bits(factor);
170 eigen_internal_assert(factorExponentBits != 0 && factorExponentBits < Binary::kExponentMask &&
171 (factorExponentBits & ~Binary::kExponentMask) == 0);
172 return numext::bit_cast<Scalar>(scale_binary_bits_by_power_of_two<Scalar>(Binary::bits(value), factorExponentBits));
173}
174
175template <typename Scalar>
176EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::complex<Scalar> scale_binary_by_power_of_two(
177 const std::complex<Scalar>& value, const Scalar& factor) {
178 return std::complex<Scalar>(scale_binary_by_power_of_two(value.real(), factor),
179 scale_binary_by_power_of_two(value.imag(), factor));
180}
181
182// Coefficient-wise value * 2^exponent through scale_binary_by_exponent(). Scalar only: the integer path is the point.
183template <typename Scalar>
184struct scale_by_exponent_op {
185 EIGEN_DEVICE_FUNC explicit scale_by_exponent_op(int exponent) : m_exponent(exponent) {}
186
187 template <typename CoeffScalar>
188 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffScalar operator()(const CoeffScalar& value) const {
189 return scale_binary_by_exponent(value, m_exponent);
190 }
191
192 int m_exponent;
193};
194
195template <typename Scalar>
196struct functor_traits<scale_by_exponent_op<Scalar>> {
197 static constexpr int Cost = 10 * NumTraits<Scalar>::MulCost;
198 static constexpr bool PacketAccess = false;
199 static constexpr bool IsRepeatable = true;
200};
201
202template <typename FactorScalar>
203struct scale_by_power_of_two_op {
204 EIGEN_DEVICE_FUNC explicit scale_by_power_of_two_op(const FactorScalar& factor) : m_factor(factor) {}
205
206 template <typename CoeffScalar>
207 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffScalar operator()(const CoeffScalar& value) const {
208 return scale_binary_by_power_of_two(value, m_factor);
209 }
210
211 FactorScalar m_factor;
212};
213
214template <typename FactorScalar>
215struct functor_traits<scale_by_power_of_two_op<FactorScalar>> {
216 // Integer significand scaling is scalar and calls an out-of-line helper.
217 static constexpr int Cost = 10 * NumTraits<FactorScalar>::MulCost;
218 static constexpr bool PacketAccess = false;
219 static constexpr bool IsRepeatable = true;
220};
221
222template <typename FactorScalar, typename CoeffScalar>
223struct use_subnormal_preserving_scaling
224 : bool_constant<
225 (std::is_same<FactorScalar, float>::value &&
226 (std::is_same<CoeffScalar, float>::value || std::is_same<CoeffScalar, std::complex<float>>::value)) ||
227 (std::is_same<FactorScalar, double>::value &&
228 (std::is_same<CoeffScalar, double>::value || std::is_same<CoeffScalar, std::complex<double>>::value))> {};
229
230// Subnormal-magnitude tests read the representation for float and double, where a comparison under DAZ
231// reads a subnormal as zero as well; the other scalars compare arithmetically.
232template <typename Scalar>
233EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_subnormal_magnitude_impl(const Scalar& value, true_type) {
234 // The smallest normal, 2^(min_exponent - 1), has frexp exponent min_exponent.
235 const int exponent = binary_frexp_exponent(value);
236 return exponent != 0 && exponent < std::numeric_limits<Scalar>::min_exponent;
237}
238template <typename Scalar>
239EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_subnormal_magnitude_impl(const Scalar& value, false_type) {
240 const Scalar magnitude = numext::abs(value);
241 return magnitude > Scalar(0) && magnitude < (std::numeric_limits<Scalar>::min)();
242}
243// 0 < |value| < the smallest normal.
244template <typename Scalar>
245EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_subnormal_magnitude(const Scalar& value) {
246 return is_subnormal_magnitude_impl(value, bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
247}
248
249// |value| < the smallest normal, zero included: a maximum that a flushing comparison may have produced. A scalar
250// reduction under FTZ/DAZ does not always return a zero: where max(a, b) is a compare and a select (Arm FZ with
251// scalar VFP or fcsel, ARMv7's double, which has no packet), every comparison is false and the running maximum keeps
252// the bits of the first operand.
253template <typename Scalar>
254EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_zero_or_subnormal_magnitude_impl(const Scalar& value, true_type) {
255 return is_below_normal_magnitude_bits<Scalar>(binary_floating_point_traits<Scalar>::magnitude(value));
256}
257template <typename Scalar>
258EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_zero_or_subnormal_magnitude_impl(const Scalar& value, false_type) {
259 return numext::abs(value) < (std::numeric_limits<Scalar>::min)();
260}
261template <typename Scalar>
262EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_zero_or_subnormal_magnitude(const Scalar& value) {
263 return is_zero_or_subnormal_magnitude_impl(value,
264 bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
265}
266
267// The larger of two non-negative values, such as two recovered maxima, read from the representation for float and
268// double: an arithmetic maximum (a compiler may lower it to fmaxnm) flushes subnormal operands to zero under FTZ/DAZ.
269template <typename Scalar>
270EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar max_preserving_subnormals_impl(const Scalar& a, const Scalar& b,
271 true_type) {
272 using Binary = binary_floating_point_traits<Scalar>;
273 return numext::bit_cast<Scalar>(larger_magnitude_bits<Scalar>(Binary::magnitude(a), Binary::magnitude(b)));
274}
275template <typename Scalar>
276EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar max_preserving_subnormals_impl(const Scalar& a, const Scalar& b,
277 false_type) {
278 return numext::maxi(a, b);
279}
280template <typename Scalar>
281EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar max_preserving_subnormals(const Scalar& a, const Scalar& b) {
282 return max_preserving_subnormals_impl(a, b, bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
283}
284
285// |value| and value < 0 of a finite real scalar, from the representation for float and double: a comparison reads a
286// negative subnormal as zero under DAZ, and an abs() that widens (MSVC 19.29 takes abs(float) through double) flushes
287// a subnormal result under FTZ when it narrows back.
288template <typename Scalar>
289EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar abs_preserving_subnormals_impl(const Scalar& value, true_type) {
290 return numext::bit_cast<Scalar>(binary_floating_point_traits<Scalar>::magnitude(value));
291}
292template <typename Scalar>
293EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar abs_preserving_subnormals_impl(const Scalar& value, false_type) {
294 return numext::abs(value);
295}
296template <typename Scalar>
297EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar abs_preserving_subnormals(const Scalar& value) {
298 return abs_preserving_subnormals_impl(value,
299 bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
300}
301
302template <typename Scalar>
303EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_negative_preserving_subnormals_impl(const Scalar& value, true_type) {
304 using Binary = binary_floating_point_traits<Scalar>;
305 return (Binary::bits(value) & Binary::kSignBit) != 0 && !is_zero_magnitude_bits<Scalar>(Binary::magnitude(value));
306}
307template <typename Scalar>
308EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_negative_preserving_subnormals_impl(const Scalar& value, false_type) {
309 return value < Scalar(0);
310}
311template <typename Scalar>
312EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool is_negative_preserving_subnormals(const Scalar& value) {
313 return is_negative_preserving_subnormals_impl(
314 value, bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
315}
316
317// The position of the largest magnitude of a real vector with at least one coefficient, read from the representation
318// and out of line like the other bit classifications.
319template <typename Scalar, typename Derived>
320EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE Index index_of_largest_magnitude(const DenseBase<Derived>& src) {
321 using Binary = binary_floating_point_traits<Scalar>;
322 using Bits = typename Binary::Bits;
323 const evaluator<Derived> coeffs(src.derived());
324 Index position = 0;
325 Bits largest = Binary::magnitude(coeffs.coeff(0));
326 for (Index i = 1; i < src.size(); ++i) {
327 const Bits magnitude = Binary::magnitude(coeffs.coeff(i));
328 if (magnitude > largest) {
329 largest = magnitude;
330 position = i;
331 }
332 }
333 return position;
334}
335
336// ldexp and frexp's exponent for a real scalar: float and double through the representation, which FTZ/DAZ cannot
337// reach, the other scalars through the C library. frexp's exponent is e with 2^(e - 1) <= |value| < 2^e, 0 for a
338// zero.
339template <typename Scalar>
340EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar ldexp_preserving_subnormals_impl(const Scalar& value, int exponent,
341 true_type) {
342 return scale_binary_by_exponent(value, exponent);
343}
344template <typename Scalar>
345EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar ldexp_preserving_subnormals_impl(const Scalar& value, int exponent,
346 false_type) {
347 return numext::ldexp(value, exponent);
348}
349template <typename Scalar>
350EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar ldexp_preserving_subnormals(const Scalar& value, int exponent) {
351 return ldexp_preserving_subnormals_impl(value, exponent,
352 bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
353}
354
355template <typename Scalar>
356EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE int frexp_exponent_preserving_subnormals_impl(const Scalar& value, true_type) {
357 return binary_frexp_exponent(value);
358}
359template <typename Scalar>
360EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE int frexp_exponent_preserving_subnormals_impl(const Scalar& value, false_type) {
361 int exponent = 0;
362 EIGEN_USING_STD(frexp);
363 frexp(value, &exponent);
364 return exponent;
365}
366template <typename Scalar>
367EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE int frexp_exponent_preserving_subnormals(const Scalar& value) {
368 return frexp_exponent_preserving_subnormals_impl(
369 value, bool_constant<use_subnormal_preserving_scaling<Scalar, Scalar>::value>());
370}
371// The significand in [0.5, 1), with the sign of value, and its exponent.
372template <typename Scalar>
373EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar frexp_preserving_subnormals(const Scalar& value, int& exponent) {
374 exponent = frexp_exponent_preserving_subnormals(value);
375 return ldexp_preserving_subnormals(value, -exponent);
376}
377
378template <typename Scalar, bool IsPowerOfTwo_>
379struct safe_scaling_operations {
380 private:
381 using Factors = safe_scaling_factors<Scalar>;
382
383 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool needs_subnormal_recovery(const Scalar& value) {
384 using Binary = binary_floating_point_traits<Scalar>;
385 using Bits = typename Binary::Bits;
386 constexpr Bits kThreshold = Bits(subnormal_recovery_exponent() + Binary::kExponentBias) << Binary::kFractionBits;
387 return Binary::magnitude(value) - Bits(1) < kThreshold - Bits(1);
388 }
389
390 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool is_identity(const Factors& factors) {
391 return factors.scale == Scalar(1) && factors.invScale == Scalar(1);
392 }
393
394 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool has_normal_reciprocal(const Scalar&, const Scalar&, false_type) {
395 return false;
396 }
397
398 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool has_normal_reciprocal(const Scalar& value, const Scalar& normalMin,
399 true_type) {
400 return value >= normalMin && value <= Scalar(1) / normalMin;
401 }
402
403 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool is_positive_finite(const Scalar& value, false_type) {
404 return value > Scalar(0) && value <= NumTraits<Scalar>::highest();
405 }
406
407 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool is_positive_finite(const Scalar& value, true_type) {
408 using Binary = binary_floating_point_traits<Scalar>;
409 const typename Binary::Bits bits = Binary::bits(value);
410 return bits > 0 && bits < Binary::kExponentMask;
411 }
412
413 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Factors select_factors(const Scalar& maxCoeff) {
414 using IsIeeeBinary = bool_constant<std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value>;
415 if (!is_positive_finite(maxCoeff, IsIeeeBinary())) return Factors{};
416 return safe_scaling<Scalar, IsPowerOfTwo_>::compute_floor_factors(maxCoeff);
417 }
418
419 template <typename Src, typename Func>
420 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void with_scaled_impl(const Src& src, const Scalar&,
421 const Factors& factors, const Func& func,
422 false_type) {
423 if (is_identity(factors)) {
424 // Keep the copy distinct from multiplication by one, which flushes subnormals under FTZ.
425 const Src* unscaled = &src;
426 EIGEN_OPTIMIZATION_BARRIER(unscaled);
427 func(*unscaled);
428 } else EIGEN_IF_CONSTEXPR (IsPowerOfTwo_)
429 func(src * factors.invScale);
430 else
431 func(src / factors.scale);
432 }
433
434 template <typename Src, typename Func>
435 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void with_scaled_impl(const Src& src, const Scalar& maxCoeff,
436 const Factors& factors, const Func& func,
437 true_type) {
438 if (!is_identity(factors) && needs_subnormal_recovery(maxCoeff))
439 func(src.unaryExpr(scale_by_power_of_two_op<Scalar>(factors.invScale)));
440 else
441 with_scaled_impl(src, maxCoeff, factors, func, false_type());
442 }
443
444 template <typename Dest, typename Src>
445 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void unscale_to_impl(Dest& dest, const Src& src, const Scalar&,
446 const Factors& factors, false_type) {
447 unscale_to(dest, src, factors);
448 }
449
450 // Below the recovery threshold, unscaling rounds coefficients that are significant relative to maxCoeff into the
451 // subnormal range, where FTZ/DAZ arithmetic would zero them.
452 template <typename Dest, typename Src>
453 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void unscale_to_impl(Dest& dest, const Src& src, const Scalar& maxCoeff,
454 const Factors& factors, true_type) {
455 if (!is_identity(factors) && needs_subnormal_recovery(maxCoeff))
456 dest = src.unaryExpr(scale_by_power_of_two_op<Scalar>(factors.scale));
457 else
458 unscale_to_impl(dest, src, maxCoeff, factors, false_type());
459 }
460
461 template <typename Src>
462 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar recover_flushed_max_coeff_impl(const Src&, const Scalar& maxCoeff,
463 false_type) {
464 return maxCoeff;
465 }
466
467 template <typename Src>
468 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar recover_flushed_max_coeff_impl(const Src& src,
469 const Scalar& maxCoeff,
470 true_type) {
471 using Binary = binary_floating_point_traits<Scalar>;
472 using Bits = typename Binary::Bits;
473 if (!is_zero_or_subnormal_magnitude(maxCoeff)) return maxCoeff;
474 // A product expression has no coefficient access; its evaluator materializes it.
475 const evaluator<Src> coeffs(src);
476 Bits maxBits = 0;
477 for (Index col = 0; col < src.cols(); ++col) {
478 for (Index row = 0; row < src.rows(); ++row) {
479 const typename Src::Scalar coeff = coeffs.coeff(row, col);
480 maxBits = numext::maxi(
481 maxBits, numext::maxi(Binary::magnitude(numext::real(coeff)), Binary::magnitude(numext::imag(coeff))));
482 }
483 }
484 return numext::bit_cast<Scalar>(maxBits);
485 }
486
487 template <typename Derived>
488 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index recover_flushed_max_coeff_index_impl(const DenseBase<Derived>&,
489 Index position, false_type) {
490 return position;
491 }
492
493 template <typename Derived>
494 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index recover_flushed_max_coeff_index_impl(const DenseBase<Derived>& src,
495 Index, true_type) {
496 return index_of_largest_magnitude<Scalar>(src);
497 }
498
499 public:
500 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Factors
501 compute_ceiling_factors_with_normal_reciprocal(const Scalar& value) {
502 Factors factors = safe_scaling<Scalar, IsPowerOfTwo_>::compute_ceiling_factors(value);
503 EIGEN_IF_CONSTEXPR (!IsPowerOfTwo_) {
504 if (factors.invScale > NumTraits<Scalar>::highest()) {
505 factors.invScale = NumTraits<Scalar>::highest();
506 factors.scale = Scalar(1) / factors.invScale;
507 }
508 }
509 return factors;
510 }
511
512 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE bool try_compute_ceiling_factors_with_normal_reciprocal(
513 const Scalar& value, const Scalar& normalMin, Factors& factors) {
514 if (!IsPowerOfTwo_ &&
515 !has_normal_reciprocal(value, normalMin, bool_constant<std::is_floating_point<Scalar>::value>()))
516 return false;
517 factors = compute_ceiling_factors_with_normal_reciprocal(value);
518 return true;
519 }
520
521 // Dispatch once per expression; ordinary scaling retains packet access. The callback consumes the lazy
522 // expression before its operands expire, without requiring a common C++ type for the different paths.
523 template <typename Src, typename Func>
524 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Factors with_scaled(const Src& src, const Scalar& maxCoeff,
525 const Func& func) {
526 const Factors factors = select_factors(maxCoeff);
527 constexpr bool kPreserveSubnormalInputs =
528 IsPowerOfTwo_ && use_subnormal_preserving_scaling<Scalar, typename Src::Scalar>::value;
529 with_scaled_impl(src, maxCoeff, factors, func, bool_constant<kPreserveSubnormalInputs>());
530 return factors;
531 }
532
533 template <typename Dest, typename Src>
534 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void unscale_to(Dest&& dest, const Src& src, const Factors& factors) {
535 if (factors.scale == Scalar(1)) {
536 const Src* unscaled = &src;
537 EIGEN_OPTIMIZATION_BARRIER(unscaled);
538 dest = *unscaled;
539 return;
540 }
541 dest = src * factors.scale;
542 }
543
544 template <typename ValueType>
545 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void unscale_in_place(ValueType& value, const Factors& factors) {
546 if (factors.scale == Scalar(1)) return;
547 value *= factors.scale;
548 }
549
550 template <typename MatrixType>
551 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void unscale_in_place(MatrixType& matrix, const Scalar& maxCoeff,
552 const Factors& factors) {
553 if (factors.scale == Scalar(1)) return;
554 unscale_to(matrix, matrix, maxCoeff, factors);
555 }
556
557 template <typename Dest, typename Src>
558 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void unscale_to(Dest&& dest, const Src& src, const Scalar& maxCoeff,
559 const Factors& factors) {
560 constexpr bool kPreserveSubnormalOutputs =
561 IsPowerOfTwo_ && use_subnormal_preserving_scaling<Scalar, typename Src::Scalar>::value;
562 unscale_to_impl(dest, src, maxCoeff, factors, bool_constant<kPreserveSubnormalOutputs>());
563 }
564
565 // Below 2^subnormal_recovery_exponent(), a coefficient within a factor 2^-kFractionBits of the largest one can be
566 // subnormal, so a scaling that must survive FTZ/DAZ goes through integer significands (scale_by_exponent_op,
567 // scale_by_power_of_two_op) instead of packet multiplication.
568 EIGEN_DEVICE_FUNC static constexpr int subnormal_recovery_exponent() {
569 return std::numeric_limits<Scalar>::min_exponent + std::numeric_limits<Scalar>::digits - 2;
570 }
571
572 // A SIMD unit that flushes subnormal inputs (ARMv7 NEON, Arm FZ, DAZ) reduces an all-subnormal matrix to a zero
573 // maximum, and a flushed scalar reduction to whichever subnormal it held first, so rescan a zero or subnormal
574 // maxCoeff from the representation. Any subnormal maximum, including the largest real or imaginary magnitude of a
575 // complex matrix, selects the same factors as the true one.
576 template <typename Src>
577 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Scalar recover_flushed_max_coeff(const Src& src,
578 const Scalar& maxCoeff) {
579 constexpr bool kPreserveSubnormalInputs =
580 IsPowerOfTwo_ && use_subnormal_preserving_scaling<Scalar, typename Src::Scalar>::value;
581 return recover_flushed_max_coeff_impl(src, maxCoeff, bool_constant<kPreserveSubnormalInputs>());
582 }
583
584 // The position of the largest magnitude in a real vector whose maxCoeff(&position) read zero or subnormal,
585 // rescanned from the representation like recover_flushed_max_coeff(); the position maxCoeff() reported where there
586 // is no representation path.
587 template <typename Derived>
588 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Index recover_flushed_max_coeff_index(const DenseBase<Derived>& src,
589 Index position) {
590 constexpr bool kPreserveSubnormalInputs =
591 IsPowerOfTwo_ && use_subnormal_preserving_scaling<Scalar, typename Derived::Scalar>::value;
592 return recover_flushed_max_coeff_index_impl(src, position, bool_constant<kPreserveSubnormalInputs>());
593 }
594
595 template <typename MatrixType>
596 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Factors scale_in_place(MatrixType& matrix, const Scalar& maxCoeff) {
597 const Factors factors = select_factors(maxCoeff);
598 scale_in_place(matrix, maxCoeff, factors);
599 return factors;
600 }
601
602 template <typename MatrixType>
603 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void scale_in_place(MatrixType& matrix, const Scalar& maxCoeff,
604 const Factors& factors) {
605 if (!is_identity(factors)) scale_to(matrix, matrix, maxCoeff, factors);
606 }
607
608 template <typename Dest, typename Src>
609 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE void scale_to(Dest& dest, const Src& src, const Scalar& maxCoeff,
610 const Factors& factors) {
611 constexpr bool kPreserveSubnormalInputs =
612 IsPowerOfTwo_ && use_subnormal_preserving_scaling<Scalar, typename Src::Scalar>::value;
613 with_scaled_impl(
614 src, maxCoeff, factors, [&](const auto& scaled) { dest = scaled; }, bool_constant<kPreserveSubnormalInputs>());
615 }
616
617 template <typename Dest, typename Src>
618 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE Factors scale_to(Dest& dest, const Src& src, const Scalar& maxCoeff) {
619 const Factors factors = select_factors(maxCoeff);
620 scale_to(dest, src, maxCoeff, factors);
621 return factors;
622 }
623};
624
625template <typename Scalar, bool>
626struct safe_scaling : safe_scaling_operations<Scalar, false> {
627 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE safe_scaling_factors<Scalar> compute_floor_factors(const Scalar& value) {
628 return {value, Scalar(1)};
629 }
630
631 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE safe_scaling_factors<Scalar> compute_ceiling_factors(
632 const Scalar& value) {
633 return {value, Scalar(1) / value};
634 }
635};
636
637template <typename Scalar>
638struct safe_scaling<Scalar, true> : safe_scaling_operations<Scalar, true> {
639 private:
640 using Binary = binary_floating_point_traits<Scalar>;
641 using Bits = typename Binary::Bits;
642
643 template <bool RoundUp>
644 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE safe_scaling_factors<Scalar> compute_power_of_two_factors(
645 const Scalar& value) {
646 constexpr Bits kMaxScale = Binary::kMaxFiniteExponentBits - Binary::kExponentUnit;
647 Bits scaleBits = RoundUp ? numext::bit_cast<Bits>(numext::ceil_power_of_two(value))
648 : Binary::bits(value) & Binary::kExponentMask;
649 if (scaleBits < Binary::kExponentUnit) scaleBits = Binary::kExponentUnit;
650 if (scaleBits > kMaxScale) scaleBits = kMaxScale;
651 const Bits invScaleBits = Bits(Binary::kMaxFiniteExponentBits - scaleBits);
652 return {numext::bit_cast<Scalar>(scaleBits), numext::bit_cast<Scalar>(invScaleBits)};
653 }
654
655 public:
656 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE safe_scaling_factors<Scalar> compute_floor_factors(const Scalar& value) {
657 return compute_power_of_two_factors<false>(value);
658 }
659
660 EIGEN_DEVICE_FUNC static EIGEN_STRONG_INLINE safe_scaling_factors<Scalar> compute_ceiling_factors(
661 const Scalar& value) {
662 return compute_power_of_two_factors<true>(value);
663 }
664};
665
666#if !defined(EIGEN_GPU_COMPILE_PHASE)
667template <>
668struct safe_scaling<long double, true> : safe_scaling_operations<long double, true> {
669 static EIGEN_STRONG_INLINE safe_scaling_factors<long double> compute_floor_factors(const long double& value) {
670 EIGEN_USING_STD(frexp);
671 int exponent = 0;
672 frexp(value, &exponent);
673 return compute_factors_from_exponent(exponent - 1);
674 }
675
676 static EIGEN_STRONG_INLINE safe_scaling_factors<long double> compute_ceiling_factors(const long double& value) {
677 EIGEN_USING_STD(frexp);
678 int exponent = 0;
679 const long double fraction = frexp(value, &exponent);
680 if (numext::equal_strict(fraction, 0.5L)) --exponent;
681 return compute_factors_from_exponent(exponent);
682 }
683
684 private:
685 static EIGEN_STRONG_INLINE safe_scaling_factors<long double> compute_factors_from_exponent(int exponent) {
686 constexpr int kMinNormalExponent = std::numeric_limits<long double>::min_exponent - 1;
687 constexpr int kMaxNormalExponent = std::numeric_limits<long double>::max_exponent - 1;
688 constexpr int kMinScaleExponent =
689 kMinNormalExponent > -kMaxNormalExponent ? kMinNormalExponent : -kMaxNormalExponent;
690 constexpr int kMaxScaleExponent =
691 kMaxNormalExponent < -kMinNormalExponent ? kMaxNormalExponent : -kMinNormalExponent;
692 if (exponent < kMinScaleExponent) exponent = kMinScaleExponent;
693 if (exponent > kMaxScaleExponent) exponent = kMaxScaleExponent;
694
695 EIGEN_USING_STD(ldexp);
696 const long double scale = ldexp(1.0L, exponent);
697 return {scale, 1.0L / scale};
698 }
699};
700#endif
701
702} // namespace internal
703} // namespace Eigen
704
705#endif // EIGEN_SAFE_SCALING_H