Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorArgMax.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2015 Eugene Brevdo <ebrevdo@gmail.com>
5// Benoit Steiner <benoit.steiner.goog@gmail.com>
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_TENSOR_TENSOR_ARG_MAX_H
13#define EIGEN_TENSOR_TENSOR_ARG_MAX_H
14
15// IWYU pragma: private
16#include "./InternalHeaderCheck.h"
17
18namespace Eigen {
19namespace internal {
20
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;
29};
30
31template <typename XprType>
32struct eval<TensorIndexPairOp<XprType>, Eigen::Dense> {
33 typedef const TensorIndexPairOp<XprType> EIGEN_DEVICE_REF type;
34};
35
36} // end namespace internal
37
43template <typename XprType>
44class TensorIndexPairOp : public TensorBase<TensorIndexPairOp<XprType>, ReadOnlyAccessors> {
45 public:
46 typedef typename Eigen::internal::traits<TensorIndexPairOp>::Scalar Scalar;
47 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
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;
52
53 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorIndexPairOp(const XprType& expr) : m_xpr(expr) {}
54
55 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename XprType::Nested>& expression() const { return m_xpr; }
56
57 protected:
58 typename XprType::Nested m_xpr;
59};
60
61// Eval as rvalue
62template <typename ArgType, typename Device>
63struct TensorEvaluator<const TensorIndexPairOp<ArgType>, Device> {
64 typedef TensorIndexPairOp<ArgType> XprType;
65 typedef typename XprType::Index Index;
66 typedef typename XprType::Scalar Scalar;
67 typedef typename XprType::CoeffReturnType CoeffReturnType;
68
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;
73
74 enum {
75 IsAligned = false,
76 PacketAccess = false,
77 BlockAccess = false,
79 CoordAccess = false, // to be implemented
80 RawAccess = false
81 };
82 static constexpr int Layout = TensorEvaluator<ArgType, Device>::Layout;
83
84 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
85 typedef internal::TensorBlockNotImplemented TensorBlock;
86 //===--------------------------------------------------------------------===//
87
88 EIGEN_STRONG_INLINE TensorEvaluator(const XprType& op, const Device& device) : m_impl(op.expression(), device) {}
89
90 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_impl.dimensions(); }
91
92 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType /*data*/) {
93 m_impl.evalSubExprsIfNeeded(nullptr);
94 return true;
95 }
96 EIGEN_STRONG_INLINE void cleanup() { m_impl.cleanup(); }
97
98 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const {
99 return CoeffReturnType(index, m_impl.coeff(index));
100 }
101
102 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool vectorized) const {
103 return m_impl.costPerCoeff(vectorized) + TensorOpCost(0, 0, 1);
104 }
105
106 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return nullptr; }
107
108 protected:
109 TensorEvaluator<ArgType, Device> m_impl;
110};
111
112namespace internal {
113
120template <typename ReduceOp, typename Dims, typename XprType>
121struct traits<TensorPairReducerOp<ReduceOp, Dims, XprType>> : public traits<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;
128};
129
130template <typename ReduceOp, typename Dims, typename XprType>
131struct eval<TensorPairReducerOp<ReduceOp, Dims, XprType>, Eigen::Dense> {
132 typedef const TensorPairReducerOp<ReduceOp, Dims, XprType> EIGEN_DEVICE_REF type;
133};
134
135} // end namespace internal
136
137template <typename ReduceOp, typename Dims, typename XprType>
138class TensorPairReducerOp : public TensorBase<TensorPairReducerOp<ReduceOp, Dims, XprType>, ReadOnlyAccessors> {
139 public:
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;
146
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) {}
150
151 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename XprType::Nested>& expression() const { return m_xpr; }
152
153 EIGEN_DEVICE_FUNC const ReduceOp& reduce_op() const { return m_reduce_op; }
154
155 EIGEN_DEVICE_FUNC const Dims& reduce_dims() const { return m_reduce_dims; }
156
157 EIGEN_DEVICE_FUNC Index return_dim() const { return m_return_dim; }
158
159 protected:
160 typename XprType::Nested m_xpr;
161 const ReduceOp m_reduce_op;
162 const Index m_return_dim;
163 const Dims m_reduce_dims;
164};
165
166// Eval as rvalue
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;
181
182 enum {
183 IsAligned = false,
184 PacketAccess = false,
185 BlockAccess = false,
186 PreferBlockAccess = TensorEvaluator<ArgType, Device>::PreferBlockAccess,
187 CoordAccess = false, // to be implemented
188 RawAccess = false
189 };
190 static constexpr int Layout =
191 TensorEvaluator<const TensorReductionOp<ReduceOp, Dims, const TensorIndexPairOp<ArgType>>, Device>::Layout;
192 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
193 typedef internal::TensorBlockNotImplemented TensorBlock;
194 //===--------------------------------------------------------------------===//
195
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;
204 } else {
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;
207 }
208 // If m_return_dim is not a valid index, returns 1 or this can crash on Windows.
209 m_stride_div =
210 ((m_return_dim >= 0) && (m_return_dim < static_cast<Index>(m_strides.size()))) ? m_strides[m_return_dim] : 1;
211 }
212
213 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_impl.dimensions(); }
214
215 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType /*data*/) {
216 m_impl.evalSubExprsIfNeeded(nullptr);
217 return true;
218 }
219 EIGEN_STRONG_INLINE void cleanup() { m_impl.cleanup(); }
220
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;
224 }
225
226 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return nullptr; }
227
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);
232 }
233
234 private:
235 EIGEN_DEVICE_FUNC void gen_strides(const InputDimensions& dims, StrideDims& strides) {
236 if (m_return_dim < 0) {
237 return; // Won't be using the strides.
238 }
239 eigen_assert(m_return_dim < NumDims && "Asking to convert index to a dimension outside of the rank");
240
241 // Calculate m_stride_div and m_stride_mod, which are used to
242 // calculate the value of an index w.r.t. the m_return_dim.
243 EIGEN_IF_CONSTEXPR (Layout == static_cast<int>(ColMajor)) {
244 strides[0] = 1;
245 for (int i = 1; i < NumDims; ++i) {
246 strides[i] = strides[i - 1] * dims[i - 1];
247 }
248 } else {
249 strides[NumDims - 1] = 1;
250 for (int i = NumDims - 2; i >= 0; --i) {
251 strides[i] = strides[i + 1] * dims[i + 1];
252 }
253 }
254 }
255
256 protected:
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;
261 Index m_stride_mod;
262 Index m_stride_div;
263};
264
265} // end namespace Eigen
266
267#endif // EIGEN_TENSOR_TENSOR_ARG_MAX_H
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