Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorTrace.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2017 Gagan Goel <gagan.nith@gmail.com>
5// Copyright (C) 2017 Benoit Steiner <benoit.steiner.goog@gmail.com>
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_TENSOR_TENSOR_TRACE_H
13#define EIGEN_TENSOR_TENSOR_TRACE_H
14
15// IWYU pragma: private
16#include "./InternalHeaderCheck.h"
17
18namespace Eigen {
19
20namespace internal {
21template <typename Dims, typename XprType>
22struct traits<TensorTraceOp<Dims, XprType> > : public traits<XprType> {
23 typedef typename XprType::Scalar Scalar;
24 typedef traits<XprType> XprTraits;
25 typedef typename XprTraits::StorageKind StorageKind;
26 typedef typename XprTraits::Index Index;
27 static constexpr int NumDimensions = XprTraits::NumDimensions - array_size<Dims>::value;
28 static constexpr int Layout = XprTraits::Layout;
29 enum {
30 // Trace is read-only.
31 Flags = traits<XprType>::Flags & ~LvalueBit
32 };
33};
34
35template <typename Dims, typename XprType>
36struct eval<TensorTraceOp<Dims, XprType>, Eigen::Dense> {
37 typedef const TensorTraceOp<Dims, XprType>& type;
38};
39
40} // end namespace internal
41
47template <typename Dims, typename XprType>
48class TensorTraceOp : public TensorBase<TensorTraceOp<Dims, XprType> > {
49 public:
50 typedef typename Eigen::internal::traits<TensorTraceOp>::Scalar Scalar;
51 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
52 typedef typename XprType::CoeffReturnType CoeffReturnType;
53 typedef typename Eigen::internal::ref_selector<TensorTraceOp>::type Nested;
54 typedef typename Eigen::internal::traits<TensorTraceOp>::StorageKind StorageKind;
55 typedef typename Eigen::internal::traits<TensorTraceOp>::Index Index;
56
57 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorTraceOp(const XprType& expr, const Dims& dims)
58 : m_xpr(expr), m_dims(dims) {}
59
60 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dims& dims() const { return m_dims; }
61
62 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const internal::remove_all_t<typename XprType::Nested>& expression() const {
63 return m_xpr;
64 }
65
66 protected:
67 typename XprType::Nested m_xpr;
68 const Dims m_dims;
69};
70
71// Eval as rvalue
72template <typename Dims, typename ArgType, typename Device>
73struct TensorEvaluator<const TensorTraceOp<Dims, ArgType>, Device> {
74 typedef TensorTraceOp<Dims, ArgType> XprType;
75 static constexpr int NumInputDims =
76 internal::array_size<typename TensorEvaluator<ArgType, Device>::Dimensions>::value;
77 static constexpr int NumReducedDims = internal::array_size<Dims>::value;
78 static constexpr int NumOutputDims = NumInputDims - NumReducedDims;
79 typedef typename XprType::Index Index;
80 typedef DSizes<Index, NumOutputDims> Dimensions;
81 typedef typename XprType::Scalar Scalar;
82 typedef typename XprType::CoeffReturnType CoeffReturnType;
83 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
84 static constexpr int PacketSize = internal::unpacket_traits<PacketReturnType>::size;
85 typedef StorageMemory<CoeffReturnType, Device> Storage;
86 typedef typename Storage::Type EvaluatorPointerType;
87
88 static constexpr int Layout = TensorEvaluator<ArgType, Device>::Layout;
89 enum {
90 IsAligned = false,
91 PacketAccess = TensorEvaluator<ArgType, Device>::PacketAccess,
92 BlockAccess = false,
94 CoordAccess = false,
95 RawAccess = false
96 };
97
98 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
99 typedef internal::TensorBlockNotImplemented TensorBlock;
100 //===--------------------------------------------------------------------===//
101
102 EIGEN_STRONG_INLINE TensorEvaluator(const XprType& op, const Device& device)
103 : m_impl(op.expression(), device), m_traceDim(1), m_device(device) {
104 EIGEN_STATIC_ASSERT((NumOutputDims >= 0), YOU_MADE_A_PROGRAMMING_MISTAKE);
105 EIGEN_STATIC_ASSERT((NumReducedDims >= 2) || ((NumReducedDims == 0) && (NumInputDims == 0)),
106 YOU_MADE_A_PROGRAMMING_MISTAKE);
107
108 for (int i = 0; i < NumInputDims; ++i) {
109 m_reduced[i] = false;
110 }
111
112 const Dims& op_dims = op.dims();
113 for (int i = 0; i < NumReducedDims; ++i) {
114 eigen_assert(op_dims[i] >= 0);
115 eigen_assert(op_dims[i] < NumInputDims);
116 m_reduced[op_dims[i]] = true;
117 }
118
119 // All the dimensions should be distinct to compute the trace
120 int num_distinct_reduce_dims = 0;
121 for (int i = 0; i < NumInputDims; ++i) {
122 if (m_reduced[i]) {
123 ++num_distinct_reduce_dims;
124 }
125 }
126
127 EIGEN_ONLY_USED_FOR_DEBUG(num_distinct_reduce_dims);
128 eigen_assert(num_distinct_reduce_dims == NumReducedDims);
129
130 // Compute the dimensions of the result.
131 const typename TensorEvaluator<ArgType, Device>::Dimensions& input_dims = m_impl.dimensions();
132
133 int output_index = 0;
134 int reduced_index = 0;
135 for (int i = 0; i < NumInputDims; ++i) {
136 if (m_reduced[i]) {
137 m_reducedDims[reduced_index] = input_dims[i];
138 if (reduced_index > 0) {
139 // All the trace dimensions must have the same size
140 eigen_assert(m_reducedDims[0] == m_reducedDims[reduced_index]);
141 }
142 ++reduced_index;
143 } else {
144 m_dimensions[output_index] = input_dims[i];
145 ++output_index;
146 }
147 }
148
149 EIGEN_IF_CONSTEXPR (NumReducedDims != 0) {
150 m_traceDim = m_reducedDims[0];
151 }
152
153 // Compute the output strides
154 EIGEN_IF_CONSTEXPR (NumOutputDims > 0) {
155 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
156 m_outputStrides[0] = 1;
157 for (int i = 1; i < NumOutputDims; ++i) {
158 m_outputStrides[i] = m_outputStrides[i - 1] * m_dimensions[i - 1];
159 }
160 } else {
161 m_outputStrides.back() = 1;
162 for (int i = NumOutputDims - 2; i >= 0; --i) {
163 m_outputStrides[i] = m_outputStrides[i + 1] * m_dimensions[i + 1];
164 }
165 }
166 }
167
168 // Compute the input strides
169 EIGEN_IF_CONSTEXPR (NumInputDims > 0) {
170 array<Index, NumInputDims> input_strides;
171 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
172 input_strides[0] = 1;
173 for (int i = 1; i < NumInputDims; ++i) {
174 input_strides[i] = input_strides[i - 1] * input_dims[i - 1];
175 }
176 } else {
177 input_strides.back() = 1;
178 for (int i = NumInputDims - 2; i >= 0; --i) {
179 input_strides[i] = input_strides[i + 1] * input_dims[i + 1];
180 }
181 }
182
183 output_index = 0;
184 reduced_index = 0;
185 for (int i = 0; i < NumInputDims; ++i) {
186 if (m_reduced[i]) {
187 m_reducedStrides[reduced_index] = input_strides[i];
188 ++reduced_index;
189 } else {
190 m_preservedStrides[output_index] = input_strides[i];
191 ++output_index;
192 }
193 }
194 }
195 }
196
197 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_dimensions; }
198
199 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType /*data*/) {
200 m_impl.evalSubExprsIfNeeded(nullptr);
201 return true;
202 }
203
204 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return nullptr; }
205
206 EIGEN_STRONG_INLINE void cleanup() { m_impl.cleanup(); }
207
208 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const {
209 // Initialize the result
210 CoeffReturnType result = internal::cast<int, CoeffReturnType>(0);
211 Index index_stride = 0;
212 for (int i = 0; i < NumReducedDims; ++i) {
213 index_stride += m_reducedStrides[i];
214 }
215
216 // If trace is requested along all dimensions, starting index would be 0
217 Index cur_index = 0;
218 EIGEN_IF_CONSTEXPR (NumOutputDims != 0) {
219 cur_index = firstInput(index);
220 }
221 for (Index i = 0; i < m_traceDim; ++i) {
222 result += m_impl.coeff(cur_index);
223 cur_index += index_stride;
224 }
225
226 return result;
227 }
228
229 template <int LoadMode>
230 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packet(Index index) const {
231 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
232
233 EIGEN_ALIGN_MAX std::remove_const_t<CoeffReturnType> values[PacketSize];
234 for (int i = 0; i < PacketSize; ++i) {
235 values[i] = coeff(index + i);
236 }
237 PacketReturnType result = internal::ploadt<PacketReturnType, LoadMode>(values);
238 return result;
239 }
240
241 protected:
242 // Given the output index, finds the first index in the input tensor used to compute the trace
243 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index firstInput(Index index) const { return firstInputImpl(index); }
244
245 template <int ND = NumOutputDims>
246 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<ND == 0, Index> firstInputImpl(Index /*index*/) const {
247 return 0;
248 }
249
250 template <int ND = NumOutputDims>
251 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE std::enable_if_t<(ND > 0), Index> firstInputImpl(Index index) const {
252 Index startInput = 0;
253 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
254 for (int i = ND - 1; i > 0; --i) {
255 const Index idx = index / m_outputStrides[i];
256 startInput += idx * m_preservedStrides[i];
257 index -= idx * m_outputStrides[i];
258 }
259 startInput += index * m_preservedStrides[0];
260 } else {
261 for (int i = 0; i < ND - 1; ++i) {
262 const Index idx = index / m_outputStrides[i];
263 startInput += idx * m_preservedStrides[i];
264 index -= idx * m_outputStrides[i];
265 }
266 startInput += index * m_preservedStrides[ND - 1];
267 }
268 return startInput;
269 }
270
271 Dimensions m_dimensions;
272 TensorEvaluator<ArgType, Device> m_impl;
273 // Initialize the size of the trace dimension
274 Index m_traceDim;
275 const Device EIGEN_DEVICE_REF m_device;
276 array<bool, NumInputDims> m_reduced;
277 array<Index, NumReducedDims> m_reducedDims;
278 array<Index, NumOutputDims> m_outputStrides;
279 array<Index, NumReducedDims> m_reducedStrides;
280 array<Index, NumOutputDims> m_preservedStrides;
281};
282
283} // End namespace Eigen
284
285#endif // EIGEN_TENSOR_TENSOR_TRACE_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
Tensor Trace class.
Definition TensorTrace.h:48
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47