11#ifndef EIGEN_FIND_COEFF_H
12#define EIGEN_FIND_COEFF_H
15#include "./InternalHeaderCheck.h"
21template <typename Scalar, int NaNPropagation, bool IsInteger = NumTraits<Scalar>::IsInteger>
22struct max_coeff_functor {
23 EIGEN_DEVICE_FUNC
inline bool compareCoeff(
const Scalar& incumbent,
const Scalar& candidate)
const {
24 return candidate > incumbent;
26 template <
typename Packet>
27 EIGEN_DEVICE_FUNC
inline Packet comparePacket(
const Packet& incumbent,
const Packet& candidate)
const {
28 return pcmp_lt(incumbent, candidate);
30 template <
typename Packet>
31 EIGEN_DEVICE_FUNC
inline Scalar predux(
const Packet& a)
const {
36template <
typename Scalar>
38 EIGEN_DEVICE_FUNC
inline bool compareCoeff(
const Scalar& incumbent,
const Scalar& candidate)
const {
39 return (candidate > incumbent) || ((candidate != candidate) && (incumbent == incumbent));
41 template <
typename Packet>
42 EIGEN_DEVICE_FUNC
inline Packet comparePacket(
const Packet& incumbent,
const Packet& candidate)
const {
43 return pandnot(pcmp_lt_or_nan(incumbent, candidate), pisnan(incumbent));
45 template <
typename Packet>
46 EIGEN_DEVICE_FUNC
inline Scalar predux(
const Packet& a)
const {
47 return predux_max<PropagateNaN>(a);
51template <
typename Scalar>
53 EIGEN_DEVICE_FUNC
inline bool compareCoeff(
const Scalar& incumbent,
const Scalar& candidate)
const {
54 return (candidate > incumbent) || ((candidate == candidate) && (incumbent != incumbent));
56 template <
typename Packet>
57 EIGEN_DEVICE_FUNC
inline Packet comparePacket(
const Packet& incumbent,
const Packet& candidate)
const {
58 return pandnot(pcmp_lt_or_nan(incumbent, candidate), pisnan(candidate));
60 template <
typename Packet>
61 EIGEN_DEVICE_FUNC
inline Scalar predux(
const Packet& a)
const {
62 return predux_max<PropagateNumbers>(a);
66template <typename Scalar, int NaNPropagation, bool IsInteger = NumTraits<Scalar>::IsInteger>
67struct min_coeff_functor {
68 EIGEN_DEVICE_FUNC
inline bool compareCoeff(
const Scalar& incumbent,
const Scalar& candidate)
const {
69 return candidate < incumbent;
71 template <
typename Packet>
72 EIGEN_DEVICE_FUNC
inline Packet comparePacket(
const Packet& incumbent,
const Packet& candidate)
const {
73 return pcmp_lt(candidate, incumbent);
75 template <
typename Packet>
76 EIGEN_DEVICE_FUNC
inline Scalar predux(
const Packet& a)
const {
81template <
typename Scalar>
83 EIGEN_DEVICE_FUNC
inline bool compareCoeff(
const Scalar& incumbent,
const Scalar& candidate)
const {
84 return (candidate < incumbent) || ((candidate != candidate) && (incumbent == incumbent));
86 template <
typename Packet>
87 EIGEN_DEVICE_FUNC
inline Packet comparePacket(
const Packet& incumbent,
const Packet& candidate)
const {
88 return pandnot(pcmp_lt_or_nan(candidate, incumbent), pisnan(incumbent));
90 template <
typename Packet>
91 EIGEN_DEVICE_FUNC
inline Scalar predux(
const Packet& a)
const {
92 return predux_min<PropagateNaN>(a);
96template <
typename Scalar>
98 EIGEN_DEVICE_FUNC
inline bool compareCoeff(
const Scalar& incumbent,
const Scalar& candidate)
const {
99 return (candidate < incumbent) || ((candidate == candidate) && (incumbent != incumbent));
101 template <
typename Packet>
102 EIGEN_DEVICE_FUNC
inline Packet comparePacket(
const Packet& incumbent,
const Packet& candidate)
const {
103 return pandnot(pcmp_lt_or_nan(candidate, incumbent), pisnan(candidate));
105 template <
typename Packet>
106 EIGEN_DEVICE_FUNC
inline Scalar predux(
const Packet& a)
const {
107 return predux_min<PropagateNumbers>(a);
113template <
typename Scalar>
114struct min_max_traits {
115 static constexpr bool PacketAccess = packet_traits<Scalar>::Vectorizable && packet_traits<Scalar>::HasCmp;
117template <
typename Scalar,
int NaNPropagation>
118struct functor_traits<max_coeff_functor<Scalar, NaNPropagation>> : min_max_traits<Scalar> {};
119template <
typename Scalar,
int NaNPropagation>
120struct functor_traits<min_coeff_functor<Scalar, NaNPropagation>> : min_max_traits<Scalar> {};
122template <
typename Evaluator,
typename Func,
bool Linear,
bool Vectorize>
123struct find_coeff_loop;
124template <
typename Evaluator,
typename Func>
125struct find_coeff_loop<Evaluator, Func, false, false> {
126 using Scalar =
typename Evaluator::Scalar;
127 static EIGEN_DEVICE_FUNC
inline void run(
const Evaluator& eval, Func& func, Scalar& res, Index& outer, Index& inner) {
128 Index outerSize = eval.outerSize();
129 Index innerSize = eval.innerSize();
136 for (Index j = 0; j < outerSize; j++) {
137 for (Index i = 0; i < innerSize; i++) {
138 Scalar xprCoeff = eval.coeffByOuterInner(j, i);
139 if (func.compareCoeff(res, xprCoeff)) {
148template <
typename Evaluator,
typename Func>
149struct find_coeff_loop<Evaluator, Func, true, false> {
150 using Scalar =
typename Evaluator::Scalar;
151 static EIGEN_DEVICE_FUNC
inline void run(
const Evaluator& eval, Func& func, Scalar& res, Index& index) {
152 Index size = eval.size();
158 for (Index k = 0; k < size; k++) {
159 Scalar xprCoeff = eval.coeff(k);
160 if (func.compareCoeff(res, xprCoeff)) {
167template <
typename Evaluator,
typename Func>
168struct find_coeff_loop<Evaluator, Func, false, true> {
169 using ScalarImpl = find_coeff_loop<Evaluator, Func, false, false>;
170 using Scalar =
typename Evaluator::Scalar;
171 using Packet =
typename Evaluator::Packet;
172 static constexpr int PacketSize = unpacket_traits<Packet>::size;
173 static EIGEN_DEVICE_FUNC
inline void run(
const Evaluator& eval, Func& func, Scalar& result, Index& outer,
175 Index outerSize = eval.outerSize();
176 Index innerSize = eval.innerSize();
177 if (innerSize < PacketSize) {
178 ScalarImpl::run(eval, func, result, outer, inner);
181 Index packetEnd = numext::round_down(innerSize, PacketSize);
188 bool checkPacket =
false;
190 for (Index j = 0; j < outerSize; j++) {
191 Packet resultPacket = pset1<Packet>(result);
192 for (Index i = 0; i < packetEnd; i += PacketSize) {
193 Packet xprPacket = eval.template packetByOuterInner<Unaligned, Packet>(j, i);
194 if (predux_any(func.comparePacket(resultPacket, xprPacket))) {
197 result = func.predux(xprPacket);
198 resultPacket = pset1<Packet>(result);
203 for (Index i = packetEnd; i < innerSize; i++) {
204 Scalar xprCoeff = eval.coeffByOuterInner(j, i);
205 if (func.compareCoeff(result, xprCoeff)) {
215 result = eval.coeffByOuterInner(outer, inner);
216 Index i_end = inner + PacketSize;
217 for (Index i = inner; i < i_end; i++) {
218 Scalar xprCoeff = eval.coeffByOuterInner(outer, i);
219 if (func.compareCoeff(result, xprCoeff)) {
227template <
typename Evaluator,
typename Func>
228struct find_coeff_loop<Evaluator, Func, true, true> {
229 using ScalarImpl = find_coeff_loop<Evaluator, Func, true, false>;
230 using Scalar =
typename Evaluator::Scalar;
231 using Packet =
typename Evaluator::Packet;
232 static constexpr int PacketSize = unpacket_traits<Packet>::size;
233 static constexpr int Alignment = Evaluator::Alignment;
235 static EIGEN_DEVICE_FUNC
inline void run(
const Evaluator& eval, Func& func, Scalar& result, Index& index) {
236 Index size = eval.size();
237 if (size < PacketSize) {
238 ScalarImpl::run(eval, func, result, index);
241 Index packetEnd = numext::round_down(size, PacketSize);
247 Packet resultPacket = pset1<Packet>(result);
248 bool checkPacket =
false;
250 for (Index k = 0; k < packetEnd; k += PacketSize) {
251 Packet xprPacket = eval.template packet<Alignment, Packet>(k);
252 if (predux_any(func.comparePacket(resultPacket, xprPacket))) {
254 result = func.predux(xprPacket);
255 resultPacket = pset1<Packet>(result);
260 for (Index k = packetEnd; k < size; k++) {
261 Scalar xprCoeff = eval.coeff(k);
262 if (func.compareCoeff(result, xprCoeff)) {
270 result = eval.coeff(index);
271 Index k_end = index + PacketSize;
272 for (Index k = index; k < k_end; k++) {
273 Scalar xprCoeff = eval.coeff(k);
274 if (func.compareCoeff(result, xprCoeff)) {
283template <
typename Derived>
284struct find_coeff_evaluator :
public evaluator<Derived> {
285 using Base = evaluator<Derived>;
286 using Scalar =
typename Derived::Scalar;
287 using Packet =
typename packet_traits<Scalar>::type;
288 static constexpr int Flags = Base::Flags;
289 static constexpr bool IsRowMajor = bool(Flags &
RowMajorBit);
290 EIGEN_DEVICE_FUNC
inline find_coeff_evaluator(
const Derived& xpr) : Base(xpr), m_xpr(xpr) {}
292 EIGEN_DEVICE_FUNC
inline Scalar coeffByOuterInner(Index outer, Index inner)
const {
293 Index row = IsRowMajor ? outer : inner;
294 Index col = IsRowMajor ? inner : outer;
295 return Base::coeff(row, col);
297 template <
int LoadMode,
typename PacketType>
298 EIGEN_DEVICE_FUNC
inline PacketType packetByOuterInner(Index outer, Index inner)
const {
299 Index row = IsRowMajor ? outer : inner;
300 Index col = IsRowMajor ? inner : outer;
301 return Base::template packet<LoadMode, PacketType>(row, col);
304 EIGEN_DEVICE_FUNC
inline Index innerSize()
const {
return m_xpr.innerSize(); }
305 EIGEN_DEVICE_FUNC
inline Index outerSize()
const {
return m_xpr.outerSize(); }
306 EIGEN_DEVICE_FUNC
inline Index size()
const {
return m_xpr.size(); }
308 const Derived& m_xpr;
311template <
typename Derived,
typename Func>
312struct find_coeff_impl {
313 using Evaluator = find_coeff_evaluator<Derived>;
314 static constexpr int Flags = Evaluator::Flags;
315 static constexpr bool IsRowMajor = Derived::IsRowMajor;
316 static constexpr int MaxInnerSizeAtCompileTime =
317 IsRowMajor ? Derived::MaxColsAtCompileTime : Derived::MaxRowsAtCompileTime;
318 static constexpr int MaxSizeAtCompileTime = Derived::MaxSizeAtCompileTime;
320 using Scalar =
typename Derived::Scalar;
321 using Packet =
typename Evaluator::Packet;
323 static constexpr int PacketSize = unpacket_traits<Packet>::size;
325 static constexpr bool DontVectorize =
326 enum_lt_not_dynamic(Linearize ? MaxSizeAtCompileTime : MaxInnerSizeAtCompileTime, PacketSize);
327 static constexpr bool Vectorize =
328 !DontVectorize && bool(Flags &
PacketAccessBit) && functor_traits<Func>::PacketAccess;
330 using Loop = find_coeff_loop<Evaluator, Func, Linearize, Vectorize>;
332 template <
bool ForwardLinearAccess = Linearize, std::enable_if_t<!ForwardLinearAccess,
bool> = true>
333 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void run(
const Derived& xpr, Func& func, Scalar& res, Index& outer,
336 Loop::run(eval, func, res, outer, inner);
338 template <
bool ForwardLinearAccess = Linearize, std::enable_if_t<ForwardLinearAccess,
bool> = true>
339 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void run(
const Derived& xpr, Func& func, Scalar& res, Index& outer,
343 run(xpr, func, res, index);
344 outer = index / xpr.innerSize();
345 inner = index % xpr.innerSize();
347 static EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void run(
const Derived& xpr, Func& func, Scalar& res, Index& index) {
349 Loop::run(eval, func, res, index);
353template <
typename Derived,
typename IndexType,
typename Func>
354EIGEN_DEVICE_FUNC
typename internal::traits<Derived>::Scalar findCoeff(
const DenseBase<Derived>& mat, Func& func,
355 IndexType* rowPtr, IndexType* colPtr) {
356 eigen_assert(mat.rows() > 0 && mat.cols() > 0 &&
"you are using an empty matrix");
358 using FindCoeffImpl = internal::find_coeff_impl<Derived, Func>;
361 Scalar res = mat.coeff(0, 0);
362 FindCoeffImpl::run(mat.derived(), func, res, outer, inner);
363 *rowPtr = internal::convert_index<IndexType>(Derived::IsRowMajor ? outer : inner);
364 if (colPtr) *colPtr = internal::convert_index<IndexType>(Derived::IsRowMajor ? inner : outer);
368template <
typename Derived,
typename IndexType,
typename Func>
369EIGEN_DEVICE_FUNC
typename internal::traits<Derived>::Scalar findCoeff(
const DenseBase<Derived>& mat, Func& func,
370 IndexType* indexPtr) {
371 eigen_assert(mat.size() > 0 &&
"you are using an empty matrix");
372 EIGEN_STATIC_ASSERT_VECTOR_ONLY(Derived)
374 using FindCoeffImpl = internal::find_coeff_impl<Derived, Func>;
376 Scalar res = mat.coeff(0);
377 FindCoeffImpl::run(mat.derived(), func, res, index);
378 *indexPtr = internal::convert_index<IndexType>(index);
397template <
typename Derived>
398template <
int NaNPropagation,
typename IndexType>
400 IndexType* colPtr)
const {
401 using Func = internal::min_coeff_functor<Scalar, NaNPropagation>;
403 return internal::findCoeff(derived(), func, rowPtr, colPtr);
419template <
typename Derived>
420template <
int NaNPropagation,
typename IndexType>
422 using Func = internal::min_coeff_functor<Scalar, NaNPropagation>;
424 return internal::findCoeff(derived(), func, indexPtr);
440template <
typename Derived>
441template <
int NaNPropagation,
typename IndexType>
443 IndexType* colPtr)
const {
444 using Func = internal::max_coeff_functor<Scalar, NaNPropagation>;
446 return internal::findCoeff(derived(), func, rowPtr, colPtr);
462template <
typename Derived>
463template <
int NaNPropagation,
typename IndexType>
465 using Func = internal::max_coeff_functor<Scalar, NaNPropagation>;
467 return internal::findCoeff(derived(), func, indexPtr);
internal::traits< Derived >::Scalar minCoeff() const
Definition Redux.h:786
internal::traits< Derived >::Scalar maxCoeff() const
Definition Redux.h:799
typename internal::traits< Derived >::Scalar Scalar
Definition DenseBase.h:63
@ PropagateNaN
Definition Constants.h:343
@ PropagateNumbers
Definition Constants.h:345
constexpr unsigned int PacketAccessBit
Definition Constants.h:98
constexpr unsigned int LinearAccessBit
Definition Constants.h:134
constexpr unsigned int RowMajorBit
Definition Constants.h:71