11#ifndef EIGEN_TENSOR_TENSOR_CONTRACTION_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_H
15#include "./InternalHeaderCheck.h"
21template <
typename Dimensions,
typename LhsXprType,
typename RhsXprType,
typename OutputKernelType>
22struct traits<TensorContractionOp<Dimensions, LhsXprType, RhsXprType, OutputKernelType>> {
24 using Scalar =
typename gebp_traits<std::remove_const_t<typename LhsXprType::Scalar>,
25 std::remove_const_t<typename RhsXprType::Scalar>>::ResScalar;
27 using StorageKind =
typename promote_storage_type<typename traits<LhsXprType>::StorageKind,
28 typename traits<RhsXprType>::StorageKind>::ret;
30 typename promote_index_type<typename traits<LhsXprType>::Index,
typename traits<RhsXprType>::Index>::type;
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;
36 std::conditional_t<Pointer_type_promotion<typename LhsXprType::Scalar, Scalar>::val,
37 typename traits<LhsXprType>::PointerType,
typename traits<RhsXprType>::PointerType>;
39 static constexpr int Flags = 0;
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>&;
47template <
typename Indices_,
typename LeftArgType_,
typename RightArgType_,
typename OutputKernelType_,
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_;
58 static constexpr int NumDimensions =
59 traits<LeftArgType_>::NumDimensions + traits<RightArgType_>::NumDimensions - 2 * array_size<Indices_>::value;
63template <
typename LhsScalar,
typename RhsScalar>
64struct TensorContractionBlockMemAllocator {
65 using BlockMemHandle =
void*;
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));
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);
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));
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));
109 template <
typename Device>
110 EIGEN_DEVICE_FUNC
static void deallocate(Device& d, BlockMemHandle handle) {
111 d.deallocate(handle);
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);
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;
156template <
typename ResScalar,
typename LhsScalar,
typename RhsScalar,
typename StorageIndex,
typename OutputMapper,
157 typename LhsMapper,
typename RhsMapper>
158struct TensorContractionKernel {
161 static constexpr bool HasBeta =
false;
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_) {}
168 using LhsBlock = LhsScalar*;
169 using RhsBlock = RhsScalar*;
172 using BlockMemAllocator = TensorContractionBlockMemAllocator<LhsScalar, RhsScalar>;
173 using BlockMemHandle =
typename BlockMemAllocator::BlockMemHandle;
175 using Traits =
typename internal::gebp_traits<LhsScalar, RhsScalar>;
177 using LhsPacker = internal::gemm_pack_lhs<LhsScalar, StorageIndex,
typename LhsMapper::SubMapper, Traits::mr,
178 Traits::LhsProgress,
typename Traits::LhsPacket4Packing,
ColMajor>;
181 internal::gemm_pack_rhs<RhsScalar, StorageIndex, typename RhsMapper::SubMapper, Traits::nr, ColMajor>;
183 using GebpKernel = internal::gebp_kernel<LhsScalar, RhsScalar, StorageIndex, OutputMapper, Traits::mr, Traits::nr,
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);
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);
198 template <
typename Device>
199 EIGEN_DEVICE_FUNC
static void deallocate(Device& d, BlockMemHandle handle) {
200 BlockMemAllocator::deallocate(d, handle);
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, 0,
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);
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) {
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 kComputeStrideFromBlockDimensions,
224 kComputeStrideFromBlockDimensions,
232 const StorageIndex m;
233 const StorageIndex k;
234 const StorageIndex n;
235 const StorageIndex bm;
236 const StorageIndex bk;
237 const StorageIndex bn;
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>{});
252 fn(bool_constant<true>{}, bool_constant<true>{}, bool_constant<false>{});
254 if (rhs_inner_dim_reordered)
255 fn(bool_constant<true>{}, bool_constant<false>{}, bool_constant<true>{});
257 fn(bool_constant<true>{}, bool_constant<false>{}, bool_constant<false>{});
260 if (rhs_inner_dim_contiguous) {
261 if (rhs_inner_dim_reordered)
262 fn(bool_constant<false>{}, bool_constant<true>{}, bool_constant<true>{});
264 fn(bool_constant<false>{}, bool_constant<true>{}, bool_constant<false>{});
266 if (rhs_inner_dim_reordered)
267 fn(bool_constant<false>{}, bool_constant<false>{}, bool_constant<true>{});
269 fn(bool_constant<false>{}, bool_constant<false>{}, bool_constant<false>{});
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)
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)
295struct TensorContractionParams {
298 bool swapped_arguments;
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 {}
332template <
typename Indices,
typename LhsXprType,
typename RhsXprType,
333 typename OutputKernelType =
const NoOpOutputKernel>
334class TensorContractionOp
335 :
public TensorBase<TensorContractionOp<Indices, LhsXprType, RhsXprType, OutputKernelType>, ReadOnlyAccessors> {
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;
344 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorContractionOp(
const LhsXprType& lhs,
const RhsXprType& rhs,
346 const OutputKernelType& output_kernel = OutputKernelType())
347 : m_lhs_xpr(lhs), m_rhs_xpr(rhs), m_indices(dims), m_output_kernel(output_kernel) {}
349 EIGEN_DEVICE_FUNC
const Indices& indices()
const {
return m_indices; }
352 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename LhsXprType::Nested>&
lhsExpression()
const {
356 EIGEN_DEVICE_FUNC
const internal::remove_all_t<typename RhsXprType::Nested>& rhsExpression()
const {
360 EIGEN_DEVICE_FUNC
const OutputKernelType& outputKernel()
const {
return m_output_kernel; }
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;
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);
379template <
int StorageOrder,
bool MatrixIsRight>
380struct GemvDirectDispatcher<false, StorageOrder, MatrixIsRight> {
381 template <
typename Evaluator,
typename Scalar>
382 static bool run(
const Evaluator*, Scalar*) {
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;
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;
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;
410 static constexpr bool RawAccess =
true;
413 using TensorBlock = internal::TensorBlockNotImplemented;
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>;
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;
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>;
436 using Dimensions = DSizes<Index, NumDims>;
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(),
442 m_rightImpl(choose(Cond<static_cast<int>(Layout) == static_cast<int>(
ColMajor)>(), op.rhsExpression(),
446 m_output_kernel(op.outputKernel()),
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);
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)) {
457 for (
int i = 0; i < LDims; i++) {
458 eval_left_dims[i] = m_leftImpl.dimensions()[i];
460 for (
int i = 0; i < RDims; i++) {
461 eval_right_dims[i] = m_rightImpl.dimensions()[i];
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;
470 for (
int i = 0; i < LDims; i++) {
471 eval_left_dims[i] = m_leftImpl.dimensions()[LDims - i - 1];
473 for (
int i = 0; i < RDims; i++) {
474 eval_right_dims[i] = m_rightImpl.dimensions()[RDims - i - 1];
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;
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]);
496 array<Index, LDims> lhs_strides;
498 for (
int i = 0; i < LDims - 1; ++i) {
499 lhs_strides[i + 1] = lhs_strides[i] * eval_left_dims[i];
502 array<Index, RDims> rhs_strides;
504 for (
int i = 0; i < RDims - 1; ++i) {
505 rhs_strides[i + 1] = rhs_strides[i] * eval_right_dims[i];
516 m_lhs_inner_dim_contiguous =
true;
518 Index nocontract_idx = 0;
520 for (
int i = 0; i < LDims; i++) {
522 bool contracting =
false;
523 for (
int j = 0; j < ContractDims; j++) {
524 if (eval_op_indices[j].first == i) {
531 m_dimensions[dim_idx] = eval_left_dims[i];
532 m_left_nocontract_strides[nocontract_idx] = lhs_strides[i];
534 m_lhs_inner_dim_contiguous =
false;
536 m_i_strides[nocontract_idx] = m_i_size;
537 m_i_size *= eval_left_dims[i];
544 for (
int i = 0; i < RDims; i++) {
545 bool contracting =
false;
547 for (
int j = 0; j < ContractDims; j++) {
548 if (eval_op_indices[j].second == i) {
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];
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;
574 Index size = eval_left_dims[left];
575 eigen_assert(size == eval_right_dims[right] &&
"Contraction axes must be same size");
577 m_k_strides[i] = m_k_size;
579 m_left_contracting_strides[i] = lhs_strides[left];
580 m_right_contracting_strides[i] = rhs_strides[right];
582 if (i > 0 && right < eval_op_indices[i - 1].second) {
583 m_rhs_inner_dim_reordered =
true;
586 m_rhs_inner_dim_contiguous =
false;
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;
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]);
616 m_tensor_contraction_params.swapped_arguments =
static_cast<int>(Layout) ==
RowMajor;
619 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE
const Dimensions& dimensions()
const {
return m_dimensions; }
621 EIGEN_STRONG_INLINE
bool evalSubExprsIfNeeded(EvaluatorPointerType data) {
622 m_leftImpl.evalSubExprsIfNeeded(
nullptr);
623 m_rightImpl.evalSubExprsIfNeeded(
nullptr);
628 m_result =
static_cast<EvaluatorPointerType
>(m_device.allocate(dimensions().TotalSize() *
sizeof(Scalar)));
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) {
640 evalToAsync(dest, [done]() { done(
false); });
642 m_result =
static_cast<EvaluatorPointerType
>(m_device.allocate(dimensions().TotalSize() *
sizeof(Scalar)));
643 evalToAsync(m_result, [done]() { done(
true); });
650 EIGEN_DEVICE_FUNC
void evalTo(Scalar* buffer)
const {
651 static_cast<const Derived*
>(
this)->
template evalProduct<Unaligned>(buffer);
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));
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) {
666 if (this->m_k_size == 0) {
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);
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>(
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)) {
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)) {
714 if (m_rhs_contracted_dims_trailing) {
715 if (internal::GemvDirectDispatcher<types_match, ColMajor, true>::run(
this, buffer)) {
721 this->
template evalGemm<lhs_inner_dim_contiguous, rhs_inner_dim_contiguous, rhs_inner_dim_reordered, Alignment>(
725 template <
bool lhs_inner_dim_contiguous,
bool rhs_inner_dim_contiguous,
bool rhs_inner_dim_reordered,
int Alignment>
726#if !defined(EIGEN_HIPCC)
730 evalGemv(Scalar* buffer)
const {
731 const Index rows = m_i_size;
732 const Index cols = m_k_size;
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>;
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>;
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);
754 const Scalar alpha(1);
755 const Index resIncr(1);
758 m_device.fill(buffer, buffer + rows, Scalar(0));
760 internal::general_matrix_vector_product<Index, LhsScalar, LhsMapper,
ColMajor,
false, RhsScalar, RhsMapper,
761 false>::run(rows, cols, lhs, rhs, buffer, resIncr, alpha);
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));
778 template <
int StorageOrder,
bool MatrixIsRight>
779#if !defined(EIGEN_HIPCC)
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>>;
789 const Index rows = MatrixIsRight ? m_j_size : m_i_size;
790 const Index cols = m_k_size;
793 const Index mat_stride = (StorageOrder ==
RowMajor) ? cols : rows;
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);
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, 1);
804 m_device.fill(buffer, buffer + rows, Scalar(0));
806 internal::general_matrix_vector_product<Index, MatScalar, LhsMapper, StorageOrder,
false, VecScalar, RhsMapper,
807 false>::run(rows, cols, lhs, rhs, buffer, 1, Scalar(1));
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));
814 template <
bool lhs_inner_dim_contiguous,
bool rhs_inner_dim_contiguous,
bool rhs_inner_dim_reordered,
int Alignment>
815#if !defined(EIGEN_HIPCC)
819 evalGemm(Scalar* buffer)
const {
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);
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 false>(buffer, k_start, k_end, num_threads);
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);
838 const Index k_slice = k_end - k_start;
841 const Index m = this->m_i_size;
844 const Index n = this->m_j_size;
846 if (m == 0 || n == 0)
return;
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);
858 using LhsScalar = std::remove_const_t<typename EvalLeftArgType::Scalar>;
859 using RhsScalar = std::remove_const_t<typename EvalRightArgType::Scalar>;
861 using LeftEvaluator = TensorEvaluator<EvalLeftArgType, Device>;
862 using RightEvaluator = TensorEvaluator<EvalRightArgType, Device>;
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;
868 internal::TensorContractionInputMapper<LhsScalar, Index, internal::Lhs, LeftEvaluator, left_nocontract_t,
869 contract_t, lhs_packet_size, lhs_inner_dim_contiguous,
false,
Unaligned>;
872 internal::TensorContractionInputMapper<RhsScalar, Index, internal::Rhs, RightEvaluator, right_nocontract_t,
873 contract_t, rhs_packet_size, rhs_inner_dim_contiguous,
876 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
878 using TensorContractionKernel =
879 internal::TensorContractionKernel<Scalar, LhsScalar, RhsScalar, Index, OutputMapper, LhsMapper, RhsMapper>;
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);
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);
888 OutputMapper output(buffer, m);
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());
897 using LhsBlock =
typename TensorContractionKernel::LhsBlock;
898 using RhsBlock =
typename TensorContractionKernel::RhsBlock;
903 TensorContractionKernel kernel(m, k_slice, n, mc, kc, nc);
905 using BlockMemHandle =
typename TensorContractionKernel::BlockMemHandle;
906 const BlockMemHandle packed_mem = kernel.allocate(this->m_device, &blockA, &blockB);
910 EIGEN_IF_CONSTEXPR (!TensorContractionKernel::HasBeta) {
911 this->m_device.fill(buffer, buffer + m * n, Scalar(0));
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) {
918 const Index actual_kc = numext::mini(k2 + kc, k_end) - k2;
919 kernel.packLhs(&blockA, lhs.getSubMapper(i2, k2), actual_kc, actual_mc);
923 const Scalar alpha = Scalar(1);
924 const Scalar beta = (TensorContractionKernel::HasBeta && k2 == k_start) ? Scalar(0) : Scalar(1);
927 for (Index j2 = 0; j2 < n; j2 += nc) {
929 const Index actual_nc = numext::mini(j2 + nc, n) - j2;
930 kernel.packRhs(&blockB, rhs.getSubMapper(k2, j2), actual_kc, actual_nc);
934 const OutputMapper output_mapper = output.getSubMapper(i2, j2);
935 kernel.invoke(output_mapper, blockA, blockB, actual_mc, actual_kc, actual_nc, alpha, beta);
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);
945 kernel.deallocate(this->m_device, packed_mem);
948 EIGEN_STRONG_INLINE
void cleanup() {
949 m_leftImpl.cleanup();
950 m_rightImpl.cleanup();
952 if (m_result !=
nullptr) {
953 m_device.deallocate(m_result);
958 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeff(Index index)
const {
return m_result[index]; }
960 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(
bool)
const {
961 return TensorOpCost(
sizeof(CoeffReturnType), 0, 0);
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);
969 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE EvaluatorPointerType data()
const {
return m_result; }
978 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements()
const {
979 return internal::TensorBlockResourceRequirements::any();
983 Dimensions m_dimensions;
985 contract_t m_k_strides{};
986 contract_t m_left_contracting_strides{};
987 contract_t m_right_contracting_strides{};
989 bool m_lhs_inner_dim_contiguous;
990 bool m_rhs_inner_dim_contiguous;
991 bool m_rhs_inner_dim_reordered;
995 bool m_lhs_contracted_dims_leading;
996 bool m_rhs_contracted_dims_leading;
997 bool m_rhs_contracted_dims_trailing;
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{};
1008 TensorContractionParams m_tensor_contraction_params;
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;
1018template <
typename Indices,
typename LeftArgType,
typename RightArgType,
typename OutputKernelType,
typename 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>;
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;
1031 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
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>;
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;
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>;
1050 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
1053 using Dimensions = DSizes<Index, NumDims>;
1055 TensorEvaluator(
const XprType& op,
const Device& device) : Base(op, device) {}
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);
1063 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered);
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