11#ifndef EIGEN_TENSOR_TENSOR_CUSTOM_OP_H
12#define EIGEN_TENSOR_TENSOR_CUSTOM_OP_H
15#include "./InternalHeaderCheck.h"
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;
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>.");
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;
41template <
typename CustomUnaryFunc,
typename XprType>
42struct eval<TensorCustomUnaryOp<CustomUnaryFunc, XprType>, Eigen::Dense> {
43 typedef const TensorCustomUnaryOp<CustomUnaryFunc, XprType> EIGEN_DEVICE_REF type;
53template <
typename CustomUnaryFunc,
typename XprType>
54class TensorCustomUnaryOp :
public TensorBase<TensorCustomUnaryOp<CustomUnaryFunc, XprType>, ReadOnlyAccessors> {
56 typedef typename internal::traits<TensorCustomUnaryOp>::Scalar Scalar;
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;
63 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCustomUnaryOp(
const XprType& expr,
const CustomUnaryFunc& func)
64 : m_expr(expr), m_func(func) {}
66 EIGEN_DEVICE_FUNC
const CustomUnaryFunc& func()
const {
return m_func; }
68 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename XprType::Nested>& expression()
const {
return m_expr; }
71 typename XprType::Nested m_expr;
72 const CustomUnaryFunc m_func;
76template <
typename CustomUnaryFunc,
typename XprType,
typename Device>
79 typedef typename internal::traits<ArgType>::Index Index;
80 static constexpr int NumDims = internal::traits<ArgType>::NumDimensions;
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;
90 static constexpr int Layout = TensorEvaluator<XprType, Device>::Layout;
93 PacketAccess = (PacketType<CoeffReturnType, Device>::size > 1),
99 BlockAccess = internal::is_arithmetic<CoeffReturnType>::value,
100 PreferBlockAccess =
false,
106 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
107 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
109 typedef typename internal::TensorMaterializedBlock<CoeffReturnType, NumDims, Layout, Index> TensorBlock;
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) {}
116 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_dimensions; }
118 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType data) {
123 m_result =
static_cast<EvaluatorPointerType
>(
124 m_device.get((CoeffReturnType*)m_device.allocate_temp(dimensions().TotalSize() *
sizeof(CoeffReturnType))));
130 EIGEN_STRONG_INLINE
void cleanup() {
132 m_device.deallocate_temp(m_result);
137 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
return m_result[index]; }
139 template <
int LoadMode>
140 EIGEN_DEVICE_FUNC PacketReturnType packet(Index index)
const {
141 return internal::ploadt<PacketReturnType, LoadMode>(m_result + index);
144 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
146 return TensorOpCost(
sizeof(CoeffReturnType), 0, 0, vectorized, PacketSize);
149 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
150 return internal::TensorBlockResourceRequirements::any();
153 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
154 bool =
false)
const {
155 eigen_assert(m_result !=
nullptr);
156 return TensorBlock::materialize(m_result, m_dimensions, desc, scratch);
159 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return m_result; }
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);
167 Dimensions m_dimensions;
169 const Device EIGEN_DEVICE_REF m_device;
170 EvaluatorPointerType m_result;
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;
189 typename promote_index_type<typename traits<LhsXprType>::Index,
typename traits<RhsXprType>::Index>::type Index;
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>.");
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>
209template <
typename CustomBinaryFunc,
typename LhsXprType,
typename RhsXprType>
210struct eval<TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>, Eigen::Dense> {
211 typedef const TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>& type;
216template <
typename CustomBinaryFunc,
typename LhsXprType,
typename RhsXprType>
217class TensorCustomBinaryOp
218 :
public TensorBase<TensorCustomBinaryOp<CustomBinaryFunc, LhsXprType, RhsXprType>, ReadOnlyAccessors> {
220 typedef typename internal::traits<TensorCustomBinaryOp>::Scalar Scalar;
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;
227 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCustomBinaryOp(
const LhsXprType& lhs,
const RhsXprType& rhs,
228 const CustomBinaryFunc& func)
230 : m_lhs_xpr(lhs), m_rhs_xpr(rhs), m_func(func) {}
232 EIGEN_DEVICE_FUNC
const CustomBinaryFunc& func()
const {
return m_func; }
234 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression()
const {
238 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression()
const {
243 typename LhsXprType::Nested m_lhs_xpr;
244 typename RhsXprType::Nested m_rhs_xpr;
245 const CustomBinaryFunc m_func;
249template <
typename CustomBinaryFunc,
typename LhsXprType,
typename RhsXprType,
typename Device>
252 typedef typename internal::traits<XprType>::Index Index;
253 static constexpr int NumDims = internal::traits<XprType>::NumDimensions;
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;
260 typedef typename Eigen::internal::traits<XprType>::PointerType TensorPointerType;
261 typedef StorageMemory<CoeffReturnType, Device> Storage;
262 typedef typename Storage::Type EvaluatorPointerType;
264 static constexpr int Layout = TensorEvaluator<LhsXprType, Device>::Layout;
267 PacketAccess = (PacketType<CoeffReturnType, Device>::size > 1),
270 BlockAccess = internal::is_arithmetic<CoeffReturnType>::value,
271 PreferBlockAccess =
false,
277 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
278 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
280 typedef typename internal::TensorMaterializedBlock<CoeffReturnType, NumDims, Layout, Index> TensorBlock;
284 EIGEN_STRONG_INLINE TensorEvaluator(
const XprType& op,
const Device& device)
285 : m_dimensions(op.func().dimensions(op.lhsExpression(), op.rhsExpression())),
290 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_dimensions; }
292 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType data) {
297 m_result =
static_cast<EvaluatorPointerType
>(
298 m_device.get((CoeffReturnType*)m_device.allocate_temp(dimensions().TotalSize() *
sizeof(CoeffReturnType))));
304 EIGEN_STRONG_INLINE
void cleanup() {
305 if (m_result !=
nullptr) {
306 m_device.deallocate_temp(m_result);
311 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
return m_result[index]; }
313 template <
int LoadMode>
314 EIGEN_DEVICE_FUNC PacketReturnType packet(Index index)
const {
315 return internal::ploadt<PacketReturnType, LoadMode>(m_result + index);
318 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
320 return TensorOpCost(
sizeof(CoeffReturnType), 0, 0, vectorized, PacketSize);
323 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
324 return internal::TensorBlockResourceRequirements::any();
327 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
328 bool =
false)
const {
329 eigen_assert(m_result !=
nullptr);
330 return TensorBlock::materialize(m_result, m_dimensions, desc, scratch);
333 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return m_result; }
336 void evalTo(EvaluatorPointerType data) {
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);
348 Dimensions m_dimensions;
350 const Device EIGEN_DEVICE_REF m_device;
351 EvaluatorPointerType m_result;
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