11#ifndef EIGEN_TENSOR_TENSOR_CONTRACTION_BLOCKING_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_BLOCKING_H
15#include "./InternalHeaderCheck.h"
20constexpr int ShardByRow = 0;
21constexpr int ShardByCol = 1;
24template <
typename ResScalar,
typename LhsScalar,
typename RhsScalar,
typename StorageIndex,
25 int ShardingType = ShardByCol>
26class TensorContractionBlocking {
43#if !defined(EIGEN_HIPCC)
46 TensorContractionBlocking(StorageIndex k, StorageIndex m, StorageIndex n, StorageIndex num_threads = 1)
47 : kc_(k), mc_(m), nc_(n) {
48 if (ShardingType == ShardByCol) {
49 computeProductBlockingSizes<LhsScalar, RhsScalar, 1>(kc_, mc_, nc_, num_threads);
51 computeProductBlockingSizes<LhsScalar, RhsScalar, 1>(kc_, nc_, mc_, num_threads);
54 const int rhs_packet_size = internal::packet_traits<RhsScalar>::size;
55 kc_ = (rhs_packet_size <= 8 || kc_ <= rhs_packet_size) ? kc_ : (kc_ / rhs_packet_size) * rhs_packet_size;
58 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE StorageIndex kc()
const {
return kc_; }
59 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE StorageIndex mc()
const {
return mc_; }
60 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE StorageIndex nc()
const {
return nc_; }
Namespace containing all symbols from the Eigen library.