Eigen  5.0.1
 
Loading...
Searching...
No Matches
Visitor.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008 Gael Guennebaud <gael.guennebaud@inria.fr>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_VISITOR_H
12#define EIGEN_VISITOR_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21// Opt in only when coefficient updates are valid without init/initpacket.
22template <typename Visitor, typename = void>
23struct visitor_already_initialized : std::false_type {};
24template <typename Visitor>
25struct visitor_already_initialized<Visitor, void_t<decltype(functor_traits<Visitor>::AlreadyInitialized)>>
26 : bool_constant<functor_traits<Visitor>::AlreadyInitialized> {};
27
28template <typename Visitor, typename Derived, int UnrollCount,
29 bool Vectorize = (Derived::PacketAccess && functor_traits<Visitor>::PacketAccess), bool LinearAccess = false,
30 bool ShortCircuitEvaluation = false>
31struct visitor_impl;
32
33template <typename Visitor, bool ShortCircuitEvaluation = false>
34struct short_circuit_eval_impl {
35 // if short circuit evaluation is not used, do nothing
36 static constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const Visitor&) { return false; }
37};
38template <typename Visitor>
39struct short_circuit_eval_impl<Visitor, true> {
40 // if short circuit evaluation is used, check the visitor
41 static constexpr EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool run(const Visitor& visitor) { return visitor.done(); }
42};
43
44// unrolled inner-outer traversal
45template <typename Visitor, typename Derived, int UnrollCount, bool Vectorize, bool ShortCircuitEvaluation>
46struct visitor_impl<Visitor, Derived, UnrollCount, Vectorize, false, ShortCircuitEvaluation> {
47 // don't use short circuit evaluation for unrolled version
48 using Scalar = typename Derived::Scalar;
49 using Packet = typename packet_traits<Scalar>::type;
50 static constexpr bool RowMajor = Derived::IsRowMajor;
51 static constexpr int RowsAtCompileTime = Derived::RowsAtCompileTime;
52 static constexpr int ColsAtCompileTime = Derived::ColsAtCompileTime;
53 static constexpr int PacketSize = packet_traits<Scalar>::size;
54 static constexpr int InnerSizeAtCompileTime = RowMajor ? ColsAtCompileTime : RowsAtCompileTime;
55 static constexpr int OuterSizeAtCompileTime = RowMajor ? RowsAtCompileTime : ColsAtCompileTime;
56 static constexpr int PacketOpsPerOuter =
57 Vectorize && InnerSizeAtCompileTime >= PacketSize ? InnerSizeAtCompileTime / PacketSize : 0;
58 static constexpr int FirstScalarInner = PacketOpsPerOuter * PacketSize;
59 static constexpr int ScalarOpsPerOuter = InnerSizeAtCompileTime - FirstScalarInner;
60 static constexpr int OpsPerOuter = PacketOpsPerOuter + ScalarOpsPerOuter;
61 static constexpr int OpCount = UnrollCount == 0 ? 0 : OuterSizeAtCompileTime * OpsPerOuter;
62
63 template <int Op>
64 static constexpr bool IsPacketOp() {
65 return OpsPerOuter != 0 && (Op % OpsPerOuter) < PacketOpsPerOuter;
66 }
67
68 template <int Op>
69 static constexpr int CoeffIndex() {
70 return OpsPerOuter == 0 ? 0
71 : (Op / OpsPerOuter) * InnerSizeAtCompileTime +
72 (IsPacketOp<Op>() ? ((Op % OpsPerOuter) * PacketSize)
73 : (FirstScalarInner + (Op % OpsPerOuter) - PacketOpsPerOuter));
74 }
75
76 template <int Op, std::enable_if_t<!IsPacketOp<Op>(), bool> = true>
77 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void visit(const Derived& mat, Visitor& visitor) {
78 constexpr int K = CoeffIndex<Op>();
79 constexpr int R = RowMajor ? (K / ColsAtCompileTime) : (K % RowsAtCompileTime);
80 constexpr int C = RowMajor ? (K % ColsAtCompileTime) : (K / RowsAtCompileTime);
81 EIGEN_IF_CONSTEXPR (Op == 0) {
82 visitor.init(mat.coeff(R, C), R, C);
83 } else {
84 visitor(mat.coeff(R, C), R, C);
85 }
86 }
87
88 template <int Op, std::enable_if_t<IsPacketOp<Op>(), bool> = true>
89 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void visit(const Derived& mat, Visitor& visitor) {
90 constexpr int K = CoeffIndex<Op>();
91 constexpr int R = RowMajor ? (K / ColsAtCompileTime) : (K % RowsAtCompileTime);
92 constexpr int C = RowMajor ? (K % ColsAtCompileTime) : (K / RowsAtCompileTime);
93 Packet P = mat.template packet<Packet>(R, C);
94 EIGEN_IF_CONSTEXPR (Op == 0) {
95 visitor.initpacket(P, R, C);
96 } else {
97 visitor.packet(P, R, C);
98 }
99 }
100
101 template <int... Ops>
102 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run_impl(const Derived& mat, Visitor& visitor,
103 std::integer_sequence<int, Ops...>) {
104 int unused[] = {0, (visit<Ops>(mat, visitor), 0)...};
105 EIGEN_UNUSED_VARIABLE(unused);
106 }
107
108 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const Derived& mat, Visitor& visitor) {
109 run_impl(mat, visitor, std::make_integer_sequence<int, OpCount>{});
110 }
111};
112
113// unrolled linear traversal
114template <typename Visitor, typename Derived, int UnrollCount, bool Vectorize, bool ShortCircuitEvaluation>
115struct visitor_impl<Visitor, Derived, UnrollCount, Vectorize, true, ShortCircuitEvaluation> {
116 // don't use short circuit evaluation for unrolled version
117 using Scalar = typename Derived::Scalar;
118 using Packet = typename packet_traits<Scalar>::type;
119 static constexpr int PacketSize = packet_traits<Scalar>::size;
120 static constexpr int PacketOps = Vectorize ? UnrollCount / PacketSize : 0;
121 static constexpr int FirstScalar = PacketOps * PacketSize;
122 static constexpr int ScalarOps = UnrollCount - FirstScalar;
123 static constexpr int OpCount = PacketOps + ScalarOps;
124
125 template <int Op>
126 static constexpr bool IsPacketOp() {
127 return Op < PacketOps;
128 }
129
130 template <int Op>
131 static constexpr int CoeffIndex() {
132 return IsPacketOp<Op>() ? (Op * PacketSize) : (FirstScalar + Op - PacketOps);
133 }
134
135 template <int Op, std::enable_if_t<!IsPacketOp<Op>(), bool> = true>
136 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void visit(const Derived& mat, Visitor& visitor) {
137 constexpr int K = CoeffIndex<Op>();
138 EIGEN_IF_CONSTEXPR (Op == 0) {
139 visitor.init(mat.coeff(K), K);
140 } else {
141 visitor(mat.coeff(K), K);
142 }
143 }
144
145 template <int Op, std::enable_if_t<IsPacketOp<Op>(), bool> = true>
146 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void visit(const Derived& mat, Visitor& visitor) {
147 constexpr int K = CoeffIndex<Op>();
148 Packet P = mat.template packet<Packet>(K);
149 EIGEN_IF_CONSTEXPR (Op == 0) {
150 visitor.initpacket(P, K);
151 } else {
152 visitor.packet(P, K);
153 }
154 }
155
156 template <int... Ops>
157 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run_impl(const Derived& mat, Visitor& visitor,
158 std::integer_sequence<int, Ops...>) {
159 int unused[] = {0, (visit<Ops>(mat, visitor), 0)...};
160 EIGEN_UNUSED_VARIABLE(unused);
161 }
162
163 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const Derived& mat, Visitor& visitor) {
164 run_impl(mat, visitor, std::make_integer_sequence<int, OpCount>{});
165 }
166};
167
168// dynamic scalar outer-inner traversal
169template <typename Visitor, typename Derived, bool ShortCircuitEvaluation>
170struct visitor_impl<Visitor, Derived, Dynamic, /*Vectorize=*/false, /*LinearAccess=*/false, ShortCircuitEvaluation> {
171 using short_circuit = short_circuit_eval_impl<Visitor, ShortCircuitEvaluation>;
172 static constexpr bool RowMajor = Derived::IsRowMajor;
173
174 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const Derived& mat, Visitor& visitor) {
175 const Index innerSize = RowMajor ? mat.cols() : mat.rows();
176 const Index outerSize = RowMajor ? mat.rows() : mat.cols();
177 if (innerSize == 0 || outerSize == 0) return;
178 {
179 visitor.init(mat.coeff(0, 0), 0, 0);
180 if (short_circuit::run(visitor)) return;
181 for (Index i = 1; i < innerSize; ++i) {
182 Index r = RowMajor ? 0 : i;
183 Index c = RowMajor ? i : 0;
184 visitor(mat.coeff(r, c), r, c);
185 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
186 }
187 }
188 for (Index j = 1; j < outerSize; j++) {
189 for (Index i = 0; i < innerSize; ++i) {
190 Index r = RowMajor ? j : i;
191 Index c = RowMajor ? i : j;
192 visitor(mat.coeff(r, c), r, c);
193 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
194 }
195 }
196 }
197};
198
199// dynamic vectorized outer-inner traversal
200template <typename Visitor, typename Derived, bool ShortCircuitEvaluation>
201struct visitor_impl<Visitor, Derived, Dynamic, /*Vectorize=*/true, /*LinearAccess=*/false, ShortCircuitEvaluation> {
202 using Scalar = typename Derived::Scalar;
203 using Packet = typename packet_traits<Scalar>::type;
204 static constexpr int PacketSize = packet_traits<Scalar>::size;
205 using short_circuit = short_circuit_eval_impl<Visitor, ShortCircuitEvaluation>;
206 static constexpr bool RowMajor = Derived::IsRowMajor;
207 static constexpr bool AlreadyInitialized = visitor_already_initialized<Visitor>::value;
208
209 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const Derived& mat, Visitor& visitor) {
210 const Index innerSize = RowMajor ? mat.cols() : mat.rows();
211 const Index outerSize = RowMajor ? mat.rows() : mat.cols();
212 if (innerSize == 0 || outerSize == 0) return;
213 EIGEN_IF_CONSTEXPR (!AlreadyInitialized) {
214 Index i = 0;
215 if (innerSize < PacketSize) {
216 visitor.init(mat.coeff(0, 0), 0, 0);
217 i = 1;
218 } else {
219 Packet p = mat.template packet<Packet>(0, 0);
220 visitor.initpacket(p, 0, 0);
221 i = PacketSize;
222 }
223 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
224 for (; i + PacketSize - 1 < innerSize; i += PacketSize) {
225 Index r = RowMajor ? 0 : i;
226 Index c = RowMajor ? i : 0;
227 Packet p = mat.template packet<Packet>(r, c);
228 visitor.packet(p, r, c);
229 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
230 }
231 for (; i < innerSize; ++i) {
232 Index r = RowMajor ? 0 : i;
233 Index c = RowMajor ? i : 0;
234 visitor(mat.coeff(r, c), r, c);
235 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
236 }
237 }
238 const Index packetEnd = innerSize - innerSize % PacketSize;
239 for (Index j = AlreadyInitialized ? 0 : 1; j < outerSize; j++) {
240 Index i = 0;
241 for (; AlreadyInitialized ? i < packetEnd : i + PacketSize - 1 < innerSize; i += PacketSize) {
242 Index r = RowMajor ? j : i;
243 Index c = RowMajor ? i : j;
244 Packet p = mat.template packet<Packet>(r, c);
245 visitor.packet(p, r, c);
246 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
247 }
248 for (; i < innerSize; ++i) {
249 Index r = RowMajor ? j : i;
250 Index c = RowMajor ? i : j;
251 visitor(mat.coeff(r, c), r, c);
252 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
253 }
254 }
255 }
256};
257
258// dynamic scalar linear traversal
259template <typename Visitor, typename Derived, bool ShortCircuitEvaluation>
260struct visitor_impl<Visitor, Derived, Dynamic, /*Vectorize=*/false, /*LinearAccess=*/true, ShortCircuitEvaluation> {
261 using short_circuit = short_circuit_eval_impl<Visitor, ShortCircuitEvaluation>;
262
263 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const Derived& mat, Visitor& visitor) {
264 const Index size = mat.size();
265 if (size == 0) return;
266 visitor.init(mat.coeff(0), 0);
267 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
268 for (Index k = 1; k < size; k++) {
269 visitor(mat.coeff(k), k);
270 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
271 }
272 }
273};
274
275// dynamic vectorized linear traversal
276template <typename Visitor, typename Derived, bool ShortCircuitEvaluation>
277struct visitor_impl<Visitor, Derived, Dynamic, /*Vectorize=*/true, /*LinearAccess=*/true, ShortCircuitEvaluation> {
278 using Scalar = typename Derived::Scalar;
279 using Packet = typename packet_traits<Scalar>::type;
280 static constexpr int PacketSize = packet_traits<Scalar>::size;
281 using short_circuit = short_circuit_eval_impl<Visitor, ShortCircuitEvaluation>;
282 static constexpr bool AlreadyInitialized = visitor_already_initialized<Visitor>::value;
283
284 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const Derived& mat, Visitor& visitor) {
285 const Index size = mat.size();
286 if (size == 0) return;
287 const Index packetEnd = size - size % PacketSize;
288 Index k = 0;
289 EIGEN_IF_CONSTEXPR (!AlreadyInitialized) {
290 if (size < PacketSize) {
291 visitor.init(mat.coeff(0), 0);
292 k = 1;
293 } else {
294 Packet p = mat.template packet<Packet>(k);
295 visitor.initpacket(p, k);
296 k = PacketSize;
297 }
298 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
299 }
300 for (; k < packetEnd; k += PacketSize) {
301 Packet p = mat.template packet<Packet>(k);
302 visitor.packet(p, k);
303 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
304 }
305 for (; k < size; k++) {
306 visitor(mat.coeff(k), k);
307 if EIGEN_PREDICT_FALSE (short_circuit::run(visitor)) return;
308 }
309 }
310};
311
312// evaluator adaptor
313template <typename XprType>
314class visitor_evaluator {
315 public:
316 using Evaluator = evaluator<XprType>;
317 using Scalar = typename XprType::Scalar;
318 using Packet = typename packet_traits<Scalar>::type;
319 using CoeffReturnType = std::remove_const_t<typename XprType::CoeffReturnType>;
320
321 static constexpr bool PacketAccess = static_cast<bool>(Evaluator::Flags & PacketAccessBit);
322 static constexpr bool LinearAccess = static_cast<bool>(Evaluator::Flags & LinearAccessBit);
323 static constexpr bool IsRowMajor = static_cast<bool>(XprType::IsRowMajor);
324 static constexpr int RowsAtCompileTime = XprType::RowsAtCompileTime;
325 static constexpr int ColsAtCompileTime = XprType::ColsAtCompileTime;
326 static constexpr int XprAlignment = Evaluator::Alignment;
327 static constexpr int CoeffReadCost = Evaluator::CoeffReadCost;
328
329 EIGEN_DEVICE_FUNC explicit visitor_evaluator(const XprType& xpr) : m_evaluator(xpr), m_xpr(xpr) {}
330
331 EIGEN_DEVICE_FUNC constexpr Index rows() const noexcept { return m_xpr.rows(); }
332 EIGEN_DEVICE_FUNC constexpr Index cols() const noexcept { return m_xpr.cols(); }
333 EIGEN_DEVICE_FUNC constexpr Index size() const noexcept { return m_xpr.size(); }
334 // outer-inner access
335 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index row, Index col) const {
336 return m_evaluator.coeff(row, col);
337 }
338 template <typename Packet, int Alignment = Unaligned>
339 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packet(Index row, Index col) const {
340 return m_evaluator.template packet<Alignment, Packet>(row, col);
341 }
342 // linear access
343 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const { return m_evaluator.coeff(index); }
344 template <typename Packet, int Alignment = XprAlignment>
345 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packet(Index index) const {
346 return m_evaluator.template packet<Alignment, Packet>(index);
347 }
348
349 protected:
350 Evaluator m_evaluator;
351 const XprType& m_xpr;
352};
353
354template <typename T, typename = void>
355struct visitor_has_linear_access : std::false_type {};
356
357template <typename T>
358struct visitor_has_linear_access<T, void_t<decltype(functor_traits<T>::LinearAccess)>>
359 : bool_constant<static_cast<bool>(functor_traits<T>::LinearAccess)> {};
360
361template <typename Derived, typename Visitor, bool ShortCircuitEvaluation>
362struct visit_impl {
363 using Evaluator = visitor_evaluator<Derived>;
364 using Scalar = typename DenseBase<Derived>::Scalar;
365
370 static constexpr int InnerSizeAtCompileTime = IsRowMajor ? ColsAtCompileTime : RowsAtCompileTime;
371 static constexpr int OuterSizeAtCompileTime = IsRowMajor ? RowsAtCompileTime : ColsAtCompileTime;
372
373 // Linear packets can make an early scalar short-circuit exit more expensive.
374 // Preinitialized visitors opt into starting directly with packet traversal.
375 static constexpr bool LinearAccess = (!ShortCircuitEvaluation || visitor_already_initialized<Visitor>::value) &&
376 Evaluator::LinearAccess && visitor_has_linear_access<Visitor>::value;
377 static constexpr bool Vectorize = Evaluator::PacketAccess && static_cast<bool>(functor_traits<Visitor>::PacketAccess);
378
379 static constexpr int PacketSize = packet_traits<Scalar>::size;
380 static constexpr int VectorOps =
381 Vectorize ? (LinearAccess ? (SizeAtCompileTime / PacketSize)
382 : (OuterSizeAtCompileTime * (InnerSizeAtCompileTime / PacketSize)))
383 : 0;
384 static constexpr int ScalarOps = SizeAtCompileTime - (VectorOps * PacketSize);
385 // treat vector op and scalar op as same cost for unroll logic
386 static constexpr int TotalOps = VectorOps + ScalarOps;
387
388 static constexpr int UnrollCost = int(Evaluator::CoeffReadCost) + int(functor_traits<Visitor>::Cost);
389 static constexpr bool Unroll = (SizeAtCompileTime != Dynamic) && ((TotalOps * UnrollCost) <= EIGEN_UNROLLING_LIMIT);
390 static constexpr int UnrollCount = Unroll ? int(SizeAtCompileTime) : Dynamic;
391
392 using impl = visitor_impl<Visitor, Evaluator, UnrollCount, Vectorize, LinearAccess, ShortCircuitEvaluation>;
393
394 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void run(const DenseBase<Derived>& mat, Visitor& visitor) {
395 Evaluator evaluator(mat.derived());
396 impl::run(evaluator, visitor);
397 }
398};
399
400} // end namespace internal
401
421template <typename Derived>
422template <typename Visitor>
423EIGEN_DEVICE_FUNC void DenseBase<Derived>::visit(Visitor& visitor) const {
424 using impl = internal::visit_impl<Derived, Visitor, /*ShortCircuitEvaluation*/ false>;
425 impl::run(derived(), visitor);
426}
427
428namespace internal {
429
430template <typename Scalar, bool Approximate>
431struct fuzzy_constant_visitor {
432 Scalar value;
433 typename NumTraits<Scalar>::Real precision;
434 bool result = true;
435
436 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void operator()(const Scalar& x, Index, Index = 0) {
437 result = result && (Approximate ? internal::isApprox(x, value, precision)
438 : internal::isMuchSmallerThan(x, Scalar(1), precision));
439 }
440 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void init(const Scalar& x, Index r, Index c = 0) { (*this)(x, r, c); }
441 EIGEN_DEVICE_FUNC bool done() const { return !result; }
442 template <typename Packet>
443 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void packet(const Packet& x, Index, Index = 0) {
444 const Packet p = pset1<Packet>(precision);
445 Packet mask;
446 EIGEN_IF_CONSTEXPR (Approximate) {
447 const Packet v = pset1<Packet>(value);
448 mask = pcmp_le(pabs(psub(x, v)), pmul(pmin(pabs(x), pabs(v)), p));
449 } else {
450 mask = pcmp_le(pabs(x), p);
451 }
452 // Reduce the comparison mask directly, without comparing its lanes to zero again.
453 result = result && !predux_any(pandnot(ptrue(mask), mask));
454 }
455 template <typename Packet>
456 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void initpacket(const Packet& x, Index r, Index c = 0) {
457 packet(x, r, c);
458 }
459};
460
461template <typename Scalar, bool Approximate>
462struct functor_traits<fuzzy_constant_visitor<Scalar, Approximate>> {
463 static constexpr bool AlreadyInitialized = true;
464 static constexpr bool LinearAccess = true;
465 static constexpr int Cost = 4 * NumTraits<Scalar>::AddCost + NumTraits<Scalar>::MulCost;
466 // ARMv7 NEON flushes subnormal float operands and results; scalar VFP does not.
467 // Explicitly flushing would change exact-zero checks; normal operands can also have a subnormal difference.
468 static constexpr bool PacketAccess =
469 (std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value) &&
470 !(EIGEN_ARCH_ARM && std::is_same<Scalar, float>::value) && packet_traits<Scalar>::HasAbs &&
471 packet_traits<Scalar>::HasCmp &&
472 (!Approximate ||
473 (packet_traits<Scalar>::HasSub && packet_traits<Scalar>::HasMin && packet_traits<Scalar>::HasMul));
474};
475
476template <typename Scalar>
477using use_fuzzy_constant_visitor =
478 bool_constant<std::is_same<Scalar, float>::value || std::is_same<Scalar, double>::value>;
479
480template <bool Approximate, typename Derived,
481 std::enable_if_t<use_fuzzy_constant_visitor<typename Derived::Scalar>::value, int> = 0>
482EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool fuzzy_constant_all(const Derived& matrix,
483 const typename Derived::Scalar& value,
484 const typename Derived::RealScalar& precision) {
485 fuzzy_constant_visitor<typename Derived::Scalar, Approximate> visitor{value, precision};
486 visit_impl<Derived, decltype(visitor), true>::run(matrix, visitor);
487 return visitor.result;
488}
489
490// Separate overloads keep custom scalars from instantiating unused comparisons in C++14.
491template <bool Approximate, typename Derived,
492 std::enable_if_t<!use_fuzzy_constant_visitor<typename Derived::Scalar>::value && Approximate, int> = 0>
493EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool fuzzy_constant_all(const Derived& matrix,
494 const typename Derived::Scalar& value,
495 const typename Derived::RealScalar& precision) {
496 for (Index j = 0; j < matrix.cols(); ++j)
497 for (Index i = 0; i < matrix.rows(); ++i)
498 if (!internal::isApprox(matrix.coeff(i, j), value, precision)) return false;
499 return true;
500}
501
502template <bool Approximate, typename Derived,
503 std::enable_if_t<!use_fuzzy_constant_visitor<typename Derived::Scalar>::value && !Approximate, int> = 0>
504EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool fuzzy_constant_all(const Derived& matrix, const typename Derived::Scalar&,
505 const typename Derived::RealScalar& precision) {
506 using Scalar = typename Derived::Scalar;
507 for (Index j = 0; j < matrix.cols(); ++j)
508 for (Index i = 0; i < matrix.rows(); ++i)
509 if (!internal::isMuchSmallerThan(matrix.coeff(i, j), Scalar(1), precision)) return false;
510 return true;
511}
512
513template <typename Scalar>
514struct all_visitor {
515 using result_type = bool;
516 using Packet = typename packet_traits<Scalar>::type;
517 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index, Index) { res = (value != Scalar(0)); }
518 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index) { res = (value != Scalar(0)); }
519 EIGEN_DEVICE_FUNC inline bool all_predux(const Packet& p) const { return predux_all(p); }
520 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index, Index) { res = all_predux(p); }
521 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index) { res = all_predux(p); }
522 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index, Index) { res = res && (value != Scalar(0)); }
523 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index) { res = res && (value != Scalar(0)); }
524 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index, Index) { res = res && all_predux(p); }
525 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index) { res = res && all_predux(p); }
526 EIGEN_DEVICE_FUNC inline bool done() const { return !res; }
527 bool res = true;
528};
529template <typename Scalar>
530struct functor_traits<all_visitor<Scalar>> {
531 enum { Cost = NumTraits<Scalar>::ReadCost, LinearAccess = true, PacketAccess = packet_traits<Scalar>::HasCmp };
532};
533
534template <typename Scalar>
535struct any_visitor {
536 using result_type = bool;
537 using Packet = typename packet_traits<Scalar>::type;
538 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index, Index) { res = (value != Scalar(0)); }
539 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index) { res = (value != Scalar(0)); }
540 EIGEN_DEVICE_FUNC inline bool any_predux(const Packet& p) const {
541 return predux_any(pandnot(ptrue(p), pcmp_eq(p, pzero(p))));
542 }
543 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index, Index) { res = any_predux(p); }
544 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index) { res = any_predux(p); }
545 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index, Index) { res = res || (value != Scalar(0)); }
546 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index) { res = res || (value != Scalar(0)); }
547 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index, Index) { res = res || any_predux(p); }
548 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index) { res = res || any_predux(p); }
549 EIGEN_DEVICE_FUNC inline bool done() const { return res; }
550 bool res = false;
551};
552template <>
553EIGEN_DEVICE_FUNC inline bool any_visitor<bool>::any_predux(const Packet& p) const {
554 return predux(p);
555}
556template <typename Scalar>
557struct functor_traits<any_visitor<Scalar>> {
558 enum { Cost = NumTraits<Scalar>::ReadCost, LinearAccess = true, PacketAccess = packet_traits<Scalar>::HasCmp };
559};
560
561template <typename Scalar>
562struct count_visitor {
563 using result_type = Index;
564 using Packet = typename packet_traits<Scalar>::type;
565 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index, Index) { res = value != Scalar(0) ? 1 : 0; }
566 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index) { res = value != Scalar(0) ? 1 : 0; }
567 EIGEN_DEVICE_FUNC inline Index count_redux(const Packet& p) const { return predux_count(p); }
568 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index, Index) { res = count_redux(p); }
569 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index) { res = count_redux(p); }
570 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index, Index) {
571 if (value != Scalar(0)) res++;
572 }
573 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index) {
574 if (value != Scalar(0)) res++;
575 }
576 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index, Index) { res += count_redux(p); }
577 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index) { res += count_redux(p); }
578 Index res = 0;
579};
580
581template <typename Scalar>
582struct functor_traits<count_visitor<Scalar>> {
583 enum {
584 Cost = NumTraits<Scalar>::AddCost,
585 LinearAccess = true,
586 PacketAccess = packet_traits<Scalar>::HasCmp && packet_traits<Scalar>::HasAdd
587 };
588};
589
590// Reduces pisfinite masks directly: their lanes are all-ones or zero, so the coefficients are all finite exactly when
591// no lane of the complement is set. predux_all, which all() on isFiniteTyped() uses, compares each lane against zero
592// instead, which a compiler can only fold away when it can prove the operand is a mask.
593template <typename Scalar>
594struct all_finite_visitor {
595 using result_type = bool;
596 using Packet = typename packet_traits<Scalar>::type;
597 EIGEN_DEVICE_FUNC inline bool finite_predux(const Packet& p) const { return !predux_any(pnot(pisfinite(p))); }
598 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index, Index) { res = (numext::isfinite)(value); }
599 EIGEN_DEVICE_FUNC inline void init(const Scalar& value, Index) { res = (numext::isfinite)(value); }
600 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index, Index) { res = finite_predux(p); }
601 EIGEN_DEVICE_FUNC inline void initpacket(const Packet& p, Index) { res = finite_predux(p); }
602 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index, Index) {
603 res = res && (numext::isfinite)(value);
604 }
605 EIGEN_DEVICE_FUNC inline void operator()(const Scalar& value, Index) { res = res && (numext::isfinite)(value); }
606 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index, Index) { res = res && finite_predux(p); }
607 EIGEN_DEVICE_FUNC inline void packet(const Packet& p, Index) { res = res && finite_predux(p); }
608 EIGEN_DEVICE_FUNC inline bool done() const { return !res; }
609 bool res = true;
610};
611template <typename Scalar>
612struct functor_traits<all_finite_visitor<Scalar>> {
613 enum {
614 Cost = NumTraits<Scalar>::ReadCost + NumTraits<Scalar>::MulCost,
615 LinearAccess = true,
616 PacketAccess = packet_traits<Scalar>::HasCmp
617 };
618};
619
620template <typename Derived, bool AlwaysTrue = NumTraits<typename traits<Derived>::Scalar>::IsInteger>
621struct all_finite_impl {
622 static EIGEN_DEVICE_FUNC inline bool run(const Derived& /*derived*/) { return true; }
623};
624#if !defined(__FINITE_MATH_ONLY__) || !(__FINITE_MATH_ONLY__)
625template <typename Derived>
626struct all_finite_impl<Derived, false> {
627 static EIGEN_DEVICE_FUNC inline bool run(const Derived& derived) {
628 using Visitor = all_finite_visitor<typename traits<Derived>::Scalar>;
629 Visitor visitor;
630 visit_impl<Derived, Visitor, /*ShortCircuitEvaluation*/ true>::run(derived, visitor);
631 return visitor.res;
632 }
633};
634#endif
635
636} // end namespace internal
637
645template <typename Derived>
646EIGEN_DEVICE_FUNC inline bool DenseBase<Derived>::all() const {
647 using Visitor = internal::all_visitor<Scalar>;
648 using impl = internal::visit_impl<Derived, Visitor, /*ShortCircuitEvaluation*/ true>;
649 Visitor visitor;
650 impl::run(derived(), visitor);
651 return visitor.res;
652}
653
658template <typename Derived>
659EIGEN_DEVICE_FUNC inline bool DenseBase<Derived>::any() const {
660 using Visitor = internal::any_visitor<Scalar>;
661 using impl = internal::visit_impl<Derived, Visitor, /*ShortCircuitEvaluation*/ true>;
662 Visitor visitor;
663 impl::run(derived(), visitor);
664 return visitor.res;
665}
666
671template <typename Derived>
672EIGEN_DEVICE_FUNC Index DenseBase<Derived>::count() const {
673 using Visitor = internal::count_visitor<Scalar>;
674 using impl = internal::visit_impl<Derived, Visitor, /*ShortCircuitEvaluation*/ false>;
675 Visitor visitor;
676 impl::run(derived(), visitor);
677 return visitor.res;
678}
679
680template <typename Derived>
681EIGEN_DEVICE_FUNC inline bool DenseBase<Derived>::hasNaN() const {
682 return derived().cwiseTypedNotEqual(derived()).any();
683}
684
689template <typename Derived>
690EIGEN_DEVICE_FUNC inline bool DenseBase<Derived>::allFinite() const {
691 return internal::all_finite_impl<Derived>::run(derived());
692}
693
695template <typename Derived>
696EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool DenseBase<Derived>::isApproxToConstant(const Scalar& val,
697 const RealScalar& prec) const {
698 typename internal::nested_eval<Derived, 1>::type self(derived());
699 return internal::fuzzy_constant_all<true>(self, val, prec);
700}
701
710template <typename Derived>
711EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE bool DenseBase<Derived>::isZero(const RealScalar& prec) const {
712 typename internal::nested_eval<Derived, 1>::type self(derived());
713 return internal::fuzzy_constant_all<false>(self, Scalar(0), prec);
714}
715
716} // end namespace Eigen
717
718#endif // EIGEN_VISITOR_H
void visit(Visitor &func) const
Definition Visitor.h:423
constexpr bool isMuchSmallerThan(const typename NumTraits< Scalar >::Real &other, const RealScalar &prec) const
Definition ApproxComparisons.h:440
typename internal::traits< ArrayWrapper< ExpressionType > >::Scalar Scalar
Definition DenseBase.h:63
Index count() const
Definition Visitor.h:672
bool any() const
Definition Visitor.h:659
bool all() const
Definition Visitor.h:646
bool allFinite() const
Definition Visitor.h:690
bool isZero(const RealScalar &prec=NumTraits< Scalar >::dummy_precision()) const
Definition Visitor.h:711
bool isApproxToConstant(const Scalar &value, const RealScalar &prec=NumTraits< Scalar >::dummy_precision()) const
Definition Visitor.h:696
@ RowMajor
Definition Constants.h:321
constexpr unsigned int PacketAccessBit
Definition Constants.h:98
constexpr unsigned int LinearAccessBit
Definition Constants.h:134