11#ifndef EIGEN_TENSOR_TENSOR_CONTRACTION_THREAD_POOL_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_THREAD_POOL_H
15#ifdef EIGEN_USE_THREADS
18#include "./InternalHeaderCheck.h"
22template <
typename Indices,
typename LeftArgType,
typename RightArgType,
typename OutputKernelType>
25 :
public TensorContractionEvaluatorBase<TensorEvaluator<
26 const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, ThreadPoolDevice>> {
27 using Device = ThreadPoolDevice;
29 using Self = TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device>;
30 using Base = TensorContractionEvaluatorBase<Self>;
32 using XprType = TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>;
33 using Scalar = std::remove_const_t<typename XprType::Scalar>;
34 using Index =
typename XprType::Index;
35 using CoeffReturnType =
typename XprType::CoeffReturnType;
36 using PacketReturnType =
typename PacketType<CoeffReturnType, Device>::type;
38 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
44 using EvalLeftArgType =
45 std::conditional_t<static_cast<int>(Layout) ==
static_cast<int>(
ColMajor), LeftArgType, RightArgType>;
46 using EvalRightArgType =
47 std::conditional_t<static_cast<int>(Layout) ==
static_cast<int>(
ColMajor), RightArgType, LeftArgType>;
49 static constexpr int LDims =
50 internal::array_size<typename TensorEvaluator<EvalLeftArgType, Device>::Dimensions>::value;
51 static constexpr int RDims =
52 internal::array_size<typename TensorEvaluator<EvalRightArgType, Device>::Dimensions>::value;
53 static constexpr int ContractDims = internal::array_size<Indices>::value;
55 using left_dim_mapper_t = array<Index, LDims>;
56 using right_dim_mapper_t = array<Index, RDims>;
58 using contract_t = array<Index, ContractDims>;
59 using left_nocontract_t = array<Index, LDims - ContractDims>;
60 using right_nocontract_t = array<Index, RDims - ContractDims>;
62 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
64 using Dimensions = DSizes<Index, NumDims>;
67 using LhsScalar = std::remove_const_t<typename EvalLeftArgType::Scalar>;
68 using RhsScalar = std::remove_const_t<typename EvalRightArgType::Scalar>;
69 using Traits =
typename internal::gebp_traits<LhsScalar, RhsScalar>;
71 using LeftEvaluator = TensorEvaluator<EvalLeftArgType, Device>;
72 using RightEvaluator = TensorEvaluator<EvalRightArgType, Device>;
74 TensorEvaluator(
const XprType& op,
const Device& device) : Base(op, device) {}
76 template <
int Alignment>
77 void evalProduct(Scalar* buffer)
const {
78 evalProductImpl<NoCallback, Alignment>(buffer, NoCallback());
81 template <
typename EvalToCallback,
int Alignment>
82 void evalProductAsync(Scalar* buffer, EvalToCallback done)
const {
83 evalProductImpl<EvalToCallback, Alignment>(buffer, std::move(done));
86 template <
typename DoneCallback,
int Alignment>
87 void evalProductImpl(Scalar* buffer, DoneCallback done)
const {
103 static constexpr bool IsEvalInSyncMode = std::is_same<DoneCallback, NoCallback>::value;
105 const Index m = this->m_i_size;
106 const Index n = this->m_j_size;
107 const Index k = this->m_k_size;
108 if (m == 0 || n == 0) {
109 EIGEN_IF_CONSTEXPR (!IsEvalInSyncMode) done();
113 internal::tensor_contraction_dispatch(
114 [&](
auto lhs_c,
auto rhs_c,
auto rhs_r) {
115 this->
template evalProductSequential<lhs_c(), rhs_c(), rhs_r(), Unaligned>(buffer);
117 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered);
118 EIGEN_IF_CONSTEXPR (!IsEvalInSyncMode) done();
145 bool shard_by_col = shardByCol(m, n, 2);
151 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByCol> blocking(k, m, n,
157 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByRow> blocking(k, m, n,
168 const TensorOpCost cost = contractionCost(m, n, bm, bn, bk, shard_by_col,
false);
170 TensorCostModel<ThreadPoolDevice>::numThreads(
static_cast<double>(n) * m, cost, this->m_device.numThreads());
171 int num_threads_by_k = numThreadsInnerDim(m, n, k);
172 if (shardByInnerDim(m, n, k, num_threads, num_threads_by_k)) {
175 EIGEN_IF_CONSTEXPR (IsEvalInSyncMode) {
176 EvalShardedByInnerDimContext<DoneCallback> ctx(
this, num_threads_by_k, buffer, m, n, k, std::move(done));
177 ctx.template run<Alignment>();
180 new EvalShardedByInnerDimContext<DoneCallback>(
this, num_threads_by_k, buffer, m, n, k, std::move(done));
181 ctx->template runAsync<Alignment>();
189 if (n == 1) num_threads = 1;
191 if (num_threads == 1) {
192 internal::tensor_contraction_dispatch(
193 [&](
auto lhs_c,
auto rhs_c,
auto rhs_r) {
194 this->
template evalProductSequential<lhs_c(), rhs_c(), rhs_r(), Unaligned>(buffer);
196 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered);
197 EIGEN_IF_CONSTEXPR (!IsEvalInSyncMode) done();
202 shard_by_col = shardByCol(m, n, num_threads);
204 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByCol> blocking(
205 k, m, n, num_threads);
210 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByRow> blocking(
211 k, m, n, num_threads);
218 Index nm0 = numext::div_ceil(m, bm);
219 Index nn0 = numext::div_ceil(n, bn);
220 Index nk = numext::div_ceil(k, bk);
231 gm = coarsenM(m, n, bm, bn, bk, gn, num_threads, shard_by_col);
232 gn = coarsenN(m, n, bm, bn, bk, gm, num_threads, shard_by_col);
234 gn = coarsenN(m, n, bm, bn, bk, gm, num_threads, shard_by_col);
235 gm = coarsenM(m, n, bm, bn, bk, gn, num_threads, shard_by_col);
238 Index nm = numext::div_ceil(nm0, gm);
239 Index nn = numext::div_ceil(nn0, gn);
245 const Index sharding_dim_tasks = shard_by_col ? nn : nm;
246 const int num_worker_threads = this->m_device.numThreadsInPool();
251 const float oversharding_factor = num_worker_threads <= 4 ? 8.0
252 : num_worker_threads <= 8 ? 4.0
253 : num_worker_threads <= 16 ? 2.0
254 : num_worker_threads <= 32 ? 1.0
255 : num_worker_threads <= 64 ? 0.8
258 const bool parallelize_by_sharding_dim_only = sharding_dim_tasks >= oversharding_factor * num_worker_threads;
267 bool parallel_pack = num_threads >= nm * nn;
269 if (m * bk * Index(
sizeof(LhsScalar)) + n * bk * Index(
sizeof(RhsScalar)) <= l2CacheSize() * num_threads)
270 parallel_pack =
true;
273 if ((shard_by_col ? nm : nn) == 1) parallel_pack =
false;
276 if (parallelize_by_sharding_dim_only) parallel_pack =
false;
278 internal::tensor_contraction_dispatch(
279 [&](
auto lhs_c,
auto rhs_c,
auto rhs_r) {
280 EIGEN_IF_CONSTEXPR (IsEvalInSyncMode) {
281 EvalParallelContext<NoCallback, lhs_c(), rhs_c(), rhs_r(), Alignment> ctx(
282 this, num_threads, buffer, m, n, k, bm, bn, bk, nm, nn, nk, gm, gn, nm0, nn0, shard_by_col,
283 parallel_pack, parallelize_by_sharding_dim_only, NoCallback());
286 auto* ctx =
new EvalParallelContext<DoneCallback, lhs_c(), rhs_c(), rhs_r(), Alignment>(
287 this, num_threads, buffer, m, n, k, bm, bn, bk, nm, nn, nk, gm, gn, nm0, nn0, shard_by_col,
288 parallel_pack, parallelize_by_sharding_dim_only, std::move(done));
292 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered);
300 void operator()()
const { eigen_assert(
false &&
"NoCallback should never be called"); }
305 template <
typename DoneCallback,
typename Context>
306 class EvalParallelNotification;
309 template <
typename Context>
310 class EvalParallelNotification<NoCallback, Context> {
312 EvalParallelNotification(Context*, NoCallback) {}
313 void Notify() { done_.Notify(); }
314 void Wait() { done_.Wait(); }
317 Eigen::Notification done_;
321 template <
typename DoneCallback,
typename Context>
322 class EvalParallelNotification {
324 EvalParallelNotification(Context* ctx, DoneCallback done) : ctx_(ctx), done_(std::move(done)) {}
330 DoneCallback done_copy = std::move(done_);
350 template <
typename DoneCallback,
bool lhs_inner_dim_contiguous,
bool rhs_inner_dim_contiguous,
351 bool rhs_inner_dim_reordered,
int Alignment>
352 class EvalParallelContext {
355 internal::TensorContractionInputMapper<LhsScalar, Index, internal::Lhs, LeftEvaluator, left_nocontract_t,
356 contract_t, internal::packet_traits<LhsScalar>::size,
357 lhs_inner_dim_contiguous,
false,
Unaligned>;
359 internal::TensorContractionInputMapper<RhsScalar, Index, internal::Rhs, RightEvaluator, right_nocontract_t,
360 contract_t, internal::packet_traits<RhsScalar>::size,
361 rhs_inner_dim_contiguous, rhs_inner_dim_reordered,
Unaligned>;
363 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
365 using TensorContractionKernel =
366 internal::TensorContractionKernel<Scalar, LhsScalar, RhsScalar, Index, OutputMapper, LhsMapper, RhsMapper>;
368 using LhsBlock =
typename TensorContractionKernel::LhsBlock;
369 using RhsBlock =
typename TensorContractionKernel::RhsBlock;
370 using BlockMemHandle =
typename TensorContractionKernel::BlockMemHandle;
372 EvalParallelContext(
const Self* self,
int num_threads, Scalar* buffer, Index tm, Index tn, Index tk, Index bm,
373 Index bn, Index bk, Index nm, Index nn, Index nk, Index gm, Index gn, Index nm0, Index nn0,
374 bool shard_by_col,
bool parallel_pack,
bool parallelize_by_sharding_dim_only, DoneCallback done)
375 : created_by_thread_id_(std::this_thread::get_id()),
376 done_(this, std::move(done)),
377 device_(self->m_device),
378 lhs_(self->m_leftImpl, self->m_left_nocontract_strides, self->m_i_strides, self->m_left_contracting_strides,
380 rhs_(self->m_rightImpl, self->m_right_nocontract_strides, self->m_j_strides,
381 self->m_right_contracting_strides, self->m_k_strides),
384 output_kernel_(self->m_output_kernel),
385 tensor_contraction_params_(self->m_tensor_contraction_params),
386 num_threads_(num_threads),
387 shard_by_col_(shard_by_col),
388 parallel_pack_(parallel_pack),
389 parallelize_by_sharding_dim_only_(parallelize_by_sharding_dim_only),
403 kernel_(m_, k_, n_, bm_, bk_, bn_),
404 num_thread_local_allocations_(0),
408 thread_local_capacity(2 * (parallelize_by_sharding_dim_only_ ? device_.numThreadsInPool() : 0)),
411 lhs_thread_local_blocks_(shard_by_col_ ? 0 : thread_local_capacity, {*
this}, {*
this}),
412 rhs_thread_local_blocks_(shard_by_col_ ? thread_local_capacity : 0, {*
this}, {*
this}) {
414 eigen_assert(!(parallel_pack && parallelize_by_sharding_dim_only));
416 for (Index x = 0; x < P; x++) {
422 x == 0 ? 1 : (parallel_pack_ ? nn_ + nm_ : (shard_by_col_ ? nn_ : nm_)) + (x == P - 1 ? nm_ * nn_ : 0);
423 state_packing_ready_[x] = parallel_pack_ ? 0 : (shard_by_col_ ? nm_ : nn_);
424 state_kernel_[x] =
new std::atomic<uint8_t>*[nm_];
425 for (Index m = 0; m < nm_; m++) {
426 state_kernel_[x][m] =
new std::atomic<uint8_t>[nn_];
430 for (Index n = 0; n < nn_; n++)
431 state_kernel_[x][m][n].store((x == 0 ? 0 : 1) + (parallel_pack_ ? 2 : 1), std::memory_order_relaxed);
436 packed_mem_ = kernel_.allocateSlices(
440 std::min<Index>(nk_, P - 1),
441 packed_lhs_, packed_rhs_);
443 if (parallelize_by_sharding_dim_only_) {
444 const int num_worker_threads = device_.numThreadsInPool();
447 can_use_thread_local_packed_ =
new std::atomic<bool>[nn_];
448 for (
int i = 0; i < nn_; ++i) can_use_thread_local_packed_[i].store(
true, std::memory_order_relaxed);
450 Index num_blocks = num_worker_threads * gn_;
451 thread_local_pre_allocated_mem_ = kernel_.allocateSlices(
456 nullptr, &rhs_thread_local_pre_allocated_);
459 can_use_thread_local_packed_ =
new std::atomic<bool>[nm_];
460 for (
int i = 0; i < nm_; ++i) can_use_thread_local_packed_[i].store(
true, std::memory_order_relaxed);
462 Index num_blocks = num_worker_threads * gm_;
463 thread_local_pre_allocated_mem_ = kernel_.allocateSlices(
467 1, &lhs_thread_local_pre_allocated_,
473 ~EvalParallelContext() {
474 for (Index x = 0; x < P; x++) {
475 for (Index m = 0; m < nm_; m++)
delete[] state_kernel_[x][m];
476 delete[] state_kernel_[x];
478 kernel_.deallocate(device_, packed_mem_);
479 if (parallelize_by_sharding_dim_only_) {
480 kernel_.deallocate(device_, thread_local_pre_allocated_mem_);
481 delete[] can_use_thread_local_packed_;
506 std::thread::id created_by_thread_id_;
510 EvalParallelNotification<DoneCallback, EvalParallelContext> done_;
512 const Device& device_;
515 Scalar*
const buffer_;
516 OutputMapper output_;
517 OutputKernelType output_kernel_;
518 TensorContractionParams tensor_contraction_params_;
519 const int num_threads_;
520 const bool shard_by_col_;
521 const bool parallel_pack_;
522 const bool parallelize_by_sharding_dim_only_;
543 TensorContractionKernel kernel_;
579 static constexpr Index P = 3;
582 BlockMemHandle packed_mem_;
583 std::vector<LhsBlock> packed_lhs_[P - 1];
584 std::vector<RhsBlock> packed_rhs_[P - 1];
604 BlockMemHandle thread_local_pre_allocated_mem_;
608 std::vector<LhsBlock> lhs_thread_local_pre_allocated_;
609 std::vector<RhsBlock> rhs_thread_local_pre_allocated_;
612 std::atomic<int> num_thread_local_allocations_;
613 const int thread_local_capacity;
621 template <
typename BlockType>
622 class ThreadLocalBlocks {
624 ThreadLocalBlocks() =
default;
626 ThreadLocalBlocks(BlockType* base,
size_t grain_size)
627 : is_pre_allocated_(true), thread_local_pre_allocated_base_(base), grain_size_(grain_size) {}
629 ThreadLocalBlocks(BlockMemHandle mem_handle, std::vector<BlockType> blocks)
630 : is_pre_allocated_(false), mem_handle_(std::move(mem_handle)), blocks_(std::move(blocks)) {}
632 BlockType& block(
int grain_index) {
633 eigen_assert(grain_index >= 0);
634 eigen_assert(
static_cast<size_t>(grain_index) < size());
635 return is_pre_allocated_ ? thread_local_pre_allocated_base_[grain_index] : blocks_[grain_index];
638 void Release(EvalParallelContext& ctx)
const {
639 if (!is_pre_allocated_) {
640 ctx.kernel_.deallocate(ctx.device_, mem_handle_);
644 size_t size()
const {
return is_pre_allocated_ ? grain_size_ : blocks_.size(); }
647 bool is_pre_allocated_;
650 BlockType* thread_local_pre_allocated_base_ =
nullptr;
651 size_t grain_size_ = 0;
654 BlockMemHandle mem_handle_{};
655 std::vector<BlockType> blocks_;
664 template <
typename BlockType,
bool is_rhs>
665 class ThreadLocalBlocksInitialize {
666 static constexpr bool kIsLhs = !is_rhs && std::is_same<BlockType, LhsBlock>::value;
667 static constexpr bool kIsRhs = is_rhs && std::is_same<BlockType, RhsBlock>::value;
668 static_assert(kIsLhs || kIsRhs,
"Unknown block type");
670 using Blocks = ThreadLocalBlocks<BlockType>;
673 ThreadLocalBlocksInitialize(EvalParallelContext& ctx)
674 : ctx_(ctx), num_worker_threads_(ctx_.device_.numThreadsInPool()) {}
676 void operator()(Blocks& blocks) {
677 const int n = ctx_.num_thread_local_allocations_.fetch_add(1, std::memory_order_relaxed);
679 if (n >= num_worker_threads_) {
680 ThreadLocalBlocksAllocator<is_rhs>::allocate(ctx_, blocks);
682 ThreadLocalBlocksAllocator<is_rhs>::reuse(ctx_, n, blocks);
689 template <
bool pack_rhs,
typename EvalCtx = EvalParallelContext>
690 struct ThreadLocalBlocksAllocator;
692 template <
typename EvalCtx>
693 struct ThreadLocalBlocksAllocator<true, EvalCtx> {
694 static void allocate(EvalCtx& ctx, Blocks& blocks) {
695 std::vector<RhsBlock> rhs_blocks;
696 BlockMemHandle mem_handle = ctx.kernel_.allocateSlices(ctx.device_,
700 nullptr, &rhs_blocks);
701 blocks = ThreadLocalBlocks<RhsBlock>(std::move(mem_handle), std::move(rhs_blocks));
704 static void reuse(EvalCtx& ctx,
int index, Blocks& blocks) {
705 RhsBlock* ptr = &ctx.rhs_thread_local_pre_allocated_[ctx.gn_ * index];
706 blocks = ThreadLocalBlocks<RhsBlock>(ptr, ctx.gn_);
710 template <
typename EvalCtx>
711 struct ThreadLocalBlocksAllocator<false, EvalCtx> {
712 static void allocate(EvalCtx& ctx, Blocks& blocks) {
713 std::vector<LhsBlock> lhs_blocks;
714 BlockMemHandle mem_handle = ctx.kernel_.allocateSlices(ctx.device_,
718 &lhs_blocks,
nullptr);
719 blocks = ThreadLocalBlocks<LhsBlock>(std::move(mem_handle), std::move(lhs_blocks));
722 static void reuse(EvalCtx& ctx,
int index, Blocks& blocks) {
723 LhsBlock* ptr = &ctx.lhs_thread_local_pre_allocated_[ctx.gm_ * index];
724 blocks = ThreadLocalBlocks<LhsBlock>(ptr, ctx.gm_);
728 EvalParallelContext& ctx_;
729 const int num_worker_threads_;
732 template <
typename BlockType>
733 class ThreadLocalBlocksRelease {
735 using Blocks = ThreadLocalBlocks<BlockType>;
736 ThreadLocalBlocksRelease(EvalParallelContext& ctx) : ctx_(ctx) {}
737 void operator()(Blocks& blocks) { blocks.Release(ctx_); }
740 EvalParallelContext& ctx_;
744 using ThreadLocalLhsInit = ThreadLocalBlocksInitialize<LhsBlock,
false>;
745 using ThreadLocalRhsInit = ThreadLocalBlocksInitialize<RhsBlock,
true>;
748 using ThreadLocalLhsRelease = ThreadLocalBlocksRelease<LhsBlock>;
749 using ThreadLocalRhsRelease = ThreadLocalBlocksRelease<RhsBlock>;
753 Eigen::ThreadLocal<ThreadLocalBlocks<LhsBlock>, ThreadLocalLhsInit, ThreadLocalLhsRelease> lhs_thread_local_blocks_;
754 Eigen::ThreadLocal<ThreadLocalBlocks<RhsBlock>, ThreadLocalRhsInit, ThreadLocalRhsRelease> rhs_thread_local_blocks_;
761 std::atomic<bool>* can_use_thread_local_packed_;
763 std::atomic<uint8_t>** state_kernel_[P];
768 std::atomic<Index> state_packing_ready_[P];
769 std::atomic<Index> state_switch_[P];
771 LhsBlock& packed_lhs(Index m, Index k, Index m1,
bool use_thread_local) {
772 if (use_thread_local) {
773 eigen_assert(!shard_by_col_);
774 ThreadLocalBlocks<LhsBlock>& blocks = lhs_thread_local_blocks_.local();
776 Index grain_index = m1 - m * gm_;
778 internal::convert_index<int>(grain_index));
780 return packed_lhs_[k % (P - 1)][m1];
784 RhsBlock& packed_rhs(Index n, Index k, Index n1,
bool use_thread_local) {
785 if (use_thread_local) {
786 eigen_assert(shard_by_col_);
787 ThreadLocalBlocks<RhsBlock>& blocks = rhs_thread_local_blocks_.local();
789 Index grain_index = n1 - n * gn_;
791 internal::convert_index<int>(grain_index));
793 return packed_rhs_[k % (P - 1)][n1];
807 void pack_lhs(Index m, Index k) {
808 bool use_thread_local =
false;
810 if (parallelize_by_sharding_dim_only_ && !shard_by_col_ &&
811 can_use_thread_local_packed_[m].load(std::memory_order_relaxed)) {
812 if (state_kernel_[k % P][m][0].load(std::memory_order_relaxed) == 1) {
813 use_thread_local =
true;
819 can_use_thread_local_packed_[m].store(
false, std::memory_order_relaxed);
823 const Index mend = m * gm_ + gm(m);
824 for (Index m1 = m * gm_; m1 < mend; m1++)
825 kernel_.packLhs(&packed_lhs(m, k, m1, use_thread_local), lhs_.getSubMapper(m1 * bm_, k * bk_), bk(k), bm(m1));
827 if (!parallel_pack_ && shard_by_col_) {
828 eigen_assert(!use_thread_local);
831 signal_switch(k + 1);
832 for (Index n = nn_ - 1; n >= 0; n--) {
833 bool sync = parallelize_by_sharding_dim_only_ || n == 0;
834 signal_kernel(m, n, k, sync, use_thread_local);
839 void pack_rhs(Index n, Index k) {
840 bool use_thread_local =
false;
842 if (parallelize_by_sharding_dim_only_ && shard_by_col_ &&
843 can_use_thread_local_packed_[n].load(std::memory_order_relaxed)) {
844 if (state_kernel_[k % P][0][n].load(std::memory_order_relaxed) == 1) {
845 use_thread_local =
true;
851 can_use_thread_local_packed_[n].store(
false, std::memory_order_relaxed);
855 const Index nend = n * gn_ + gn(n);
856 for (Index n1 = n * gn_; n1 < nend; n1++) {
857 EIGEN_IF_CONSTEXPR (!TensorContractionKernel::HasBeta) {
868 std::fill_n(buffer_ + n1 * bn_ * m_, bn(n1) * m_, Scalar(0));
871 kernel_.packRhs(&packed_rhs(n, k, n1, use_thread_local), rhs_.getSubMapper(k * bk_, n1 * bn_), bk(k), bn(n1));
874 if (parallel_pack_ || shard_by_col_) {
875 signal_switch(k + 1);
876 for (Index m = nm_ - 1; m >= 0; m--) {
877 bool sync = parallelize_by_sharding_dim_only_ || m == 0;
878 signal_kernel(m, n, k, sync, use_thread_local);
881 eigen_assert(!use_thread_local);
886 void kernel(Index m, Index n, Index k,
bool use_thread_local) {
890 const Index nend = n * gn_ + gn(n);
891 const Index mend = m * gm_ + gm(m);
894 const Scalar alpha = Scalar(1);
895 const Scalar beta = (TensorContractionKernel::HasBeta && k == 0) ? Scalar(0) : Scalar(1);
898 for (Index n1 = n * gn_; n1 < nend; n1++) {
899 for (Index m1 = m * gm_; m1 < mend; m1++) {
900 const auto output_mapper = output_.getSubMapper(m1 * bm_, n1 * bn_);
901 kernel_.invoke(output_mapper, packed_lhs(m, k, m1, !shard_by_col_ && use_thread_local),
902 packed_rhs(n, k, n1, shard_by_col_ && use_thread_local), bm(m1), bk(k), bn(n1), alpha, beta);
906 output_kernel_(output_mapper, tensor_contraction_params_, m1 * bm_, n1 * bn_, bm(m1), bn(n1));
911 for (Index m1 = m * gm_; m1 < mend; m1++)
912 for (Index n1 = n * gn_; n1 < nend; n1++) {
913 const auto output_mapper = output_.getSubMapper(m1 * bm_, n1 * bn_);
914 kernel_.invoke(output_mapper, packed_lhs(m, k, m1, !shard_by_col_ && use_thread_local),
915 packed_rhs(n, k, n1, shard_by_col_ && use_thread_local), bm(m1), bk(k), bn(n1), alpha, beta);
919 output_kernel_(output_mapper, tensor_contraction_params_, m1 * bm_, n1 * bn_, bm(m1), bn(n1));
923 signal_kernel(m, n, k + 1,
false,
false);
924 signal_switch(k + 2);
927 void signal_packing(Index k) {
928 eigen_assert(!parallel_pack_);
929 Index s = state_packing_ready_[k % P].fetch_sub(1);
932 state_packing_ready_[k % P] = shard_by_col_ ? nm_ : nn_;
933 enqueue_packing(k, shard_by_col_);
936 void signal_kernel(Index m, Index n, Index k,
bool sync,
bool use_thread_local) {
937 std::atomic<uint8_t>* state = &state_kernel_[k % P][m][n];
938 Index s = state->load();
940 if (s != 1 && state->fetch_sub(1) != 1) {
941 eigen_assert(!use_thread_local);
944 state->store(parallel_pack_ ? 3 : 2, std::memory_order_relaxed);
946 kernel(m, n, k, use_thread_local);
948 eigen_assert(!use_thread_local);
949 device_.enqueue([
this, m, n, k, use_thread_local]() { kernel(m, n, k, use_thread_local); });
953 void signal_switch(Index k, Index v = 1) {
954 Index s = state_switch_[k % P].fetch_sub(v);
955 eigen_assert(s >= v);
960 state_switch_[k % P] = (parallel_pack_ ? nm_ + nn_ : (shard_by_col_ ? nn_ : nm_)) + nm_ * nn_;
964 if (parallel_pack_) {
965 enqueue_packing(k, !shard_by_col_);
966 enqueue_packing(k, shard_by_col_);
967 }
else if (shard_by_col_) {
968 enqueue_packing(k,
false);
970 enqueue_packing(k,
true);
978 }
else if (k == nk_) {
979 signal_switch(k + 1, parallel_pack_ ? nm_ + nn_ : (shard_by_col_ ? nn_ : nm_));
986 void enqueue_packing(Index k,
bool rhs) { enqueue_packing_helper(0, rhs ? nn_ : nm_, k, rhs); }
988 void enqueue_packing_helper(Index start, Index end, Index k,
bool rhs) {
989 if (end - start == 1) {
995 while (end - start > 1) {
996 Index mid = (start + end) / 2;
997 device_.enqueue([
this, mid, end, k, rhs]() { enqueue_packing_helper(mid, end, k, rhs); });
1009 bool pack_async = (start == 0) && (parallelize_by_sharding_dim_only_ && shard_by_col_ == rhs) &&
1010 (k > 0 || std::this_thread::get_id() == created_by_thread_id_);
1013 device_.enqueue([
this, start, end, k, rhs]() { enqueue_packing_helper(start, end, k, rhs); });
1015 enqueue_packing_helper(start, end, k, rhs);
1021 Index bm(Index m)
const {
return m + 1 < nm0_ ? bm_ : m_ + bm_ - bm_ * nm0_; }
1022 Index bn(Index n)
const {
return n + 1 < nn0_ ? bn_ : n_ + bn_ - bn_ * nn0_; }
1023 Index bk(Index k)
const {
return k + 1 < nk_ ? bk_ : k_ + bk_ - bk_ * nk_; }
1025 Index gm(Index m)
const {
return m + 1 < nm_ ? gm_ : nm0_ + gm_ - gm_ * nm_; }
1026 Index gn(Index n)
const {
return n + 1 < nn_ ? gn_ : nn0_ + gn_ - gn_ * nn_; }
1028 EvalParallelContext(
const EvalParallelContext&) =
delete;
1029 void operator=(
const EvalParallelContext&) =
delete;
1039 template <
typename DoneCallback>
1040 struct EvalShardedByInnerDimContext {
1041 EvalShardedByInnerDimContext(
const Self* self,
int num_threads, Scalar* result_buffer, Index m_size, Index n_size,
1042 Index k_size, DoneCallback done_callback)
1044 m_lhs_inner_dim_contiguous(evaluator->m_lhs_inner_dim_contiguous),
1045 m_rhs_inner_dim_contiguous(evaluator->m_rhs_inner_dim_contiguous),
1046 m_rhs_inner_dim_reordered(evaluator->m_rhs_inner_dim_reordered),
1047 result(result_buffer),
1051 done(std::move(done_callback)),
1052 buffer_size_bytes(m * n * sizeof(Scalar)),
1053 block_size(blockSize(k, num_threads)),
1054 num_blocks(numext::div_ceil<Index>(k, block_size)),
1055 num_pending_blocks(internal::convert_index<int>(num_blocks)),
1056 l0_ranges(numext::div_ceil<Index>(num_blocks, l0_size)),
1057 l0_state(l0_ranges),
1058 block_buffers(num_blocks) {
1060 for (
int i = 0; i < l0_ranges; ++i) {
1061 const Index num_pending_tasks = actualRangeSize(l0_ranges, l0_size, i);
1062 l0_state.emplace_back(internal::convert_index<int>(num_pending_tasks));
1066 for (Index block_idx = 0; block_idx < num_blocks; ++block_idx) {
1067 Scalar* buf = block_idx == 0 ? result :
static_cast<Scalar*
>(evaluator->m_device.allocate(buffer_size_bytes));
1068 block_buffers.emplace_back(buf);
1072 ~EvalShardedByInnerDimContext() {
1073 for (Index i = 1; i < num_blocks; ++i) {
1074 evaluator->m_device.deallocate(block_buffers[i]);
1078 template <
int Alignment>
1080 Barrier barrier(internal::convert_index<int>(num_blocks));
1081 eval<Alignment>(barrier, 0, num_blocks);
1085 aggregateL0Blocks<Alignment>();
1088 applyOutputKernel();
1091 template <
int Alignment>
1093 evalAsync<Alignment>(0, num_blocks);
1099 static constexpr Index packet_size = internal::packet_traits<RhsScalar>::size;
1101 const Self* evaluator;
1104 bool m_lhs_inner_dim_contiguous;
1105 bool m_rhs_inner_dim_contiguous;
1106 bool m_rhs_inner_dim_reordered;
1120 Index buffer_size_bytes;
1126 std::atomic<int> num_pending_blocks;
1144 static constexpr Index l0_size = 4;
1148 MaxSizeVector<std::atomic<int>> l0_state;
1151 MaxSizeVector<Scalar*> block_buffers;
1153 template <
int Alignment>
1154 void processBlock(Index block_idx, Index begin, Index end) {
1155 Scalar* buf = block_buffers[block_idx];
1157 internal::tensor_contraction_dispatch(
1158 [&](
auto lhs_c,
auto rhs_c,
auto rhs_r) {
1159 evaluator->template evalGemmPartialWithoutOutputKernel<lhs_c(), rhs_c(), rhs_r(), Alignment>(
1160 buf, begin, end, internal::convert_index<int>(num_blocks));
1162 m_lhs_inner_dim_contiguous, m_rhs_inner_dim_contiguous, m_rhs_inner_dim_reordered);
1165 const Index l0_index = block_idx / l0_size;
1166 const int v = l0_state[l0_index].fetch_sub(1);
1167 eigen_assert(v >= 1);
1172 const Index rng_size = actualRangeSize(l0_ranges, l0_size, l0_index);
1173 const Index dst_block_idx = l0_index * l0_size;
1175 if (rng_size == l0_size) {
1176 addAllToBuffer<Alignment>(m * n,
1177 block_buffers[dst_block_idx + 1],
1178 block_buffers[dst_block_idx + 2],
1179 block_buffers[dst_block_idx + 3],
1180 block_buffers[dst_block_idx]);
1183 for (
int i = 1; i < rng_size; ++i) {
1184 addToBuffer<Alignment>(m * n,
1185 block_buffers[dst_block_idx + i],
1186 block_buffers[dst_block_idx]);
1193 template <
int Alignment>
1194 void aggregateL0Blocks()
const {
1197 for (; l0_index + 2 < l0_ranges; l0_index += 3) {
1198 addAllToBuffer<Alignment>(m * n,
1199 block_buffers[(l0_index + 0) * l0_size],
1200 block_buffers[(l0_index + 1) * l0_size],
1201 block_buffers[(l0_index + 2) * l0_size],
1205 for (; l0_index < l0_ranges; ++l0_index) {
1206 addToBuffer<Alignment>(m * n, block_buffers[l0_index * l0_size], block_buffers[0]);
1210 void applyOutputKernel()
const {
1211 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
1212 evaluator->m_output_kernel(OutputMapper(result, m), evaluator->m_tensor_contraction_params,
1213 static_cast<Eigen::Index
>(0),
static_cast<Eigen::Index
>(0), m, n);
1217 Index actualBlockSize(Index block_idx)
const {
1218 return block_idx + 1 < num_blocks ? block_size : k + block_size - block_size * num_blocks;
1222 Index actualRangeSize(Index num_ranges, Index range_size, Index range_idx)
const {
1223 eigen_assert(range_idx < num_ranges);
1224 return range_idx + 1 < num_ranges ? range_size : num_blocks + range_size - range_size * num_ranges;
1227 template <
int Alignment>
1228 EIGEN_STRONG_INLINE
static void addToBuffer(
size_t n,
const Scalar* src_buf, Scalar* tgt_buf) {
1229 const int output_packet_size = internal::unpacket_traits<PacketReturnType>::size;
1231 const size_t num_packets = n / output_packet_size;
1232 for (; i < output_packet_size * num_packets; i += output_packet_size) {
1233 const PacketReturnType src_val = internal::pload<PacketReturnType>(src_buf + i);
1234 const PacketReturnType tgt_val = internal::ploadt<PacketReturnType, Alignment>(tgt_buf + i);
1235 const PacketReturnType sum = internal::padd(src_val, tgt_val);
1236 internal::pstoret<Scalar, PacketReturnType, Alignment>(tgt_buf + i, sum);
1238 for (; i < n; ++i) {
1239 tgt_buf[i] += src_buf[i];
1243 template <
int Alignment>
1244 EIGEN_STRONG_INLINE
static void addAllToBuffer(
size_t n,
const Scalar* src_buf0,
const Scalar* src_buf1,
1245 const Scalar* src_buf2, Scalar* dst_buf) {
1246 using ::Eigen::internal::padd;
1247 using ::Eigen::internal::pload;
1248 using ::Eigen::internal::ploadt;
1249 using ::Eigen::internal::pstoret;
1251 const int output_packet_size = internal::unpacket_traits<PacketReturnType>::size;
1254 const size_t num_packets = n / output_packet_size;
1255 for (; i < output_packet_size * num_packets; i += output_packet_size) {
1256 const auto src_val0 = pload<PacketReturnType>(src_buf0 + i);
1257 const auto src_val1 = pload<PacketReturnType>(src_buf1 + i);
1258 const auto src_val2 = pload<PacketReturnType>(src_buf2 + i);
1260 const auto dst_val = ploadt<PacketReturnType, Alignment>(dst_buf + i);
1261 const auto sum = padd(padd(dst_val, src_val0), padd(src_val1, src_val2));
1263 pstoret<Scalar, PacketReturnType, Alignment>(dst_buf + i, sum);
1265 for (; i < n; ++i) {
1266 dst_buf[i] += src_buf0[i] + src_buf1[i] + src_buf2[i];
1270 template <
int Alignment>
1271 void eval(Barrier& barrier, Index start_block_idx, Index end_block_idx) {
1272 while (end_block_idx - start_block_idx > 1) {
1273 Index mid_block_idx = (start_block_idx + end_block_idx) / 2;
1274 evaluator->m_device.enqueue([
this, &barrier, mid_block_idx, end_block_idx]() {
1275 eval<Alignment>(barrier, mid_block_idx, end_block_idx);
1277 end_block_idx = mid_block_idx;
1280 Index block_idx = start_block_idx;
1281 Index block_start = block_idx * block_size;
1282 Index block_end = block_start + actualBlockSize(block_idx);
1284 processBlock<Alignment>(block_idx, block_start, block_end);
1288 template <
int Alignment>
1289 void evalAsync(Index start_block_idx, Index end_block_idx) {
1290 while (end_block_idx - start_block_idx > 1) {
1291 Index mid_block_idx = (start_block_idx + end_block_idx) / 2;
1292 evaluator->m_device.enqueue(
1293 [
this, mid_block_idx, end_block_idx]() { evalAsync<Alignment>(mid_block_idx, end_block_idx); });
1294 end_block_idx = mid_block_idx;
1297 Index block_idx = start_block_idx;
1299 Index block_start = block_idx * block_size;
1300 Index block_end = block_start + actualBlockSize(block_idx);
1302 processBlock<Alignment>(block_idx, block_start, block_end);
1304 int v = num_pending_blocks.fetch_sub(1);
1305 eigen_assert(v >= 1);
1309 aggregateL0Blocks<Alignment>();
1312 applyOutputKernel();
1319 DoneCallback done_copy = std::move(done);
1332 static Index blockSize(Index k,
int num_threads) {
1333 const auto round_up = [=](Index index) -> Index {
1334 const Index kmultiple = packet_size <= 8 ? 8 : packet_size;
1335 return numext::div_ceil<Index>(index, kmultiple) * kmultiple;
1338 const Index target_block_size = round_up(numext::div_ceil<Index>(k, num_threads));
1339 const Index desired_min_block_size = 12 * packet_size;
1341 return numext::mini<Index>(k, numext::maxi<Index>(desired_min_block_size, target_block_size));
1344 EvalShardedByInnerDimContext(
const EvalShardedByInnerDimContext&) =
delete;
1345 void operator=(
const EvalShardedByInnerDimContext&) =
delete;
1354 static bool shardByCol(Index m, Index n, Index num_threads) {
1361 if (m / num_threads >= Traits::nr &&
1363 (n / num_threads < Traits::nr ||
1366 (n / num_threads < 4 * Traits::nr && (n % (num_threads * Traits::nr)) != 0 &&
1368 ((m % (num_threads * Traits::nr)) == 0 ||
1376 if (n / num_threads < 16 * Traits::nr && m > n * 32)
return false;
1380 Index coarsenM(Index m, Index n, Index bm, Index bn, Index bk, Index gn,
int num_threads,
bool shard_by_col)
const {
1383 Index nm0 = numext::div_ceil(m, bm);
1389 while (gm1 <= nm0 && nm1 == numext::div_ceil(nm0, gm1)) gm1++;
1390 if (gm1 > nm0)
break;
1392 int res = checkGrain(m, n, bm, bn, bk, gm1, gn, gm, gn, num_threads, shard_by_col);
1394 nm1 = numext::div_ceil(nm0, gm1);
1395 if (res == 0)
continue;
1402 Index coarsenN(Index m, Index n, Index bm, Index bn, Index bk, Index gm,
int num_threads,
bool shard_by_col)
const {
1405 Index nn0 = numext::div_ceil(n, bn);
1408 while (gn1 <= nn0 && nn1 == numext::div_ceil(nn0, gn1)) gn1++;
1409 if (gn1 > nn0)
break;
1410 int res = checkGrain(m, n, bm, bn, bk, gm, gn1, gm, gn, num_threads, shard_by_col);
1412 nn1 = numext::div_ceil(nn0, gn1);
1413 if (res == 0)
continue;
1421 int checkGrain(Index m, Index n, Index bm, Index bn, Index bk, Index gm, Index gn, Index oldgm, Index oldgn,
1422 int num_threads,
bool shard_by_col)
const {
1423 const TensorOpCost cost = contractionCost(bm * gm, bn * gn, bm, bn, bk, shard_by_col,
true);
1424 double taskSize = TensorCostModel<ThreadPoolDevice>::taskSize(
static_cast<double>(bm) * gm * bn * gn, cost);
1427 if (taskSize < 1)
return 1;
1429 if (taskSize > 2)
return -1;
1436 Index nm0 = numext::div_ceil(m, bm);
1437 Index nn0 = numext::div_ceil(n, bn);
1438 Index new_tasks = numext::div_ceil(nm0, gm) * numext::div_ceil(nn0, gn);
1439 double new_parallelism =
1440 static_cast<double>(new_tasks) / (numext::div_ceil<Index>(new_tasks, num_threads) * num_threads);
1441 Index old_tasks = numext::div_ceil(nm0, oldgm) * numext::div_ceil(nn0, oldgn);
1442 double old_parallelism =
1443 static_cast<double>(old_tasks) / (numext::div_ceil<Index>(old_tasks, num_threads) * num_threads);
1444 if (new_parallelism > old_parallelism || new_parallelism == 1)
return 1;
1448 TensorOpCost contractionCost(Index m, Index n, Index bm, Index bn, Index bk,
bool shard_by_col,
1449 bool prepacked)
const {
1450 const int packed_size = std::min<int>(PacketType<LhsScalar, Device>::size, PacketType<RhsScalar, Device>::size);
1451 const int output_packet_size = internal::unpacket_traits<PacketReturnType>::size;
1452 const double kd =
static_cast<double>(bk);
1453 double compute_bandwidth = computeBandwidth(
false, bm, bn, bk);
1455 TensorOpCost cost = TensorOpCost(0, 0, kd * compute_bandwidth,
true, packed_size);
1457 cost += TensorOpCost(0,
sizeof(CoeffReturnType), 0,
true, output_packet_size);
1465 TensorOpCost lhsCost = this->m_leftImpl.costPerCoeff(
true) * (kd / n);
1466 TensorOpCost rhsCost = this->m_rightImpl.costPerCoeff(
true) * (kd / m);
1470 lhsCost.dropMemoryCost();
1472 rhsCost.dropMemoryCost();
1473 return cost + lhsCost + rhsCost;
1478 static bool shardByInnerDim(Index m, Index n, Index k,
int num_threads,
int num_threads_by_k) {
1479 std::ptrdiff_t bufsize = m * n *
sizeof(Scalar);
1480 bool shard_by_k =
false;
1482 num_threads_by_k < 2 ||
1483 num_threads_by_k < num_threads ||
1484 bufsize > l3CacheSize() / num_threads_by_k ||
1486 k / num_threads_by_k < 2 * Traits::nr) {
1488 }
else if (numext::maxi(m, n) / num_threads < Traits::nr ||
1490 (k / num_threads_by_k > 8 * Traits::nr &&
1493 (numext::mini(m, n) < 2 * Traits::nr || num_threads_by_k > num_threads))) {
1499 TensorOpCost contractionCostPerInnerDim(Index m, Index n, Index k)
const {
1501 const int output_packet_size = internal::unpacket_traits<PacketReturnType>::size;
1502 TensorOpCost cost(0, 0, (computeBandwidth(
true, m, n, k) * m) * n,
true, output_packet_size);
1504 cost += TensorOpCost(0,
sizeof(CoeffReturnType), 0,
true, output_packet_size);
1505 TensorOpCost lhsCost = this->m_leftImpl.costPerCoeff(
true) * m;
1506 TensorOpCost rhsCost = this->m_rightImpl.costPerCoeff(
true) * n;
1509 lhsCost.dropMemoryCost();
1510 return cost + lhsCost + rhsCost;
1513 int numThreadsInnerDim(Index m, Index n, Index k)
const {
1514 const int output_packet_size = internal::unpacket_traits<PacketReturnType>::size;
1515 TensorOpCost cost = contractionCostPerInnerDim(m, n, k);
1516 double total_parallel_cost = TensorCostModel<ThreadPoolDevice>::totalCost(k, cost);
1519 double reduction_cost =
1520 TensorCostModel<ThreadPoolDevice>::totalCost(m * n, TensorOpCost(2, 1, 1,
true, output_packet_size));
1521 int num_threads = 1;
1522 double min_cost = total_parallel_cost;
1523 double kPerThreadOverHead = 3000;
1524 double kFixedOverHead = 20000;
1525 for (
int nt = 2; nt <= this->m_device.numThreads(); nt += 2) {
1526 double sequential_cost = kFixedOverHead + nt * (reduction_cost + kPerThreadOverHead);
1527 double parallel_cost = total_parallel_cost / nt + sequential_cost;
1528 if (parallel_cost < min_cost) {
1530 min_cost = parallel_cost;
1536 double computeBandwidth(
bool shard_by_col, Index bm, Index bn, Index bk)
const {
1540 double computeBandwidth = bk == 1 ? 4.0
1541 : (shard_by_col ? bn : bm) < Traits::nr || (shard_by_col ? bm : bn) < Traits::mr ? 2.0
1543#ifndef EIGEN_VECTORIZE_FMA
1548 if (computeBandwidth == 0.5) computeBandwidth = 1.0;
1550 return computeBandwidth;
Definition TensorContraction.h:335
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47