Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
SpecialFunctionsFunctors.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2016 Eugene Brevdo <ebrevdo@gmail.com>
5// Copyright (C) 2016 Gael Guennebaud <gael.guennebaud@inria.fr>
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_SPECIALFUNCTIONS_FUNCTORS_H
13#define EIGEN_SPECIALFUNCTIONS_FUNCTORS_H
14
15// IWYU pragma: private
16#include "./InternalHeaderCheck.h"
17
18namespace Eigen {
19
20namespace internal {
21
27template <typename Scalar>
28struct scalar_igamma_op : binary_op_base<Scalar, Scalar> {
29 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a, const Scalar& x) const {
30 using numext::igamma;
31 return igamma(a, x);
32 }
33 template <typename Packet>
34 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& a, const Packet& x) const {
35 return internal::pigamma(a, x);
36 }
37};
38template <typename Scalar>
39struct functor_traits<scalar_igamma_op<Scalar> > {
40 enum {
41 // Guesstimate
42 Cost = 20 * NumTraits<Scalar>::MulCost + 10 * NumTraits<Scalar>::AddCost,
43 PacketAccess = packet_traits<Scalar>::HasIGamma
44 };
45};
46
53template <typename Scalar>
54struct scalar_igamma_der_a_op {
55 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a, const Scalar& x) const {
56 using numext::igamma_der_a;
57 return igamma_der_a(a, x);
58 }
59 template <typename Packet>
60 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& a, const Packet& x) const {
61 return internal::pigamma_der_a(a, x);
62 }
63};
64template <typename Scalar>
65struct functor_traits<scalar_igamma_der_a_op<Scalar> > {
66 enum {
67 // 2x the cost of igamma
68 Cost = 40 * NumTraits<Scalar>::MulCost + 20 * NumTraits<Scalar>::AddCost,
69 PacketAccess = packet_traits<Scalar>::HasIGammaDerA
70 };
71};
72
80template <typename Scalar>
81struct scalar_gamma_sample_der_alpha_op {
82 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& alpha, const Scalar& sample) const {
83 using numext::gamma_sample_der_alpha;
84 return gamma_sample_der_alpha(alpha, sample);
85 }
86 template <typename Packet>
87 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& alpha, const Packet& sample) const {
88 return internal::pgamma_sample_der_alpha(alpha, sample);
89 }
90};
91template <typename Scalar>
92struct functor_traits<scalar_gamma_sample_der_alpha_op<Scalar> > {
93 enum {
94 // 2x the cost of igamma, minus the lgamma cost (the lgamma cancels out)
95 Cost = 30 * NumTraits<Scalar>::MulCost + 15 * NumTraits<Scalar>::AddCost,
96 PacketAccess = packet_traits<Scalar>::HasGammaSampleDerAlpha
97 };
98};
99
105template <typename Scalar>
106struct scalar_igammac_op : binary_op_base<Scalar, Scalar> {
107 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a, const Scalar& x) const {
108 using numext::igammac;
109 return igammac(a, x);
110 }
111 template <typename Packet>
112 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& a, const Packet& x) const {
113 return internal::pigammac(a, x);
114 }
115};
116template <typename Scalar>
117struct functor_traits<scalar_igammac_op<Scalar> > {
118 enum {
119 // Guesstimate
120 Cost = 20 * NumTraits<Scalar>::MulCost + 10 * NumTraits<Scalar>::AddCost,
121 PacketAccess = packet_traits<Scalar>::HasIGammac
122 };
123};
124
129template <typename Scalar>
130struct scalar_betainc_op {
131 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& x, const Scalar& a,
132 const Scalar& b) const {
133 using numext::betainc;
134 return betainc(x, a, b);
135 }
136 template <typename Packet>
137 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Packet packetOp(const Packet& x, const Packet& a, const Packet& b) const {
138 return internal::pbetainc(x, a, b);
139 }
140};
141template <typename Scalar>
142struct functor_traits<scalar_betainc_op<Scalar> > {
143 enum {
144 // Guesstimate
145 Cost = 400 * NumTraits<Scalar>::MulCost + 400 * NumTraits<Scalar>::AddCost,
146 PacketAccess = packet_traits<Scalar>::HasBetaInc
147 };
148};
149
155template <typename Scalar>
156struct scalar_lgamma_op {
157 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a) const {
158 using numext::lgamma;
159 return lgamma(a);
160 }
161 typedef typename packet_traits<Scalar>::type Packet;
162 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const { return internal::plgamma(a); }
163};
164template <typename Scalar>
165struct functor_traits<scalar_lgamma_op<Scalar> > {
166 enum {
167 // Guesstimate
168 Cost = 10 * NumTraits<Scalar>::MulCost + 5 * NumTraits<Scalar>::AddCost,
169 PacketAccess = packet_traits<Scalar>::HasLGamma
170 };
171};
172
177template <typename Scalar>
178struct scalar_digamma_op {
179 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a) const {
180 using numext::digamma;
181 return digamma(a);
182 }
183 typedef typename packet_traits<Scalar>::type Packet;
184 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const { return internal::pdigamma(a); }
185};
186template <typename Scalar>
187struct functor_traits<scalar_digamma_op<Scalar> > {
188 enum {
189 // Guesstimate
190 Cost = 10 * NumTraits<Scalar>::MulCost + 5 * NumTraits<Scalar>::AddCost,
191 PacketAccess = packet_traits<Scalar>::HasDiGamma
192 };
193};
194
199template <typename Scalar>
200struct scalar_zeta_op {
201 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& x, const Scalar& q) const {
202 using numext::zeta;
203 return zeta(x, q);
204 }
205 typedef typename packet_traits<Scalar>::type Packet;
206 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& x, const Packet& q) const {
207 return internal::pzeta(x, q);
208 }
209};
210template <typename Scalar>
211struct functor_traits<scalar_zeta_op<Scalar> > {
212 enum {
213 // Guesstimate
214 Cost = 10 * NumTraits<Scalar>::MulCost + 5 * NumTraits<Scalar>::AddCost,
215 PacketAccess = packet_traits<Scalar>::HasZeta
216 };
217};
218
223template <typename Scalar>
224struct scalar_polygamma_op {
225 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& n, const Scalar& x) const {
226 using numext::polygamma;
227 return polygamma(n, x);
228 }
229 typedef typename packet_traits<Scalar>::type Packet;
230 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& n, const Packet& x) const {
231 return internal::ppolygamma(n, x);
232 }
233};
234template <typename Scalar>
235struct functor_traits<scalar_polygamma_op<Scalar> > {
236 enum {
237 // Guesstimate
238 Cost = 10 * NumTraits<Scalar>::MulCost + 5 * NumTraits<Scalar>::AddCost,
239 PacketAccess = packet_traits<Scalar>::HasPolygamma
240 };
241};
242
247template <typename Scalar>
248struct scalar_erf_op {
249 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a) const { return numext::erf(a); }
250 template <typename Packet>
251 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& x) const {
252 return perf(x);
253 }
254};
255template <typename Scalar>
256struct functor_traits<scalar_erf_op<Scalar> > {
257 enum {
258 PacketAccess = packet_traits<Scalar>::HasErf,
259 Cost = (PacketAccess
260#ifdef EIGEN_VECTORIZE_FMA
261 // TODO(rmlarsen): Move the FMA cost model to a central location.
262 // Haswell can issue 2 add/mul/madd per cycle.
263 // 10 pmadd, 2 pmul, 1 div, 2 other
264 ? (2 * NumTraits<Scalar>::AddCost + 7 * NumTraits<Scalar>::MulCost +
265 scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value)
266#else
267 ? (12 * NumTraits<Scalar>::AddCost + 12 * NumTraits<Scalar>::MulCost +
268 scalar_div_cost<Scalar, packet_traits<Scalar>::HasDiv>::value)
269#endif
270 // Assume for simplicity that this is as expensive as an exp().
271 : (functor_traits<scalar_exp_op<Scalar> >::Cost))
272 };
273};
274
280template <typename Scalar>
281struct scalar_erfc_op {
282 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a) const {
283 using numext::erfc;
284 return erfc(a);
285 }
286 typedef typename packet_traits<Scalar>::type Packet;
287 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const { return internal::perfc(a); }
288};
289template <typename Scalar>
290struct functor_traits<scalar_erfc_op<Scalar> > {
291 enum {
292 // Guesstimate
293 Cost = 10 * NumTraits<Scalar>::MulCost + 5 * NumTraits<Scalar>::AddCost,
294 PacketAccess = packet_traits<Scalar>::HasErfc
295 };
296};
297
303template <typename Scalar>
304struct scalar_ndtri_op {
305 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(const Scalar& a) const {
306 using numext::ndtri;
307 return ndtri(a);
308 }
309 typedef typename packet_traits<Scalar>::type Packet;
310 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(const Packet& a) const { return internal::pndtri(a); }
311};
312template <typename Scalar>
313struct functor_traits<scalar_ndtri_op<Scalar> > {
314 enum {
315 // On average, we are evaluating rational functions with degree N=9 in the
316 // numerator and denominator. This results in 2*N additions and 2*N
317 // multiplications.
318 Cost = 18 * NumTraits<Scalar>::MulCost + 18 * NumTraits<Scalar>::AddCost,
319 PacketAccess = packet_traits<Scalar>::HasNdtri
320 };
321};
322
323} // end namespace internal
324
325} // end namespace Eigen
326
327#endif // EIGEN_SPECIALFUNCTIONS_FUNCTORS_H
Namespace containing all symbols from the Eigen library.
const Eigen::CwiseBinaryOp< Eigen::internal::scalar_igammac_op< typename Derived::Scalar >, const Derived, const ExponentDerived > igammac(const Eigen::ArrayBase< Derived > &a, const Eigen::ArrayBase< ExponentDerived > &x)
Definition SpecialFunctionsArrayAPI.h:86
const Eigen::CwiseBinaryOp< Eigen::internal::scalar_igamma_der_a_op< typename Derived::Scalar >, const Derived, const ExponentDerived > igamma_der_a(const Eigen::ArrayBase< Derived > &a, const Eigen::ArrayBase< ExponentDerived > &x)
Definition SpecialFunctionsArrayAPI.h:49
const Eigen::CwiseTernaryOp< Eigen::internal::scalar_betainc_op< typename ArgXDerived::Scalar >, const ArgADerived, const ArgBDerived, const ArgXDerived > betainc(const Eigen::ArrayBase< ArgADerived > &a, const Eigen::ArrayBase< ArgBDerived > &b, const Eigen::ArrayBase< ArgXDerived > &x)
Definition SpecialFunctionsArrayAPI.h:120
const Eigen::CwiseBinaryOp< Eigen::internal::scalar_gamma_sample_der_alpha_op< typename AlphaDerived::Scalar >, const AlphaDerived, const SampleDerived > gamma_sample_der_alpha(const Eigen::ArrayBase< AlphaDerived > &alpha, const Eigen::ArrayBase< SampleDerived > &sample)
Definition SpecialFunctionsArrayAPI.h:69
const Eigen::CwiseBinaryOp< Eigen::internal::scalar_polygamma_op< typename DerivedX::Scalar >, const DerivedN, const DerivedX > polygamma(const Eigen::ArrayBase< DerivedN > &n, const Eigen::ArrayBase< DerivedX > &x)
Definition SpecialFunctionsArrayAPI.h:103
const Eigen::CwiseBinaryOp< Eigen::internal::scalar_igamma_op< typename Derived::Scalar >, const Derived, const ExponentDerived > igamma(const Eigen::ArrayBase< Derived > &a, const Eigen::ArrayBase< ExponentDerived > &x)
Definition SpecialFunctionsArrayAPI.h:31
const Eigen::CwiseBinaryOp< Eigen::internal::scalar_zeta_op< typename DerivedX::Scalar >, const DerivedX, const DerivedQ > zeta(const Eigen::ArrayBase< DerivedX > &x, const Eigen::ArrayBase< DerivedQ > &q)
Definition SpecialFunctionsArrayAPI.h:141