12#ifndef EIGEN_TENSOR_TENSOR_ARG_MAX_H
13#define EIGEN_TENSOR_TENSOR_ARG_MAX_H
16#include "./InternalHeaderCheck.h"
21template <
typename XprType>
22struct traits<TensorIndexPairOp<XprType>> :
public traits<XprType> {
23 typedef traits<XprType> XprTraits;
24 typedef typename XprTraits::StorageKind StorageKind;
25 typedef typename XprTraits::Index Index;
26 typedef Pair<Index, typename XprTraits::Scalar> Scalar;
27 static constexpr int NumDimensions = XprTraits::NumDimensions;
28 static constexpr int Layout = XprTraits::Layout;
31template <
typename XprType>
32struct eval<TensorIndexPairOp<XprType>, Eigen::Dense> {
33 typedef const TensorIndexPairOp<XprType> EIGEN_DEVICE_REF type;
43template <
typename XprType>
44class TensorIndexPairOp :
public TensorBase<TensorIndexPairOp<XprType>, ReadOnlyAccessors> {
46 typedef typename Eigen::internal::traits<TensorIndexPairOp>::Scalar Scalar;
48 typedef typename Eigen::internal::ref_selector<TensorIndexPairOp>::type Nested;
49 typedef typename Eigen::internal::traits<TensorIndexPairOp>::StorageKind StorageKind;
50 typedef typename Eigen::internal::traits<TensorIndexPairOp>::Index Index;
51 typedef Pair<Index, typename XprType::CoeffReturnType> CoeffReturnType;
53 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorIndexPairOp(
const XprType& expr) : m_xpr(expr) {}
55 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename XprType::Nested>& expression()
const {
return m_xpr; }
58 typename XprType::Nested m_xpr;
62template <
typename ArgType,
typename Device>
65 typedef typename XprType::Index Index;
66 typedef typename XprType::Scalar Scalar;
69 typedef typename TensorEvaluator<ArgType, Device>::Dimensions
Dimensions;
70 static constexpr int NumDims = internal::array_size<Dimensions>::value;
71 typedef StorageMemory<CoeffReturnType, Device> Storage;
72 typedef typename Storage::Type EvaluatorPointerType;
82 static constexpr int Layout = TensorEvaluator<ArgType, Device>::Layout;
85 typedef internal::TensorBlockNotImplemented TensorBlock;
88 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device) : m_impl(op.expression(), device) {}
90 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_impl.dimensions(); }
92 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType ) {
93 m_impl.evalSubExprsIfNeeded(
nullptr);
96 EIGEN_STRONG_INLINE
void cleanup() { m_impl.cleanup(); }
98 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
99 return CoeffReturnType(index, m_impl.coeff(index));
102 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
103 return m_impl.costPerCoeff(vectorized) + TensorOpCost(0, 0, 1);
106 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return nullptr; }
109 TensorEvaluator<ArgType, Device> m_impl;
120template <
typename ReduceOp,
typename Dims,
typename XprType>
122 typedef traits<XprType> XprTraits;
123 typedef typename XprTraits::StorageKind StorageKind;
124 typedef typename XprTraits::Index Index;
125 typedef Index Scalar;
126 static constexpr int NumDimensions = XprTraits::NumDimensions - array_size<Dims>::value;
127 static constexpr int Layout = XprTraits::Layout;
130template <
typename ReduceOp,
typename Dims,
typename XprType>
137template <
typename ReduceOp,
typename Dims,
typename XprType>
138class TensorPairReducerOp :
public TensorBase<TensorPairReducerOp<ReduceOp, Dims, XprType>, ReadOnlyAccessors> {
140 typedef typename Eigen::internal::traits<TensorPairReducerOp>::Scalar Scalar;
141 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
142 typedef typename Eigen::internal::ref_selector<TensorPairReducerOp>::type Nested;
143 typedef typename Eigen::internal::traits<TensorPairReducerOp>::StorageKind StorageKind;
144 typedef typename Eigen::internal::traits<TensorPairReducerOp>::Index Index;
145 typedef Index CoeffReturnType;
147 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorPairReducerOp(
const XprType& expr,
const ReduceOp& reduce_op,
148 const Index return_dim,
const Dims& reduce_dims)
149 : m_xpr(expr), m_reduce_op(reduce_op), m_return_dim(return_dim), m_reduce_dims(reduce_dims) {}
151 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename XprType::Nested>& expression()
const {
return m_xpr; }
153 EIGEN_DEVICE_FUNC
const ReduceOp& reduce_op()
const {
return m_reduce_op; }
155 EIGEN_DEVICE_FUNC
const Dims& reduce_dims()
const {
return m_reduce_dims; }
157 EIGEN_DEVICE_FUNC Index return_dim()
const {
return m_return_dim; }
160 typename XprType::Nested m_xpr;
161 const ReduceOp m_reduce_op;
162 const Index m_return_dim;
163 const Dims m_reduce_dims;
167template <
typename ReduceOp,
typename Dims,
typename ArgType,
typename Device>
168struct TensorEvaluator<const TensorPairReducerOp<ReduceOp, Dims, ArgType>, Device> {
169 typedef TensorPairReducerOp<ReduceOp, Dims, ArgType> XprType;
170 typedef typename XprType::Index Index;
171 typedef typename XprType::Scalar Scalar;
172 typedef typename XprType::CoeffReturnType CoeffReturnType;
173 typedef typename TensorIndexPairOp<ArgType>::CoeffReturnType PairType;
174 typedef typename TensorEvaluator<const TensorReductionOp<ReduceOp, Dims, const TensorIndexPairOp<ArgType>>,
175 Device>::Dimensions Dimensions;
176 typedef typename TensorEvaluator<const TensorIndexPairOp<ArgType>, Device>::Dimensions InputDimensions;
177 static constexpr int NumDims = internal::array_size<InputDimensions>::value;
178 typedef array<Index, NumDims> StrideDims;
179 typedef StorageMemory<CoeffReturnType, Device> Storage;
180 typedef typename Storage::Type EvaluatorPointerType;
184 PacketAccess =
false,
186 PreferBlockAccess = TensorEvaluator<ArgType, Device>::PreferBlockAccess,
190 static constexpr int Layout =
191 TensorEvaluator<const TensorReductionOp<ReduceOp, Dims, const TensorIndexPairOp<ArgType>>, Device>::Layout;
193 typedef internal::TensorBlockNotImplemented TensorBlock;
196 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device)
197 : m_orig_impl(op.expression(), device),
198 m_impl(op.expression().index_pairs().reduce(op.reduce_dims(), op.reduce_op()), device),
199 m_return_dim(op.return_dim()) {
200 gen_strides(m_orig_impl.dimensions(), m_strides);
201 EIGEN_IF_CONSTEXPR (Layout ==
static_cast<int>(
ColMajor)) {
202 const Index total_size = internal::array_prod(m_orig_impl.dimensions());
203 m_stride_mod = (m_return_dim < NumDims - 1) ? m_strides[m_return_dim + 1] : total_size;
205 const Index total_size = internal::array_prod(m_orig_impl.dimensions());
206 m_stride_mod = (m_return_dim > 0) ? m_strides[m_return_dim - 1] : total_size;
210 ((m_return_dim >= 0) && (m_return_dim < static_cast<Index>(m_strides.size()))) ? m_strides[m_return_dim] : 1;
213 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_impl.dimensions(); }
215 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType ) {
216 m_impl.evalSubExprsIfNeeded(
nullptr);
219 EIGEN_STRONG_INLINE
void cleanup() { m_impl.cleanup(); }
221 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
222 const PairType v = m_impl.coeff(index);
223 return (m_return_dim < 0) ? v.first : (v.first % m_stride_mod) / m_stride_div;
226 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return nullptr; }
228 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
229 const double compute_cost =
230 1.0 + (m_return_dim < 0 ? 0.0 : (TensorOpCost::ModCost<Index>() + TensorOpCost::DivCost<Index>()));
231 return m_orig_impl.costPerCoeff(vectorized) + m_impl.costPerCoeff(vectorized) + TensorOpCost(0, 0, compute_cost);
235 EIGEN_DEVICE_FUNC
void gen_strides(
const InputDimensions& dims, StrideDims& strides) {
236 if (m_return_dim < 0) {
239 eigen_assert(m_return_dim < NumDims &&
"Asking to convert index to a dimension outside of the rank");
243 EIGEN_IF_CONSTEXPR (Layout ==
static_cast<int>(
ColMajor)) {
245 for (
int i = 1; i < NumDims; ++i) {
246 strides[i] = strides[i - 1] * dims[i - 1];
249 strides[NumDims - 1] = 1;
250 for (
int i = NumDims - 2; i >= 0; --i) {
251 strides[i] = strides[i + 1] * dims[i + 1];
257 TensorEvaluator<const TensorIndexPairOp<ArgType>, Device> m_orig_impl;
258 TensorEvaluator<const TensorReductionOp<ReduceOp, Dims, const TensorIndexPairOp<ArgType>>, Device> m_impl;
259 const Index m_return_dim;
260 StrideDims m_strides;
The tensor base class.
Definition TensorForwardDeclarations.h:69
Tensor + Index Pair class.
Definition TensorArgMax.h:44
Converts to Tensor<Pair<Index, Scalar> > and reduces to Tensor<Index>.
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47