11#ifndef EIGEN_TENSOR_TENSOR_REF_H
12#define EIGEN_TENSOR_TENSOR_REF_H
15#include "./InternalHeaderCheck.h"
21template <
typename Dimensions,
typename Scalar>
22class TensorLazyBaseEvaluator {
24 TensorLazyBaseEvaluator() =
default;
25 virtual ~TensorLazyBaseEvaluator() =
default;
27 EIGEN_DEVICE_FUNC
virtual const Dimensions& dimensions()
const = 0;
28 EIGEN_DEVICE_FUNC
virtual const Scalar* data()
const = 0;
30 EIGEN_DEVICE_FUNC
virtual const Scalar coeff(DenseIndex index)
const = 0;
31 EIGEN_DEVICE_FUNC
virtual Scalar& coeffRef(DenseIndex index) = 0;
33 void incrRefCount() { ++m_refcount; }
34 void decrRefCount() { --m_refcount; }
35 int refCount()
const {
return m_refcount; }
38 TensorLazyBaseEvaluator(
const TensorLazyBaseEvaluator& other) =
delete;
39 TensorLazyBaseEvaluator& operator=(
const TensorLazyBaseEvaluator& other) =
delete;
44template <
typename Dimensions,
typename Expr,
typename Device>
45class TensorLazyEvaluatorReadOnly
46 :
public TensorLazyBaseEvaluator<Dimensions, typename TensorEvaluator<Expr, Device>::Scalar> {
48 EIGEN_MAKE_ALIGNED_OPERATOR_NEW
49 typedef typename TensorEvaluator<Expr, Device>::Scalar Scalar;
50 typedef StorageMemory<Scalar, Device> Storage;
51 typedef typename Storage::Type EvaluatorPointerType;
52 typedef TensorEvaluator<Expr, Device> EvalType;
54 TensorLazyEvaluatorReadOnly(
const Expr& expr,
const Device& device) : m_impl(expr, device), m_dummy() {
56 internal::array_size<Dimensions>::value == internal::array_size<typename EvalType::Dimensions>::value,
57 "Dimension sizes must match.");
58 const auto& other_dims = m_impl.dimensions();
59 for (std::size_t i = 0; i < m_dims.size(); ++i) {
60 m_dims[i] = other_dims[i];
62 m_impl.evalSubExprsIfNeeded(
nullptr);
64 virtual ~TensorLazyEvaluatorReadOnly() { m_impl.cleanup(); }
66 EIGEN_DEVICE_FUNC
virtual const Dimensions& dimensions()
const {
return m_dims; }
67 EIGEN_DEVICE_FUNC
virtual const Scalar* data()
const {
return m_impl.data(); }
69 EIGEN_DEVICE_FUNC
virtual const Scalar coeff(DenseIndex index)
const {
70 return m_impl.coeff(internal::convert_index<typename EvalType::Index>(index));
72 EIGEN_DEVICE_FUNC
virtual Scalar& coeffRef(DenseIndex ) {
73 eigen_assert(
false &&
"can't reference the coefficient of a rvalue");
78 TensorEvaluator<Expr, Device> m_impl;
83template <
typename Dimensions,
typename Expr,
typename Device>
84class TensorLazyEvaluatorWritable :
public TensorLazyEvaluatorReadOnly<Dimensions, Expr, Device> {
86 typedef TensorLazyEvaluatorReadOnly<Dimensions, Expr, Device> Base;
87 typedef typename Base::Scalar Scalar;
88 typedef StorageMemory<Scalar, Device> Storage;
89 typedef typename Storage::Type EvaluatorPointerType;
91 TensorLazyEvaluatorWritable(
const Expr& expr,
const Device& device) : Base(expr, device) {}
92 virtual ~TensorLazyEvaluatorWritable() =
default;
94 EIGEN_DEVICE_FUNC
virtual Scalar& coeffRef(DenseIndex index) {
95 return this->m_impl.coeffRef(internal::convert_index<typename Base::EvalType::Index>(index));
99template <
typename Dimensions,
typename Expr,
typename Device,
bool IsWritable>
100class TensorLazyEvaluator :
public std::conditional_t<IsWritable, TensorLazyEvaluatorWritable<Dimensions, Expr, Device>,
101 TensorLazyEvaluatorReadOnly<Dimensions, const Expr, Device>> {
103 typedef std::conditional_t<IsWritable, TensorLazyEvaluatorWritable<Dimensions, Expr, Device>,
104 TensorLazyEvaluatorReadOnly<Dimensions, const Expr, Device>>
106 typedef typename Base::Scalar Scalar;
108 TensorLazyEvaluator(
const Expr& expr,
const Device& device) : Base(expr, device) {}
109 virtual ~TensorLazyEvaluator() =
default;
112template <
typename Derived>
113class TensorRefBase :
public TensorBase<Derived> {
115 typedef typename traits<Derived>::PlainObjectType PlainObjectType;
116 typedef typename PlainObjectType::Base Base;
117 typedef typename Eigen::internal::ref_selector<Derived>::type Nested;
118 typedef typename traits<PlainObjectType>::StorageKind StorageKind;
119 typedef typename traits<PlainObjectType>::Index Index;
120 typedef typename traits<PlainObjectType>::Scalar Scalar;
121 typedef typename NumTraits<Scalar>::Real RealScalar;
122 typedef typename Base::CoeffReturnType CoeffReturnType;
123 typedef Scalar* PointerType;
124 typedef PointerType PointerArgType;
126 static constexpr Index NumIndices = PlainObjectType::NumIndices;
127 typedef typename PlainObjectType::Dimensions Dimensions;
129 static constexpr int Layout = PlainObjectType::Layout;
132 PacketAccess =
false,
134 PreferBlockAccess =
false,
140 typedef TensorBlockNotImplemented TensorBlock;
143 EIGEN_STRONG_INLINE TensorRefBase() =
default;
145 TensorRefBase(
const TensorRefBase& other) : TensorBase<Derived>(other), m_evaluator(other.m_evaluator) {
146 eigen_assert(m_evaluator->refCount() > 0);
147 m_evaluator->incrRefCount();
150 TensorRefBase& operator=(
const TensorRefBase& other) {
151 if (
this != &other) {
153 m_evaluator = other.m_evaluator;
154 eigen_assert(m_evaluator->refCount() > 0);
155 m_evaluator->incrRefCount();
160 template <
typename Expression,
161 typename EnableIf = std::enable_if_t<!std::is_same<std::decay_t<Expression>, Derived>::value>>
162 EIGEN_STRONG_INLINE TensorRefBase(
const Expression& expr)
163 : m_evaluator(new TensorLazyEvaluator<Dimensions, Expression, DefaultDevice,
164 !std::is_const<PlainObjectType>::value &&
165 bool(is_lvalue<Expression>::value)>(expr, DefaultDevice())) {
166 m_evaluator->incrRefCount();
169 template <
typename Expression,
170 typename EnableIf = std::enable_if_t<!std::is_same<std::decay_t<Expression>, Derived>::value>>
171 EIGEN_STRONG_INLINE TensorRefBase& operator=(
const Expression& expr) {
173 m_evaluator =
new TensorLazyEvaluator < Dimensions, Expression, DefaultDevice,
174 !std::is_const<PlainObjectType>::value&& bool(is_lvalue<Expression>::value) >
175 (expr, DefaultDevice());
176 m_evaluator->incrRefCount();
180 ~TensorRefBase() { unrefEvaluator(); }
182 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index rank()
const {
return m_evaluator->dimensions().size(); }
183 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index dimension(Index n)
const {
return m_evaluator->dimensions()[n]; }
184 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_evaluator->dimensions(); }
185 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index size()
const {
return m_evaluator->dimensions().TotalSize(); }
186 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar* data()
const {
return m_evaluator->data(); }
188 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar operator()(Index index)
const {
return m_evaluator->coeff(index); }
190 template <
typename... IndexTypes>
191 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar operator()(Index firstIndex, IndexTypes... otherIndices)
const {
192 eigen_assert(internal::indices_fit<Index>(otherIndices...));
193 const std::size_t num_indices =
sizeof...(otherIndices) + 1;
194 const array<Index, num_indices> indices{{firstIndex,
static_cast<Index
>(otherIndices)...}};
195 return coeff(indices);
198 template <std::
size_t NumIndices>
199 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar coeff(
const array<Index, NumIndices>& indices)
const {
200 const Dimensions& dims = this->dimensions();
202 EIGEN_IF_CONSTEXPR (PlainObjectType::Options &
RowMajor) {
204 for (
size_t i = 1; i < NumIndices; ++i) {
205 index = index * dims[i] + indices[i];
208 index += indices[NumIndices - 1];
209 for (
int i = NumIndices - 2; i >= 0; --i) {
210 index = index * dims[i] + indices[i];
213 return m_evaluator->coeff(index);
216 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Scalar coeff(Index index)
const {
return m_evaluator->coeff(index); }
218 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) {
return m_evaluator->coeffRef(index); }
221 TensorLazyBaseEvaluator<Dimensions, Scalar>* evaluator() {
return m_evaluator; }
224 EIGEN_STRONG_INLINE
void unrefEvaluator() {
226 m_evaluator->decrRefCount();
227 if (m_evaluator->refCount() == 0) {
233 TensorLazyBaseEvaluator<Dimensions, Scalar>* m_evaluator =
nullptr;
245template <
typename PlainObjectType>
246class TensorRef :
public internal::TensorRefBase<TensorRef<PlainObjectType>> {
247 typedef internal::TensorRefBase<TensorRef<PlainObjectType>> Base;
250 using Scalar =
typename Base::Scalar;
251 using Dimensions =
typename Base::Dimensions;
254 using Index =
typename Base::Index;
256 EIGEN_STRONG_INLINE TensorRef() =
default;
258 EIGEN_STRONG_INLINE TensorRef(
const TensorRef& other) =
default;
260 template <
typename Expression>
261 EIGEN_STRONG_INLINE TensorRef(
const Expression& expr) : Base(expr) {
262 EIGEN_STATIC_ASSERT(internal::is_lvalue<Expression>::value,
263 "Expression must be mutable to create a mutable TensorRef<Expression>. Did you mean "
264 "TensorRef<const Expression>?)");
267 TensorRef& operator=(
const TensorRef& other) {
return Base::operator=(other).derived(); }
269 template <
typename Expression>
270 EIGEN_STRONG_INLINE TensorRef& operator=(
const Expression& expr) {
271 EIGEN_STATIC_ASSERT(internal::is_lvalue<Expression>::value,
272 "Expression must be mutable to create a mutable TensorRef<Expression>. Did you mean "
273 "TensorRef<const Expression>?)");
274 return Base::operator=(expr).derived();
277 template <
typename... IndexTypes>
278 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index firstIndex, IndexTypes... otherIndices) {
279 eigen_assert(internal::indices_fit<Index>(otherIndices...));
280 const std::size_t num_indices =
sizeof...(otherIndices) + 1;
281 const array<Index, num_indices> indices{{firstIndex,
static_cast<Index
>(otherIndices)...}};
282 return coeffRef(indices);
285 template <std::
size_t NumIndices>
286 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(
const array<Index, NumIndices>& indices) {
287 const Dimensions& dims = this->dimensions();
289 EIGEN_IF_CONSTEXPR (PlainObjectType::Options &
RowMajor) {
291 for (
size_t i = 1; i < NumIndices; ++i) {
292 index = index * dims[i] + indices[i];
295 index += indices[NumIndices - 1];
296 for (
int i = NumIndices - 2; i >= 0; --i) {
297 index = index * dims[i] + indices[i];
300 return Base::evaluator()->coeffRef(index);
303 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) {
return Base::evaluator()->coeffRef(index); }
313template <
typename PlainObjectType>
314class TensorRef<const PlainObjectType> :
public internal::TensorRefBase<TensorRef<const PlainObjectType>> {
315 typedef internal::TensorRefBase<TensorRef<const PlainObjectType>> Base;
318 EIGEN_STRONG_INLINE TensorRef() =
default;
320 EIGEN_STRONG_INLINE TensorRef(
const TensorRef& other) =
default;
322 template <
typename Expression>
323 EIGEN_STRONG_INLINE TensorRef(
const Expression& expr) : Base(expr) {}
325 TensorRef& operator=(
const TensorRef& other) {
return Base::operator=(other).derived(); }
327 template <
typename Expression>
328 EIGEN_STRONG_INLINE TensorRef& operator=(
const Expression& expr) {
329 return Base::operator=(expr).derived();
334template <
typename Derived,
typename Device>
336 typedef typename Derived::Index Index;
337 typedef typename Derived::Scalar Scalar;
339 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
340 typedef typename Derived::Dimensions
Dimensions;
341 typedef StorageMemory<CoeffReturnType, Device> Storage;
342 typedef typename Storage::Type EvaluatorPointerType;
344 static constexpr int Layout = TensorRef<Derived>::Layout;
347 PacketAccess =
false,
349 PreferBlockAccess =
false,
355 typedef internal::TensorBlockNotImplemented TensorBlock;
358 EIGEN_STRONG_INLINE TensorEvaluator(
const TensorRef<Derived>& m,
const Device&) : m_ref(m) {}
360 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_ref.dimensions(); }
362 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType) {
return true; }
364 EIGEN_STRONG_INLINE
void cleanup() {}
366 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
return m_ref.coeff(index); }
368 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) {
return m_ref.coeffRef(index); }
370 EIGEN_DEVICE_FUNC
const Scalar* data()
const {
return m_ref.data(); }
373 TensorRef<Derived> m_ref;
377template <
typename Derived,
typename Device>
379 typedef typename Derived::Index Index;
380 typedef typename Derived::Scalar Scalar;
381 typedef typename Derived::Scalar CoeffReturnType;
382 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
383 typedef typename Derived::Dimensions Dimensions;
385 typedef TensorEvaluator<const TensorRef<Derived>, Device> Base;
387 enum { IsAligned =
false, PacketAccess =
false, BlockAccess =
false, PreferBlockAccess =
false, RawAccess =
false };
390 typedef internal::TensorBlockNotImplemented TensorBlock;
393 EIGEN_STRONG_INLINE TensorEvaluator(TensorRef<Derived>& m,
const Device& d) : Base(m, d) {}
395 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) {
return this->m_ref.coeffRef(index); }
A reference to a tensor expression The expression will be evaluated lazily (as much as possible).
Definition TensorRef.h:246
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47