Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorCustomOp.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014 Benoit Steiner <benoit.steiner.goog@gmail.com>
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_TENSOR_TENSOR_CUSTOM_OP_H
12#define EIGEN_TENSOR_TENSOR_CUSTOM_OP_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20template <typename CustomUnaryFunc, typename XprType>
21struct traits<TensorCustomUnaryOp<CustomUnaryFunc, XprType> > {
22 typedef typename XprType::Scalar Scalar;
23 typedef typename XprType::StorageKind StorageKind;
24 typedef typename XprType::Index Index;
25 // The functor's dimensions() determines the output shape, so the rank of the
26 // result may differ from the rank of the input. The argument is spelled
27 // exactly as in the evaluator's call so both resolve to the same overload.
28 using CustomDimensions = remove_all_t<decltype(std::declval<const CustomUnaryFunc&>().dimensions(
29 std::declval<const remove_all_t<typename XprType::Nested>&>()))>;
30 static constexpr ptrdiff_t CustomRank = array_size<CustomDimensions>::value;
31 static_assert(CustomRank >= 0,
32 "The dimensions() method of a custom tensor functor must return a fixed-rank "
33 "array-like type such as DSizes<Index, Rank>.");
34 // Clamped so a failed assertion doesn't cascade into DSizes<Index, -1> errors.
35 static constexpr int NumDimensions = CustomRank < 0 ? 1 : static_cast<int>(CustomRank);
36 static constexpr int Layout = traits<XprType>::Layout;
37 typedef typename traits<XprType>::PointerType PointerType;
38 enum { Flags = 0 };
39};
40
41template <typename CustomUnaryFunc, typename XprType>
42struct eval<TensorCustomUnaryOp<CustomUnaryFunc, XprType>, Eigen::Dense> {
43 typedef const TensorCustomUnaryOp<CustomUnaryFunc, XprType> EIGEN_DEVICE_REF type;
44};
45
46} // end namespace internal
47
53template <typename CustomUnaryFunc, typename XprType>
54class TensorCustomUnaryOp : public TensorBase<TensorCustomUnaryOp<CustomUnaryFunc, XprType>, ReadOnlyAccessors> {
55 public:
56 typedef typename internal::traits<TensorCustomUnaryOp>::Scalar Scalar;
57 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
58 typedef typename XprType::CoeffReturnType CoeffReturnType;
59 typedef typename internal::ref_selector<TensorCustomUnaryOp>::non_const_type Nested;
60 typedef typename internal::traits<TensorCustomUnaryOp>::StorageKind StorageKind;
61 typedef typename internal::traits<TensorCustomUnaryOp>::Index Index;
62
63 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCustomUnaryOp(const XprType& expr, const CustomUnaryFunc& func)
64 : m_expr(expr), m_func(func) {}
65
66 EIGEN_DEVICE_FUNC const CustomUnaryFunc& func() const { return m_func; }
67
68 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename XprType::Nested>& expression() const { return m_expr; }
69
70 protected:
71 typename XprType::Nested m_expr;
72 const CustomUnaryFunc m_func;
73};
74
75// Eval as rvalue
76template <typename CustomUnaryFunc, typename XprType, typename Device>
77struct TensorEvaluator<const TensorCustomUnaryOp<CustomUnaryFunc, XprType>, Device> {
79 typedef typename internal::traits<ArgType>::Index Index;
80 static constexpr int NumDims = internal::traits<ArgType>::NumDimensions;
81 typedef DSizes<Index, NumDims> Dimensions;
82 typedef std::remove_const_t<typename ArgType::Scalar> Scalar;
83 typedef std::remove_const_t<typename XprType::CoeffReturnType> CoeffReturnType;
84 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
85 static constexpr int PacketSize = PacketType<CoeffReturnType, Device>::size;
86 typedef typename Eigen::internal::traits<XprType>::PointerType TensorPointerType;
87 typedef StorageMemory<CoeffReturnType, Device> Storage;
88 typedef typename Storage::Type EvaluatorPointerType;
89
90 static constexpr int Layout = TensorEvaluator<XprType, Device>::Layout;
91 enum {
92 IsAligned = false,
93 PacketAccess = (PacketType<CoeffReturnType, Device>::size > 1),
94 // The custom op is eagerly evaluated into a dense buffer (m_result), so
95 // blocks and raw storage can be served straight from it, exactly like
96 // TensorForcedEvalOp. Without these flags a custom op disables tiled
97 // evaluation for any expression containing it and hides its buffer from
98 // consumers with data()-based fast paths.
99 BlockAccess = internal::is_arithmetic<CoeffReturnType>::value,
100 PreferBlockAccess = false,
101 CoordAccess = false, // to be implemented
102 RawAccess = true
103 };
104
105 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
106 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
107 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
108
109 typedef typename internal::TensorMaterializedBlock<CoeffReturnType, NumDims, Layout, Index> TensorBlock;
110 //===--------------------------------------------------------------------===//
111
112 // The functor's dimensions() may return an index type that promotes to Index.
113 EIGEN_STRONG_INLINE TensorEvaluator(const ArgType& op, const Device& device)
114 : m_dimensions(op.func().dimensions(op.expression())), m_op(op), m_device(device), m_result(nullptr) {}
115
116 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_dimensions; }
117
118 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType data) {
119 if (data) {
120 evalTo(data);
121 return false;
122 } else {
123 m_result = static_cast<EvaluatorPointerType>(
124 m_device.get((CoeffReturnType*)m_device.allocate_temp(dimensions().TotalSize() * sizeof(CoeffReturnType))));
125 evalTo(m_result);
126 return true;
127 }
128 }
129
130 EIGEN_STRONG_INLINE void cleanup() {
131 if (m_result) {
132 m_device.deallocate_temp(m_result);
133 m_result = nullptr;
134 }
135 }
136
137 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const { return m_result[index]; }
138
139 template <int LoadMode>
140 EIGEN_DEVICE_FUNC PacketReturnType packet(Index index) const {
141 return internal::ploadt<PacketReturnType, LoadMode>(m_result + index);
142 }
143
144 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool vectorized) const {
145 // TODO(rmlarsen): Extend CustomOp API to return its cost estimate.
146 return TensorOpCost(sizeof(CoeffReturnType), 0, 0, vectorized, PacketSize);
147 }
148
149 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements() const {
150 return internal::TensorBlockResourceRequirements::any();
151 }
152
153 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
154 bool /*root_of_expr_ast*/ = false) const {
155 eigen_assert(m_result != nullptr);
156 return TensorBlock::materialize(m_result, m_dimensions, desc, scratch);
157 }
158
159 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return m_result; }
160
161 protected:
162 void evalTo(EvaluatorPointerType data) {
163 TensorMap<Tensor<CoeffReturnType, NumDims, Layout, Index> > result(m_device.get(data), m_dimensions);
164 m_op.func().eval(m_op.expression(), result, m_device);
165 }
166
167 Dimensions m_dimensions;
168 const ArgType m_op;
169 const Device EIGEN_DEVICE_REF m_device;
170 EvaluatorPointerType m_result;
171};
172
180namespace internal {
181template <typename CustomBinaryFunc, typename LhsXprType, typename RhsXprType>
182struct traits<TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType> > {
183 typedef typename internal::promote_storage_type<typename LhsXprType::Scalar, typename RhsXprType::Scalar>::ret Scalar;
184 typedef typename internal::promote_storage_type<typename LhsXprType::CoeffReturnType,
185 typename RhsXprType::CoeffReturnType>::ret CoeffReturnType;
186 typedef typename promote_storage_type<typename traits<LhsXprType>::StorageKind,
187 typename traits<RhsXprType>::StorageKind>::ret StorageKind;
188 typedef
189 typename promote_index_type<typename traits<LhsXprType>::Index, typename traits<RhsXprType>::Index>::type Index;
190 // The functor's dimensions() determines the output shape, so the rank of the
191 // result may differ from the ranks of the inputs. The arguments are spelled
192 // exactly as in the evaluator's call so both resolve to the same overload.
193 using CustomDimensions = remove_all_t<decltype(std::declval<const CustomBinaryFunc&>().dimensions(
194 std::declval<const remove_all_t<typename LhsXprType::Nested>&>(),
195 std::declval<const remove_all_t<typename RhsXprType::Nested>&>()))>;
196 static constexpr ptrdiff_t CustomRank = array_size<CustomDimensions>::value;
197 static_assert(CustomRank >= 0,
198 "The dimensions() method of a custom tensor functor must return a fixed-rank "
199 "array-like type such as DSizes<Index, Rank>.");
200 // Clamped so a failed assertion doesn't cascade into DSizes<Index, -1> errors.
201 static constexpr int NumDimensions = CustomRank < 0 ? 1 : static_cast<int>(CustomRank);
202 static constexpr int Layout = traits<LhsXprType>::Layout;
203 typedef std::conditional_t<Pointer_type_promotion<typename LhsXprType::Scalar, Scalar>::val,
204 typename traits<LhsXprType>::PointerType, typename traits<RhsXprType>::PointerType>
205 PointerType;
206 enum { Flags = 0 };
207};
208
209template <typename CustomBinaryFunc, typename LhsXprType, typename RhsXprType>
210struct eval<TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>, Eigen::Dense> {
211 typedef const TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>& type;
212};
213
214} // end namespace internal
215
216template <typename CustomBinaryFunc, typename LhsXprType, typename RhsXprType>
217class TensorCustomBinaryOp
218 : public TensorBase<TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>, ReadOnlyAccessors> {
219 public:
220 typedef typename internal::traits<TensorCustomBinaryOp>::Scalar Scalar;
221 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
222 typedef typename internal::traits<TensorCustomBinaryOp>::CoeffReturnType CoeffReturnType;
223 typedef typename internal::ref_selector<TensorCustomBinaryOp>::non_const_type Nested;
224 typedef typename internal::traits<TensorCustomBinaryOp>::StorageKind StorageKind;
225 typedef typename internal::traits<TensorCustomBinaryOp>::Index Index;
226
227 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCustomBinaryOp(const LhsXprType& lhs, const RhsXprType& rhs,
228 const CustomBinaryFunc& func)
229
230 : m_lhs_xpr(lhs), m_rhs_xpr(rhs), m_func(func) {}
231
232 EIGEN_DEVICE_FUNC const CustomBinaryFunc& func() const { return m_func; }
233
234 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression() const {
235 return m_lhs_xpr;
236 }
237
238 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression() const {
239 return m_rhs_xpr;
240 }
241
242 protected:
243 typename LhsXprType::Nested m_lhs_xpr;
244 typename RhsXprType::Nested m_rhs_xpr;
245 const CustomBinaryFunc m_func;
246};
247
248// Eval as rvalue
249template <typename CustomBinaryFunc, typename LhsXprType, typename RhsXprType, typename Device>
250struct TensorEvaluator<const TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>, Device> {
252 typedef typename internal::traits<XprType>::Index Index;
253 static constexpr int NumDims = internal::traits<XprType>::NumDimensions;
254 typedef DSizes<Index, NumDims> Dimensions;
255 typedef typename XprType::Scalar Scalar;
256 typedef std::remove_const_t<typename XprType::CoeffReturnType> CoeffReturnType;
257 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
258 static constexpr int PacketSize = PacketType<CoeffReturnType, Device>::size;
259
260 typedef typename Eigen::internal::traits<XprType>::PointerType TensorPointerType;
261 typedef StorageMemory<CoeffReturnType, Device> Storage;
262 typedef typename Storage::Type EvaluatorPointerType;
263
264 static constexpr int Layout = TensorEvaluator<LhsXprType, Device>::Layout;
265 enum {
266 IsAligned = false,
267 PacketAccess = (PacketType<CoeffReturnType, Device>::size > 1),
268 // See the unary evaluator above: serve blocks and raw storage from the
269 // eagerly materialized buffer, like TensorForcedEvalOp.
270 BlockAccess = internal::is_arithmetic<CoeffReturnType>::value,
271 PreferBlockAccess = false,
272 CoordAccess = false, // to be implemented
273 RawAccess = true
274 };
275
276 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
277 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
278 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
279
280 typedef typename internal::TensorMaterializedBlock<CoeffReturnType, NumDims, Layout, Index> TensorBlock;
281 //===--------------------------------------------------------------------===//
282
283 // The functor's dimensions() may return an index type that promotes to Index.
284 EIGEN_STRONG_INLINE TensorEvaluator(const XprType& op, const Device& device)
285 : m_dimensions(op.func().dimensions(op.lhsExpression(), op.rhsExpression())),
286 m_op(op),
287 m_device(device),
288 m_result(nullptr) {}
289
290 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_dimensions; }
291
292 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType data) {
293 if (data) {
294 evalTo(data);
295 return false;
296 } else {
297 m_result = static_cast<EvaluatorPointerType>(
298 m_device.get((CoeffReturnType*)m_device.allocate_temp(dimensions().TotalSize() * sizeof(CoeffReturnType))));
299 evalTo(m_result);
300 return true;
301 }
302 }
303
304 EIGEN_STRONG_INLINE void cleanup() {
305 if (m_result != nullptr) {
306 m_device.deallocate_temp(m_result);
307 m_result = nullptr;
308 }
309 }
310
311 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const { return m_result[index]; }
312
313 template <int LoadMode>
314 EIGEN_DEVICE_FUNC PacketReturnType packet(Index index) const {
315 return internal::ploadt<PacketReturnType, LoadMode>(m_result + index);
316 }
317
318 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool vectorized) const {
319 // TODO(rmlarsen): Extend CustomOp API to return its cost estimate.
320 return TensorOpCost(sizeof(CoeffReturnType), 0, 0, vectorized, PacketSize);
321 }
322
323 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements() const {
324 return internal::TensorBlockResourceRequirements::any();
325 }
326
327 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
328 bool /*root_of_expr_ast*/ = false) const {
329 eigen_assert(m_result != nullptr);
330 return TensorBlock::materialize(m_result, m_dimensions, desc, scratch);
331 }
332
333 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return m_result; }
334
335 protected:
336 void evalTo(EvaluatorPointerType data) {
337 // The Output type handed to eval() is a compatibility surface: functors are
338 // compiled against a DenseIndex-typed map, so widen its index type only
339 // when the expressions' promoted Index is strictly wider than DenseIndex.
340 // DenseIndex must stay the first argument: promote_index_type keeps that
341 // one on a tie, which preserves the map type for equal-width distinct
342 // index types such as long long versus long.
343 using MapIndex = typename internal::promote_index_type<DenseIndex, Index>::type;
344 TensorMap<Tensor<CoeffReturnType, NumDims, Layout, MapIndex> > result(m_device.get(data), m_dimensions);
345 m_op.func().eval(m_op.lhsExpression(), m_op.rhsExpression(), result, m_device);
346 }
347
348 Dimensions m_dimensions;
349 const XprType m_op;
350 const Device EIGEN_DEVICE_REF m_device;
351 EvaluatorPointerType m_result;
352};
353
354} // end namespace Eigen
355
356#endif // EIGEN_TENSOR_TENSOR_CUSTOM_OP_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
Tensor custom class.
Definition TensorCustomOp.h:218
Tensor custom class.
Definition TensorCustomOp.h:54
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47