Eigen  5.0.1
 
Loading...
Searching...
No Matches
ApproxComparisons.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2006-2008 Benoit Jacob <jacob.benoit.1@gmail.com>
5// Copyright (C) 2008 Gael Guennebaud <gael.guennebaud@inria.fr>
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_APPROX_COMPARISONS_H
13#define EIGEN_APPROX_COMPARISONS_H
14
15// IWYU pragma: private
16#include "./InternalHeaderCheck.h"
17
18namespace Eigen {
19
20namespace internal {
21
22// Keep custom scalars on the algebraic path: exponent scaling requires binary floating-point arithmetic.
23template <typename Scalar>
24struct use_scaled_comparison
25 : bool_constant<(std::is_floating_point<Scalar>::value && std::numeric_limits<Scalar>::radix == 2) ||
26 std::is_same<Scalar, half>::value || std::is_same<Scalar, bfloat16>::value> {};
27
28template <typename RealScalar>
29struct use_scaled_comparison<std::complex<RealScalar>> : use_scaled_comparison<RealScalar> {};
30
31// A nonnegative magnitude fraction * 2^exponent, including norms larger than the scalar range. Positive finite
32// magnitudes order lexicographically by (exponent, fraction), the member order.
33template <typename RealScalar>
34struct comparison_magnitude {
35 int exponent = 0;
36 RealScalar fraction;
37
38 EIGEN_DEVICE_FUNC explicit comparison_magnitude(const RealScalar& value) : fraction(value) {
39 if (value > RealScalar(0) && value <= NumTraits<RealScalar>::highest()) {
40 EIGEN_USING_STD(frexp);
41 fraction = frexp(value, &exponent);
42 }
43 }
44
45 template <typename OtherRealScalar>
46 EIGEN_DEVICE_FUNC explicit comparison_magnitude(const comparison_magnitude<OtherRealScalar>& value)
47 : exponent(value.exponent), fraction(RealScalar(value.fraction)) {}
48
49 EIGEN_DEVICE_FUNC void multiply(const RealScalar& value) {
50 const comparison_magnitude factor(value);
51 const comparison_magnitude product(fraction * factor.fraction);
52 fraction = product.fraction;
53 exponent += factor.exponent + product.exponent;
54 }
55
56 EIGEN_DEVICE_FUNC bool isFinite() const { return fraction <= RealScalar(1); }
57};
58
59template <typename RealScalar>
60EIGEN_DEVICE_FUNC bool operator<=(const comparison_magnitude<RealScalar>& x,
61 const comparison_magnitude<RealScalar>& y) {
62 if (!x.isFinite()) return false;
63 // Zero and non-finite magnitudes carry no exponent.
64 if (!y.isFinite() || x.fraction == RealScalar(0) || y.fraction == RealScalar(0)) return x.fraction <= y.fraction;
65 return x.exponent < y.exponent || (x.exponent == y.exponent && x.fraction <= y.fraction);
66}
67
68template <typename RealScalar, typename OtherRealScalar>
69EIGEN_DEVICE_FUNC bool operator<=(const comparison_magnitude<RealScalar>& x,
70 const comparison_magnitude<OtherRealScalar>& y) {
71 using Common = std::common_type_t<RealScalar, OtherRealScalar>;
72 return comparison_magnitude<Common>(x) <= comparison_magnitude<Common>(y);
73}
74
75// Types whose data(), rowStride() and colStride() address every coefficient. DirectAccessBit alone does not promise
76// this: a single-row IndexedView of a column-major matrix and the RealView of a strided complex block report strides
77// that do not.
78template <typename X>
79struct scaled_comparison_viewable : std::false_type {};
80
81template <typename Scalar, int Rows, int Cols, int Options, int MaxRows, int MaxCols>
82struct scaled_comparison_viewable<Matrix<Scalar, Rows, Cols, Options, MaxRows, MaxCols>> : std::true_type {};
83
84template <typename Scalar, int Rows, int Cols, int Options, int MaxRows, int MaxCols>
85struct scaled_comparison_viewable<Array<Scalar, Rows, Cols, Options, MaxRows, MaxCols>> : std::true_type {};
86
87template <typename PlainObjectType, int Options, typename StrideType>
88struct scaled_comparison_viewable<Map<PlainObjectType, Options, StrideType>> : std::true_type {};
89
90template <typename PlainObjectType, int Options, typename StrideType>
91struct scaled_comparison_viewable<Ref<PlainObjectType, Options, StrideType>> : std::true_type {};
92
93template <typename XprType, int BlockRows, int BlockCols, bool InnerPanel>
94struct scaled_comparison_viewable<Block<XprType, BlockRows, BlockCols, InnerPanel>>
95 : scaled_comparison_viewable<std::remove_const_t<XprType>> {};
96
97template <typename XprType>
98struct scaled_comparison_viewable<Transpose<XprType>> : scaled_comparison_viewable<std::remove_const_t<XprType>> {};
99
100template <typename XprType>
101struct scaled_comparison_viewable<ArrayWrapper<XprType>> : scaled_comparison_viewable<std::remove_const_t<XprType>> {};
102
103template <typename XprType>
104struct scaled_comparison_viewable<MatrixWrapper<XprType>> : scaled_comparison_viewable<std::remove_const_t<XprType>> {};
105
106// The scaled path reduces each operand several times, so it runs on at most three view types per scalar type rather
107// than being instantiated for every operand expression. A viewable operand with direct access is viewed in place,
108// keeping packet access when its inner stride is one; any other operand, including a lazy expression or a shape that
109// hides operator*(Scalar) as Homogeneous does, is evaluated first, which allocates only for dynamic sizes and for
110// fixed sizes above EIGEN_STACK_ALLOCATION_LIMIT.
111template <typename X, bool DirectAccess = has_direct_access<X>::value && scaled_comparison_viewable<X>::value,
112 bool UnitInnerStride = inner_stride_at_compile_time<X>::value == 1>
113struct scaled_comparison_operand {
114 using View =
115 Map<const Matrix<typename X::Scalar, Dynamic, Dynamic, X::IsRowMajor ? RowMajor : ColMajor>, 0, OuterStride<>>;
116 EIGEN_DEVICE_FUNC explicit scaled_comparison_operand(const X& x)
117 : view(x.data(), x.rows(), x.cols(), OuterStride<>(x.outerStride())) {}
118 View view;
119};
120
121template <typename X>
122struct scaled_comparison_operand<X, true, false> {
123 using View = Map<const Matrix<typename X::Scalar, Dynamic, Dynamic>, 0, Stride<Dynamic, Dynamic>>;
124 EIGEN_DEVICE_FUNC explicit scaled_comparison_operand(const X& x)
125 : view(x.data(), x.rows(), x.cols(), Stride<Dynamic, Dynamic>(x.colStride(), x.rowStride())) {}
126 View view;
127};
128
129template <typename X, bool UnitInnerStride>
130struct scaled_comparison_operand<X, false, UnitInnerStride> {
131 // A fixed capacity the stack limit would reject makes the comparison fail to compile for an operand the caller
132 // never stores, so the temporary takes it only where DenseStorage accepts it.
133 static constexpr bool FixedCapacity =
134 EIGEN_STACK_ALLOCATION_LIMIT == 0 ||
135 (X::MaxSizeAtCompileTime != Dynamic &&
136 std::ptrdiff_t(X::MaxSizeAtCompileTime) * std::ptrdiff_t(sizeof(typename X::Scalar)) <=
137 std::ptrdiff_t(EIGEN_STACK_ALLOCATION_LIMIT));
138 // Column-major unless it holds at most one row, where Matrix requires row-major storage.
139 using Plain =
140 Matrix<typename X::Scalar, Dynamic, Dynamic,
141 X::MaxRowsAtCompileTime == 1 && X::MaxColsAtCompileTime != 1 ? RowMajor : ColMajor,
142 FixedCapacity ? X::MaxRowsAtCompileTime : Dynamic, FixedCapacity ? X::MaxColsAtCompileTime : Dynamic>;
143 using View = typename scaled_comparison_operand<Plain>::View;
144 EIGEN_DEVICE_FUNC explicit scaled_comparison_operand(const X& x)
145 : value(x.matrix()), view(scaled_comparison_operand<Plain>(value).view) {}
146 Plain value;
147 View view;
148};
149
150template <typename Components>
151EIGEN_DEVICE_FUNC typename Components::Scalar scaled_comparison_max_coeff(const Components& components) {
152 using RealScalar = typename Components::Scalar;
153 if (components.size() == 0) return RealScalar(0);
154 return safe_scaling<RealScalar>::recover_flushed_max_coeff(components,
155 components.cwiseAbs().template maxCoeff<PropagateNaN>());
156}
157
158template <typename Derived>
159EIGEN_DEVICE_FUNC comparison_magnitude<typename stable_norm_accumulator<typename Derived::RealScalar>::type>
160scaled_comparison_norm_impl(const MatrixBase<Derived>& matrix) {
161 using RealScalar = typename stable_norm_accumulator<typename Derived::RealScalar>::type;
162 const auto& realComponents = matrix.realView();
163 const auto& components = realComponents.template cast<RealScalar>();
164 const RealScalar scale = scaled_comparison_max_coeff(components);
165 // Classify first so NaNs do not reach the ordered comparison.
166 if (!(numext::isfinite)(scale) || !(scale > RealScalar(0))) return comparison_magnitude<RealScalar>(scale);
167 RealScalar squaredNorm = RealScalar(0);
168 const auto factors = safe_scaling<RealScalar>::with_scaled(
169 components, scale, [&](const auto& scaled) { squaredNorm = scaled.squaredNorm(); });
170 comparison_magnitude<RealScalar> result(numext::sqrt(squaredNorm));
171 result.multiply(factors.scale);
172 return result;
173}
174
175template <typename X, typename Y>
176EIGEN_DEVICE_FUNC comparison_magnitude<typename stable_norm_accumulator<typename X::RealScalar>::type>
177scaled_comparison_distance_impl(const MatrixBase<X>& matrixX, const MatrixBase<Y>& matrixY) {
178 using Accumulator = typename stable_norm_accumulator<typename X::RealScalar>::type;
179 using WideScalar =
180 std::conditional_t<NumTraits<typename X::Scalar>::IsComplex || NumTraits<typename Y::Scalar>::IsComplex,
181 std::complex<Accumulator>, Accumulator>;
182 const auto& wideX = matrixX.template cast<WideScalar>();
183 const auto& wideY = matrixY.template cast<WideScalar>();
184 // FTZ flushes a subnormal difference of normal operands. Scaling a maximum M < 1 up by a power of two first loses
185 // only differences below M * min; scaling larger operands down could underflow their small components.
186 const Accumulator maxCoeff = numext::mini(
187 Accumulator(1),
188 numext::maxi(scaled_comparison_max_coeff(wideX.realView()), scaled_comparison_max_coeff(wideY.realView())));
189 const safe_scaling_factors<Accumulator> factors = supports_power_of_two_scaling<Accumulator>::value
190 ? safe_scaling<Accumulator>::compute_floor_factors(maxCoeff)
191 : safe_scaling_factors<Accumulator>();
192 // Both calls share one expression type, so scaled_comparison_norm_impl is instantiated once.
193 auto difference = scaled_comparison_norm_impl(wideX * factors.invScale - wideY * factors.invScale);
194 if (!difference.isFinite()) {
195 // Finite operands can overflow on subtraction; halving first keeps every component representable.
196 const Accumulator halfInvScale = factors.invScale * Accumulator(0.5);
197 difference = scaled_comparison_norm_impl(wideX * halfInvScale - wideY * halfInvScale);
198 difference.multiply(Accumulator(2));
199 }
200 difference.multiply(factors.scale);
201 return difference;
202}
203
204template <typename Derived>
205EIGEN_DEVICE_FUNC comparison_magnitude<typename stable_norm_accumulator<typename Derived::RealScalar>::type>
206scaled_comparison_norm(const Derived& x) {
207 return scaled_comparison_norm_impl(scaled_comparison_operand<Derived>(x).view);
208}
209
210template <typename X, typename Y>
211EIGEN_DEVICE_FUNC comparison_magnitude<typename stable_norm_accumulator<typename X::RealScalar>::type>
212scaled_comparison_distance(const X& x, const Y& y) {
213 return scaled_comparison_distance_impl(scaled_comparison_operand<X>(x).view, scaled_comparison_operand<Y>(y).view);
214}
215
216// Coefficients widened to the stable-norm accumulator, in which ordinary comparisons square and sum. The cast is the
217// identity for float and double; half and bfloat16 widen exactly to float, where no half square underflows.
218template <typename X>
219using approx_comparison_wide_t =
220 std::conditional_t<NumTraits<typename X::Scalar>::IsComplex,
221 std::complex<typename stable_norm_accumulator<typename X::RealScalar>::type>,
222 typename stable_norm_accumulator<typename X::RealScalar>::type>;
223
224template <typename Scalar, bool = use_scaled_comparison<Scalar>::value>
225struct approx_comparison_impl {
226 using RealScalar = typename NumTraits<Scalar>::Real;
227
228 template <typename X, typename Y>
229 EIGEN_DEVICE_FUNC static bool isApprox(const X& x, const Y& y, const RealScalar& prec) {
230 return (x.matrix() - y.matrix()).cwiseAbs2().sum() <=
231 prec * prec * numext::mini(x.cwiseAbs2().sum(), y.cwiseAbs2().sum());
232 }
233
234 template <typename X, typename Y>
235 EIGEN_DEVICE_FUNC static bool isMuchSmallerThan(const X& x, const Y& y, const RealScalar& prec) {
236 return x.cwiseAbs2().sum() <= numext::abs2(prec) * y.cwiseAbs2().sum();
237 }
238
239 template <typename X>
240 EIGEN_DEVICE_FUNC static bool isMuchSmallerThan(const X& x, const RealScalar& y, const RealScalar& prec) {
241 return x.cwiseAbs2().sum() <= numext::abs2(prec * y);
242 }
243};
244
245template <typename Scalar>
246struct approx_comparison_impl<Scalar, true> {
247 using RealScalar = typename NumTraits<Scalar>::Real;
248 using Accumulator = typename stable_norm_accumulator<RealScalar>::type;
249 template <typename Y>
250 using CommonAccumulator =
251 std::common_type_t<Accumulator, typename stable_norm_accumulator<typename Y::RealScalar>::type>;
252
253 template <typename ValueScalar>
254 EIGEN_DEVICE_FUNC static typename stable_norm_accumulator<ValueScalar>::type squared_norm_lower_bound(Index size) {
255 using ValueAccumulator = typename stable_norm_accumulator<ValueScalar>::type;
256 // Squares accumulate in ValueAccumulator; below n * min / epsilon, flushed squares can affect the comparison.
257 return stable_normalization_normal_min<ValueAccumulator, ValueAccumulator>::run() /
258 NumTraits<ValueAccumulator>::epsilon() * ValueAccumulator(size);
259 }
260
261 template <typename BoundScalar>
262 EIGEN_DEVICE_FUNC static bool safe_squared_norm(const BoundScalar& value, Index size) {
263 using Common = std::common_type_t<Accumulator, BoundScalar>;
264 return Common(value) >= Common(squared_norm_lower_bound<RealScalar>(size)) &&
265 Common(value) <= Common(NumTraits<Accumulator>::highest());
266 }
267
268 template <typename X, typename Y>
269 EIGEN_DEVICE_FUNC static bool isApprox(const X& x, const Y& y, const RealScalar& prec) {
270 // Widening must not admit operands that cannot be subtracted, such as half and float.
271 EIGEN_CHECK_BINARY_COMPATIBILITY(scalar_difference_op<typename X::Scalar EIGEN_COMMA typename Y::Scalar>,
272 typename X::Scalar, typename Y::Scalar)
273 const auto& wideX = x.template cast<approx_comparison_wide_t<X>>();
274 const auto& wideY = y.template cast<approx_comparison_wide_t<Y>>();
275 const Accumulator x2 = wideX.cwiseAbs2().sum();
276 const Accumulator y2 = wideY.cwiseAbs2().sum();
277 const Accumulator minimum = numext::mini(x2, y2);
278 const Accumulator precision2 = Accumulator(prec) * Accumulator(prec);
279 const Accumulator bound = precision2 * minimum;
280 // Only the smaller norm enters the bound; overflow of the larger norm is harmless.
281 if (safe_squared_norm(bound, x.size()) && minimum >= squared_norm_lower_bound<RealScalar>(x.size()) &&
282 precision2 >= squared_norm_lower_bound<RealScalar>(1))
283 return (wideX.matrix() - wideY.matrix()).cwiseAbs2().sum() <= bound;
284
285 return isApprox_scaled(x, y, prec);
286 }
287
288 template <typename X, typename Y>
289 EIGEN_DEVICE_FUNC static bool isMuchSmallerThan(const X& x, const Y& y, const RealScalar& prec) {
290 using Common = CommonAccumulator<Y>;
291 typename nested_eval<X, 2>::type nested(x);
292 typename nested_eval<Y, 2>::type otherNested(y);
293 const Accumulator x2 = nested.template cast<approx_comparison_wide_t<X>>().cwiseAbs2().sum();
294 const auto y2 = otherNested.template cast<approx_comparison_wide_t<Y>>().cwiseAbs2().sum();
295 const Accumulator precision2 = numext::abs2(Accumulator(prec));
296 const Common bound = Common(precision2) * Common(y2);
297 // A finite bound above the flushing error makes overflow/underflow of x2 harmless.
298 if (safe_squared_norm(bound, x.size()) &&
299 Common(y2) >= Common(squared_norm_lower_bound<typename Y::RealScalar>(y.size())) &&
300 precision2 >= squared_norm_lower_bound<RealScalar>(1))
301 return Common(x2) <= bound;
302 return isMuchSmallerThan_scaled(nested, otherNested, prec);
303 }
304
305 template <typename X>
306 EIGEN_DEVICE_FUNC static bool isMuchSmallerThan(const X& x, const RealScalar& y, const RealScalar& prec) {
307 typename nested_eval<X, 2>::type nested(x);
308 const Accumulator x2 = nested.template cast<approx_comparison_wide_t<X>>().cwiseAbs2().sum();
309 const Accumulator bound = numext::abs2(Accumulator(prec) * Accumulator(y));
310 if (safe_squared_norm(bound, x.size())) return x2 <= bound;
311 return isMuchSmallerThan_scaled(nested, y, prec);
312 }
313
314 private:
315 // Keep exponent scaling from inhibiting inlining of ordinary comparisons.
316 template <typename X, typename Y>
317 EIGEN_DEVICE_FUNC static EIGEN_DONT_INLINE bool isApprox_scaled(const X& xExpr, const Y& yExpr,
318 const RealScalar& prec) {
319 const scaled_comparison_operand<X> x(xExpr);
320 const scaled_comparison_operand<Y> y(yExpr);
321 const auto nx = scaled_comparison_norm_impl(x.view);
322 const auto ny = scaled_comparison_norm_impl(y.view);
323 if (!nx.isFinite() || !ny.isFinite()) return false;
324 auto tolerance = nx <= ny ? nx : ny;
325 tolerance.multiply(numext::abs(Accumulator(prec)));
326 return scaled_comparison_distance_impl(x.view, y.view) <= tolerance;
327 }
328
329 template <typename X, typename Y>
330 EIGEN_DEVICE_FUNC static EIGEN_DONT_INLINE bool isMuchSmallerThan_scaled(const X& x, const Y& y,
331 const RealScalar& prec) {
332 using Common = CommonAccumulator<Y>;
333 comparison_magnitude<Common> tolerance(scaled_comparison_norm(y));
334 tolerance.multiply(numext::abs(Common(prec)));
335 return scaled_comparison_norm(x) <= tolerance;
336 }
337
338 template <typename X>
339 EIGEN_DEVICE_FUNC static EIGEN_DONT_INLINE bool isMuchSmallerThan_scaled(const X& x, const RealScalar& y,
340 const RealScalar& prec) {
341 comparison_magnitude<Accumulator> tolerance(numext::abs(Accumulator(y)));
342 tolerance.multiply(numext::abs(Accumulator(prec)));
343 return scaled_comparison_norm(x) <= tolerance;
344 }
345};
346
347// Exponent scaling needs binary floating-point on both sides.
348template <typename Derived, typename OtherDerived>
349using approx_comparison_impl_t =
350 approx_comparison_impl<typename Derived::Scalar, use_scaled_comparison<typename Derived::Scalar>::value &&
351 use_scaled_comparison<typename OtherDerived::Scalar>::value>;
352
353template <typename Derived, typename OtherDerived, bool is_integer = NumTraits<typename Derived::Scalar>::IsInteger>
354struct isApprox_selector {
355 EIGEN_DEVICE_FUNC static bool run(const Derived& x, const OtherDerived& y, const typename Derived::RealScalar& prec) {
356 typename internal::nested_eval<Derived, 2>::type nested(x);
357 typename internal::nested_eval<OtherDerived, 2>::type otherNested(y);
358 return approx_comparison_impl_t<Derived, OtherDerived>::isApprox(nested, otherNested, prec);
359 }
360};
361
362template <typename Derived, typename OtherDerived>
363struct isApprox_selector<Derived, OtherDerived, true> {
364 EIGEN_DEVICE_FUNC static bool run(const Derived& x, const OtherDerived& y, const typename Derived::RealScalar&) {
365 return x.matrix() == y.matrix();
366 }
367};
368
369template <typename Derived, typename OtherDerived, bool is_integer = NumTraits<typename Derived::Scalar>::IsInteger>
370struct isMuchSmallerThan_object_selector {
371 EIGEN_DEVICE_FUNC static bool run(const Derived& x, const OtherDerived& y, const typename Derived::RealScalar& prec) {
372 return approx_comparison_impl_t<Derived, OtherDerived>::isMuchSmallerThan(x, y, prec);
373 }
374};
375
376template <typename Derived, typename OtherDerived>
377struct isMuchSmallerThan_object_selector<Derived, OtherDerived, true> {
378 EIGEN_DEVICE_FUNC static bool run(const Derived& x, const OtherDerived&, const typename Derived::RealScalar&) {
379 return x.matrix() == Derived::Zero(x.rows(), x.cols()).matrix();
380 }
381};
382
383template <typename Derived, bool is_integer = NumTraits<typename Derived::Scalar>::IsInteger>
384struct isMuchSmallerThan_scalar_selector {
385 EIGEN_DEVICE_FUNC static bool run(const Derived& x, const typename Derived::RealScalar& y,
386 const typename Derived::RealScalar& prec) {
387 return approx_comparison_impl<typename Derived::Scalar>::isMuchSmallerThan(x, y, prec);
388 }
389};
390
391template <typename Derived>
392struct isMuchSmallerThan_scalar_selector<Derived, true>
393 : isMuchSmallerThan_object_selector<Derived, typename Derived::RealScalar, true> {};
394
395} // end namespace internal
396
419template <typename Derived>
420template <typename OtherDerived>
421EIGEN_DEVICE_FUNC constexpr bool DenseBase<Derived>::isApprox(const DenseBase<OtherDerived>& other,
422 const RealScalar& prec) const {
423 return internal::isApprox_selector<Derived, OtherDerived>::run(derived(), other.derived(), prec);
424}
425
439template <typename Derived>
440EIGEN_DEVICE_FUNC constexpr bool DenseBase<Derived>::isMuchSmallerThan(const typename NumTraits<Scalar>::Real& other,
441 const RealScalar& prec) const {
442 return internal::isMuchSmallerThan_scalar_selector<Derived>::run(derived(), other, prec);
443}
444
455template <typename Derived>
456template <typename OtherDerived>
457EIGEN_DEVICE_FUNC constexpr bool DenseBase<Derived>::isMuchSmallerThan(const DenseBase<OtherDerived>& other,
458 const RealScalar& prec) const {
459 return internal::isMuchSmallerThan_object_selector<Derived, OtherDerived>::run(derived(), other.derived(), prec);
460}
461
462} // end namespace Eigen
463
464#endif // EIGEN_APPROX_COMPARISONS_H
constexpr bool isApprox(const DenseBase< OtherDerived > &other, const RealScalar &prec=NumTraits< Scalar >::dummy_precision()) const
Definition ApproxComparisons.h:421
constexpr DenseBase()=default
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321