Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorContractionMapper.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_CONTRACTION_MAPPER_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_MAPPER_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21constexpr int Rhs = 0;
22constexpr int Lhs = 1;
23
24/*
25 * Implementation of the Eigen blas_data_mapper class for tensors.
26 */
29template <typename Tensor, bool HasRawAccess, template <class> class MakePointer_ = MakePointer>
30struct CoeffLoader;
31
32template <typename Scalar, typename Index, int side, typename Tensor, typename nocontract_t, typename contract_t,
33 int packet_size, bool inner_dim_contiguous, bool inner_dim_reordered, int Alignment,
34 template <class> class MakePointer_ = MakePointer>
35class BaseTensorContractionMapper;
36
37template <typename Tensor, bool HasRawAccess, template <class> class MakePointer_>
38struct CoeffLoader {
39 static constexpr bool DirectOffsets = false;
40
41 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE CoeffLoader(const Tensor& tensor) : m_tensor(tensor) {}
42
43 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void offsetBuffer(typename Tensor::Index) {
44 eigen_assert(false && "unsupported");
45 }
46
47 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE const typename MakePointer_<const typename Tensor::Scalar>::Type data() const {
48 eigen_assert(false && "unsupported");
49 return nullptr;
50 }
51
52 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE typename Tensor::Scalar coeff(typename Tensor::Index index) const {
53 return m_tensor.coeff(index);
54 }
55
56 template <int LoadMode>
57 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE typename Tensor::PacketReturnType packet(typename Tensor::Index index) const {
58 return m_tensor.template packet<LoadMode>(index);
59 }
60
61 private:
62 const Tensor m_tensor;
63};
64
65template <typename Tensor, template <class> class MakePointer_>
66struct CoeffLoader<Tensor, true, MakePointer_> {
67 static constexpr bool DirectOffsets = true;
68
69 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE CoeffLoader(const Tensor& tensor) : m_data(tensor.data()) {}
70
71 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void offsetBuffer(typename Tensor::Index offset) { m_data += offset; }
72
73 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE const typename MakePointer_<const typename Tensor::Scalar>::Type data() const {
74 return m_data;
75 }
76
77 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE typename Tensor::Scalar coeff(typename Tensor::Index index) const {
78 return loadConstant(m_data + index);
79 }
80
81 template <int LoadMode>
82 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE typename Tensor::PacketReturnType packet(typename Tensor::Index index) const {
83 return internal::ploadt_ro<typename Tensor::PacketReturnType, LoadMode>(m_data + index);
84 }
85
86 private:
87 using Scalar = typename Tensor::Scalar;
88
89 typename MakePointer_<const Scalar>::Type m_data;
90};
91
92template <typename Scalar, typename Index, int side, typename Tensor, typename nocontract_t, typename contract_t,
93 int packet_size, bool inner_dim_contiguous, int Alignment, template <class> class MakePointer_ = MakePointer>
94class SimpleTensorContractionMapper {
95 public:
96 EIGEN_DEVICE_FUNC SimpleTensorContractionMapper(const Tensor& tensor, const nocontract_t& nocontract_strides,
97 const nocontract_t& ij_strides, const contract_t& contract_strides,
98 const contract_t& k_strides)
99 : m_tensor(tensor),
100 m_nocontract_strides(nocontract_strides),
101 m_ij_strides(ij_strides),
102 m_contract_strides(contract_strides),
103 m_k_strides(k_strides) {}
104
105 static constexpr bool DirectOffsets = CoeffLoader<Tensor, Tensor::RawAccess, MakePointer_>::DirectOffsets;
106
107 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void offsetBuffer(typename Tensor::Index offset) {
108 m_tensor.offsetBuffer(offset);
109 }
110
111 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE void prefetch(Index /*i*/) {}
112
113 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(Index row) const {
114 // column major assumption
115 return operator()(row, 0);
116 }
117
118 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator()(Index row, Index col) const {
119 return m_tensor.coeff(computeIndex(row, col));
120 }
121
122#ifdef EIGEN_MULTIDIMENSIONAL_SUBSCRIPT
123 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Scalar operator[](Index row, Index col) const { return operator()(row, col); }
124#endif
125
126 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index computeIndex(Index row, Index col) const {
127 const bool left = (side == Lhs);
128 Index nocontract_val = left ? row : col;
129 Index linidx = 0;
130 EIGEN_UNROLL_LOOP
131 for (int i = static_cast<int>(array_size<nocontract_t>::value) - 1; i > 0; i--) {
132 const Index idx = nocontract_val / m_ij_strides[i];
133 linidx += idx * m_nocontract_strides[i];
134 nocontract_val -= idx * m_ij_strides[i];
135 }
136 EIGEN_IF_CONSTEXPR (array_size<typename Tensor::Dimensions>::value > array_size<contract_t>::value) {
137 EIGEN_IF_CONSTEXPR (side == Lhs && inner_dim_contiguous) {
138 eigen_assert(m_nocontract_strides[0] == 1);
139 linidx += nocontract_val;
140 } else {
141 linidx += nocontract_val * m_nocontract_strides[0];
142 }
143 }
144
145 Index contract_val = left ? col : row;
146 EIGEN_IF_CONSTEXPR (array_size<contract_t>::value > 0) {
147 EIGEN_UNROLL_LOOP
148 for (int i = static_cast<int>(array_size<contract_t>::value) - 1; i > 0; i--) {
149 const Index idx = contract_val / m_k_strides[i];
150 linidx += idx * m_contract_strides[i];
151 contract_val -= idx * m_k_strides[i];
152 }
153
154 EIGEN_IF_CONSTEXPR (side == Rhs && inner_dim_contiguous) {
155 eigen_assert(m_contract_strides[0] == 1);
156 linidx += contract_val;
157 } else {
158 linidx += contract_val * m_contract_strides[0];
159 }
160 }
161
162 return linidx;
163 }
164
165 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE IndexPair<Index> computeIndexPair(Index row, Index col,
166 const Index distance) const {
167 const bool left = (side == Lhs);
168 Index nocontract_val[2] = {left ? row : col, left ? row + distance : col};
169 Index linidx[2] = {0, 0};
170 EIGEN_IF_CONSTEXPR (array_size<typename Tensor::Dimensions>::value > array_size<contract_t>::value) {
171 EIGEN_UNROLL_LOOP
172 for (int i = static_cast<int>(array_size<nocontract_t>::value) - 1; i > 0; i--) {
173 const Index idx0 = nocontract_val[0] / m_ij_strides[i];
174 const Index idx1 = nocontract_val[1] / m_ij_strides[i];
175 linidx[0] += idx0 * m_nocontract_strides[i];
176 linidx[1] += idx1 * m_nocontract_strides[i];
177 nocontract_val[0] -= idx0 * m_ij_strides[i];
178 nocontract_val[1] -= idx1 * m_ij_strides[i];
179 }
180 EIGEN_IF_CONSTEXPR (side == Lhs && inner_dim_contiguous) {
181 eigen_assert(m_nocontract_strides[0] == 1);
182 linidx[0] += nocontract_val[0];
183 linidx[1] += nocontract_val[1];
184 } else {
185 linidx[0] += nocontract_val[0] * m_nocontract_strides[0];
186 linidx[1] += nocontract_val[1] * m_nocontract_strides[0];
187 }
188 }
189
190 Index contract_val[2] = {left ? col : row, left ? col : row + distance};
191 EIGEN_IF_CONSTEXPR (array_size<contract_t>::value > 0) {
192 EIGEN_UNROLL_LOOP
193 for (int i = static_cast<int>(array_size<contract_t>::value) - 1; i > 0; i--) {
194 const Index idx0 = contract_val[0] / m_k_strides[i];
195 const Index idx1 = contract_val[1] / m_k_strides[i];
196 linidx[0] += idx0 * m_contract_strides[i];
197 linidx[1] += idx1 * m_contract_strides[i];
198 contract_val[0] -= idx0 * m_k_strides[i];
199 contract_val[1] -= idx1 * m_k_strides[i];
200 }
201
202 EIGEN_IF_CONSTEXPR (side == Rhs && inner_dim_contiguous) {
203 eigen_assert(m_contract_strides[0] == 1);
204 linidx[0] += contract_val[0];
205 linidx[1] += contract_val[1];
206 } else {
207 linidx[0] += contract_val[0] * m_contract_strides[0];
208 linidx[1] += contract_val[1] * m_contract_strides[0];
209 }
210 }
211 return IndexPair<Index>(linidx[0], linidx[1]);
212 }
213
214 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Index firstAligned(Index size) const {
215 // Only claim alignment when we can compute the actual stride (ie when we're
216 // dealing with the lhs with inner_dim_contiguous). This is because the
217 // matrix-vector product relies on the stride when dealing with aligned inputs.
218 return (Alignment == Aligned) && (side == Lhs) && inner_dim_contiguous ? 0 : size;
219 }
220 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Index stride() const {
221 return ((side == Lhs) && inner_dim_contiguous && array_size<contract_t>::value > 0) ? m_contract_strides[0] : 1;
222 }
223
224 const CoeffLoader<Tensor, Tensor::RawAccess, MakePointer_>& tensor() const { return m_tensor; }
225
226 const nocontract_t& nocontract_strides() const { return m_nocontract_strides; }
227 const nocontract_t& ij_strides() const { return m_ij_strides; }
228 const contract_t& contract_strides() const { return m_contract_strides; }
229 const contract_t& k_strides() const { return m_k_strides; }
230
231 protected:
232 CoeffLoader<Tensor, Tensor::RawAccess, MakePointer_> m_tensor;
233 const nocontract_t m_nocontract_strides;
234 const nocontract_t m_ij_strides;
235 const contract_t m_contract_strides;
236 const contract_t m_k_strides;
237};
238
239template <typename Scalar, typename Index, int side, typename Tensor, typename nocontract_t, typename contract_t,
240 int packet_size, bool inner_dim_contiguous, bool inner_dim_reordered, int Alignment,
241 template <class> class MakePointer_>
242class BaseTensorContractionMapper
243 : public SimpleTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, packet_size,
244 inner_dim_contiguous, Alignment, MakePointer_> {
245 public:
246 using ParentMapper = SimpleTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, packet_size,
247 inner_dim_contiguous, Alignment, MakePointer_>;
248
249 EIGEN_DEVICE_FUNC BaseTensorContractionMapper(const Tensor& tensor, const nocontract_t& nocontract_strides,
250 const nocontract_t& ij_strides, const contract_t& contract_strides,
251 const contract_t& k_strides)
252 : ParentMapper(tensor, nocontract_strides, ij_strides, contract_strides, k_strides) {}
253
254 template <typename PacketT, int AlignmentType>
255 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
256 std::enable_if_t<internal::unpacket_traits<PacketT>::size == packet_size, PacketT>
257 load(Index i, Index j) const {
258 // whole method makes column major assumption
259
260 // don't need to add offsets for now (because operator handles that)
261 // current code assumes packet size must be a multiple of 2
262 EIGEN_STATIC_ASSERT(packet_size % 2 == 0, YOU_MADE_A_PROGRAMMING_MISTAKE);
263
264 EIGEN_IF_CONSTEXPR (Tensor::PacketAccess && inner_dim_contiguous && !inner_dim_reordered) {
265 const Index index = this->computeIndex(i, j);
266 eigen_assert(this->computeIndex(i + packet_size - 1, j) == index + packet_size - 1);
267 return this->m_tensor.template packet<AlignmentType>(index);
268 }
269
270 const IndexPair<Index> indexPair = this->computeIndexPair(i, j, packet_size - 1);
271 const Index first = indexPair.first;
272 const Index lastIdx = indexPair.second;
273
274 // We can always do optimized packet reads from left hand side right now, because
275 // the vertical matrix dimension on the left hand side is never contracting.
276 // On the right hand side we need to check if the contracting dimensions may have
277 // been shuffled first.
278 EIGEN_IF_CONSTEXPR (Tensor::PacketAccess &&
279 (side == Lhs || internal::array_size<contract_t>::value <= 1 || !inner_dim_reordered)) {
280 if ((lastIdx - first) == (packet_size - 1)) {
281 return this->m_tensor.template packet<AlignmentType>(first);
282 }
283 }
284
285 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketT>::alignment) Scalar data[packet_size];
286
287 data[0] = this->m_tensor.coeff(first);
288 EIGEN_UNROLL_LOOP
289 for (Index k = 1; k < packet_size - 1; k += 2) {
290 const IndexPair<Index> internal_pair = this->computeIndexPair(i + k, j, 1);
291 data[k] = this->m_tensor.coeff(internal_pair.first);
292 data[k + 1] = this->m_tensor.coeff(internal_pair.second);
293 }
294 data[packet_size - 1] = this->m_tensor.coeff(lastIdx);
295
296 return pload<PacketT>(data);
297 }
298
299 template <typename PacketT, int AlignmentType>
300 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
301 std::enable_if_t<internal::unpacket_traits<PacketT>::size != packet_size, PacketT>
302 load(Index i, Index j) const {
303 const Index requested_packet_size = internal::unpacket_traits<PacketT>::size;
304 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketT>::alignment) Scalar data[requested_packet_size];
305
306 const IndexPair<Index> indexPair = this->computeIndexPair(i, j, requested_packet_size - 1);
307 const Index first = indexPair.first;
308 const Index lastIdx = indexPair.second;
309
310 data[0] = this->m_tensor.coeff(first);
311 for (Index k = 1; k < requested_packet_size - 1; k += 2) {
312 const IndexPair<Index> internal_pair = this->computeIndexPair(i + k, j, 1);
313 data[k] = this->m_tensor.coeff(internal_pair.first);
314 data[k + 1] = this->m_tensor.coeff(internal_pair.second);
315 }
316 data[requested_packet_size - 1] = this->m_tensor.coeff(lastIdx);
317
318 return pload<PacketT>(data);
319 }
320
321 template <typename PacketT, int AlignmentType>
322 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketT loadPacket(Index i, Index j) const {
323 return this->load<PacketT, AlignmentType>(i, j);
324 }
325};
326
327template <typename Scalar, typename Index, int side, typename Tensor, typename nocontract_t, typename contract_t,
328 bool inner_dim_contiguous, bool inner_dim_reordered, int Alignment, template <class> class MakePointer_>
329class BaseTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, 1, inner_dim_contiguous,
330 inner_dim_reordered, Alignment, MakePointer_>
331 : public SimpleTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, 1,
332 inner_dim_contiguous, Alignment, MakePointer_> {
333 public:
334 using ParentMapper = SimpleTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, 1,
335 inner_dim_contiguous, Alignment, MakePointer_>;
336
337 EIGEN_DEVICE_FUNC BaseTensorContractionMapper(const Tensor& tensor, const nocontract_t& nocontract_strides,
338 const nocontract_t& ij_strides, const contract_t& contract_strides,
339 const contract_t& k_strides)
340 : ParentMapper(tensor, nocontract_strides, ij_strides, contract_strides, k_strides) {}
341
342 template <typename PacketT, int>
343 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketT loadPacket(Index i, Index j) const {
344 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketT>::alignment) Scalar data[1];
345 data[0] = this->m_tensor.coeff(this->computeIndex(i, j));
346 return pload<PacketT>(data);
347 }
348 template <typename PacketT, int>
349 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketT load(Index i, Index j) const {
350 EIGEN_ALIGN_TO_BOUNDARY(unpacket_traits<PacketT>::alignment) Scalar data[1];
351 data[0] = this->m_tensor.coeff(this->computeIndex(i, j));
352 return pload<PacketT>(data);
353 }
354};
355
356template <typename Scalar, typename Index, int side, typename Tensor, typename nocontract_t, typename contract_t,
357 int packet_size, bool inner_dim_contiguous, bool inner_dim_reordered, int Alignment,
358 template <class> class MakePointer_ = MakePointer>
359class TensorContractionSubMapper {
360 public:
361 using ParentMapper = BaseTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, packet_size,
362 inner_dim_contiguous, inner_dim_reordered, Alignment, MakePointer_>;
363 using Self = TensorContractionSubMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, packet_size,
364 inner_dim_contiguous, inner_dim_reordered, Alignment, MakePointer_>;
365 using LinearMapper = Self;
366 using SubMapper = Self;
367
368 // We can use direct offsets iff the parent mapper supports them and we can compute the strides.
369 // TODO: we should also enable direct offsets for the Rhs case.
370 static constexpr bool UseDirectOffsets =
371 ParentMapper::DirectOffsets && (side == Lhs) && inner_dim_contiguous && (array_size<contract_t>::value > 0);
372
373 EIGEN_DEVICE_FUNC TensorContractionSubMapper(const ParentMapper& base_mapper, Index vert_offset, Index horiz_offset)
374 : m_base_mapper(base_mapper), m_vert_offset(vert_offset), m_horiz_offset(horiz_offset) {
375 // Bake the offsets into the buffer used by the base mapper whenever possible. This avoids the need to recompute
376 // this offset every time we attempt to access a coefficient.
377 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
378 Index stride = m_base_mapper.stride();
379 m_base_mapper.offsetBuffer(vert_offset + horiz_offset * stride);
380 }
381 }
382
383 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Scalar operator()(Index i) const {
384 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
385 return m_base_mapper(i, 0);
386 }
387 return m_base_mapper(i + m_vert_offset, m_horiz_offset);
388 }
389 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE Scalar operator()(Index i, Index j) const {
390 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
391 return m_base_mapper(i, j);
392 }
393 return m_base_mapper(i + m_vert_offset, j + m_horiz_offset);
394 }
395
396 template <typename PacketT>
397 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT loadPacket(Index i) const {
398 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
399 return m_base_mapper.template loadPacket<PacketT, Alignment>(i, 0);
400 }
401 return m_base_mapper.template loadPacket<PacketT, Alignment>(i + m_vert_offset, m_horiz_offset);
402 }
403
404 template <typename PacketT>
405 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT loadPacket(Index i, Index j) const {
406 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
407 return m_base_mapper.template loadPacket<PacketT, Alignment>(i, j);
408 }
409 return m_base_mapper.template loadPacket<PacketT, Alignment>(i + m_vert_offset, j + m_horiz_offset);
410 }
411
412 template <typename PacketT>
413 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT loadPacketPartial(Index i, Index j, Index, Index = 0) const {
414 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
415 return m_base_mapper.template loadPacket<PacketT, Alignment>(i, j);
416 }
417 return m_base_mapper.template loadPacket<PacketT, Alignment>(i + m_vert_offset, j + m_horiz_offset);
418 }
419
420 template <typename PacketT, int AlignmentType>
421 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT loadPacket(Index i, Index j) const {
422 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
423 return m_base_mapper.template load<PacketT, AlignmentType>(i, j);
424 }
425 return m_base_mapper.template loadPacket<PacketT, AlignmentType>(i + m_vert_offset, j + m_horiz_offset);
426 }
427
428 template <typename PacketT>
429 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE void storePacket(Index i, const PacketT& p) const {
430 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
431 m_base_mapper.storePacket(i, 0, p);
432 } else {
433 m_base_mapper.storePacket(i + m_vert_offset, m_horiz_offset, p);
434 }
435 }
436
437 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE LinearMapper getLinearMapper(Index i, Index j) const {
438 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
439 return LinearMapper(m_base_mapper, i, j);
440 }
441 return LinearMapper(m_base_mapper, i + m_vert_offset, j + m_horiz_offset);
442 }
443
444 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE SubMapper getSubMapper(Index i, Index j) const {
445 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
446 return SubMapper(m_base_mapper, i, j);
447 }
448 return SubMapper(m_base_mapper, i + m_vert_offset, j + m_horiz_offset);
449 }
450
451 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE const Index stride() const { return m_base_mapper.stride(); }
452
453 template <typename PacketT, int AlignmentType>
454 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketT load(Index i) const {
455 static_assert(std::is_same<PacketT, PacketT>::value, "YOU_MADE_A_PROGRAMMING_MISTAKE");
456 constexpr int ActualAlignment = (AlignmentType == Aligned) && (Alignment == Aligned) ? Aligned : Unaligned;
457 EIGEN_IF_CONSTEXPR (UseDirectOffsets) {
458 return m_base_mapper.template loadPacket<PacketT, ActualAlignment>(i, 0);
459 }
460 return m_base_mapper.template loadPacket<PacketT, ActualAlignment>(i + m_vert_offset, m_horiz_offset);
461 }
462
463 template <typename PacketT>
464 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE bool aligned(Index) const {
465 return false;
466 }
467
468 const ParentMapper& base_mapper() const { return m_base_mapper; }
469 Index vert_offset() const { return m_vert_offset; }
470 Index horiz_offset() const { return m_horiz_offset; }
471
472 private:
473 ParentMapper m_base_mapper;
474 const Index m_vert_offset;
475 const Index m_horiz_offset;
476};
477
478template <typename Scalar_, typename Index, int side, typename Tensor, typename nocontract_t, typename contract_t,
479 int packet_size, bool inner_dim_contiguous, bool inner_dim_reordered, int Alignment,
480 template <class> class MakePointer_ = MakePointer>
481class TensorContractionInputMapper
482 : public BaseTensorContractionMapper<Scalar_, Index, side, Tensor, nocontract_t, contract_t, packet_size,
483 inner_dim_contiguous, inner_dim_reordered, Alignment, MakePointer_> {
484 public:
485 using Scalar = Scalar_;
486 using Base = BaseTensorContractionMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, packet_size,
487 inner_dim_contiguous, inner_dim_reordered, Alignment, MakePointer_>;
488 using SubMapper = TensorContractionSubMapper<Scalar, Index, side, Tensor, nocontract_t, contract_t, packet_size,
489 inner_dim_contiguous, inner_dim_reordered, Alignment, MakePointer_>;
490 using VectorMapper = SubMapper;
491 using LinearMapper = SubMapper;
492
493 EIGEN_DEVICE_FUNC TensorContractionInputMapper(const Tensor& tensor, const nocontract_t& nocontract_strides,
494 const nocontract_t& ij_strides, const contract_t& contract_strides,
495 const contract_t& k_strides)
496 : Base(tensor, nocontract_strides, ij_strides, contract_strides, k_strides) {}
497
498 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE SubMapper getSubMapper(Index i, Index j) const {
499 return SubMapper(*this, i, j);
500 }
501
502 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE LinearMapper getLinearMapper(Index i, Index j) const {
503 return LinearMapper(*this, i, j);
504 }
505
506 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE VectorMapper getVectorMapper(Index i, Index j) const {
507 return VectorMapper(*this, i, j);
508 }
509
510 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE const CoeffLoader<Tensor, Tensor::RawAccess, MakePointer_>& get_tensor() const {
511 return Base::m_tensor;
512 }
513};
514
515template <typename T>
516struct TensorContractionInputMapperTrait;
517
518template <typename Scalar_, typename Index_, int side_, typename Tensor_, typename nocontract_t_, typename contract_t_,
519 int packet_size_, bool inner_dim_contiguous_, bool inner_dim_reordered_, int Alignment_,
520 template <class> class MakePointer_>
521struct TensorContractionInputMapperTrait<
522 TensorContractionInputMapper<Scalar_, Index_, side_, Tensor_, nocontract_t_, contract_t_, packet_size_,
523 inner_dim_contiguous_, inner_dim_reordered_, Alignment_, MakePointer_> > {
524 using XprType = Tensor_;
525 static constexpr bool inner_dim_contiguous = inner_dim_contiguous_;
526 static constexpr bool inner_dim_reordered = inner_dim_reordered_;
527};
528
529} // end namespace internal
530} // end namespace Eigen
531
532#endif // EIGEN_TENSOR_TENSOR_CONTRACTION_MAPPER_H
The tensor class.
Definition Tensor.h:69
Namespace containing all symbols from the Eigen library.
Definition TensorContractionMapper.h:38