Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorRef.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_REF_H
12#define EIGEN_TENSOR_TENSOR_REF_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename Dimensions, typename Scalar>
22class TensorLazyBaseEvaluator {
23 public:
24 TensorLazyBaseEvaluator() = default;
25 virtual ~TensorLazyBaseEvaluator() = default;
26
27 EIGEN_DEVICE_FUNC virtual const Dimensions& dimensions() const = 0;
28 EIGEN_DEVICE_FUNC virtual const Scalar* data() const = 0;
29
30 EIGEN_DEVICE_FUNC virtual const Scalar coeff(DenseIndex index) const = 0;
31 EIGEN_DEVICE_FUNC virtual Scalar& coeffRef(DenseIndex index) = 0;
32
33 void incrRefCount() { ++m_refcount; }
34 void decrRefCount() { --m_refcount; }
35 int refCount() const { return m_refcount; }
36
37 private:
38 TensorLazyBaseEvaluator(const TensorLazyBaseEvaluator& other) = delete;
39 TensorLazyBaseEvaluator& operator=(const TensorLazyBaseEvaluator& other) = delete;
40
41 int m_refcount = 0;
42};
43
44template <typename Dimensions, typename Expr, typename Device>
45class TensorLazyEvaluatorReadOnly
46 : public TensorLazyBaseEvaluator<Dimensions, typename TensorEvaluator<Expr, Device>::Scalar> {
47 public:
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;
53
54 TensorLazyEvaluatorReadOnly(const Expr& expr, const Device& device) : m_impl(expr, device), m_dummy() {
55 EIGEN_STATIC_ASSERT(
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];
61 }
62 m_impl.evalSubExprsIfNeeded(nullptr);
63 }
64 virtual ~TensorLazyEvaluatorReadOnly() { m_impl.cleanup(); }
65
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(); }
68
69 EIGEN_DEVICE_FUNC virtual const Scalar coeff(DenseIndex index) const {
70 return m_impl.coeff(internal::convert_index<typename EvalType::Index>(index));
71 }
72 EIGEN_DEVICE_FUNC virtual Scalar& coeffRef(DenseIndex /*index*/) {
73 eigen_assert(false && "can't reference the coefficient of a rvalue");
74 return m_dummy;
75 }
76
77 protected:
78 TensorEvaluator<Expr, Device> m_impl;
79 Dimensions m_dims;
80 Scalar m_dummy;
81};
82
83template <typename Dimensions, typename Expr, typename Device>
84class TensorLazyEvaluatorWritable : public TensorLazyEvaluatorReadOnly<Dimensions, Expr, Device> {
85 public:
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;
90
91 TensorLazyEvaluatorWritable(const Expr& expr, const Device& device) : Base(expr, device) {}
92 virtual ~TensorLazyEvaluatorWritable() = default;
93
94 EIGEN_DEVICE_FUNC virtual Scalar& coeffRef(DenseIndex index) {
95 return this->m_impl.coeffRef(internal::convert_index<typename Base::EvalType::Index>(index));
96 }
97};
98
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>> {
102 public:
103 typedef std::conditional_t<IsWritable, TensorLazyEvaluatorWritable<Dimensions, Expr, Device>,
104 TensorLazyEvaluatorReadOnly<Dimensions, const Expr, Device>>
105 Base;
106 typedef typename Base::Scalar Scalar;
107
108 TensorLazyEvaluator(const Expr& expr, const Device& device) : Base(expr, device) {}
109 virtual ~TensorLazyEvaluator() = default;
110};
111
112template <typename Derived>
113class TensorRefBase : public TensorBase<Derived> {
114 public:
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;
125
126 static constexpr Index NumIndices = PlainObjectType::NumIndices;
127 typedef typename PlainObjectType::Dimensions Dimensions;
128
129 static constexpr int Layout = PlainObjectType::Layout;
130 enum {
131 IsAligned = false,
132 PacketAccess = false,
133 BlockAccess = false,
134 PreferBlockAccess = false,
135 CoordAccess = false, // to be implemented
136 RawAccess = false
137 };
138
139 //===- Tensor block evaluation strategy (see TensorBlock.h) -----------===//
140 typedef TensorBlockNotImplemented TensorBlock;
141 //===------------------------------------------------------------------===//
142
143 EIGEN_STRONG_INLINE TensorRefBase() = default;
144
145 TensorRefBase(const TensorRefBase& other) : TensorBase<Derived>(other), m_evaluator(other.m_evaluator) {
146 eigen_assert(m_evaluator->refCount() > 0);
147 m_evaluator->incrRefCount();
148 }
149
150 TensorRefBase& operator=(const TensorRefBase& other) {
151 if (this != &other) {
152 unrefEvaluator();
153 m_evaluator = other.m_evaluator;
154 eigen_assert(m_evaluator->refCount() > 0);
155 m_evaluator->incrRefCount();
156 }
157 return *this;
158 }
159
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 /*IsWritable=*/!std::is_const<PlainObjectType>::value &&
165 bool(is_lvalue<Expression>::value)>(expr, DefaultDevice())) {
166 m_evaluator->incrRefCount();
167 }
168
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) {
172 unrefEvaluator();
173 m_evaluator = new TensorLazyEvaluator < Dimensions, Expression, DefaultDevice,
174 /*IsWritable=*/!std::is_const<PlainObjectType>::value&& bool(is_lvalue<Expression>::value) >
175 (expr, DefaultDevice());
176 m_evaluator->incrRefCount();
177 return *this;
178 }
179
180 ~TensorRefBase() { unrefEvaluator(); }
181
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(); }
187
188 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar operator()(Index index) const { return m_evaluator->coeff(index); }
189
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);
196 }
197
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();
201 Index index = 0;
202 EIGEN_IF_CONSTEXPR (PlainObjectType::Options & RowMajor) {
203 index += indices[0];
204 for (size_t i = 1; i < NumIndices; ++i) {
205 index = index * dims[i] + indices[i];
206 }
207 } else {
208 index += indices[NumIndices - 1];
209 for (int i = NumIndices - 2; i >= 0; --i) {
210 index = index * dims[i] + indices[i];
211 }
212 }
213 return m_evaluator->coeff(index);
214 }
215
216 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Scalar coeff(Index index) const { return m_evaluator->coeff(index); }
217
218 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) { return m_evaluator->coeffRef(index); }
219
220 protected:
221 TensorLazyBaseEvaluator<Dimensions, Scalar>* evaluator() { return m_evaluator; }
222
223 private:
224 EIGEN_STRONG_INLINE void unrefEvaluator() {
225 if (m_evaluator) {
226 m_evaluator->decrRefCount();
227 if (m_evaluator->refCount() == 0) {
228 delete m_evaluator;
229 }
230 }
231 }
232
233 TensorLazyBaseEvaluator<Dimensions, Scalar>* m_evaluator = nullptr;
234};
235
236} // namespace internal
237
245template <typename PlainObjectType>
246class TensorRef : public internal::TensorRefBase<TensorRef<PlainObjectType>> {
247 typedef internal::TensorRefBase<TensorRef<PlainObjectType>> Base;
248
249 public:
250 using Scalar = typename Base::Scalar;
251 using Dimensions = typename Base::Dimensions;
252 // Without this, unqualified Index below does not find the dependent base's typedef and resolves to Eigen::Index,
253 // giving the accessors a different index type than the rest of the class.
254 using Index = typename Base::Index;
255
256 EIGEN_STRONG_INLINE TensorRef() = default;
257
258 EIGEN_STRONG_INLINE TensorRef(const TensorRef& other) = default;
259
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>?)");
265 }
266
267 TensorRef& operator=(const TensorRef& other) { return Base::operator=(other).derived(); }
268
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();
275 }
276
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);
283 }
284
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();
288 Index index = 0;
289 EIGEN_IF_CONSTEXPR (PlainObjectType::Options & RowMajor) {
290 index += indices[0];
291 for (size_t i = 1; i < NumIndices; ++i) {
292 index = index * dims[i] + indices[i];
293 }
294 } else {
295 index += indices[NumIndices - 1];
296 for (int i = NumIndices - 2; i >= 0; --i) {
297 index = index * dims[i] + indices[i];
298 }
299 }
300 return Base::evaluator()->coeffRef(index);
301 }
302
303 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) { return Base::evaluator()->coeffRef(index); }
304};
305
313template <typename PlainObjectType>
314class TensorRef<const PlainObjectType> : public internal::TensorRefBase<TensorRef<const PlainObjectType>> {
315 typedef internal::TensorRefBase<TensorRef<const PlainObjectType>> Base;
316
317 public:
318 EIGEN_STRONG_INLINE TensorRef() = default;
319
320 EIGEN_STRONG_INLINE TensorRef(const TensorRef& other) = default;
321
322 template <typename Expression>
323 EIGEN_STRONG_INLINE TensorRef(const Expression& expr) : Base(expr) {}
324
325 TensorRef& operator=(const TensorRef& other) { return Base::operator=(other).derived(); }
326
327 template <typename Expression>
328 EIGEN_STRONG_INLINE TensorRef& operator=(const Expression& expr) {
329 return Base::operator=(expr).derived();
330 }
331};
332
333// evaluator for rvalues
334template <typename Derived, typename Device>
335struct TensorEvaluator<const TensorRef<Derived>, Device> {
336 typedef typename Derived::Index Index;
337 typedef typename Derived::Scalar Scalar;
338 typedef typename Derived::Scalar CoeffReturnType;
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;
343
344 static constexpr int Layout = TensorRef<Derived>::Layout;
345 enum {
346 IsAligned = false,
347 PacketAccess = false,
348 BlockAccess = false,
349 PreferBlockAccess = false,
350 CoordAccess = false, // to be implemented
351 RawAccess = false
352 };
353
354 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
355 typedef internal::TensorBlockNotImplemented TensorBlock;
356 //===--------------------------------------------------------------------===//
357
358 EIGEN_STRONG_INLINE TensorEvaluator(const TensorRef<Derived>& m, const Device&) : m_ref(m) {}
359
360 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_ref.dimensions(); }
361
362 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType) { return true; }
363
364 EIGEN_STRONG_INLINE void cleanup() {}
365
366 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const { return m_ref.coeff(index); }
367
368 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) { return m_ref.coeffRef(index); }
369
370 EIGEN_DEVICE_FUNC const Scalar* data() const { return m_ref.data(); }
371
372 protected:
373 TensorRef<Derived> m_ref;
374};
375
376// evaluator for lvalues
377template <typename Derived, typename Device>
378struct TensorEvaluator<TensorRef<Derived>, Device> : public TensorEvaluator<const TensorRef<Derived>, 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;
384
385 typedef TensorEvaluator<const TensorRef<Derived>, Device> Base;
386
387 enum { IsAligned = false, PacketAccess = false, BlockAccess = false, PreferBlockAccess = false, RawAccess = false };
388
389 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
390 typedef internal::TensorBlockNotImplemented TensorBlock;
391 //===--------------------------------------------------------------------===//
392
393 EIGEN_STRONG_INLINE TensorEvaluator(TensorRef<Derived>& m, const Device& d) : Base(m, d) {}
394
395 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar& coeffRef(Index index) { return this->m_ref.coeffRef(index); }
396};
397
398} // end namespace Eigen
399
400#endif // EIGEN_TENSOR_TENSOR_REF_H
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