11#ifndef EIGEN_TENSOR_TENSOR_ASSIGN_H
12#define EIGEN_TENSOR_TENSOR_ASSIGN_H
15#include "./InternalHeaderCheck.h"
20template <
typename LhsXprType,
typename RhsXprType>
21struct traits<TensorAssignOp<LhsXprType, RhsXprType> > {
22 typedef typename LhsXprType::Scalar Scalar;
23 typedef typename traits<LhsXprType>::StorageKind StorageKind;
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;
33template <
typename LhsXprType,
typename RhsXprType>
34struct eval<TensorAssignOp<LhsXprType, RhsXprType>, Eigen::Dense> {
35 typedef const TensorAssignOp<LhsXprType, RhsXprType>& type;
46template <
typename LhsXprType,
typename RhsXprType>
47class TensorAssignOp :
public TensorBase<TensorAssignOp<LhsXprType, RhsXprType> > {
49 typedef typename Eigen::internal::traits<TensorAssignOp>::Scalar Scalar;
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;
56 static constexpr int NumDims = Eigen::internal::traits<TensorAssignOp>::NumDimensions;
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)
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);
69 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression()
const {
74 internal::remove_all_t<typename LhsXprType::Nested>& m_lhs_xpr;
75 const internal::remove_all_t<typename RhsXprType::Nested>& m_rhs_xpr;
78template <
typename LeftArgType,
typename RightArgType,
typename 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;
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;
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
108 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
109 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
111 typedef typename TensorEvaluator<const RightArgType, Device>::TensorBlock RightTensorBlock;
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);
121 EIGEN_DEVICE_FUNC
const Dimensions& dimensions()
const {
125 return m_rightImpl.dimensions();
128 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType) {
129 eigen_assert(dimensions_match(m_leftImpl.dimensions(), m_rightImpl.dimensions()));
130 m_leftImpl.evalSubExprsIfNeeded(
nullptr);
135 return m_rightImpl.evalSubExprsIfNeeded(m_leftImpl.data());
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); });
147 EIGEN_STRONG_INLINE
void cleanup() {
148 m_leftImpl.cleanup();
149 m_rightImpl.cleanup();
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));
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)));
163 EIGEN_DEVICE_FUNC CoeffReturnType coeff(Index index)
const {
164 return m_leftImpl.coeff(internal::convert_index<LeftIndex>(index));
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));
171 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool vectorized)
const {
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);
182 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
183 return internal::TensorBlockResourceRequirements::merge(m_leftImpl.getResourceRequirements(),
184 m_rightImpl.getResourceRequirements());
187 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void evalBlock(TensorBlockDesc& desc, TensorBlockScratch& scratch) {
188 if (TensorEvaluator<LeftArgType, Device>::RawAccess && m_leftImpl.data() !=
nullptr) {
191 desc.template AddDestinationBuffer<Layout>(
192 m_leftImpl.data() + desc.offset(),
193 internal::strides<Layout>(m_leftImpl.dimensions()));
196 RightTensorBlock block = m_rightImpl.block(desc, scratch,
true);
198 if (block.kind() != internal::TensorBlockKind::kMaterializedInOutput) {
199 m_leftImpl.writeBlock(desc, block);
204 EIGEN_DEVICE_FUNC EvaluatorPointerType data()
const {
return m_leftImpl.data(); }
207 TensorEvaluator<LeftArgType, Device> m_leftImpl;
208 TensorEvaluator<RightArgType, Device> m_rightImpl;
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