Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorExpr.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_EXPR_H
12#define EIGEN_TENSOR_TENSOR_EXPR_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20template <typename NullaryOp, typename XprType>
21struct traits<TensorCwiseNullaryOp<NullaryOp, XprType> > : traits<XprType> {
22 typedef traits<XprType> XprTraits;
23 typedef typename XprType::Scalar Scalar;
24 static constexpr int NumDimensions = XprTraits::NumDimensions;
25 static constexpr int Layout = XprTraits::Layout;
26 typedef typename XprTraits::PointerType PointerType;
27 enum { Flags = 0 };
28};
29
30} // end namespace internal
31
40template <typename NullaryOp, typename XprType>
41class TensorCwiseNullaryOp : public TensorBase<TensorCwiseNullaryOp<NullaryOp, XprType>, ReadOnlyAccessors> {
42 public:
43 typedef typename Eigen::internal::traits<TensorCwiseNullaryOp>::Scalar Scalar;
44 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
45 typedef typename XprType::CoeffReturnType CoeffReturnType;
46 typedef typename Eigen::internal::ref_selector<TensorCwiseNullaryOp>::non_const_type Nested;
47 typedef typename Eigen::internal::traits<TensorCwiseNullaryOp>::StorageKind StorageKind;
48 typedef typename Eigen::internal::traits<TensorCwiseNullaryOp>::Index Index;
49
50 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCwiseNullaryOp(const XprType& xpr, const NullaryOp& func = NullaryOp())
51 : m_xpr(xpr), m_functor(func) {}
52
53 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename XprType::Nested>& nestedExpression() const { return m_xpr; }
54
55 EIGEN_DEVICE_FUNC const NullaryOp& functor() const { return m_functor; }
56
57 protected:
58 typename XprType::Nested m_xpr;
59 const NullaryOp m_functor;
60};
61
62namespace internal {
63template <typename UnaryOp, typename XprType>
64struct traits<TensorCwiseUnaryOp<UnaryOp, XprType> > : traits<XprType> {
65 // TODO(phli): Add InputScalar, InputPacket. Check references to
66 // current Scalar/Packet to see if the intent is Input or Output.
67 typedef typename result_of<UnaryOp(typename XprType::Scalar)>::type Scalar;
68 typedef traits<XprType> XprTraits;
69 static constexpr int NumDimensions = XprTraits::NumDimensions;
70 static constexpr int Layout = XprTraits::Layout;
71 typedef typename TypeConversion<Scalar, typename XprTraits::PointerType>::type PointerType;
72};
73
74template <typename UnaryOp, typename XprType>
75struct eval<TensorCwiseUnaryOp<UnaryOp, XprType>, Eigen::Dense> {
76 typedef const TensorCwiseUnaryOp<UnaryOp, XprType>& type;
77};
78
79} // end namespace internal
80
89template <typename UnaryOp, typename XprType>
90class TensorCwiseUnaryOp : public TensorBase<TensorCwiseUnaryOp<UnaryOp, XprType>, ReadOnlyAccessors> {
91 public:
92 // TODO(phli): Add InputScalar, InputPacket. Check references to
93 // current Scalar/Packet to see if the intent is Input or Output.
94 typedef typename Eigen::internal::traits<TensorCwiseUnaryOp>::Scalar Scalar;
95 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
96 typedef Scalar CoeffReturnType;
97 typedef typename Eigen::internal::ref_selector<TensorCwiseUnaryOp>::type Nested;
98 typedef typename Eigen::internal::traits<TensorCwiseUnaryOp>::StorageKind StorageKind;
99 typedef typename Eigen::internal::traits<TensorCwiseUnaryOp>::Index Index;
100
101 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCwiseUnaryOp(const XprType& xpr, const UnaryOp& func = UnaryOp())
102 : m_xpr(xpr), m_functor(func) {}
103
104 EIGEN_DEVICE_FUNC const UnaryOp& functor() const { return m_functor; }
105
107 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename XprType::Nested>& nestedExpression() const { return m_xpr; }
108
109 protected:
110 typename XprType::Nested m_xpr;
111 const UnaryOp m_functor;
112};
113
114namespace internal {
115template <typename BinaryOp, typename LhsXprType, typename RhsXprType>
116struct traits<TensorCwiseBinaryOp<BinaryOp, LhsXprType, RhsXprType> > {
117 // Type promotion to handle the case where the types of the lhs and the rhs
118 // are different.
119 // TODO(phli): Add Lhs/RhsScalar, Lhs/RhsPacket. Check references to
120 // current Scalar/Packet to see if the intent is Inputs or Output.
121 typedef typename result_of<BinaryOp(typename LhsXprType::Scalar, typename RhsXprType::Scalar)>::type Scalar;
122 typedef traits<LhsXprType> XprTraits;
123 typedef typename promote_storage_type<typename traits<LhsXprType>::StorageKind,
124 typename traits<RhsXprType>::StorageKind>::ret StorageKind;
125 typedef
126 typename promote_index_type<typename traits<LhsXprType>::Index, typename traits<RhsXprType>::Index>::type Index;
127 static constexpr int NumDimensions = XprTraits::NumDimensions;
128 static constexpr int Layout = XprTraits::Layout;
129 typedef typename TypeConversion<Scalar,
130 std::conditional_t<Pointer_type_promotion<typename LhsXprType::Scalar, Scalar>::val,
131 typename traits<LhsXprType>::PointerType,
132 typename traits<RhsXprType>::PointerType> >::type PointerType;
133 enum { Flags = 0 };
134};
135
136template <typename BinaryOp, typename LhsXprType, typename RhsXprType>
137struct eval<TensorCwiseBinaryOp<BinaryOp, LhsXprType, RhsXprType>, Eigen::Dense> {
138 typedef const TensorCwiseBinaryOp<BinaryOp, LhsXprType, RhsXprType>& type;
139};
140
141} // end namespace internal
142
151template <typename BinaryOp, typename LhsXprType, typename RhsXprType>
152class TensorCwiseBinaryOp
153 : public TensorBase<TensorCwiseBinaryOp<BinaryOp, LhsXprType, RhsXprType>, ReadOnlyAccessors> {
154 public:
155 // TODO(phli): Add Lhs/RhsScalar, Lhs/RhsPacket. Check references to
156 // current Scalar/Packet to see if the intent is Inputs or Output.
157 typedef typename Eigen::internal::traits<TensorCwiseBinaryOp>::Scalar Scalar;
158 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
159 typedef Scalar CoeffReturnType;
160 typedef typename Eigen::internal::ref_selector<TensorCwiseBinaryOp>::type Nested;
161 typedef typename Eigen::internal::traits<TensorCwiseBinaryOp>::StorageKind StorageKind;
162 typedef typename Eigen::internal::traits<TensorCwiseBinaryOp>::Index Index;
163
164 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCwiseBinaryOp(const LhsXprType& lhs, const RhsXprType& rhs,
165 const BinaryOp& func = BinaryOp())
166 : m_lhs_xpr(lhs), m_rhs_xpr(rhs), m_functor(func) {}
167
168 EIGEN_DEVICE_FUNC const BinaryOp& functor() const { return m_functor; }
169
171 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression() const {
172 return m_lhs_xpr;
173 }
174
175 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression() const {
176 return m_rhs_xpr;
177 }
178
179 protected:
180 typename LhsXprType::Nested m_lhs_xpr;
181 typename RhsXprType::Nested m_rhs_xpr;
182 const BinaryOp m_functor;
183};
184
185namespace internal {
186template <typename TernaryOp, typename Arg1XprType, typename Arg2XprType, typename Arg3XprType>
187struct traits<TensorCwiseTernaryOp<TernaryOp, Arg1XprType, Arg2XprType, Arg3XprType> > {
188 // Type promotion to handle the case where the types of the args are different.
189 typedef typename result_of<TernaryOp(typename Arg1XprType::Scalar, typename Arg2XprType::Scalar,
190 typename Arg3XprType::Scalar)>::type Scalar;
191 typedef traits<Arg1XprType> XprTraits;
192 typedef typename traits<Arg1XprType>::StorageKind StorageKind;
193 typedef typename traits<Arg1XprType>::Index Index;
194 static constexpr int NumDimensions = XprTraits::NumDimensions;
195 static constexpr int Layout = XprTraits::Layout;
196 typedef typename TypeConversion<Scalar,
197 std::conditional_t<Pointer_type_promotion<typename Arg2XprType::Scalar, Scalar>::val,
198 typename traits<Arg2XprType>::PointerType,
199 typename traits<Arg3XprType>::PointerType> >::type PointerType;
200 enum { Flags = 0 };
201};
202
203template <typename TernaryOp, typename Arg1XprType, typename Arg2XprType, typename Arg3XprType>
204struct eval<TensorCwiseTernaryOp<TernaryOp, Arg1XprType, Arg2XprType, Arg3XprType>, Eigen::Dense> {
205 typedef const TensorCwiseTernaryOp<TernaryOp, Arg1XprType, Arg2XprType, Arg3XprType>& type;
206};
207
208} // end namespace internal
209
210template <typename TernaryOp, typename Arg1XprType, typename Arg2XprType, typename Arg3XprType>
211class TensorCwiseTernaryOp
212 : public TensorBase<TensorCwiseTernaryOp<TernaryOp, Arg1XprType, Arg2XprType, Arg3XprType>, ReadOnlyAccessors> {
213 public:
214 typedef typename Eigen::internal::traits<TensorCwiseTernaryOp>::Scalar Scalar;
215 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
216 typedef Scalar CoeffReturnType;
217 typedef typename Eigen::internal::ref_selector<TensorCwiseTernaryOp>::type Nested;
218 typedef typename Eigen::internal::traits<TensorCwiseTernaryOp>::StorageKind StorageKind;
219 typedef typename Eigen::internal::traits<TensorCwiseTernaryOp>::Index Index;
220
221 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorCwiseTernaryOp(const Arg1XprType& arg1, const Arg2XprType& arg2,
222 const Arg3XprType& arg3,
223 const TernaryOp& func = TernaryOp())
224 : m_arg1_xpr(arg1), m_arg2_xpr(arg2), m_arg3_xpr(arg3), m_functor(func) {}
225
226 EIGEN_DEVICE_FUNC const TernaryOp& functor() const { return m_functor; }
227
229 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename Arg1XprType::Nested>& arg1Expression() const {
230 return m_arg1_xpr;
231 }
232
233 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename Arg2XprType::Nested>& arg2Expression() const {
234 return m_arg2_xpr;
235 }
236
237 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename Arg3XprType::Nested>& arg3Expression() const {
238 return m_arg3_xpr;
239 }
240
241 protected:
242 typename Arg1XprType::Nested m_arg1_xpr;
243 typename Arg2XprType::Nested m_arg2_xpr;
244 typename Arg3XprType::Nested m_arg3_xpr;
245 const TernaryOp m_functor;
246};
247
248namespace internal {
249template <typename IfXprType, typename ThenXprType, typename ElseXprType>
250struct traits<TensorSelectOp<IfXprType, ThenXprType, ElseXprType> > : traits<ThenXprType> {
251 typedef typename traits<ThenXprType>::Scalar Scalar;
252 typedef traits<ThenXprType> XprTraits;
253 typedef typename promote_storage_type<typename traits<ThenXprType>::StorageKind,
254 typename traits<ElseXprType>::StorageKind>::ret StorageKind;
255 typedef
256 typename promote_index_type<typename traits<ElseXprType>::Index, typename traits<ThenXprType>::Index>::type Index;
257 static constexpr int NumDimensions = XprTraits::NumDimensions;
258 static constexpr int Layout = XprTraits::Layout;
259 typedef std::conditional_t<Pointer_type_promotion<typename ThenXprType::Scalar, Scalar>::val,
260 typename traits<ThenXprType>::PointerType, typename traits<ElseXprType>::PointerType>
261 PointerType;
262};
263
264template <typename IfXprType, typename ThenXprType, typename ElseXprType>
265struct eval<TensorSelectOp<IfXprType, ThenXprType, ElseXprType>, Eigen::Dense> {
266 typedef const TensorSelectOp<IfXprType, ThenXprType, ElseXprType>& type;
267};
268
269} // end namespace internal
270
271template <typename IfXprType, typename ThenXprType, typename ElseXprType>
272class TensorSelectOp : public TensorBase<TensorSelectOp<IfXprType, ThenXprType, ElseXprType>, ReadOnlyAccessors> {
273 public:
274 typedef typename Eigen::internal::traits<TensorSelectOp>::Scalar Scalar;
275 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
276 typedef typename internal::promote_storage_type<typename ThenXprType::CoeffReturnType,
277 typename ElseXprType::CoeffReturnType>::ret CoeffReturnType;
278 typedef typename Eigen::internal::ref_selector<TensorSelectOp>::type Nested;
279 typedef typename Eigen::internal::traits<TensorSelectOp>::StorageKind StorageKind;
280 typedef typename Eigen::internal::traits<TensorSelectOp>::Index Index;
281
282 EIGEN_DEVICE_FUNC TensorSelectOp(const IfXprType& a_condition, const ThenXprType& a_then, const ElseXprType& a_else)
283 : m_condition(a_condition), m_then(a_then), m_else(a_else) {}
284
285 EIGEN_DEVICE_FUNC const IfXprType& ifExpression() const { return m_condition; }
286
287 EIGEN_DEVICE_FUNC const ThenXprType& thenExpression() const { return m_then; }
288
289 EIGEN_DEVICE_FUNC const ElseXprType& elseExpression() const { return m_else; }
290
291 protected:
292 typename IfXprType::Nested m_condition;
293 typename ThenXprType::Nested m_then;
294 typename ElseXprType::Nested m_else;
295};
296
297} // end namespace Eigen
298
299#endif // EIGEN_TENSOR_TENSOR_EXPR_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
const internal::remove_all_t< typename LhsXprType::Nested > & lhsExpression() const
Definition TensorExpr.h:171
Tensor unary expression.
Definition TensorExpr.h:90
const internal::remove_all_t< typename XprType::Nested > & nestedExpression() const
Definition TensorExpr.h:107
Namespace containing all symbols from the Eigen library.