10#ifndef EIGEN_CXX11_TENSOR_TENSOR_ASSIGN_H
11#define EIGEN_CXX11_TENSOR_TENSOR_ASSIGN_H
16template<
typename LhsXprType,
typename RhsXprType>
17struct traits<TensorAssignOp<LhsXprType, RhsXprType> >
19 typedef typename LhsXprType::Scalar Scalar;
20 typedef typename traits<LhsXprType>::StorageKind StorageKind;
21 typedef typename promote_index_type<typename traits<LhsXprType>::Index,
22 typename traits<RhsXprType>::Index>::type
Index;
23 typedef typename LhsXprType::Nested LhsNested;
24 typedef typename RhsXprType::Nested RhsNested;
25 typedef typename remove_reference<LhsNested>::type _LhsNested;
26 typedef typename remove_reference<RhsNested>::type _RhsNested;
27 static const std::size_t NumDimensions = internal::traits<LhsXprType>::NumDimensions;
28 static const int Layout = internal::traits<LhsXprType>::Layout;
35template<
typename LhsXprType,
typename RhsXprType>
36struct eval<TensorAssignOp<LhsXprType, RhsXprType>, Eigen::Dense>
38 typedef const TensorAssignOp<LhsXprType, RhsXprType>& type;
41template<
typename LhsXprType,
typename RhsXprType>
42struct nested<TensorAssignOp<LhsXprType, RhsXprType>, 1, typename eval<TensorAssignOp<LhsXprType, RhsXprType> >::type>
44 typedef TensorAssignOp<LhsXprType, RhsXprType> type;
55template <
typename LhsXprType,
typename RhsXprType>
56class TensorAssignOp :
public TensorBase<TensorAssignOp<LhsXprType, RhsXprType> > {
58 typedef typename Eigen::internal::traits<TensorAssignOp>::Scalar Scalar;
60 typedef typename LhsXprType::CoeffReturnType CoeffReturnType;
61 typedef typename Eigen::internal::nested<TensorAssignOp>::type Nested;
62 typedef typename Eigen::internal::traits<TensorAssignOp>::StorageKind StorageKind;
63 typedef typename Eigen::internal::traits<TensorAssignOp>::Index Index;
65 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorAssignOp(LhsXprType& lhs,
const RhsXprType& rhs)
66 : m_lhs_xpr(lhs), m_rhs_xpr(rhs) {}
70 typename internal::remove_all<typename LhsXprType::Nested>::type&
71 lhsExpression()
const {
return *((
typename internal::remove_all<typename LhsXprType::Nested>::type*)&m_lhs_xpr); }
74 const typename internal::remove_all<typename RhsXprType::Nested>::type&
75 rhsExpression()
const {
return m_rhs_xpr; }
78 typename internal::remove_all<typename LhsXprType::Nested>::type& m_lhs_xpr;
79 const typename internal::remove_all<typename RhsXprType::Nested>::type& m_rhs_xpr;
83template<
typename LeftArgType,
typename RightArgType,
typename Device>
86 typedef TensorAssignOp<LeftArgType, RightArgType> XprType;
87 typedef typename XprType::Index Index;
88 typedef typename XprType::Scalar Scalar;
89 typedef typename XprType::CoeffReturnType CoeffReturnType;
90 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
91 typedef typename TensorEvaluator<RightArgType, Device>::Dimensions Dimensions;
92 static const int PacketSize = internal::unpacket_traits<PacketReturnType>::size;
95 IsAligned = TensorEvaluator<LeftArgType, Device>::IsAligned & TensorEvaluator<RightArgType, Device>::IsAligned,
96 PacketAccess = TensorEvaluator<LeftArgType, Device>::PacketAccess & TensorEvaluator<RightArgType, Device>::PacketAccess,
97 Layout = TensorEvaluator<LeftArgType, Device>::Layout,
98 RawAccess = TensorEvaluator<LeftArgType, Device>::RawAccess
101 EIGEN_DEVICE_FUNC TensorEvaluator(
const XprType& op,
const Device&
device) :
102 m_leftImpl(op.lhsExpression(),
device),
103 m_rightImpl(op.rhsExpression(),
device)
105 EIGEN_STATIC_ASSERT((
static_cast<int>(TensorEvaluator<LeftArgType, Device>::Layout) ==
static_cast<int>(TensorEvaluator<RightArgType, Device>::Layout)), YOU_MADE_A_PROGRAMMING_MISTAKE);
108 EIGEN_DEVICE_FUNC
const Dimensions& dimensions()
const
113 return m_rightImpl.dimensions();
116 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(Scalar*) {
117 eigen_assert(dimensions_match(m_leftImpl.dimensions(), m_rightImpl.dimensions()));
118 m_leftImpl.evalSubExprsIfNeeded(NULL);
123 return m_rightImpl.evalSubExprsIfNeeded(m_leftImpl.data());
125 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void cleanup() {
126 m_leftImpl.cleanup();
127 m_rightImpl.cleanup();
130 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void evalScalar(Index i) {
131 m_leftImpl.coeffRef(i) = m_rightImpl.coeff(i);
133 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
void evalPacket(Index i) {
134 const int LhsStoreMode = TensorEvaluator<LeftArgType, Device>::IsAligned ?
Aligned :
Unaligned;
135 const int RhsLoadMode = TensorEvaluator<RightArgType, Device>::IsAligned ?
Aligned :
Unaligned;
136 m_leftImpl.template writePacket<LhsStoreMode>(i, m_rightImpl.template packet<RhsLoadMode>(i));
138 EIGEN_DEVICE_FUNC CoeffReturnType coeff(Index index)
const
140 return m_leftImpl.coeff(index);
142 template<
int LoadMode>
143 EIGEN_DEVICE_FUNC PacketReturnType packet(Index index)
const
145 return m_leftImpl.template packet<LoadMode>(index);
148 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost
149 costPerCoeff(
bool vectorized)
const {
153 TensorOpCost left = m_leftImpl.costPerCoeff(vectorized);
154 return m_rightImpl.costPerCoeff(vectorized) +
156 numext::maxi(0.0, left.bytes_loaded() -
sizeof(CoeffReturnType)),
157 left.bytes_stored(), left.compute_cycles()) +
158 TensorOpCost(0,
sizeof(CoeffReturnType), 0, vectorized, PacketSize);
162 const TensorEvaluator<LeftArgType, Device>& left_impl()
const {
return m_leftImpl; }
164 const TensorEvaluator<RightArgType, Device>& right_impl()
const {
return m_rightImpl; }
166 EIGEN_DEVICE_FUNC CoeffReturnType* data()
const {
return m_leftImpl.data(); }
169 TensorEvaluator<LeftArgType, Device> m_leftImpl;
170 TensorEvaluator<RightArgType, Device> m_rightImpl;
Definition TensorAssign.h:56
internal::remove_all< typenameLhsXprType::Nested >::type & lhsExpression() const
Definition TensorAssign.h:71
The tensor base class.
Definition TensorForwardDeclarations.h:29
Namespace containing all symbols from the Eigen library.
EIGEN_DEFAULT_DENSE_INDEX_TYPE Index
The tensor evaluator class.
Definition TensorEvaluator.h:27
const Device & device() const
required by sycl in order to construct sycl buffer from raw pointer
Definition TensorEvaluator.h:112