Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorContraction.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_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
21template <typename Dimensions, typename LhsXprType, typename RhsXprType, typename OutputKernelType>
22struct traits<TensorContractionOp<Dimensions, LhsXprType, RhsXprType, OutputKernelType>> {
23 // Type promotion to handle the case where the types of the lhs and the rhs are different.
24 using Scalar = typename gebp_traits<std::remove_const_t<typename LhsXprType::Scalar>,
25 std::remove_const_t<typename RhsXprType::Scalar>>::ResScalar;
26
27 using StorageKind = typename promote_storage_type<typename traits<LhsXprType>::StorageKind,
28 typename traits<RhsXprType>::StorageKind>::ret;
29 using Index =
30 typename promote_index_type<typename traits<LhsXprType>::Index, typename traits<RhsXprType>::Index>::type;
31 // From NumDims below.
32 static constexpr int NumDimensions =
33 traits<LhsXprType>::NumDimensions + traits<RhsXprType>::NumDimensions - 2 * array_size<Dimensions>::value;
34 static constexpr int Layout = traits<LhsXprType>::Layout;
35 using PointerType =
36 std::conditional_t<Pointer_type_promotion<typename LhsXprType::Scalar, Scalar>::val,
37 typename traits<LhsXprType>::PointerType, typename traits<RhsXprType>::PointerType>;
38
39 static constexpr int Flags = 0;
40};
41
42template <typename Dimensions, typename LhsXprType, typename RhsXprType, typename OutputKernelType>
43struct eval<TensorContractionOp<Dimensions, LhsXprType, RhsXprType, OutputKernelType>, Eigen::Dense> {
44 using type = const TensorContractionOp<Dimensions, LhsXprType, RhsXprType, OutputKernelType>&;
45};
46
47template <typename Indices_, typename LeftArgType_, typename RightArgType_, typename OutputKernelType_,
48 typename Device_>
49struct traits<
50 TensorEvaluator<const TensorContractionOp<Indices_, LeftArgType_, RightArgType_, OutputKernelType_>, Device_>> {
51 using Indices = Indices_;
52 using LeftArgType = LeftArgType_;
53 using RightArgType = RightArgType_;
54 using OutputKernelType = OutputKernelType_;
55 using Device = Device_;
56
57 // From NumDims below.
58 static constexpr int NumDimensions =
59 traits<LeftArgType_>::NumDimensions + traits<RightArgType_>::NumDimensions - 2 * array_size<Indices_>::value;
60};
61
62// Helper class to allocate and deallocate temporary memory for packed buffers.
63template <typename LhsScalar, typename RhsScalar>
64struct TensorContractionBlockMemAllocator {
65 using BlockMemHandle = void*;
66
67 template <typename Device>
68 EIGEN_DEVICE_FUNC static BlockMemHandle allocate(Device& d, const Index bm, const Index bk, const Index bn,
69 LhsScalar** lhs_block, RhsScalar** rhs_block) {
70 eigen_assert(lhs_block);
71 eigen_assert(rhs_block);
72 BlockSizes sz = ComputeLhsRhsBlockSizes(bm, bk, bn);
73 char* block_mem = static_cast<char*>(d.allocate(sz.lhs_size + sz.rhs_size));
74 *lhs_block = static_cast<LhsScalar*>(static_cast<void*>(block_mem));
75 *rhs_block = static_cast<RhsScalar*>(static_cast<void*>(block_mem + sz.lhs_size));
76 return block_mem;
77 }
78
79 template <typename Device>
80 EIGEN_DEVICE_FUNC static BlockMemHandle allocateSlices(Device& d, const Index bm, const Index bk, const Index bn,
81 const Index num_lhs, const Index num_rhs,
82 const Index num_slices, std::vector<LhsScalar*>* lhs_blocks,
83 std::vector<RhsScalar*>* rhs_blocks) {
84 eigen_assert(num_slices > 0);
85 eigen_assert(num_lhs >= 0 && num_rhs >= 0);
86 eigen_assert(num_lhs == 0 || lhs_blocks);
87 eigen_assert(num_rhs == 0 || rhs_blocks);
88 BlockSizes sz = ComputeLhsRhsBlockSizes(bm, bk, bn);
89 void* block_mem = d.allocate((num_lhs * sz.lhs_size + num_rhs * sz.rhs_size) * num_slices);
90 eigen_assert(block_mem);
91 char* mem = static_cast<char*>(block_mem);
92
93 for (Index x = 0; x < num_slices; x++) {
94 if (num_lhs > 0) lhs_blocks[x].resize(num_lhs);
95 for (Index m = 0; m < num_lhs; m++) {
96 lhs_blocks[x][m] = static_cast<LhsScalar*>(static_cast<void*>(mem));
97 mem += sz.lhs_size;
98 }
99 if (num_rhs > 0) rhs_blocks[x].resize(num_rhs);
100 for (Index n = 0; n < num_rhs; n++) {
101 rhs_blocks[x][n] = static_cast<RhsScalar*>(static_cast<void*>(mem));
102 mem += sz.rhs_size;
103 }
104 }
105
106 return block_mem;
107 }
108
109 template <typename Device>
110 EIGEN_DEVICE_FUNC static void deallocate(Device& d, BlockMemHandle handle) {
111 d.deallocate(handle);
112 }
113
114 private:
115 struct BlockSizes {
116 Index lhs_size;
117 Index rhs_size;
118 };
119 EIGEN_DEVICE_FUNC static BlockSizes ComputeLhsRhsBlockSizes(const Index bm, const Index bk, const Index bn) {
120 Index align = numext::maxi(EIGEN_MAX_ALIGN_BYTES, 1);
121 BlockSizes sz;
122 sz.lhs_size = numext::div_ceil<Index>(bm * bk * sizeof(LhsScalar), align) * align;
123 sz.rhs_size = numext::div_ceil<Index>(bn * bk * sizeof(RhsScalar), align) * align;
124 return sz;
125 }
126};
127
128// WARNING: In this code we assume that Lhs and Rhs tensor expressions are in
129// ColMajor storage order. This property is guaranteed by the
130// TensorContractionOp evaluator. TensorContractionKernel specifies how we pack
131// blocks of Lhs and Rhs tensor expressions, and how we invoke matrix
132// multiplication for these blocks. Default tensor contraction uses
133// gemm_pack_rhs, gemm_pack_lhs and gebp_kernel from Eigen Core (see
134// GeneralBlockPanelKernel.h for details).
135//
136// By specializing contraction kernels we can use other low level libraries to
137// perform matrix multiplication, and still rely on Eigen contraction evaluator.
138// This also includes full support in TensorContractionThreadPool, assuming that
139// underlying gemm does not use its own threading.
140//
141// - ResScalar/LhsScalar/RhsScalar - scalar type for the result of
142// multiplication, lhs tensor and rhs tensor respectively.
143//
144// - StorageIndex - index type for the tensor expressions. In practice almost
145// always is Eigen::Index.
146//
147// - OutputMapper provides access to the memory of the output matrix. In
148// practice it's always column major blas_data_mapper (it must be of ResScalar
149// type).
150//
151// - LhsMapper/RhsMapper similarly to blas_data_mapper provide a two dimensional
152// view into the Lhs/Rhs tensor expressions. In practice it's
153// TensorContractionInputMapper, or some specialization of it based on the
154// type of tensor expression (e.g. TensorImagePatchOp has optimized input
155// mapper).
156template <typename ResScalar, typename LhsScalar, typename RhsScalar, typename StorageIndex, typename OutputMapper,
157 typename LhsMapper, typename RhsMapper>
158struct TensorContractionKernel {
159 // True if `invoke()` supports `beta` in `C <- alpha * A * B + beta * C`
160 // (otherwise beta should be always equal to 1).
161 static constexpr bool HasBeta = false;
162
163 EIGEN_DEVICE_FUNC TensorContractionKernel(StorageIndex m_, StorageIndex k_, StorageIndex n_, StorageIndex bm_,
164 StorageIndex bk_, StorageIndex bn_)
165 : m(m_), k(k_), n(n_), bm(bm_), bk(bk_), bn(bn_) {}
166
167 // Pack blocks of Lhs and Rhs into contiguous blocks in memory.
168 using LhsBlock = LhsScalar*;
169 using RhsBlock = RhsScalar*;
170
171 // Packed Lhs/Rhs block memory allocator.
172 using BlockMemAllocator = TensorContractionBlockMemAllocator<LhsScalar, RhsScalar>;
173 using BlockMemHandle = typename BlockMemAllocator::BlockMemHandle;
174
175 using Traits = typename internal::gebp_traits<LhsScalar, RhsScalar>;
176
177 using LhsPacker = internal::gemm_pack_lhs<LhsScalar, StorageIndex, typename LhsMapper::SubMapper, Traits::mr,
178 Traits::LhsProgress, typename Traits::LhsPacket4Packing, ColMajor>;
179
180 using RhsPacker =
181 internal::gemm_pack_rhs<RhsScalar, StorageIndex, typename RhsMapper::SubMapper, Traits::nr, ColMajor>;
182
183 using GebpKernel = internal::gebp_kernel<LhsScalar, RhsScalar, StorageIndex, OutputMapper, Traits::mr, Traits::nr,
184 /*ConjugateLhs*/ false, /*ConjugateRhs*/ false>;
185
186 template <typename Device>
187 EIGEN_DEVICE_FUNC BlockMemHandle allocate(Device& d, LhsBlock* lhs_block, RhsBlock* rhs_block) {
188 return BlockMemAllocator::allocate(d, bm, bk, bn, lhs_block, rhs_block);
189 }
190
191 template <typename Device>
192 EIGEN_DEVICE_FUNC BlockMemHandle allocateSlices(Device& d, const StorageIndex num_lhs, const StorageIndex num_rhs,
193 const StorageIndex num_slices, std::vector<LhsBlock>* lhs_blocks,
194 std::vector<RhsBlock>* rhs_blocks) {
195 return BlockMemAllocator::allocateSlices(d, bm, bk, bn, num_lhs, num_rhs, num_slices, lhs_blocks, rhs_blocks);
196 }
197
198 template <typename Device>
199 EIGEN_DEVICE_FUNC static void deallocate(Device& d, BlockMemHandle handle) {
200 BlockMemAllocator::deallocate(d, handle);
201 }
202
203 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE void packLhs(LhsBlock* lhsBlock, const typename LhsMapper::SubMapper& data_mapper,
204 const StorageIndex depth, const StorageIndex rows) {
205 LhsPacker()(*lhsBlock, data_mapper, depth, rows, /*stride*/ 0,
206 /*offset*/ 0);
207 }
208
209 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE void packRhs(RhsBlock* rhsBlock, const typename RhsMapper::SubMapper& data_mapper,
210 const StorageIndex depth, const StorageIndex cols) {
211 RhsPacker()(*rhsBlock, data_mapper, depth, cols);
212 }
213
214 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE void invoke(const OutputMapper& output_mapper, const LhsBlock& lhsBlock,
215 const RhsBlock& rhsBlock, const StorageIndex rows,
216 const StorageIndex depth, const StorageIndex cols,
217 const ResScalar alpha, const ResScalar beta) {
218 // Default GEBP kernel does not support beta.
219 EIGEN_ONLY_USED_FOR_DEBUG(beta);
220 eigen_assert(beta == ResScalar(1));
221 static constexpr int kComputeStrideFromBlockDimensions = -1;
222 GebpKernel()(output_mapper, lhsBlock, rhsBlock, rows, depth, cols, alpha,
223 /*strideA*/ kComputeStrideFromBlockDimensions,
224 /*strideB*/ kComputeStrideFromBlockDimensions,
225 /*offsetA*/ 0, /*offsetB*/ 0);
226 }
227
228 private:
229 // These are dimensions of the original Tensors, and selected block sizes. The
230 // actual block sizes passed to all function above might be smaller because of
231 // the partial blocks at the end.
232 const StorageIndex m;
233 const StorageIndex k;
234 const StorageIndex n;
235 const StorageIndex bm;
236 const StorageIndex bk;
237 const StorageIndex bn;
238};
239
240// Dispatches a contraction operation over all 8 combinations of the three
241// runtime boolean parameters (lhs_inner_dim_contiguous, rhs_inner_dim_contiguous,
242// rhs_inner_dim_reordered), passing them as compile-time bool_constant
243// tags to the callable `fn`.
244template <typename Func>
245EIGEN_STRONG_INLINE void tensor_contraction_dispatch(Func&& fn, bool lhs_inner_dim_contiguous,
246 bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered) {
247 if (lhs_inner_dim_contiguous) {
248 if (rhs_inner_dim_contiguous) {
249 if (rhs_inner_dim_reordered)
250 fn(bool_constant<true>{}, bool_constant<true>{}, bool_constant<true>{});
251 else
252 fn(bool_constant<true>{}, bool_constant<true>{}, bool_constant<false>{});
253 } else {
254 if (rhs_inner_dim_reordered)
255 fn(bool_constant<true>{}, bool_constant<false>{}, bool_constant<true>{});
256 else
257 fn(bool_constant<true>{}, bool_constant<false>{}, bool_constant<false>{});
258 }
259 } else {
260 if (rhs_inner_dim_contiguous) {
261 if (rhs_inner_dim_reordered)
262 fn(bool_constant<false>{}, bool_constant<true>{}, bool_constant<true>{});
263 else
264 fn(bool_constant<false>{}, bool_constant<true>{}, bool_constant<false>{});
265 } else {
266 if (rhs_inner_dim_reordered)
267 fn(bool_constant<false>{}, bool_constant<false>{}, bool_constant<true>{});
268 else
269 fn(bool_constant<false>{}, bool_constant<false>{}, bool_constant<false>{});
270 }
271 }
272}
273
274} // end namespace internal
275
276// Legacy macros kept for backward compatibility with code that overrides them
277// (e.g. TensorFlow Lite restricts template instantiations for binary size).
278// New Eigen code should use internal::tensor_contraction_dispatch() instead.
279#ifndef TENSOR_CONTRACTION_DISPATCH
280#define TENSOR_CONTRACTION_DISPATCH(METHOD, ALIGNMENT, ARGS) \
281 ::Eigen::internal::tensor_contraction_dispatch( \
282 [&](auto lhs_c, auto rhs_c, auto rhs_r) { METHOD<lhs_c(), rhs_c(), rhs_r(), ALIGNMENT> ARGS; }, \
283 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered)
284#endif
285
286#ifndef TENSOR_CONTRACTION_ASYNC_DISPATCH
287#define TENSOR_CONTRACTION_ASYNC_DISPATCH(METHOD, DONE, ALIGNMENT, ARGS, FN) \
288 ::Eigen::internal::tensor_contraction_dispatch( \
289 [&](auto lhs_c, auto rhs_c, auto rhs_r) { (new METHOD<DONE, lhs_c(), rhs_c(), rhs_r(), ALIGNMENT> ARGS)->FN; }, \
290 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered)
291#endif
292
293// Tensor contraction params that should enable to get from output matrix
294// 2-dimensional coordinates to the output tensor dimensions.
295struct TensorContractionParams {
296 // TensorContraction evaluator assumes that both tensors are in ColMajor
297 // layout, if tensors are in RowMajor evaluator swap lhs with rhs.
298 bool swapped_arguments;
299};
300
301// Output kernel allows to fuse operations into the tensor contraction.
302//
303// Examples:
304// 1. Elementwise Relu transformation following Conv2D.
305// 2. AddBias to the Conv2D output channels dimension.
306//
307// The NoOpOutputKernel implements an output kernel that does absolutely nothing.
308struct NoOpOutputKernel {
324 template <typename Index, typename Scalar>
325 EIGEN_ALWAYS_INLINE void operator()(const internal::blas_data_mapper<Scalar, Index, ColMajor>&,
326 const TensorContractionParams&, Index, Index, Index, Index) const {}
327};
328
332template <typename Indices, typename LhsXprType, typename RhsXprType,
333 typename OutputKernelType = const NoOpOutputKernel>
334class TensorContractionOp
335 : public TensorBase<TensorContractionOp<Indices, LhsXprType, RhsXprType, OutputKernelType>, ReadOnlyAccessors> {
336 public:
337 using Scalar = typename Eigen::internal::traits<TensorContractionOp>::Scalar;
338 using CoeffReturnType = typename internal::gebp_traits<typename LhsXprType::CoeffReturnType,
339 typename RhsXprType::CoeffReturnType>::ResScalar;
340 using Nested = typename Eigen::internal::ref_selector<TensorContractionOp>::type;
341 using StorageKind = typename Eigen::internal::traits<TensorContractionOp>::StorageKind;
342 using Index = typename Eigen::internal::traits<TensorContractionOp>::Index;
343
344 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorContractionOp(const LhsXprType& lhs, const RhsXprType& rhs,
345 const Indices& dims,
346 const OutputKernelType& output_kernel = OutputKernelType())
347 : m_lhs_xpr(lhs), m_rhs_xpr(rhs), m_indices(dims), m_output_kernel(output_kernel) {}
348
349 EIGEN_DEVICE_FUNC const Indices& indices() const { return m_indices; }
350
352 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename LhsXprType::Nested>& lhsExpression() const {
353 return m_lhs_xpr;
354 }
355
356 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression() const {
357 return m_rhs_xpr;
358 }
359
360 EIGEN_DEVICE_FUNC const OutputKernelType& outputKernel() const { return m_output_kernel; }
361
362 protected:
363 typename LhsXprType::Nested m_lhs_xpr;
364 typename RhsXprType::Nested m_rhs_xpr;
365 const Indices m_indices;
366 const OutputKernelType m_output_kernel;
367};
368
369namespace internal {
370template <bool TypesMatch, int StorageOrder, bool MatrixIsRight>
371struct GemvDirectDispatcher {
372 template <typename Evaluator, typename Scalar>
373 static bool run(const Evaluator* self, Scalar* buffer) {
374 self->template evalGemvDirect<StorageOrder, MatrixIsRight>(buffer);
375 return true;
376 }
377};
378
379template <int StorageOrder, bool MatrixIsRight>
380struct GemvDirectDispatcher<false, StorageOrder, MatrixIsRight> {
381 template <typename Evaluator, typename Scalar>
382 static bool run(const Evaluator*, Scalar*) {
383 return false;
384 }
385};
386} // namespace internal
387
388template <typename Derived>
389struct TensorContractionEvaluatorBase {
390 using Indices = typename internal::traits<Derived>::Indices;
391 using LeftArgType = typename internal::traits<Derived>::LeftArgType;
392 using RightArgType = typename internal::traits<Derived>::RightArgType;
393 using OutputKernelType = typename internal::traits<Derived>::OutputKernelType;
394 using Device = typename internal::traits<Derived>::Device;
395
396 using XprType = TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>;
397 using Scalar = std::remove_const_t<typename XprType::Scalar>;
398 using Index = typename XprType::Index;
399 using CoeffReturnType = typename XprType::CoeffReturnType;
400 using PacketReturnType = typename PacketType<CoeffReturnType, Device>::type;
401 using Storage = StorageMemory<Scalar, Device>;
402 using EvaluatorPointerType = typename Storage::Type;
403
404 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
405 static constexpr bool IsAligned = true;
406 static constexpr bool PacketAccess = (PacketType<CoeffReturnType, Device>::size > 1);
407 static constexpr bool BlockAccess = false;
408 static constexpr bool PreferBlockAccess = false;
409 static constexpr bool CoordAccess = false; // to be implemented
410 static constexpr bool RawAccess = true;
411
412 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
413 using TensorBlock = internal::TensorBlockNotImplemented;
414 //===--------------------------------------------------------------------===//
415
416 // Most of the code is assuming that both input tensors are ColMajor. If the
417 // inputs are RowMajor, we will "cheat" by swapping the LHS and RHS:
418 // If we want to compute A * B = C, where A is LHS and B is RHS, the code
419 // will pretend B is LHS and A is RHS.
420 using EvalLeftArgType =
421 std::conditional_t<static_cast<int>(Layout) == static_cast<int>(ColMajor), LeftArgType, RightArgType>;
422 using EvalRightArgType =
423 std::conditional_t<static_cast<int>(Layout) == static_cast<int>(ColMajor), RightArgType, LeftArgType>;
424
425 static constexpr int LDims =
426 internal::array_size<typename TensorEvaluator<EvalLeftArgType, Device>::Dimensions>::value;
427 static constexpr int RDims =
428 internal::array_size<typename TensorEvaluator<EvalRightArgType, Device>::Dimensions>::value;
429 static constexpr int ContractDims = internal::array_size<Indices>::value;
430 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
431
432 using contract_t = array<Index, ContractDims>;
433 using left_nocontract_t = array<Index, LDims - ContractDims>;
434 using right_nocontract_t = array<Index, RDims - ContractDims>;
435
436 using Dimensions = DSizes<Index, NumDims>;
437
438 EIGEN_STRONG_INLINE TensorContractionEvaluatorBase(const XprType& op, const Device& device)
439 : m_leftImpl(choose(Cond<static_cast<int>(Layout) == static_cast<int>(ColMajor)>(), op.lhsExpression(),
440 op.rhsExpression()),
441 device),
442 m_rightImpl(choose(Cond<static_cast<int>(Layout) == static_cast<int>(ColMajor)>(), op.rhsExpression(),
443 op.lhsExpression()),
444 device),
445 m_device(device),
446 m_output_kernel(op.outputKernel()),
447 m_result(nullptr) {
448 EIGEN_STATIC_ASSERT((static_cast<int>(TensorEvaluator<LeftArgType, Device>::Layout) ==
449 static_cast<int>(TensorEvaluator<RightArgType, Device>::Layout)),
450 YOU_MADE_A_PROGRAMMING_MISTAKE);
451
452 DSizes<Index, LDims> eval_left_dims;
453 DSizes<Index, RDims> eval_right_dims;
454 array<IndexPair<Index>, ContractDims> eval_op_indices;
455 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
456 // For ColMajor, we keep using the existing dimensions
457 for (int i = 0; i < LDims; i++) {
458 eval_left_dims[i] = m_leftImpl.dimensions()[i];
459 }
460 for (int i = 0; i < RDims; i++) {
461 eval_right_dims[i] = m_rightImpl.dimensions()[i];
462 }
463 // We keep the pairs of contracting indices.
464 for (int i = 0; i < ContractDims; i++) {
465 eval_op_indices[i].first = op.indices()[i].first;
466 eval_op_indices[i].second = op.indices()[i].second;
467 }
468 } else {
469 // For RowMajor, we need to reverse the existing dimensions
470 for (int i = 0; i < LDims; i++) {
471 eval_left_dims[i] = m_leftImpl.dimensions()[LDims - i - 1];
472 }
473 for (int i = 0; i < RDims; i++) {
474 eval_right_dims[i] = m_rightImpl.dimensions()[RDims - i - 1];
475 }
476 // We need to flip all the pairs of contracting indices as well as
477 // reversing the dimensions.
478 for (int i = 0; i < ContractDims; i++) {
479 eval_op_indices[i].first = LDims - 1 - op.indices()[ContractDims - 1 - i].second;
480 eval_op_indices[i].second = RDims - 1 - op.indices()[ContractDims - 1 - i].first;
481 }
482 }
483
484 // Check for duplicate axes and make sure the first index in eval_op_indices
485 // is increasing. Using O(n^2) sorting is OK since ContractDims is small
486 for (int i = 0; i < ContractDims; i++) {
487 for (int j = i + 1; j < ContractDims; j++) {
488 eigen_assert(eval_op_indices[j].first != eval_op_indices[i].first &&
489 eval_op_indices[j].second != eval_op_indices[i].second && "contraction axes should be unique");
490 if (eval_op_indices[j].first < eval_op_indices[i].first) {
491 numext::swap(eval_op_indices[j], eval_op_indices[i]);
492 }
493 }
494 }
495
496 array<Index, LDims> lhs_strides;
497 lhs_strides[0] = 1;
498 for (int i = 0; i < LDims - 1; ++i) {
499 lhs_strides[i + 1] = lhs_strides[i] * eval_left_dims[i];
500 }
501
502 array<Index, RDims> rhs_strides;
503 rhs_strides[0] = 1;
504 for (int i = 0; i < RDims - 1; ++i) {
505 rhs_strides[i + 1] = rhs_strides[i] * eval_right_dims[i];
506 }
507
508 m_i_size = 1;
509 m_j_size = 1;
510 m_k_size = 1;
511
512 // To compute the dimension, we simply concatenate the non-contracting
513 // dimensions of the left and then the right tensor. Additionally, we also
514 // compute the strides corresponding to the left non-contracting
515 // dimensions and right non-contracting dimensions.
516 m_lhs_inner_dim_contiguous = true;
517 int dim_idx = 0;
518 Index nocontract_idx = 0;
519
520 for (int i = 0; i < LDims; i++) {
521 // find if we are contracting on index i of left tensor
522 bool contracting = false;
523 for (int j = 0; j < ContractDims; j++) {
524 if (eval_op_indices[j].first == i) {
525 contracting = true;
526 break;
527 }
528 }
529 if (!contracting) {
530 // add dimension size to output dimensions
531 m_dimensions[dim_idx] = eval_left_dims[i];
532 m_left_nocontract_strides[nocontract_idx] = lhs_strides[i];
533 if (dim_idx != i) {
534 m_lhs_inner_dim_contiguous = false;
535 }
536 m_i_strides[nocontract_idx] = m_i_size;
537 m_i_size *= eval_left_dims[i];
538 dim_idx++;
539 nocontract_idx++;
540 }
541 }
542
543 nocontract_idx = 0;
544 for (int i = 0; i < RDims; i++) {
545 bool contracting = false;
546 // find if we are contracting on index i of right tensor
547 for (int j = 0; j < ContractDims; j++) {
548 if (eval_op_indices[j].second == i) {
549 contracting = true;
550 break;
551 }
552 }
553 if (!contracting) {
554 m_dimensions[dim_idx] = eval_right_dims[i];
555 m_j_strides[nocontract_idx] = m_j_size;
556 m_j_size *= eval_right_dims[i];
557 m_right_nocontract_strides[nocontract_idx] = rhs_strides[i];
558 dim_idx++;
559 nocontract_idx++;
560 }
561 }
562
563 // Now compute the strides corresponding to the contracting dimensions. We
564 // assumed above that non-contracting axes are represented in the same order
565 // in the matrix as they are in the tensor. This is not the case for
566 // contracting axes. As the contracting axes must be of the same size in
567 // each tensor, we'll only look at the first tensor here.
568 m_rhs_inner_dim_contiguous = true;
569 m_rhs_inner_dim_reordered = false;
570 for (int i = 0; i < ContractDims; i++) {
571 Index left = eval_op_indices[i].first;
572 Index right = eval_op_indices[i].second;
573
574 Index size = eval_left_dims[left];
575 eigen_assert(size == eval_right_dims[right] && "Contraction axes must be same size");
576
577 m_k_strides[i] = m_k_size;
578 m_k_size *= size;
579 m_left_contracting_strides[i] = lhs_strides[left];
580 m_right_contracting_strides[i] = rhs_strides[right];
581
582 if (i > 0 && right < eval_op_indices[i - 1].second) {
583 m_rhs_inner_dim_reordered = true;
584 }
585 if (right != i) {
586 m_rhs_inner_dim_contiguous = false;
587 }
588 }
589
590 // After the sort above, eval_op_indices[*].first is in ascending order. The
591 // contracted LHS dims are the contiguous prefix of storage iff that prefix
592 // is exactly {0, 1, ..., ContractDims-1}. Similarly for the RHS — and we
593 // also detect the "contracted dims at the trailing end" case, which
594 // matters when LHS/RHS were swapped at the type level for RowMajor inputs.
595 m_lhs_contracted_dims_leading = true;
596 m_rhs_contracted_dims_leading = true;
597 m_rhs_contracted_dims_trailing = true;
598 const int rhs_trail_start = RDims - ContractDims;
599 for (int i = 0; i < ContractDims; i++) {
600 if (eval_op_indices[i].first != i) m_lhs_contracted_dims_leading = false;
601 if (eval_op_indices[i].second != i) m_rhs_contracted_dims_leading = false;
602 if (eval_op_indices[i].second != rhs_trail_start + i) m_rhs_contracted_dims_trailing = false;
603 }
604
605 // If the layout is RowMajor, we need to reverse the m_dimensions
606 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(RowMajor)) {
607 for (int i = 0, j = NumDims - 1; i < j; i++, j--) {
608 numext::swap(m_dimensions[i], m_dimensions[j]);
609 }
610 }
611
612 // A set of parameters that will allow output kernel to get from output
613 // tensor dimensions (i, j) into the original tensor dimensions.
614 // TODO(ezhulenev): Add parameters required to infer output tensor index for
615 // more complex contractions than 2x2 on internal dimension.
616 m_tensor_contraction_params.swapped_arguments = static_cast<int>(Layout) == RowMajor;
617 }
618
619 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_dimensions; }
620
621 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType data) {
622 m_leftImpl.evalSubExprsIfNeeded(nullptr);
623 m_rightImpl.evalSubExprsIfNeeded(nullptr);
624 if (data) {
625 evalTo(data);
626 return false;
627 } else {
628 m_result = static_cast<EvaluatorPointerType>(m_device.allocate(dimensions().TotalSize() * sizeof(Scalar)));
629 evalTo(m_result);
630 return true;
631 }
632 }
633
634#ifdef EIGEN_USE_THREADS
635 template <typename EvalSubExprsCallback>
636 EIGEN_STRONG_INLINE void evalSubExprsIfNeededAsync(EvaluatorPointerType dest, EvalSubExprsCallback done) {
637 m_leftImpl.evalSubExprsIfNeededAsync(nullptr, [this, done, dest](bool) {
638 m_rightImpl.evalSubExprsIfNeededAsync(nullptr, [this, done, dest](bool) {
639 if (dest) {
640 evalToAsync(dest, [done]() { done(false); });
641 } else {
642 m_result = static_cast<EvaluatorPointerType>(m_device.allocate(dimensions().TotalSize() * sizeof(Scalar)));
643 evalToAsync(m_result, [done]() { done(true); });
644 }
645 });
646 });
647 }
648#endif // EIGEN_USE_THREADS
649
650 EIGEN_DEVICE_FUNC void evalTo(Scalar* buffer) const {
651 static_cast<const Derived*>(this)->template evalProduct<Unaligned>(buffer);
652 }
653
654#ifdef EIGEN_USE_THREADS
655 template <typename EvalToCallback>
656 void evalToAsync(Scalar* buffer, EvalToCallback done) const {
657 static_cast<const Derived*>(this)->template evalProductAsync<EvalToCallback, Unaligned>(buffer, std::move(done));
658 }
659#endif // EIGEN_USE_THREADS
660
661 template <bool lhs_inner_dim_contiguous, bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered, int Alignment>
662 void evalProductSequential(Scalar* buffer) const {
663 if (this->m_i_size == 0 || this->m_j_size == 0) {
664 return;
665 }
666 if (this->m_k_size == 0) {
667 // Contraction over dimension of size zero results in an all-zero output.
668 this->m_device.fill(buffer, buffer + this->m_i_size * this->m_j_size, Scalar(0));
669 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
670 this->m_output_kernel(OutputMapper(buffer, this->m_i_size), this->m_tensor_contraction_params,
671 static_cast<Index>(0), static_cast<Index>(0), this->m_i_size, this->m_j_size);
672 return;
673 }
674 // Gemv-shape contractions (output is a vector) get a direct GEMV kernel
675 // call when the operand layout admits one. There are four shapes:
676 //
677 // m_j_size == 1: ColMajor inputs. m_leftImpl is the matrix.
678 // (A) lhs_inner_dim_contiguous: contracted dims at non-contiguous end
679 // of LHS. Existing evalGemv via TensorContractionInputMapper.
680 // (B) m_lhs_contracted_dims_leading: contracted dims at contiguous end.
681 // RowMajor view of LHS via const_blas_data_mapper.
682 //
683 // m_i_size == 1: RowMajor inputs (lines 431-434 swap LHS/RHS at the type
684 // level, so the matrix is m_rightImpl).
685 // (C) m_rhs_contracted_dims_leading: contracted dims at contiguous end
686 // of (eval-)RHS. RowMajor view of m_rightImpl.
687 // (D) m_rhs_contracted_dims_trailing: contracted dims at non-contiguous
688 // end. ColMajor view of m_rightImpl.
689 //
690 // Cases B/C/D require direct memory (data() != nullptr) on both impls.
691 // Anything else falls back to evalGemm.
692 constexpr bool types_match = std::is_same<typename EvalLeftArgType::Scalar, Scalar>::value &&
693 std::is_same<typename EvalRightArgType::Scalar, Scalar>::value;
694 if (this->m_j_size == 1) {
695 if (lhs_inner_dim_contiguous) {
696 this->template evalGemv<lhs_inner_dim_contiguous, rhs_inner_dim_contiguous, rhs_inner_dim_reordered, Alignment>(
697 buffer);
698 return;
699 }
700 EIGEN_IF_CONSTEXPR (types_match) {
701 if (m_lhs_contracted_dims_leading && m_leftImpl.data() != nullptr && m_rightImpl.data() != nullptr) {
702 if (internal::GemvDirectDispatcher<types_match, RowMajor, false>::run(this, buffer)) {
703 return;
704 }
705 }
706 }
707 } else if (this->m_i_size == 1 && m_leftImpl.data() != nullptr && m_rightImpl.data() != nullptr) {
708 EIGEN_IF_CONSTEXPR (types_match) {
709 if (m_rhs_contracted_dims_leading) {
710 if (internal::GemvDirectDispatcher<types_match, RowMajor, true>::run(this, buffer)) {
711 return;
712 }
713 }
714 if (m_rhs_contracted_dims_trailing) {
715 if (internal::GemvDirectDispatcher<types_match, ColMajor, true>::run(this, buffer)) {
716 return;
717 }
718 }
719 }
720 }
721 this->template evalGemm<lhs_inner_dim_contiguous, rhs_inner_dim_contiguous, rhs_inner_dim_reordered, Alignment>(
722 buffer);
723 }
724
725 template <bool lhs_inner_dim_contiguous, bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered, int Alignment>
726#if !defined(EIGEN_HIPCC)
727 EIGEN_DEVICE_FUNC
728#endif
729 void
730 evalGemv(Scalar* buffer) const {
731 const Index rows = m_i_size;
732 const Index cols = m_k_size;
733
734 using LhsScalar = std::remove_const_t<typename EvalLeftArgType::Scalar>;
735 using RhsScalar = std::remove_const_t<typename EvalRightArgType::Scalar>;
736 using LeftEvaluator = TensorEvaluator<EvalLeftArgType, Device>;
737 using RightEvaluator = TensorEvaluator<EvalRightArgType, Device>;
738 const int lhs_packet_size = internal::unpacket_traits<typename LeftEvaluator::PacketReturnType>::size;
739 const int rhs_packet_size = internal::unpacket_traits<typename RightEvaluator::PacketReturnType>::size;
740 constexpr int lhs_alignment = LeftEvaluator::IsAligned ? Aligned : Unaligned;
741 constexpr int rhs_alignment = RightEvaluator::IsAligned ? Aligned : Unaligned;
742 using LhsMapper = internal::TensorContractionInputMapper<LhsScalar, Index, internal::Lhs, LeftEvaluator,
743 left_nocontract_t, contract_t, lhs_packet_size,
744 lhs_inner_dim_contiguous, false, lhs_alignment>;
745
746 using RhsMapper =
747 internal::TensorContractionInputMapper<RhsScalar, Index, internal::Rhs, RightEvaluator, right_nocontract_t,
748 contract_t, rhs_packet_size, rhs_inner_dim_contiguous,
749 rhs_inner_dim_reordered, rhs_alignment>;
750
751 LhsMapper lhs(m_leftImpl, m_left_nocontract_strides, m_i_strides, m_left_contracting_strides, m_k_strides);
752 RhsMapper rhs(m_rightImpl, m_right_nocontract_strides, m_j_strides, m_right_contracting_strides, m_k_strides);
753
754 const Scalar alpha(1);
755 const Index resIncr(1);
756
757 // zero out the result buffer (which must be of size at least rows * sizeof(Scalar)
758 m_device.fill(buffer, buffer + rows, Scalar(0));
759
760 internal::general_matrix_vector_product<Index, LhsScalar, LhsMapper, ColMajor, false, RhsScalar, RhsMapper,
761 false>::run(rows, cols, lhs, rhs, buffer, resIncr, alpha);
762
763 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
764 m_output_kernel(OutputMapper(buffer, rows), m_tensor_contraction_params, static_cast<Index>(0),
765 static_cast<Index>(0), rows, static_cast<Index>(1));
766 }
767
768 // Direct gemv fast path: drives the GEMV kernel against raw tensor memory via
769 // const_blas_data_mapper, bypassing TensorContractionInputMapper (which only
770 // delivers vectorized loads along the free dim) and the GEBP packing overhead
771 // of evalGemm.
772 //
773 // StorageOrder is the layout of the (rows × cols) matrix view in memory.
774 // MatrixIsRight selects which TensorEvaluator holds the matrix and which
775 // holds the vector — needed because the constructor swaps LHS/RHS for
776 // RowMajor inputs (lines 431-434), so the gemv shape lives in m_rightImpl
777 // there. Caller guarantees both impls expose direct memory.
778 template <int StorageOrder, bool MatrixIsRight>
779#if !defined(EIGEN_HIPCC)
780 EIGEN_DEVICE_FUNC
781#endif
782 void
783 evalGemvDirect(Scalar* buffer) const {
784 using MatScalar = std::remove_const_t<
785 std::conditional_t<MatrixIsRight, typename EvalRightArgType::Scalar, typename EvalLeftArgType::Scalar>>;
786 using VecScalar = std::remove_const_t<
787 std::conditional_t<MatrixIsRight, typename EvalLeftArgType::Scalar, typename EvalRightArgType::Scalar>>;
788
789 const Index rows = MatrixIsRight ? m_j_size : m_i_size;
790 const Index cols = m_k_size;
791 // For a row-major (rows × cols) view the row stride is cols; for a
792 // column-major view the column stride is rows.
793 const Index mat_stride = (StorageOrder == RowMajor) ? cols : rows;
794
795 const MatScalar* mat_data = MatrixIsRight ? m_rightImpl.data() : m_leftImpl.data();
796 const VecScalar* vec_data = MatrixIsRight ? m_leftImpl.data() : m_rightImpl.data();
797 eigen_assert(mat_data != nullptr && vec_data != nullptr);
798
799 using LhsMapper = internal::const_blas_data_mapper<MatScalar, Index, StorageOrder>;
800 using RhsMapper = internal::const_blas_data_mapper<VecScalar, Index, ColMajor>;
801 LhsMapper lhs(mat_data, mat_stride);
802 RhsMapper rhs(vec_data, /*stride=*/1);
803
804 m_device.fill(buffer, buffer + rows, Scalar(0));
805
806 internal::general_matrix_vector_product<Index, MatScalar, LhsMapper, StorageOrder, false, VecScalar, RhsMapper,
807 false>::run(rows, cols, lhs, rhs, buffer, /*resIncr=*/1, Scalar(1));
808
809 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
810 m_output_kernel(OutputMapper(buffer, rows), m_tensor_contraction_params, static_cast<Index>(0),
811 static_cast<Index>(0), rows, static_cast<Index>(1));
812 }
813
814 template <bool lhs_inner_dim_contiguous, bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered, int Alignment>
815#if !defined(EIGEN_HIPCC)
816 EIGEN_DEVICE_FUNC
817#endif
818 void
819 evalGemm(Scalar* buffer) const {
820 // columns in left side, rows in right side
821 const Index k = this->m_k_size;
822 this->template evalGemmPartial<lhs_inner_dim_contiguous, rhs_inner_dim_contiguous, rhs_inner_dim_reordered,
823 Alignment, true>(buffer, 0, k, 1);
824 }
825
826 template <bool lhs_inner_dim_contiguous, bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered, int Alignment>
827 EIGEN_DEVICE_FUNC void evalGemmPartialWithoutOutputKernel(Scalar* buffer, Index k_start, Index k_end,
828 int num_threads) const {
829 evalGemmPartial<lhs_inner_dim_contiguous, rhs_inner_dim_contiguous, rhs_inner_dim_reordered, Alignment,
830 /*use_output_kernel*/ false>(buffer, k_start, k_end, num_threads);
831 }
832
833 template <bool lhs_inner_dim_contiguous, bool rhs_inner_dim_contiguous, bool rhs_inner_dim_reordered, int Alignment,
834 bool use_output_kernel>
835 EIGEN_DEVICE_FUNC void evalGemmPartial(Scalar* buffer, Index k_start, Index k_end, int num_threads) const {
836 eigen_assert(k_end >= k_start && k_start >= 0 && k_end <= this->m_k_size);
837 // columns in slice on left side, rows on right side
838 const Index k_slice = k_end - k_start;
839
840 // rows in left side
841 const Index m = this->m_i_size;
842
843 // columns in right side
844 const Index n = this->m_j_size;
845
846 if (m == 0 || n == 0) return;
847 if (k_slice == 0) {
848 this->m_device.fill(buffer, buffer + m * n, Scalar(0));
849 if (use_output_kernel) {
850 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
851 m_output_kernel(OutputMapper(buffer, m), m_tensor_contraction_params, static_cast<Index>(0),
852 static_cast<Index>(0), m, n);
853 }
854 return;
855 }
856
857 // define data mappers for Lhs and Rhs
858 using LhsScalar = std::remove_const_t<typename EvalLeftArgType::Scalar>;
859 using RhsScalar = std::remove_const_t<typename EvalRightArgType::Scalar>;
860
861 using LeftEvaluator = TensorEvaluator<EvalLeftArgType, Device>;
862 using RightEvaluator = TensorEvaluator<EvalRightArgType, Device>;
863
864 const int lhs_packet_size = internal::unpacket_traits<typename LeftEvaluator::PacketReturnType>::size;
865 const int rhs_packet_size = internal::unpacket_traits<typename RightEvaluator::PacketReturnType>::size;
866
867 using LhsMapper =
868 internal::TensorContractionInputMapper<LhsScalar, Index, internal::Lhs, LeftEvaluator, left_nocontract_t,
869 contract_t, lhs_packet_size, lhs_inner_dim_contiguous, false, Unaligned>;
870
871 using RhsMapper =
872 internal::TensorContractionInputMapper<RhsScalar, Index, internal::Rhs, RightEvaluator, right_nocontract_t,
873 contract_t, rhs_packet_size, rhs_inner_dim_contiguous,
874 rhs_inner_dim_reordered, Unaligned>;
875
876 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
877
878 using TensorContractionKernel =
879 internal::TensorContractionKernel<Scalar, LhsScalar, RhsScalar, Index, OutputMapper, LhsMapper, RhsMapper>;
880
881 // initialize data mappers
882 LhsMapper lhs(this->m_leftImpl, this->m_left_nocontract_strides, this->m_i_strides,
883 this->m_left_contracting_strides, this->m_k_strides);
884
885 RhsMapper rhs(this->m_rightImpl, this->m_right_nocontract_strides, this->m_j_strides,
886 this->m_right_contracting_strides, this->m_k_strides);
887
888 OutputMapper output(buffer, m);
889
890 // Sizes of the blocks to load in cache. See the Goto paper for details.
891 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByCol> blocking(
892 k_slice, m, n, num_threads);
893 const Index kc = blocking.kc();
894 const Index mc = numext::mini(m, blocking.mc());
895 const Index nc = numext::mini(n, blocking.nc());
896
897 using LhsBlock = typename TensorContractionKernel::LhsBlock;
898 using RhsBlock = typename TensorContractionKernel::RhsBlock;
899
900 LhsBlock blockA;
901 RhsBlock blockB;
902
903 TensorContractionKernel kernel(m, k_slice, n, mc, kc, nc);
904
905 using BlockMemHandle = typename TensorContractionKernel::BlockMemHandle;
906 const BlockMemHandle packed_mem = kernel.allocate(this->m_device, &blockA, &blockB);
907
908 // If a contraction kernel does not support beta, explicitly initialize
909 // output buffer with zeroes.
910 EIGEN_IF_CONSTEXPR (!TensorContractionKernel::HasBeta) {
911 this->m_device.fill(buffer, buffer + m * n, Scalar(0));
912 }
913
914 for (Index i2 = 0; i2 < m; i2 += mc) {
915 const Index actual_mc = numext::mini(i2 + mc, m) - i2;
916 for (Index k2 = k_start; k2 < k_end; k2 += kc) {
917 // make sure we don't overshoot right edge of left matrix, then pack vertical panel
918 const Index actual_kc = numext::mini(k2 + kc, k_end) - k2;
919 kernel.packLhs(&blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
920
921 // If kernel supports beta, there is no need to initialize output
922 // buffer with zeroes.
923 const Scalar alpha = Scalar(1);
924 const Scalar beta = (TensorContractionKernel::HasBeta && k2 == k_start) ? Scalar(0) : Scalar(1);
925
926 // series of horizontal blocks
927 for (Index j2 = 0; j2 < n; j2 += nc) {
928 // make sure we don't overshoot right edge of right matrix, then pack block
929 const Index actual_nc = numext::mini(j2 + nc, n) - j2;
930 kernel.packRhs(&blockB, rhs.getSubMapper(k2, j2), actual_kc, actual_nc);
931
932 // call gebp (matrix kernel)
933 // The parameters here are copied from Eigen's GEMM implementation
934 const OutputMapper output_mapper = output.getSubMapper(i2, j2);
935 kernel.invoke(output_mapper, blockA, blockB, actual_mc, actual_kc, actual_nc, alpha, beta);
936
937 // We are done with this [i2, j2] output block.
938 if (use_output_kernel && k2 + kc >= k_end) {
939 m_output_kernel(output_mapper, m_tensor_contraction_params, i2, j2, actual_mc, actual_nc);
940 }
941 }
942 }
943 }
944
945 kernel.deallocate(this->m_device, packed_mem);
946 }
947
948 EIGEN_STRONG_INLINE void cleanup() {
949 m_leftImpl.cleanup();
950 m_rightImpl.cleanup();
951
952 if (m_result != nullptr) {
953 m_device.deallocate(m_result);
954 m_result = nullptr;
955 }
956 }
957
958 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index) const { return m_result[index]; }
959
960 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool) const {
961 return TensorOpCost(sizeof(CoeffReturnType), 0, 0);
962 }
963
964 template <int LoadMode>
965 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packet(Index index) const {
966 return internal::ploadt<PacketReturnType, LoadMode>(m_result + index);
967 }
968
969 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE EvaluatorPointerType data() const { return m_result; }
970
971 // Required so a contraction can be composed with operators whose own
972 // getResourceRequirements() forwards into m_impl (TensorPaddingOp,
973 // TensorBroadcastingOp, etc.). Without this, e.g. an expression like
974 // `Tensor B = A.contract(C, dims).pad(p)` fails to compile because
975 // Pad's BlockAccess is gated on m_impl.RawAccess (which is true here)
976 // and instantiating Pad's getResourceRequirements then requires this
977 // method on the operand evaluator.
978 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements() const {
979 return internal::TensorBlockResourceRequirements::any();
980 }
981
982 protected:
983 Dimensions m_dimensions;
984
985 contract_t m_k_strides{};
986 contract_t m_left_contracting_strides{};
987 contract_t m_right_contracting_strides{};
988
989 bool m_lhs_inner_dim_contiguous;
990 bool m_rhs_inner_dim_contiguous;
991 bool m_rhs_inner_dim_reordered;
992 // Set in the constructor; describe whether the contracted LHS/RHS dims
993 // form a contiguous block at the leading or trailing end of storage. Used
994 // by evalProductSequential to pick a direct GEMV path when applicable.
995 bool m_lhs_contracted_dims_leading;
996 bool m_rhs_contracted_dims_leading;
997 bool m_rhs_contracted_dims_trailing;
998
999 left_nocontract_t m_i_strides{};
1000 right_nocontract_t m_j_strides{};
1001 left_nocontract_t m_left_nocontract_strides{};
1002 right_nocontract_t m_right_nocontract_strides{};
1003
1004 Index m_i_size;
1005 Index m_j_size;
1006 Index m_k_size;
1007
1008 TensorContractionParams m_tensor_contraction_params;
1009
1010 TensorEvaluator<EvalLeftArgType, Device> m_leftImpl;
1011 TensorEvaluator<EvalRightArgType, Device> m_rightImpl;
1012 const Device EIGEN_DEVICE_REF m_device;
1013 OutputKernelType m_output_kernel;
1014 EvaluatorPointerType m_result;
1015};
1016
1017// evaluator for default device
1018template <typename Indices, typename LeftArgType, typename RightArgType, typename OutputKernelType, typename Device>
1019struct TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device>
1020 : public TensorContractionEvaluatorBase<
1021 TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device>> {
1022 using Self = TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device>;
1023 using Base = TensorContractionEvaluatorBase<Self>;
1024
1025 using XprType = TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>;
1026 using Scalar = std::remove_const_t<typename XprType::Scalar>;
1027 using Index = typename XprType::Index;
1028 using CoeffReturnType = typename XprType::CoeffReturnType;
1029 using PacketReturnType = typename PacketType<CoeffReturnType, Device>::type;
1030
1031 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
1032
1033 // Most of the code is assuming that both input tensors are ColMajor. If the
1034 // inputs are RowMajor, we will "cheat" by swapping the LHS and RHS:
1035 // If we want to compute A * B = C, where A is LHS and B is RHS, the code
1036 // will pretend B is LHS and A is RHS.
1037 using EvalLeftArgType = std::conditional_t<Layout == static_cast<int>(ColMajor), LeftArgType, RightArgType>;
1038 using EvalRightArgType = std::conditional_t<Layout == static_cast<int>(ColMajor), RightArgType, LeftArgType>;
1039
1040 static constexpr int LDims =
1041 internal::array_size<typename TensorEvaluator<EvalLeftArgType, Device>::Dimensions>::value;
1042 static constexpr int RDims =
1043 internal::array_size<typename TensorEvaluator<EvalRightArgType, Device>::Dimensions>::value;
1044 static constexpr int ContractDims = internal::array_size<Indices>::value;
1045
1046 using contract_t = array<Index, ContractDims>;
1047 using left_nocontract_t = array<Index, LDims - ContractDims>;
1048 using right_nocontract_t = array<Index, RDims - ContractDims>;
1049
1050 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
1051
1052 // Could we use NumDimensions here?
1053 using Dimensions = DSizes<Index, NumDims>;
1054
1055 TensorEvaluator(const XprType& op, const Device& device) : Base(op, device) {}
1056
1057 template <int Alignment>
1058 void evalProduct(Scalar* buffer) const {
1059 internal::tensor_contraction_dispatch(
1060 [&](auto lhs_c, auto rhs_c, auto rhs_r) {
1061 this->template evalProductSequential<lhs_c(), rhs_c(), rhs_r(), Alignment>(buffer);
1062 },
1063 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered);
1064 }
1065};
1066
1067} // end namespace Eigen
1068
1069#endif // EIGEN_TENSOR_TENSOR_CONTRACTION_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
Definition TensorContraction.h:335
const internal::remove_all_t< typename LhsXprType::Nested > & lhsExpression() const
Definition TensorContraction.h:352
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47