Eigen  5.0.1
 
Loading...
Searching...
No Matches
UnaryFunctors.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2008-2016 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_UNARY_FUNCTORS_H
12#define EIGEN_UNARY_FUNCTORS_H
13
14// IWYU pragma: private
15#include "../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
26template <typename Scalar>
27struct scalar_opposite_op {
28 EIGEN_STATIC_ASSERT((!std::is_same<Scalar, bool>::value),
29 BOOLEAN_NEGATION_IS_NOT_SUPPORTED__CAST_TO_A_SIGNED_INTEGER_TYPE)
30 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::negate(a); }
31 template <typename Packet>
32 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
33 return internal::pnegate(a);
34 }
35};
36template <typename Scalar>
37struct functor_traits<scalar_opposite_op<Scalar>> {
38 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasNegate };
39};
40
46template <typename Scalar>
47struct scalar_abs_op {
48 using result_type = typename NumTraits<Scalar>::Real;
49 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a) const {
50 return numext::abs(a);
51 }
52 template <typename Packet>
53 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
54 return internal::pabs(a);
55 }
56};
57template <typename Scalar>
58struct functor_traits<scalar_abs_op<Scalar>> {
59 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasAbs };
60};
61
67template <typename Scalar>
68struct scalar_score_coeff_op : scalar_abs_op<Scalar> {
69 using Score_is_abs = void;
70};
71template <typename Scalar>
72struct functor_traits<scalar_score_coeff_op<Scalar>> : functor_traits<scalar_abs_op<Scalar>> {};
73
74/* Avoid recomputing abs when we know the score and they are the same. Not a true Eigen functor. */
75template <typename Scalar, typename = void>
76struct abs_knowing_score {
77 using result_type = typename NumTraits<Scalar>::Real;
78 template <typename Score>
79 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a, const Score&) const {
80 return numext::abs(a);
81 }
82};
83template <typename Scalar>
84struct abs_knowing_score<Scalar, typename scalar_score_coeff_op<Scalar>::Score_is_abs> {
85 using result_type = typename NumTraits<Scalar>::Real;
86 template <typename Scal>
87 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scal&, const result_type& a) const {
88 return a;
89 }
90};
91
97template <typename Scalar>
98struct scalar_abs2_op {
99 using result_type = typename NumTraits<Scalar>::Real;
100 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a) const {
101 return numext::abs2(a);
102 }
103 template <typename Packet>
104 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
105 return internal::pmul(a, a);
106 }
107};
108template <typename Scalar>
109struct functor_traits<scalar_abs2_op<Scalar>> {
110 enum {
111 Cost = NumTraits<Scalar>::MulCost,
112 PacketAccess = packet_traits<Scalar>::HasMul && !NumTraits<Scalar>::IsComplex
113 };
114};
115
121template <typename Scalar>
122struct scalar_conjugate_op {
123 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::conj(a); }
124 template <typename Packet>
125 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
126 return internal::pconj(a);
127 }
128};
129template <typename Scalar>
130struct functor_traits<scalar_conjugate_op<Scalar>> {
131 enum {
132 Cost = 0,
133 // Yes the cost is zero even for complexes because in most cases for which
134 // the cost is used, conjugation turns to be a no-op. Some examples:
135 // cost(a*conj(b)) == cost(a*b)
136 // cost(a+conj(b)) == cost(a+b)
137 // <etc.
138 // If we don't set it to zero, then:
139 // A.conjugate().lazyProduct(B.conjugate())
140 // will bake its operands. We definitely don't want that!
141 PacketAccess = packet_traits<Scalar>::HasConj
142 };
143};
144
150template <typename Scalar>
151struct scalar_arg_op {
152 using result_type = typename NumTraits<Scalar>::Real;
153 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a) const {
154 return numext::arg(a);
155 }
156 template <typename Packet>
157 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
158 return internal::parg(a);
159 }
160};
161template <typename Scalar>
162struct functor_traits<scalar_arg_op<Scalar>> {
163 enum {
164 Cost = NumTraits<Scalar>::IsComplex ? 5 * NumTraits<Scalar>::MulCost : NumTraits<Scalar>::AddCost,
165 PacketAccess = packet_traits<Scalar>::HasArg
166 };
167};
168
174template <typename Scalar>
175struct scalar_carg_op {
176 using result_type = Scalar;
177 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
178 return Scalar(numext::arg(a));
179 }
180 template <typename Packet>
181 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
182 return pcarg(a);
183 }
184};
185template <typename Scalar>
186struct functor_traits<scalar_carg_op<Scalar>> {
187 using RealScalar = typename NumTraits<Scalar>::Real;
188 enum {
189 Cost = functor_traits<scalar_atan2_op<RealScalar>>::Cost,
190 // The generic pcarg lowers to patan2, whose quotient-based reduction needs pdiv.
191 PacketAccess = packet_traits<RealScalar>::HasATan && packet_traits<RealScalar>::HasDiv
192 };
193};
194
200template <typename Scalar, typename NewType>
201struct scalar_cast_op {
202 using result_type = NewType;
203 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE NewType operator()(const Scalar& a) const {
204 return cast<Scalar, NewType>(a);
205 }
206};
207
208template <typename Scalar, typename NewType>
209struct functor_traits<scalar_cast_op<Scalar, NewType>> {
210 enum { Cost = std::is_same<Scalar, NewType>::value ? 0 : NumTraits<NewType>::AddCost, PacketAccess = false };
211};
212
219template <typename SrcType, typename DstType>
220struct core_cast_op : scalar_cast_op<SrcType, DstType> {};
221
222template <typename SrcType, typename DstType>
223struct functor_traits<core_cast_op<SrcType, DstType>> {
224 using CastingTraits = type_casting_traits<SrcType, DstType>;
225 enum {
226 Cost = std::is_same<SrcType, DstType>::value ? 0 : NumTraits<DstType>::AddCost,
227 PacketAccess = CastingTraits::VectorizedCast && (CastingTraits::SrcCoeffRatio <= 8)
228 };
229};
230
236template <typename Scalar, int N>
237struct scalar_arithmetic_shift_right_op {
238 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
239 return numext::arithmetic_shift_right(a, N);
240 }
241 template <typename Packet>
242 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
243 return internal::parithmetic_shift_right<N>(a);
244 }
245};
246template <typename Scalar, int N>
247struct functor_traits<scalar_arithmetic_shift_right_op<Scalar, N>> {
248 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasShift };
249};
250
256template <typename Scalar, int N>
257struct scalar_logical_shift_right_op {
258 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
259 return numext::logical_shift_right(a, N);
260 }
261 template <typename Packet>
262 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
263 return internal::plogical_shift_right<N>(a);
264 }
265};
266template <typename Scalar, int N>
267struct functor_traits<scalar_logical_shift_right_op<Scalar, N>> {
268 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasShift };
269};
270
276template <typename Scalar, int N>
277struct scalar_logical_shift_left_op {
278 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
279 return numext::logical_shift_left(a, N);
280 }
281 template <typename Packet>
282 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
283 return internal::plogical_shift_left<N>(a);
284 }
285};
286template <typename Scalar, int N>
287struct functor_traits<scalar_logical_shift_left_op<Scalar, N>> {
288 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasShift };
289};
290
296template <typename Scalar>
297struct scalar_real_op {
298 using result_type = typename NumTraits<Scalar>::Real;
299 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a) const {
300 return numext::real(a);
301 }
302};
303template <typename Scalar>
304struct functor_traits<scalar_real_op<Scalar>> {
305 enum { Cost = 0, PacketAccess = false };
306};
307
313template <typename Scalar>
314struct scalar_imag_op {
315 using result_type = typename NumTraits<Scalar>::Real;
316 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a) const {
317 return numext::imag(a);
318 }
319};
320template <typename Scalar>
321struct functor_traits<scalar_imag_op<Scalar>> {
322 enum { Cost = 0, PacketAccess = false };
323};
324
330template <typename Scalar>
331struct scalar_real_ref_op {
332 using result_type = typename NumTraits<Scalar>::Real;
333 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE const result_type& operator()(const Scalar& a) const {
334 return numext::real_ref(a);
335 }
336 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type& operator()(Scalar& a) const {
337 return numext::real_ref(a);
338 }
339};
340template <typename Scalar>
341struct functor_traits<scalar_real_ref_op<Scalar>> {
342 enum { Cost = 0, PacketAccess = false };
343};
344
350template <typename Scalar>
351struct scalar_imag_ref_op {
352 // For a real Scalar, numext::imag_ref returns the temporary RealScalar(0) by value; binding it to the
353 // reference return types below would dangle (issue #3096). Real-valued objects get a read-only zero
354 // expression from imag() instead (see NonConstImagReturnType in CommonCwiseUnaryOps.inc).
355 static_assert(NumTraits<Scalar>::IsComplex,
356 "THE IMAGINARY PART OF A REAL-VALUED OBJECT IS NOT AN LVALUE. USE THE READ-ONLY imag() OVERLOAD.");
357 using result_type = typename NumTraits<Scalar>::Real;
358 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type& operator()(Scalar& a) const {
359 return numext::imag_ref(a);
360 }
361 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE const result_type& operator()(const Scalar& a) const {
362 return numext::imag_ref(a);
363 }
364};
365template <typename Scalar>
366struct functor_traits<scalar_imag_ref_op<Scalar>> {
367 enum { Cost = 0, PacketAccess = false };
368};
369
376template <typename Scalar>
377struct scalar_exp_op {
378 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return internal::pexp(a); }
379 template <typename Packet>
380 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
381 return internal::pexp(a);
382 }
383};
384template <typename Scalar>
385struct functor_traits<scalar_exp_op<Scalar>> {
386 enum {
387 PacketAccess = packet_traits<Scalar>::HasExp,
388 // The following numbers are based on the AVX implementation.
389#ifdef EIGEN_VECTORIZE_FMA
390 // Haswell can issue 2 add/mul/madd per cycle.
391 Cost = (sizeof(Scalar) == 4
392 // float: 8 pmadd, 4 pmul, 2 padd/psub, 6 other
393 ? (8 * NumTraits<Scalar>::AddCost + 6 * NumTraits<Scalar>::MulCost)
394 // double: 7 pmadd, 5 pmul, 3 padd/psub, 1 div, 13 other
395 : (14 * NumTraits<Scalar>::AddCost + 6 * NumTraits<Scalar>::MulCost +
396 scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value))
397#else
398 Cost = (sizeof(Scalar) == 4
399 // float: 7 pmadd, 6 pmul, 4 padd/psub, 10 other
400 ? (21 * NumTraits<Scalar>::AddCost + 13 * NumTraits<Scalar>::MulCost)
401 // double: 7 pmadd, 5 pmul, 3 padd/psub, 1 div, 13 other
402 : (23 * NumTraits<Scalar>::AddCost + 12 * NumTraits<Scalar>::MulCost +
403 scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value))
404#endif
405 };
406};
407
408template <typename Scalar>
409struct scalar_exp2_op {
410 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return internal::pexp2(a); }
411 template <typename Packet>
412 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
413 return internal::pexp2(a);
414 }
415};
416template <typename Scalar>
417struct functor_traits<scalar_exp2_op<Scalar>> {
418 enum {
419 // There is no complex pexp2; complex arguments take numext::exp2.
420 PacketAccess = packet_traits<Scalar>::HasExp && !NumTraits<Scalar>::IsComplex,
421 Cost = functor_traits<scalar_exp_op<Scalar>>::Cost // TODO: measure cost of exp2
422 };
423};
424
431template <typename Scalar>
432struct scalar_ldexp_op {
433 static_assert(!NumTraits<Scalar>::IsComplex && !NumTraits<Scalar>::IsInteger,
434 "ldexp is only defined for real floating-point scalar types");
435 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE explicit scalar_ldexp_op(int exponent) : m_exponent(exponent) {}
436 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
437 return numext::ldexp(a, m_exponent);
438 }
439 template <typename Packet>
440 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
441 // pldexp clamps exponents, so a saturating conversion to Scalar is safe.
442 return internal::pldexp(a, pset1<Packet>(static_cast<Scalar>(m_exponent)));
443 }
444
445 private:
446 const int m_exponent;
447};
448template <typename Scalar>
449struct functor_traits<scalar_ldexp_op<Scalar>> {
450 enum {
451 // HasExp packets already require a tested pldexp implementation.
452 PacketAccess = packet_traits<Scalar>::HasExp,
453 Cost = 4 * NumTraits<Scalar>::MulCost
454 };
455};
456
463template <typename Scalar>
464struct scalar_expm1_op {
465 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::expm1(a); }
466 template <typename Packet>
467 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
468 return internal::pexpm1(a);
469 }
470};
471template <typename Scalar>
472struct functor_traits<scalar_expm1_op<Scalar>> {
473 enum {
474 PacketAccess = packet_traits<Scalar>::HasExpm1,
475 Cost = functor_traits<scalar_exp_op<Scalar>>::Cost // TODO: measure cost of expm1
476 };
477};
478
485template <typename Scalar>
486struct scalar_log_op {
487 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::log(a); }
488 template <typename Packet>
489 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
490 return internal::plog(a);
491 }
492};
493template <typename Scalar>
494struct functor_traits<scalar_log_op<Scalar>> {
495 enum {
496 PacketAccess = packet_traits<Scalar>::HasLog,
497 Cost = (PacketAccess
498 // The following numbers are based on the AVX implementation.
499#ifdef EIGEN_VECTORIZE_FMA
500 // 8 pmadd, 6 pmul, 8 padd/psub, 16 other, can issue 2 add/mul/madd per cycle.
501 ? (20 * NumTraits<Scalar>::AddCost + 7 * NumTraits<Scalar>::MulCost)
502#else
503 // 8 pmadd, 6 pmul, 8 padd/psub, 20 other
504 ? (36 * NumTraits<Scalar>::AddCost + 14 * NumTraits<Scalar>::MulCost)
505#endif
506 // Measured cost of std::log.
507 : sizeof(Scalar) == 4 ? 40 : 85)
508 };
509};
510
517template <typename Scalar>
518struct scalar_log1p_op {
519 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::log1p(a); }
520 template <typename Packet>
521 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
522 return internal::plog1p(a);
523 }
524};
525template <typename Scalar>
526struct functor_traits<scalar_log1p_op<Scalar>> {
527 enum {
528 PacketAccess = packet_traits<Scalar>::HasLog1p,
529 Cost = functor_traits<scalar_log_op<Scalar>>::Cost // TODO: measure cost of log1p
530 };
531};
532
539template <typename Scalar>
540struct scalar_log10_op {
541 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const {
542 EIGEN_USING_STD(log10) return log10(a);
543 }
544 template <typename Packet>
545 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
546 return internal::plog10(a);
547 }
548};
549template <typename Scalar>
550struct functor_traits<scalar_log10_op<Scalar>> {
551 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasLog10 };
552};
553
560template <typename Scalar>
561struct scalar_log2_op {
562 using RealScalar = typename NumTraits<Scalar>::Real;
563 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const {
564 return internal::mul(RealScalar(EIGEN_LOG2E), numext::log(a));
565 }
566 template <typename Packet>
567 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
568 return internal::plog2(a);
569 }
570};
571template <typename Scalar>
572struct functor_traits<scalar_log2_op<Scalar>> {
573 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasLog };
574};
575
580template <typename Scalar>
581struct scalar_sqrt_op {
582 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::sqrt(a); }
583 template <typename Packet>
584 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
585 return internal::psqrt(a);
586 }
587};
588template <typename Scalar>
589struct functor_traits<scalar_sqrt_op<Scalar>> {
590 enum {
591#if EIGEN_FAST_MATH
592 // The following numbers are based on the AVX implementation.
593 Cost = (sizeof(Scalar) == 8 ? 28
594 // 4 pmul, 1 pmadd, 3 other
595 : (3 * NumTraits<Scalar>::AddCost + 5 * NumTraits<Scalar>::MulCost)),
596#else
597 // The following numbers are based on min VSQRT throughput on Haswell.
598 Cost = (sizeof(Scalar) == 8 ? 28 : 14),
599#endif
600 PacketAccess = packet_traits<Scalar>::HasSqrt
601 };
602};
603
604// Boolean specialization to eliminate -Wimplicit-conversion-floating-point-to-bool warnings.
605template <>
606struct scalar_sqrt_op<bool> {
607 EIGEN_DEPRECATED EIGEN_DEVICE_FUNC inline bool operator()(const bool& a) const { return a; }
608 template <typename Packet>
609 EIGEN_DEPRECATED EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
610 return a;
611 }
612};
613template <>
614struct functor_traits<scalar_sqrt_op<bool>> {
615 enum { Cost = 1, PacketAccess = packet_traits<bool>::Vectorizable };
616};
617
622template <typename Scalar>
623struct scalar_cbrt_op {
624 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::cbrt(a); }
625 template <typename Packet>
626 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
627 return internal::pcbrt(a);
628 }
629};
630
631template <typename Scalar>
632struct functor_traits<scalar_cbrt_op<Scalar>> {
633 enum { Cost = 20 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasCbrt };
634};
635
640template <typename Scalar>
641struct scalar_rsqrt_op {
642 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::rsqrt(a); }
643 template <typename Packet>
644 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
645 return internal::prsqrt(a);
646 }
647};
648
649template <typename Scalar>
650struct functor_traits<scalar_rsqrt_op<Scalar>> {
651 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasRsqrt };
652};
653
658template <typename Scalar>
659struct scalar_cos_op {
660 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::cos(a); }
661 template <typename Packet>
662 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
663 return internal::pcos(a);
664 }
665};
666template <typename Scalar>
667struct functor_traits<scalar_cos_op<Scalar>> {
668 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasCos };
669};
670
675template <typename Scalar>
676struct scalar_sin_op {
677 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::sin(a); }
678 template <typename Packet>
679 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
680 return internal::psin(a);
681 }
682};
683template <typename Scalar>
684struct functor_traits<scalar_sin_op<Scalar>> {
685 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasSin };
686};
687
692template <typename Scalar>
693struct scalar_tan_op {
694 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::tan(a); }
695 template <typename Packet>
696 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
697 return internal::ptan(a);
698 }
699};
700template <typename Scalar>
701struct functor_traits<scalar_tan_op<Scalar>> {
702 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasTan };
703};
704
709template <typename Scalar>
710struct scalar_acos_op {
711 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::acos(a); }
712 template <typename Packet>
713 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
714 return internal::pacos(a);
715 }
716};
717template <typename Scalar>
718struct functor_traits<scalar_acos_op<Scalar>> {
719 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasACos };
720};
721
726template <typename Scalar>
727struct scalar_asin_op {
728 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::asin(a); }
729 template <typename Packet>
730 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
731 return internal::pasin(a);
732 }
733};
734template <typename Scalar>
735struct functor_traits<scalar_asin_op<Scalar>> {
736 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasASin };
737};
738
743template <typename Scalar>
744struct scalar_atan_op {
745 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::atan(a); }
746 template <typename Packet>
747 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
748 return internal::patan(a);
749 }
750};
751template <typename Scalar>
752struct functor_traits<scalar_atan_op<Scalar>> {
753 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasATan };
754};
755
760template <typename Scalar>
761struct scalar_tanh_op {
762 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::tanh(a); }
763 template <typename Packet>
764 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& x) const {
765 return ptanh(x);
766 }
767};
768
769template <typename Scalar>
770struct functor_traits<scalar_tanh_op<Scalar>> {
771 enum {
772 PacketAccess = packet_traits<Scalar>::HasTanh,
773 Cost = ((EIGEN_FAST_MATH && std::is_same<Scalar, float>::value)
774// The following numbers are based on the AVX implementation,
775#ifdef EIGEN_VECTORIZE_FMA
776 // Haswell can issue 2 add/mul/madd per cycle.
777 // 9 pmadd, 2 pmul, 1 div, 2 other
778 ? (2 * NumTraits<Scalar>::AddCost + 6 * NumTraits<Scalar>::MulCost +
779 scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value)
780#else
781 ? (11 * NumTraits<Scalar>::AddCost + 11 * NumTraits<Scalar>::MulCost +
782 scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value)
783#endif
784 // This number assumes a naive implementation of tanh
785 : (6 * NumTraits<Scalar>::AddCost + 3 * NumTraits<Scalar>::MulCost +
786 2 * scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value +
787 functor_traits<scalar_exp_op<Scalar>>::Cost))
788 };
789};
790
795template <typename Scalar>
796struct scalar_atanh_op {
797 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::atanh(a); }
798 template <typename Packet>
799 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& x) const {
800 return patanh(x);
801 }
802};
803
804template <typename Scalar>
805struct functor_traits<scalar_atanh_op<Scalar>> {
806 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasATanh };
807};
808
813template <typename Scalar>
814struct scalar_sinh_op {
815 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::sinh(a); }
816 template <typename Packet>
817 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
818 return internal::psinh(a);
819 }
820};
821template <typename Scalar>
822struct functor_traits<scalar_sinh_op<Scalar>> {
823 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasSinh };
824};
825
830template <typename Scalar>
831struct scalar_asinh_op {
832 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::asinh(a); }
833 template <typename Packet>
834 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
835 return internal::pasinh(a);
836 }
837};
838
839template <typename Scalar>
840struct functor_traits<scalar_asinh_op<Scalar>> {
841 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasASinh };
842};
843
848template <typename Scalar>
849struct scalar_cosh_op {
850 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::cosh(a); }
851 template <typename Packet>
852 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
853 return internal::pcosh(a);
854 }
855};
856template <typename Scalar>
857struct functor_traits<scalar_cosh_op<Scalar>> {
858 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasCosh };
859};
860
865template <typename Scalar>
866struct scalar_acosh_op {
867 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::acosh(a); }
868 template <typename Packet>
869 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
870 return internal::pacosh(a);
871 }
872};
873
874template <typename Scalar>
875struct functor_traits<scalar_acosh_op<Scalar>> {
876 enum { Cost = 5 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasACosh };
877};
878
883template <typename Scalar>
884struct scalar_inverse_op {
885 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return Scalar(1) / a; }
886 template <typename Packet>
887 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
888 return internal::preciprocal(a);
889 }
890};
891template <typename Scalar>
892struct functor_traits<scalar_inverse_op<Scalar>> {
893 enum {
894 PacketAccess = packet_traits<Scalar>::HasDiv,
895 // If packet_traits<Scalar>::HasReciprocal then the Estimated cost is that
896 // of computing an approximation plus a single Newton-Raphson step, which
897 // consists of 1 pmul + 1 pmadd.
898 Cost = (packet_traits<Scalar>::HasReciprocal ? 4 * NumTraits<Scalar>::MulCost
899 : scalar_div_cost<Scalar, PacketAccess>::value)
900 };
901};
902
907template <typename Scalar>
908struct scalar_square_op {
909 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return internal::mul(a, a); }
910 template <typename Packet>
911 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
912 return internal::pmul(a, a);
913 }
914};
915template <typename Scalar>
916struct functor_traits<scalar_square_op<Scalar>> {
917 enum { Cost = NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasMul };
918};
919
920// Boolean specialization to avoid -Wint-in-bool-context warnings on GCC.
921template <>
922struct scalar_square_op<bool> {
923 EIGEN_DEPRECATED EIGEN_DEVICE_FUNC inline bool operator()(const bool& a) const { return a; }
924 template <typename Packet>
925 EIGEN_DEPRECATED EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
926 return a;
927 }
928};
929template <>
930struct functor_traits<scalar_square_op<bool>> {
931 enum { Cost = 0, PacketAccess = packet_traits<bool>::Vectorizable };
932};
933
938template <typename Scalar>
939struct scalar_cube_op {
940 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const {
941 return internal::mul(a, internal::mul(a, a));
942 }
943 template <typename Packet>
944 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
945 return internal::pmul(a, pmul(a, a));
946 }
947};
948template <typename Scalar>
949struct functor_traits<scalar_cube_op<Scalar>> {
950 enum { Cost = 2 * NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasMul };
951};
952
953// Boolean specialization to avoid -Wint-in-bool-context warnings on GCC.
954template <>
955struct scalar_cube_op<bool> {
956 EIGEN_DEPRECATED EIGEN_DEVICE_FUNC inline bool operator()(const bool& a) const { return a; }
957 template <typename Packet>
958 EIGEN_DEPRECATED EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
959 return a;
960 }
961};
962template <>
963struct functor_traits<scalar_cube_op<bool>> {
964 enum { Cost = 0, PacketAccess = packet_traits<bool>::Vectorizable };
965};
966
971template <typename Scalar>
972struct scalar_round_op {
973 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::round(a); }
974 template <typename Packet>
975 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
976 return internal::pround(a);
977 }
978};
979template <typename Scalar>
980struct functor_traits<scalar_round_op<Scalar>> {
981 enum {
982 Cost = NumTraits<Scalar>::MulCost,
983 PacketAccess = packet_traits<Scalar>::HasRound || NumTraits<Scalar>::IsInteger
984 };
985};
986
991template <typename Scalar>
992struct scalar_floor_op {
993 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::floor(a); }
994 template <typename Packet>
995 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
996 return internal::pfloor(a);
997 }
998};
999template <typename Scalar>
1000struct functor_traits<scalar_floor_op<Scalar>> {
1001 enum {
1002 Cost = NumTraits<Scalar>::MulCost,
1003 PacketAccess = packet_traits<Scalar>::HasRound || NumTraits<Scalar>::IsInteger
1004 };
1005};
1006
1011template <typename Scalar>
1012struct scalar_rint_op {
1013 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::rint(a); }
1014 template <typename Packet>
1015 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1016 return internal::print(a);
1017 }
1018};
1019template <typename Scalar>
1020struct functor_traits<scalar_rint_op<Scalar>> {
1021 enum {
1022 Cost = NumTraits<Scalar>::MulCost,
1023 PacketAccess = packet_traits<Scalar>::HasRound || NumTraits<Scalar>::IsInteger
1024 };
1025};
1026
1031template <typename Scalar>
1032struct scalar_ceil_op {
1033 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::ceil(a); }
1034 template <typename Packet>
1035 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1036 return internal::pceil(a);
1037 }
1038};
1039template <typename Scalar>
1040struct functor_traits<scalar_ceil_op<Scalar>> {
1041 enum {
1042 Cost = NumTraits<Scalar>::MulCost,
1043 PacketAccess = packet_traits<Scalar>::HasRound || NumTraits<Scalar>::IsInteger
1044 };
1045};
1046
1051template <typename Scalar>
1052struct scalar_trunc_op {
1053 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const { return numext::trunc(a); }
1054 template <typename Packet>
1055 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1056 return internal::ptrunc(a);
1057 }
1058};
1059template <typename Scalar>
1060struct functor_traits<scalar_trunc_op<Scalar>> {
1061 enum {
1062 Cost = NumTraits<Scalar>::MulCost,
1063 PacketAccess = packet_traits<Scalar>::HasRound || NumTraits<Scalar>::IsInteger
1064 };
1065};
1066
1071template <typename Scalar, bool UseTypedPredicate = false>
1072struct scalar_isnan_op {
1073 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE bool operator()(const Scalar& a) const {
1074#if defined(SYCL_DEVICE_ONLY)
1075 return numext::isnan(a);
1076#else
1077 return numext::isnan EIGEN_NOT_A_MACRO(a);
1078#endif
1079 }
1080};
1081
1082template <typename Scalar>
1083struct scalar_isnan_op<Scalar, true> {
1084 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1085#if defined(SYCL_DEVICE_ONLY)
1086 return (numext::isnan(a) ? ptrue(a) : pzero(a));
1087#else
1088 return (numext::isnan EIGEN_NOT_A_MACRO(a) ? ptrue(a) : pzero(a));
1089#endif
1090 }
1091 template <typename Packet>
1092 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1093 return pisnan(a);
1094 }
1095};
1096
1097template <typename Scalar, bool UseTypedPredicate>
1098struct functor_traits<scalar_isnan_op<Scalar, UseTypedPredicate>> {
1099 enum { Cost = NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasCmp && UseTypedPredicate };
1100};
1101
1106template <typename Scalar, bool UseTypedPredicate = false>
1107struct scalar_isinf_op {
1108 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE bool operator()(const Scalar& a) const {
1109#if defined(SYCL_DEVICE_ONLY)
1110 return numext::isinf(a);
1111#else
1112 return (numext::isinf)(a);
1113#endif
1114 }
1115};
1116
1117template <typename Scalar>
1118struct scalar_isinf_op<Scalar, true> {
1119 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1120#if defined(SYCL_DEVICE_ONLY)
1121 return (numext::isinf(a) ? ptrue(a) : pzero(a));
1122#else
1123 return (numext::isinf EIGEN_NOT_A_MACRO(a) ? ptrue(a) : pzero(a));
1124#endif
1125 }
1126 template <typename Packet>
1127 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1128 return pisinf(a);
1129 }
1130};
1131template <typename Scalar, bool UseTypedPredicate>
1132struct functor_traits<scalar_isinf_op<Scalar, UseTypedPredicate>> {
1133 enum { Cost = NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasCmp && UseTypedPredicate };
1134};
1135
1140template <typename Scalar, bool UseTypedPredicate = false>
1141struct scalar_isfinite_op {
1142 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE bool operator()(const Scalar& a) const {
1143#if defined(SYCL_DEVICE_ONLY)
1144 return numext::isfinite(a);
1145#else
1146 return (numext::isfinite)(a);
1147#endif
1148 }
1149};
1150
1151template <typename Scalar>
1152struct scalar_isfinite_op<Scalar, true> {
1153 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1154#if defined(SYCL_DEVICE_ONLY)
1155 return (numext::isfinite(a) ? ptrue(a) : pzero(a));
1156#else
1157 return (numext::isfinite EIGEN_NOT_A_MACRO(a) ? ptrue(a) : pzero(a));
1158#endif
1159 }
1160 template <typename Packet>
1161 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1162 return pisfinite(a);
1163 }
1164};
1165template <typename Scalar, bool UseTypedPredicate>
1166struct functor_traits<scalar_isfinite_op<Scalar, UseTypedPredicate>> {
1167 enum { Cost = NumTraits<Scalar>::MulCost, PacketAccess = packet_traits<Scalar>::HasCmp && UseTypedPredicate };
1168};
1169
1175template <typename Scalar>
1176struct scalar_boolean_not_op {
1177 using result_type = Scalar;
1178 // `false` any value `a` that satisfies `a == Scalar(0)`
1179 // `true` is the complement of `false`
1180 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1181 return a == Scalar(0) ? Scalar(1) : Scalar(0);
1182 }
1183 template <typename Packet>
1184 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
1185 const Packet cst_one = pset1<Packet>(Scalar(1));
1186 EIGEN_IF_CONSTEXPR ((std::is_same<Scalar, bool>::value)) {
1187 // Boolean packet lanes are canonical, so logical NOT is 1 & ~a.
1188 return pandnot(cst_one, a);
1189 } else {
1190 Packet not_a = pcmp_eq(a, pzero(a));
1191 return pand(not_a, cst_one);
1192 }
1193 }
1194};
1195template <typename Scalar>
1196struct functor_traits<scalar_boolean_not_op<Scalar>> {
1197 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = packet_traits<Scalar>::HasCmp };
1198};
1199
1200template <typename Scalar, bool IsComplex = NumTraits<Scalar>::IsComplex>
1201struct bitwise_unary_impl {
1202 static constexpr size_t Size = sizeof(Scalar);
1203 using uint_t = typename numext::get_integer_by_size<Size>::unsigned_type;
1204 static EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar run_not(const Scalar& a) {
1205 uint_t a_as_uint = numext::bit_cast<uint_t, Scalar>(a);
1206 uint_t result = ~a_as_uint;
1207 return numext::bit_cast<Scalar, uint_t>(result);
1208 }
1209};
1210
1211template <typename Scalar>
1212struct bitwise_unary_impl<Scalar, true> {
1213 using Real = typename NumTraits<Scalar>::Real;
1214 static EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar run_not(const Scalar& a) {
1215 Real real_result = bitwise_unary_impl<Real>::run_not(numext::real(a));
1216 Real imag_result = bitwise_unary_impl<Real>::run_not(numext::imag(a));
1217 return Scalar(real_result, imag_result);
1218 }
1219};
1220
1226template <typename Scalar>
1227struct scalar_bitwise_not_op {
1228 EIGEN_STATIC_ASSERT(!NumTraits<Scalar>::RequireInitialization,
1229 BITWISE OPERATIONS MAY ONLY BE PERFORMED ON PLAIN DATA TYPES)
1230 EIGEN_STATIC_ASSERT((!std::is_same<Scalar, bool>::value), DONT USE BITWISE OPS ON BOOLEAN TYPES)
1231 using result_type = Scalar;
1232 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1233 return bitwise_unary_impl<Scalar>::run_not(a);
1234 }
1235 template <typename Packet>
1236 EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const {
1237 return pandnot(ptrue(a), a);
1238 }
1239};
1240template <typename Scalar>
1241struct functor_traits<scalar_bitwise_not_op<Scalar>> {
1242 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess = true };
1243};
1244
1249template <typename Scalar>
1250struct scalar_sign_op {
1251 EIGEN_DEVICE_FUNC constexpr inline Scalar operator()(const Scalar& a) const { return numext::sign(a); }
1252
1253 template <typename Packet>
1254 EIGEN_DEVICE_FUNC inline Packet packetOp(const Packet& a) const {
1255 return internal::psign(a);
1256 }
1257};
1258
1259template <typename Scalar>
1260struct functor_traits<scalar_sign_op<Scalar>> {
1261 enum {
1262 Cost = NumTraits<Scalar>::IsComplex ? (8 * NumTraits<Scalar>::MulCost) // roughly
1263 : (3 * NumTraits<Scalar>::AddCost),
1264 PacketAccess = packet_traits<Scalar>::HasSign && packet_traits<Scalar>::Vectorizable
1265 };
1266};
1267
1268// Real-valued implementation.
1269template <typename T, typename EnableIf = void>
1270struct scalar_logistic_op_impl {
1271 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE T operator()(const T& x) const { return packetOp(x); }
1272
1273 template <typename Packet>
1274 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& x) const {
1275 const Packet one = pset1<Packet>(T(1));
1276 const Packet inf = pset1<Packet>(NumTraits<T>::infinity());
1277 const Packet e = pexp(x);
1278 const Packet inf_mask = pcmp_eq(e, inf);
1279 return pselect(inf_mask, one, pdiv(e, padd(one, e)));
1280 }
1281};
1282
1283// Complex-valued implementation.
1284template <typename T>
1285struct scalar_logistic_op_impl<T, std::enable_if_t<NumTraits<T>::IsComplex>> {
1286 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE T operator()(const T& x) const {
1287 const T e = numext::exp(x);
1288 return (numext::isinf)(numext::real(e)) ? T(1) : e / (e + T(1));
1289 }
1290};
1291
1296template <typename T>
1297struct scalar_logistic_op : scalar_logistic_op_impl<T> {};
1298
1299// TODO(rmlarsen): Enable the following on host when integer_packet is defined
1300// for the relevant packet types.
1301#ifndef EIGEN_GPUCC
1302
1318template <>
1319struct scalar_logistic_op<float> {
1320 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE float operator()(const float& x) const {
1321 // Truncate at the first point where the interpolant is exactly one.
1322 const float cst_exp_hi = 16.6355324f;
1323 const float e = numext::exp(numext::mini(x, cst_exp_hi));
1324 return e / (1.0f + e);
1325 }
1326
1327 template <typename Packet>
1328 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& _x) const {
1329 const Packet cst_zero = pset1<Packet>(0.0f);
1330 const Packet cst_one = pset1<Packet>(1.0f);
1331 const Packet cst_half = pset1<Packet>(0.5f);
1332 // Truncate at the first point where the interpolant is exactly one.
1333 const Packet cst_exp_hi = pset1<Packet>(16.6355324f);
1334 const Packet cst_exp_lo = pset1<Packet>(-104.f);
1335
1336 // Clamp x to the non-trivial range where S(x). Outside this
1337 // interval the correctly rounded value of S(x) is either zero
1338 // or one.
1339 Packet zero_mask = pcmp_lt(_x, cst_exp_lo);
1340 Packet x = pmin(_x, cst_exp_hi);
1341
1342 // 1. Multiplicative range reduction:
1343 // Reduce the range of x by a factor of 2. This avoids having
1344 // to compute exp(x) accurately where the result is a denormalized
1345 // value.
1346 x = pmul(x, cst_half);
1347
1348 // 2. Subtractive range reduction:
1349 // Express exp(x) as exp(m*ln(2) + r) = 2^m*exp(r), start by extracting
1350 // m = floor(x/ln(2) + 0.5), such that x = m*ln(2) + r.
1351 const Packet cst_cephes_LOG2EF = pset1<Packet>(1.44269504088896341f);
1352 Packet m = pfloor(pmadd(x, cst_cephes_LOG2EF, cst_half));
1353 // Get r = x - m*ln(2). We use a trick from Cephes where the term
1354 // m*ln(2) is subtracted out in two parts, m*C1+m*C2 = m*ln(2),
1355 // to avoid accumulating truncation errors.
1356 const Packet cst_cephes_exp_C1 = pset1<Packet>(-0.693359375f);
1357 const Packet cst_cephes_exp_C2 = pset1<Packet>(2.12194440e-4f);
1358 Packet r = pmadd(m, cst_cephes_exp_C1, x);
1359 r = pmadd(m, cst_cephes_exp_C2, r);
1360
1361 // 3. Compute an approximation to exp(r) using a degree 5 minimax polynomial.
1362 // We compute even and odd terms separately to increase instruction level
1363 // parallelism.
1364 Packet r2 = pmul(r, r);
1365 const Packet cst_p2 = pset1<Packet>(0.49999141693115234375f);
1366 const Packet cst_p3 = pset1<Packet>(0.16666877269744873046875f);
1367 const Packet cst_p4 = pset1<Packet>(4.1898667812347412109375e-2f);
1368 const Packet cst_p5 = pset1<Packet>(8.33471305668354034423828125e-3f);
1369
1370 const Packet p_even = pmadd(r2, cst_p4, cst_p2);
1371 const Packet p_odd = pmadd(r2, cst_p5, cst_p3);
1372 const Packet p_low = padd(r, cst_one);
1373 Packet p = pmadd(r, p_odd, p_even);
1374 p = pmadd(r2, p, p_low);
1375
1376 // 4. Undo subtractive range reduction exp(m*ln(2) + r) = 2^m * exp(r).
1377 Packet e = pldexp_fast(p, m);
1378
1379 // 5. Undo multiplicative range reduction by using exp(r) = exp(r/2)^2.
1380 e = pmul(e, e);
1381
1382 // Return exp(x) / (1 + exp(x))
1383 return pselect(zero_mask, cst_zero, pdiv(e, padd(cst_one, e)));
1384 }
1385};
1386#endif // #ifndef EIGEN_GPUCC
1387
1388template <typename T>
1389struct functor_traits<scalar_logistic_op<T>> {
1390 enum {
1391 // The cost estimate for float here is for the common(?) case where
1392 // all arguments are greater than -9.
1393 Cost = scalar_div_cost<T, packet_traits<T>::HasDiv>::value +
1394 (std::is_same<T, float>::value ? NumTraits<T>::AddCost * 15 + NumTraits<T>::MulCost * 11
1395 : NumTraits<T>::AddCost * 2 + functor_traits<scalar_exp_op<T>>::Cost),
1396 // Both packet paths branch with pcmp_*/pselect.
1397 PacketAccess = !NumTraits<T>::IsComplex && packet_traits<T>::HasAdd && packet_traits<T>::HasDiv &&
1398 packet_traits<T>::HasCmp &&
1399 (std::is_same<T, float>::value
1400 ? packet_traits<T>::HasMul && packet_traits<T>::HasMax && packet_traits<T>::HasMin
1401 : packet_traits<T>::HasNegate && packet_traits<T>::HasExp)
1402 };
1403};
1404
1405template <typename Scalar, typename ExponentScalar, bool IsBaseInteger = NumTraits<Scalar>::IsInteger,
1406 bool IsExponentInteger = NumTraits<ExponentScalar>::IsInteger,
1407 bool IsBaseComplex = NumTraits<Scalar>::IsComplex,
1408 bool IsExponentComplex = NumTraits<ExponentScalar>::IsComplex>
1409struct scalar_unary_pow_op {
1410 using PromotedExponent = typename internal::promote_scalar_arg<
1411 Scalar, ExponentScalar,
1412 internal::has_ReturnType<ScalarBinaryOpTraits<Scalar, ExponentScalar, scalar_unary_pow_op>>::value>::type;
1413 using result_type = typename ScalarBinaryOpTraits<Scalar, PromotedExponent, scalar_unary_pow_op>::ReturnType;
1414 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE scalar_unary_pow_op(const ExponentScalar& exponent)
1415 : m_exponent(exponent) {}
1416 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE result_type operator()(const Scalar& a) const {
1417 EIGEN_USING_STD(pow);
1418 return static_cast<result_type>(pow(a, m_exponent));
1419 }
1420
1421 private:
1422 const ExponentScalar m_exponent;
1423 scalar_unary_pow_op() {}
1424};
1425
1426template <typename T>
1427constexpr int exponent_digits() {
1428 return CHAR_BIT * sizeof(T) - NumTraits<T>::digits() - NumTraits<T>::IsSigned;
1429}
1430
1431template <typename From, typename To>
1432struct is_floating_exactly_representable {
1433 // TODO(rmlarsen): Add radix to NumTraits and enable this check.
1434 // (NumTraits<To>::radix == NumTraits<From>::radix) &&
1435 static constexpr bool value =
1436 (exponent_digits<To>() >= exponent_digits<From>() && NumTraits<To>::digits() >= NumTraits<From>::digits());
1437};
1438
1439// Specialization for real, non-integer types, non-complex types.
1440template <typename Scalar, typename ExponentScalar>
1441struct scalar_unary_pow_op<Scalar, ExponentScalar, false, false, false, false> {
1442 template <bool IsExactlyRepresentable = is_floating_exactly_representable<ExponentScalar, Scalar>::value>
1443 std::enable_if_t<IsExactlyRepresentable, void> check_is_representable() const {}
1444
1445 // Issue a deprecation warning if we do a narrowing conversion on the exponent.
1446 template <bool IsExactlyRepresentable = is_floating_exactly_representable<ExponentScalar, Scalar>::value>
1447 EIGEN_DEPRECATED std::enable_if_t<!IsExactlyRepresentable, void> check_is_representable() const {}
1448
1449 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE scalar_unary_pow_op(const ExponentScalar& exponent)
1450 : m_exponent(static_cast<Scalar>(exponent)) {
1451 check_is_representable();
1452 }
1453
1454 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1455 EIGEN_USING_STD(pow);
1456 return static_cast<Scalar>(pow(a, m_exponent));
1457 }
1458 template <typename Packet>
1459 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet packetOp(const Packet& a) const {
1460 return unary_pow_impl<Packet, Scalar>::run(a, m_exponent);
1461 }
1462
1463 private:
1464 const Scalar m_exponent;
1465 scalar_unary_pow_op() {}
1466};
1467
1468// Specialization for a complex base and a real, non-integer exponent type: integer-valued exponents take the
1469// repeated-squaring path, which is exact where the power is representable; the rest defer to pow. There is no
1470// vectorized complex pow, so this stays scalar.
1471template <typename Scalar, typename ExponentScalar>
1472struct scalar_unary_pow_op<Scalar, ExponentScalar, false, false, true, false> {
1473 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE scalar_unary_pow_op(const ExponentScalar& exponent)
1474 : m_exponent(exponent),
1475 m_use_repeated_squaring((numext::isfinite)(exponent) && numext::round(exponent) == exponent &&
1476 unary_pow::use_repeated_squaring<Scalar>(exponent)) {}
1477 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1478 if (m_use_repeated_squaring) return unary_pow::int_pow(a, m_exponent);
1479 EIGEN_USING_STD(pow);
1480 return static_cast<Scalar>(pow(a, m_exponent));
1481 }
1482
1483 private:
1484 const ExponentScalar m_exponent;
1485 const bool m_use_repeated_squaring;
1486 scalar_unary_pow_op() {}
1487};
1488
1489template <typename Scalar, typename ExponentScalar, bool BaseIsInteger, bool BaseIsComplex>
1490struct scalar_unary_pow_op<Scalar, ExponentScalar, BaseIsInteger, true, BaseIsComplex, false> {
1491 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE scalar_unary_pow_op(const ExponentScalar& exponent)
1492 : m_exponent(exponent) {}
1493 EIGEN_DEVICE_FUNC constexpr EIGEN_STRONG_INLINE Scalar operator()(const Scalar& a) const {
1494 return unary_pow_impl<Scalar, ExponentScalar>::run(a, m_exponent);
1495 }
1496 template <typename Packet>
1497 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Packet packetOp(const Packet& a) const {
1498 return unary_pow_impl<Packet, ExponentScalar>::run(a, m_exponent);
1499 }
1500
1501 private:
1502 const ExponentScalar m_exponent;
1503 scalar_unary_pow_op() {}
1504};
1505
1506template <typename Scalar, typename ExponentScalar>
1507struct functor_traits<scalar_unary_pow_op<Scalar, ExponentScalar>> {
1508 enum {
1509 GenPacketAccess = functor_traits<scalar_pow_op<Scalar, ExponentScalar>>::PacketAccess,
1510 // Only the real-exponent specializations define packetOp. Every path multiplies packets of the base type and
1511 // divides them for a negative exponent unless the base is an integer; the compares and selects run on the
1512 // real packet, which for a complex base is the component view of its packet. A complex base is vectorized
1513 // only through the double-word path, so it also needs that path's packet support.
1514 IntPacketAccess =
1515 !NumTraits<ExponentScalar>::IsComplex && packet_traits<Scalar>::HasMul &&
1516 (packet_traits<Scalar>::HasDiv || NumTraits<Scalar>::IsInteger) &&
1517 packet_traits<typename NumTraits<Scalar>::Real>::HasCmp &&
1518 (!NumTraits<Scalar>::IsComplex || unary_pow::use_double_word<typename packet_traits<Scalar>::type>::value),
1519 PacketAccess = NumTraits<ExponentScalar>::IsInteger ? IntPacketAccess : (IntPacketAccess && GenPacketAccess),
1520 Cost = functor_traits<scalar_pow_op<Scalar, ExponentScalar>>::Cost
1521 };
1522};
1523
1524} // end namespace internal
1525
1526} // end namespace Eigen
1527
1528#endif // EIGEN_UNARY_FUNCTORS_H