Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorAssign.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_ASSIGN_H
12#define EIGEN_TENSOR_TENSOR_ASSIGN_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20template <typename LhsXprType, typename RhsXprType>
21struct traits<TensorAssignOp<LhsXprType, RhsXprType> > {
22 typedef typename LhsXprType::Scalar Scalar;
23 typedef typename traits<LhsXprType>::StorageKind StorageKind;
24 typedef
25 typename promote_index_type<typename traits<LhsXprType>::Index, typename traits<RhsXprType>::Index>::type Index;
26 static constexpr std::size_t NumDimensions = internal::traits<LhsXprType>::NumDimensions;
27 static constexpr int Layout = internal::traits<LhsXprType>::Layout;
28 typedef typename traits<LhsXprType>::PointerType PointerType;
29
30 enum { Flags = 0 };
31};
32
33template <typename LhsXprType, typename RhsXprType>
34struct eval<TensorAssignOp<LhsXprType, RhsXprType>, Eigen::Dense> {
35 typedef const TensorAssignOp<LhsXprType, RhsXprType>& type;
36};
37
38} // end namespace internal
39
46template <typename LhsXprType, typename RhsXprType>
47class TensorAssignOp : public TensorBase<TensorAssignOp<LhsXprType, RhsXprType> > {
48 public:
49 typedef typename Eigen::internal::traits<TensorAssignOp>::Scalar Scalar;
50 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
51 typedef typename LhsXprType::CoeffReturnType CoeffReturnType;
52 typedef typename Eigen::internal::ref_selector<TensorAssignOp>::type Nested;
53 typedef typename Eigen::internal::traits<TensorAssignOp>::StorageKind StorageKind;
54 typedef typename Eigen::internal::traits<TensorAssignOp>::Index Index;
55
56 static constexpr int NumDims = Eigen::internal::traits<TensorAssignOp>::NumDimensions;
57
58 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorAssignOp(LhsXprType& lhs, const RhsXprType& rhs)
59 : m_lhs_xpr(lhs), m_rhs_xpr(rhs) {
60 EIGEN_STATIC_ASSERT((internal::traits<LhsXprType>::NumDimensions == internal::traits<RhsXprType>::NumDimensions),
61 Number_of_dimensions_must_match)
62 }
63
65 EIGEN_DEVICE_FUNC internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression() const {
66 return *((internal::remove_all_t<typename LhsXprType::Nested>*)&m_lhs_xpr);
67 }
68
69 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression() const {
70 return m_rhs_xpr;
71 }
72
73 protected:
74 internal::remove_all_t<typename LhsXprType::Nested>& m_lhs_xpr;
75 const internal::remove_all_t<typename RhsXprType::Nested>& m_rhs_xpr;
76};
77
78template <typename LeftArgType, typename RightArgType, typename Device>
79struct TensorEvaluator<const TensorAssignOp<LeftArgType, RightArgType>, Device> {
80 typedef TensorAssignOp<LeftArgType, RightArgType> XprType;
81 typedef typename XprType::Index Index;
82 using LeftIndex = typename TensorEvaluator<LeftArgType, Device>::Index;
83 using RightIndex = typename TensorEvaluator<RightArgType, Device>::Index;
84 typedef typename XprType::Scalar Scalar;
85 typedef typename XprType::CoeffReturnType CoeffReturnType;
86 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
87 typedef typename TensorEvaluator<RightArgType, Device>::Dimensions Dimensions;
88 typedef StorageMemory<CoeffReturnType, Device> Storage;
89 typedef typename Storage::Type EvaluatorPointerType;
90
91 static constexpr int PacketSize = PacketType<CoeffReturnType, Device>::size;
92 static constexpr int NumDims = XprType::NumDims;
93 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
94
95 enum {
96 IsAligned =
97 int(TensorEvaluator<LeftArgType, Device>::IsAligned) & int(TensorEvaluator<RightArgType, Device>::IsAligned),
98 PacketAccess = int(TensorEvaluator<LeftArgType, Device>::PacketAccess) &
99 int(TensorEvaluator<RightArgType, Device>::PacketAccess),
100 BlockAccess = int(TensorEvaluator<LeftArgType, Device>::BlockAccess) &
101 int(TensorEvaluator<RightArgType, Device>::BlockAccess),
102 PreferBlockAccess = int(TensorEvaluator<LeftArgType, Device>::PreferBlockAccess) |
103 int(TensorEvaluator<RightArgType, Device>::PreferBlockAccess),
104 RawAccess = TensorEvaluator<LeftArgType, Device>::RawAccess
105 };
106
107 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
108 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
109 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
110
111 typedef typename TensorEvaluator<const RightArgType, Device>::TensorBlock RightTensorBlock;
112 //===--------------------------------------------------------------------===//
113
114 TensorEvaluator(const XprType& op, const Device& device)
115 : m_leftImpl(op.lhsExpression(), device), m_rightImpl(op.rhsExpression(), device) {
116 EIGEN_STATIC_ASSERT((static_cast<int>(TensorEvaluator<LeftArgType, Device>::Layout) ==
117 static_cast<int>(TensorEvaluator<RightArgType, Device>::Layout)),
118 YOU_MADE_A_PROGRAMMING_MISTAKE);
119 }
120
121 EIGEN_DEVICE_FUNC const Dimensions& dimensions() const {
122 // The dimensions of the lhs and the rhs tensors should be equal to prevent
123 // overflows and ensure the result is fully initialized.
124 // TODO: use left impl instead if right impl dimensions are known at compile time.
125 return m_rightImpl.dimensions();
126 }
127
128 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType) {
129 eigen_assert(dimensions_match(m_leftImpl.dimensions(), m_rightImpl.dimensions()));
130 m_leftImpl.evalSubExprsIfNeeded(nullptr);
131 // If the lhs provides raw access to its storage area (i.e. if m_leftImpl.data() returns a non
132 // null value), attempt to evaluate the rhs expression in place. Returns true iff in place
133 // evaluation isn't supported and the caller still needs to manually assign the values generated
134 // by the rhs to the lhs.
135 return m_rightImpl.evalSubExprsIfNeeded(m_leftImpl.data());
136 }
137
138#ifdef EIGEN_USE_THREADS
139 template <typename EvalSubExprsCallback>
140 EIGEN_STRONG_INLINE void evalSubExprsIfNeededAsync(EvaluatorPointerType, EvalSubExprsCallback done) {
141 m_leftImpl.evalSubExprsIfNeededAsync(nullptr, [this, done](bool) {
142 m_rightImpl.evalSubExprsIfNeededAsync(m_leftImpl.data(), [done](bool need_assign) { done(need_assign); });
143 });
144 }
145#endif // EIGEN_USE_THREADS
146
147 EIGEN_STRONG_INLINE void cleanup() {
148 m_leftImpl.cleanup();
149 m_rightImpl.cleanup();
150 }
151
152 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void evalScalar(Index i) const {
153 m_leftImpl.coeffRef(internal::convert_index<LeftIndex>(i)) =
154 m_rightImpl.coeff(internal::convert_index<RightIndex>(i));
155 }
156 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void evalPacket(Index i) const {
157 constexpr int LhsStoreMode = TensorEvaluator<LeftArgType, Device>::IsAligned ? Aligned : Unaligned;
158 constexpr int RhsLoadMode = TensorEvaluator<RightArgType, Device>::IsAligned ? Aligned : Unaligned;
159 m_leftImpl.template writePacket<LhsStoreMode>(
160 internal::convert_index<LeftIndex>(i),
161 m_rightImpl.template packet<RhsLoadMode>(internal::convert_index<RightIndex>(i)));
162 }
163 EIGEN_DEVICE_FUNC CoeffReturnType coeff(Index index) const {
164 return m_leftImpl.coeff(internal::convert_index<LeftIndex>(index));
165 }
166 template <int LoadMode>
167 EIGEN_DEVICE_FUNC PacketReturnType packet(Index index) const {
168 return m_leftImpl.template packet<LoadMode>(internal::convert_index<LeftIndex>(index));
169 }
170
171 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool vectorized) const {
172 // We assume that evalPacket or evalScalar is called to perform the
173 // assignment and account for the cost of the write here, but reduce left
174 // cost by one load because we are using m_leftImpl.coeffRef.
175 TensorOpCost left = m_leftImpl.costPerCoeff(vectorized);
176 return m_rightImpl.costPerCoeff(vectorized) +
177 TensorOpCost(numext::maxi(0.0, left.bytes_loaded() - sizeof(CoeffReturnType)), left.bytes_stored(),
178 left.compute_cycles()) +
179 TensorOpCost(0, sizeof(CoeffReturnType), 0, vectorized, PacketSize);
180 }
181
182 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements() const {
183 return internal::TensorBlockResourceRequirements::merge(m_leftImpl.getResourceRequirements(),
184 m_rightImpl.getResourceRequirements());
185 }
186
187 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void evalBlock(TensorBlockDesc& desc, TensorBlockScratch& scratch) {
188 if (TensorEvaluator<LeftArgType, Device>::RawAccess && m_leftImpl.data() != nullptr) {
189 // If destination has raw data access, we pass it as a potential
190 // destination for a block descriptor evaluation.
191 desc.template AddDestinationBuffer<Layout>(
192 /*dst_base=*/m_leftImpl.data() + desc.offset(),
193 /*dst_strides=*/internal::strides<Layout>(m_leftImpl.dimensions()));
194 }
195
196 RightTensorBlock block = m_rightImpl.block(desc, scratch, /*root_of_expr_ast=*/true);
197 // If block was evaluated into a destination, there is no need to do assignment.
198 if (block.kind() != internal::TensorBlockKind::kMaterializedInOutput) {
199 m_leftImpl.writeBlock(desc, block);
200 }
201 block.cleanup();
202 }
203
204 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return m_leftImpl.data(); }
205
206 private:
207 TensorEvaluator<LeftArgType, Device> m_leftImpl;
208 TensorEvaluator<RightArgType, Device> m_rightImpl;
209};
210
211} // namespace Eigen
212
213#endif // EIGEN_TENSOR_TENSOR_ASSIGN_H
Definition TensorAssign.h:47
internal::remove_all_t< typename LhsXprType::Nested > & lhsExpression() const
Definition TensorAssign.h:65
The tensor base class.
Definition TensorForwardDeclarations.h:69
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47