11#ifndef EIGEN_NULLARY_FUNCTORS_H
12#define EIGEN_NULLARY_FUNCTORS_H
15#include "../InternalHeaderCheck.h"
21template <
typename Scalar>
22struct scalar_constant_op {
23 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE scalar_constant_op(
const Scalar& other) : m_other(other) {}
24 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()()
const {
return m_other; }
25 template <
typename PacketType>
26 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetOp()
const {
27 return internal::pset1<PacketType>(m_other);
31template <
typename Scalar>
32struct functor_traits<scalar_constant_op<Scalar>> {
35 PacketAccess = packet_traits<Scalar>::Vectorizable,
40template <
typename Scalar>
41struct scalar_zero_op {
42 EIGEN_DEVICE_FUNC scalar_zero_op() =
default;
43 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()()
const {
return Scalar(0); }
44 template <
typename PacketType>
45 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketType packetOp()
const {
46 return internal::pzero<PacketType>(PacketType());
49template <
typename Scalar>
50struct functor_traits<scalar_zero_op<Scalar>> : functor_traits<scalar_constant_op<Scalar>> {};
52template <
typename Scalar>
53struct scalar_identity_op {
54 template <
typename IndexType>
55 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()(IndexType row, IndexType col)
const {
56 return row == col ? Scalar(1) : Scalar(0);
59template <
typename Scalar>
60struct functor_traits<scalar_identity_op<Scalar>> {
61 enum { Cost = NumTraits<Scalar>::AddCost, PacketAccess =
false, IsRepeatable =
true };
64template <
typename Scalar,
bool IsInteger>
65struct linspaced_op_impl;
67template <
typename Scalar>
68struct linspaced_op_impl<Scalar, false> {
69 using RealScalar =
typename NumTraits<Scalar>::Real;
71 EIGEN_DEVICE_FUNC
constexpr linspaced_op_impl(
const Scalar& low,
const Scalar& high, Index num_steps)
74 m_size1(num_steps == 1 ? 1 : num_steps - 1),
75 m_step(num_steps == 1 ? Scalar() : Scalar((high - low) / RealScalar(num_steps - 1))),
76 m_flip(numext::abs(high) < numext::abs(low)) {}
78 template <
typename IndexType>
79 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()(IndexType i)
const {
81 return (i == 0) ? m_low : Scalar(m_high - RealScalar(m_size1 - i) * m_step);
83 return (i == m_size1) ? m_high : Scalar(m_low + RealScalar(i) * m_step);
86 template <
typename Packet,
typename IndexType>
87 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(IndexType i)
const {
90 Packet low = pset1<Packet>(m_low);
91 Packet high = pset1<Packet>(m_high);
92 Packet step = pset1<Packet>(m_step);
94 Packet pi = plset<Packet>(Scalar(i - m_size1));
95 Packet res = pmadd(step, pi, high);
96 Packet mask = pcmp_lt(pzero(res), plset<Packet>(Scalar(i)));
97 return pselect<Packet>(mask, res, low);
99 Packet pi = plset<Packet>(Scalar(i));
100 Packet res = pmadd(step, pi, low);
101 Packet mask = pcmp_lt(pi, pset1<Packet>(Scalar(m_size1)));
102 return pselect<Packet>(mask, res, high);
113template <
typename Scalar>
114struct linspaced_op_impl<Scalar, true> {
115 EIGEN_DEVICE_FUNC
constexpr linspaced_op_impl(
const Scalar& low,
const Scalar& high, Index num_steps)
117 m_multiplier((high - low) / convert_index<Scalar>(num_steps <= 1 ? 1 : num_steps - 1)),
118 m_divisor(convert_index<Scalar>((high >= low ? num_steps : -num_steps) + (high - low)) /
119 ((numext::abs(high - low) + 1) == 0 ? 1 : (numext::abs(high - low) + 1))),
120 m_use_divisor(num_steps > 1 && (numext::abs(high - low) + 1) < num_steps) {}
122 template <
typename IndexType>
123 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()(IndexType i)
const {
125 return m_low + convert_index<Scalar>(i) / m_divisor;
127 return m_low + convert_index<Scalar>(i) * m_multiplier;
131 const Scalar m_multiplier;
132 const Scalar m_divisor;
133 const bool m_use_divisor;
141template <
typename Scalar>
143template <
typename Scalar>
144struct functor_traits<linspaced_op<Scalar>> {
147 PacketAccess = (!NumTraits<Scalar>::IsInteger) && packet_traits<Scalar>::HasSetLinear,
153template <
typename Scalar>
155 EIGEN_DEVICE_FUNC
constexpr linspaced_op(
const Scalar& low,
const Scalar& high, Index num_steps)
156 : impl((num_steps == 1 ? high : low), high, num_steps) {}
158 template <
typename IndexType>
159 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()(IndexType i)
const {
163 template <
typename Packet,
typename IndexType>
164 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(IndexType i)
const {
165 return impl.template packetOp<Packet>(i);
170 const linspaced_op_impl<Scalar, NumTraits<Scalar>::IsInteger> impl;
173template <
typename Scalar>
174struct equalspaced_op {
175 EIGEN_DEVICE_FUNC
constexpr equalspaced_op(
const Scalar& start,
const Scalar& step) : m_start(start), m_step(step) {}
176 template <
typename IndexType>
177 EIGEN_DEVICE_FUNC
constexpr EIGEN_STRONG_INLINE Scalar operator()(IndexType i)
const {
178 return m_start + m_step *
static_cast<Scalar
>(i);
180 template <
typename Packet,
typename IndexType>
181 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Packet packetOp(IndexType i)
const {
182 const Packet cst_start = pset1<Packet>(m_start);
183 const Packet cst_step = pset1<Packet>(m_step);
184 const Packet cst_lin0 = plset<Packet>(Scalar(0));
185 const Packet cst_offset = pmadd(cst_lin0, cst_step, cst_start);
187 Packet i_packet = pset1<Packet>(
static_cast<Scalar
>(i));
188 return pmadd(i_packet, cst_step, cst_offset);
190 const Scalar m_start;
194template <
typename Scalar>
195struct functor_traits<equalspaced_op<Scalar>> {
197 Cost = NumTraits<Scalar>::AddCost + NumTraits<Scalar>::MulCost,
199 packet_traits<Scalar>::HasSetLinear && packet_traits<Scalar>::HasMul && packet_traits<Scalar>::HasAdd,
208template <
typename Functor>
209using functor_has_linear_access = bool_constant<!has_binary_operator<Functor>::value>;