Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorContractionBlocking.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_BLOCKING_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_BLOCKING_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18namespace internal {
19
20constexpr int ShardByRow = 0;
21constexpr int ShardByCol = 1;
22
23// Default Blocking Strategy
24template <typename ResScalar, typename LhsScalar, typename RhsScalar, typename StorageIndex,
25 int ShardingType = ShardByCol>
26class TensorContractionBlocking {
27 public:
28 /*
29 adding EIGEN_DEVICE_FUNC unconditionally to 'TensorContractionBlocking' constructor in `TensorContractionBlocking.h`
30 requires adding EIGEN_DEVICE_FUNC to `computeProductBlockingSizes` in `GeneralBlockPanelKernel.h`
31 which in turn, requires adding EIGEN_DEVICE_FUNC to `evaluateProductBlockingSizesHeuristic` in
32 `GeneralBlockPanelKernel.h` which in turn, requires adding EIGEN_DEVICE_FUNC to `manage_caching_sizes` in
33 `GeneralBlockPanelKernel.h` (else HIPCC will error out)
34
35 However adding EIGEN_DEVICE_FUNC to `manage_caching_sizes` in `GeneralBlockPanelKernel.h`
36 results in NVCC erroring out with the following error
37
38 ../Eigen/src/Core/products/GeneralBlockPanelKernel.h(57): error #2901:
39 dynamic initialization is not supported for function-scope static variables within a __device__/__global__
40 function
41 */
42
43#if !defined(EIGEN_HIPCC)
44 EIGEN_DEVICE_FUNC
45#endif
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);
50 } else {
51 computeProductBlockingSizes<LhsScalar, RhsScalar, 1>(kc_, nc_, mc_, num_threads);
52 }
53
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;
56 }
57
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_; }
61
62 private:
63 StorageIndex kc_;
64 StorageIndex mc_;
65 StorageIndex nc_;
66};
67
68} // end namespace internal
69} // end namespace Eigen
70
71#endif // EIGEN_TENSOR_TENSOR_CONTRACTION_BLOCKING_H
Namespace containing all symbols from the Eigen library.