Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorContractionThreadPool.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_THREAD_POOL_H
12#define EIGEN_TENSOR_TENSOR_CONTRACTION_THREAD_POOL_H
13
14// evaluator for thread pool device
15#ifdef EIGEN_USE_THREADS
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20namespace Eigen {
21
22template <typename Indices, typename LeftArgType, typename RightArgType, typename OutputKernelType>
23struct TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>,
24 ThreadPoolDevice>
25 : public TensorContractionEvaluatorBase<TensorEvaluator<
26 const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, ThreadPoolDevice>> {
27 using Device = ThreadPoolDevice;
28
29 using Self = TensorEvaluator<const TensorContractionOp<Indices, LeftArgType, RightArgType, OutputKernelType>, Device>;
30 using Base = TensorContractionEvaluatorBase<Self>;
31
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;
37
38 static constexpr int Layout = TensorEvaluator<LeftArgType, Device>::Layout;
39
40 // Most of the code is assuming that both input tensors are ColMajor. If the
41 // inputs are RowMajor, we will "cheat" by swapping the LHS and RHS:
42 // If we want to compute A * B = C, where A is LHS and B is RHS, the code
43 // will pretend B is LHS and A is RHS.
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>;
48
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;
54
55 using left_dim_mapper_t = array<Index, LDims>;
56 using right_dim_mapper_t = array<Index, RDims>;
57
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>;
61
62 static constexpr int NumDims = LDims + RDims - 2 * ContractDims;
63
64 using Dimensions = DSizes<Index, NumDims>;
65
66 // typedefs needed in evalTo
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>;
70
71 using LeftEvaluator = TensorEvaluator<EvalLeftArgType, Device>;
72 using RightEvaluator = TensorEvaluator<EvalRightArgType, Device>;
73
74 TensorEvaluator(const XprType& op, const Device& device) : Base(op, device) {}
75
76 template <int Alignment>
77 void evalProduct(Scalar* buffer) const {
78 evalProductImpl<NoCallback, Alignment>(buffer, NoCallback());
79 }
80
81 template <typename EvalToCallback, int Alignment>
82 void evalProductAsync(Scalar* buffer, EvalToCallback done) const {
83 evalProductImpl<EvalToCallback, Alignment>(buffer, std::move(done));
84 }
85
86 template <typename DoneCallback, int Alignment>
87 void evalProductImpl(Scalar* buffer, DoneCallback done) const {
88 // This function computes a lot of heuristics in multiple steps, and it
89 // also has multiple exit points. To keep it sane, readable and all in one
90 // place, sync/async execution decision is made at runtime at the very end.
91 //
92 // (1) In sync mode we allocate Context on the stack, submit computations
93 // to the device thread pool, and block on a barrier until it is
94 // completed.
95 //
96 // (2) In async mode we allocate Context on the heap, and after all tasks
97 // are finished, we call the provided done callback, and delete a
98 // context from the heap.
99 //
100 // (*) EvalParallelContext & EvalShardedByInnerDimContext owns all the state
101 // and temporary buffers, required for executing the tensor contraction.
102 // They are responsible for cleaning it up after contraction is done.
103 static constexpr bool IsEvalInSyncMode = std::is_same<DoneCallback, NoCallback>::value;
104
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();
110 return;
111 }
112 if (k == 0) {
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);
116 },
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();
119 return;
120 }
121
122 // Compute a set of algorithm parameters:
123 // - kernel block sizes (bm, bn, bk)
124 // - task grain sizes (number of kernels executed per task: gm, gn)
125 // - number of threads
126 // - sharding by row/column
127 // - parallel packing or first lhs then rhs
128 // and some derived parameters:
129 // - number of tasks (nm, nn, nk)
130 // - number of kernels (nm0, nn0)
131 // Unfortunately, all these parameters are tightly interdependent.
132 // So in some cases we first compute approximate values, then compute other
133 // values based on these approximations and then refine the approximations.
134
135 // There are lots of heuristics here. There is some reasoning behind them,
136 // but ultimately they are just tuned on contraction benchmarks for
137 // different input configurations, thread counts and instruction sets.
138 // So feel free to question any of them.
139
140 // Compute whether we want to shard by row or by column.
141 // This is a first approximation, it will be refined later. Since we don't
142 // know number of threads yet we use 2, because what we are most
143 // interested in at this point is whether it makes sense to use
144 // parallelization at all or not.
145 bool shard_by_col = shardByCol(m, n, 2);
146
147 // First approximation of kernel blocking sizes.
148 // Again, we don't know number of threads yet, so we use 2.
149 Index bm, bn, bk;
150 if (shard_by_col) {
151 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByCol> blocking(k, m, n,
152 2);
153 bm = blocking.mc();
154 bn = blocking.nc();
155 bk = blocking.kc();
156 } else {
157 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByRow> blocking(k, m, n,
158 2);
159 bm = blocking.mc();
160 bn = blocking.nc();
161 bk = blocking.kc();
162 }
163
164 // Compute optimal number of threads.
165 // Note: we use bk instead of k here because we are interested in amount of
166 // _parallelizable_ computations, and computations are not parallelizable
167 // across k dimension.
168 const TensorOpCost cost = contractionCost(m, n, bm, bn, bk, shard_by_col, false);
169 int num_threads =
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)) {
173 // We are in the scenario where it is more effective to shard by the
174 // inner dimension.
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>();
178 } else {
179 auto* ctx =
180 new EvalShardedByInnerDimContext<DoneCallback>(this, num_threads_by_k, buffer, m, n, k, std::move(done));
181 ctx->template runAsync<Alignment>();
182 }
183
184 return;
185 }
186
187 // TODO(dvyukov): this is a stop-gap to prevent regressions while the cost
188 // model is not tuned. Remove this when the cost model is tuned.
189 if (n == 1) num_threads = 1;
190
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);
195 },
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();
198 return;
199 }
200
201 // Now that we know number of threads, recalculate sharding and blocking.
202 shard_by_col = shardByCol(m, n, num_threads);
203 if (shard_by_col) {
204 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByCol> blocking(
205 k, m, n, num_threads);
206 bm = blocking.mc();
207 bn = blocking.nc();
208 bk = blocking.kc();
209 } else {
210 internal::TensorContractionBlocking<Scalar, LhsScalar, RhsScalar, Index, internal::ShardByRow> blocking(
211 k, m, n, num_threads);
212 bm = blocking.mc();
213 bn = blocking.nc();
214 bk = blocking.kc();
215 }
216
217 // Number of kernels for each dimension.
218 Index nm0 = numext::div_ceil(m, bm);
219 Index nn0 = numext::div_ceil(n, bn);
220 Index nk = numext::div_ceil(k, bk);
221
222 // Calculate task grain size (number of kernels executed per task).
223 // This task size coarsening serves two purposes:
224 // 1. It reduces per-task overheads including synchronization overheads.
225 // 2. It allows to use caches better (reuse the same packed rhs in several
226 // consecutive kernels).
227 Index gm = 1;
228 Index gn = 1;
229 // If we are sharding by column, then we prefer to reduce rows first.
230 if (shard_by_col) {
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);
233 } else {
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);
236 }
237 // Number of tasks in each dimension.
238 Index nm = numext::div_ceil(nm0, gm);
239 Index nn = numext::div_ceil(nn0, gn);
240
241 // If there is enough concurrency in the sharding dimension, we choose not
242 // to parallelize by the other dimension, and execute all kernels in sync
243 // mode. This reduces parallelism from the nm x nn down to nn
244 // (shard_by_col==true) or nm (shard_by_col==false).
245 const Index sharding_dim_tasks = shard_by_col ? nn : nm;
246 const int num_worker_threads = this->m_device.numThreadsInPool();
247
248 // With small number of threads we want to make sure that we do not reduce
249 // parallelism too much. With large number of threads we trade maximum
250 // parallelism for better memory locality.
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
256 : /* num_worker_threads > 64 */ 0.6;
257
258 const bool parallelize_by_sharding_dim_only = sharding_dim_tasks >= oversharding_factor * num_worker_threads;
259
260 // Last by not least, decide whether we want to issue both lhs and rhs
261 // packing in parallel; or issue lhs packing first, and then issue rhs
262 // packing when lhs packing completes (for !shard_by_col lhs and rhs are
263 // swapped). Parallel packing allows more parallelism (for both packing and
264 // kernels), while sequential packing provides better locality (once
265 // a thread finishes rhs packing it proceed to kernels with that rhs).
266 // First, we are interested in parallel packing if there are few tasks.
267 bool parallel_pack = num_threads >= nm * nn;
268 // Also do parallel packing if all data fits into L2$.
269 if (m * bk * Index(sizeof(LhsScalar)) + n * bk * Index(sizeof(RhsScalar)) <= l2CacheSize() * num_threads)
270 parallel_pack = true;
271 // But don't do it if we will use each rhs only once. Locality seems to be
272 // more important in this case.
273 if ((shard_by_col ? nm : nn) == 1) parallel_pack = false;
274 // Also don't get in the way of parallelize_by_sharding_dim_only
275 // optimization.
276 if (parallelize_by_sharding_dim_only) parallel_pack = false;
277
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());
284 ctx.run();
285 } else {
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));
289 ctx->run();
290 }
291 },
292 this->m_lhs_inner_dim_contiguous, this->m_rhs_inner_dim_contiguous, this->m_rhs_inner_dim_reordered);
293 }
294
295 // ------------------------------------------------------------------------ //
296
297 // Dummy struct to represent an empty DoneCallback.
298
299 struct NoCallback {
300 void operator()() const { eigen_assert(false && "NoCallback should never be called"); }
301 };
302
303 // ------------------------------------------------------------------------ //
304
305 template <typename DoneCallback, typename Context>
306 class EvalParallelNotification;
307
308 // Synchronous evaluation notification that blocks caller thread in Wait().
309 template <typename Context>
310 class EvalParallelNotification<NoCallback, Context> {
311 public:
312 EvalParallelNotification(Context*, NoCallback) {}
313 void Notify() { done_.Notify(); }
314 void Wait() { done_.Wait(); }
315
316 private:
317 Eigen::Notification done_;
318 };
319
320 // Asynchronous evaluation notification that does not block in Wait().
321 template <typename DoneCallback, typename Context>
322 class EvalParallelNotification {
323 public:
324 EvalParallelNotification(Context* ctx, DoneCallback done) : ctx_(ctx), done_(std::move(done)) {}
325
326 void Notify() {
327 // Make a copy of done callback, because it will be destructed when we
328 // will delete context in the next line (EvalParallelNotification is a
329 // data member of EvalParallelContext class).
330 DoneCallback done_copy = std::move(done_);
331
332 // Delete parallel evaluation context.
333 delete ctx_;
334
335 // Now safely call the done callback.
336 done_copy();
337 }
338
339 void Wait() {}
340
341 private:
342 Context* ctx_;
343 DoneCallback done_;
344 };
345
346 // Context orchestrates sync/async parallel contraction evaluation. When it is
347 // executed in asynchronous mode, it owns all the shared state that might be
348 // accessible by block packing and kernel tasks.
349
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 {
353 public:
354 using LhsMapper =
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>;
358 using RhsMapper =
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>;
362
363 using OutputMapper = internal::blas_data_mapper<Scalar, Index, ColMajor>;
364
365 using TensorContractionKernel =
366 internal::TensorContractionKernel<Scalar, LhsScalar, RhsScalar, Index, OutputMapper, LhsMapper, RhsMapper>;
367
368 using LhsBlock = typename TensorContractionKernel::LhsBlock;
369 using RhsBlock = typename TensorContractionKernel::RhsBlock;
370 using BlockMemHandle = typename TensorContractionKernel::BlockMemHandle;
371
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,
379 self->m_k_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),
382 buffer_(buffer),
383 output_(buffer, tm),
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),
390 m_(tm),
391 n_(tn),
392 k_(tk),
393 bm_(bm),
394 bn_(bn),
395 bk_(bk),
396 nm_(nm),
397 nn_(nn),
398 nk_(nk),
399 gm_(gm),
400 gn_(gn),
401 nm0_(nm0),
402 nn0_(nn0),
403 kernel_(m_, k_, n_, bm_, bk_, bn_),
404 num_thread_local_allocations_(0),
405 // We reserve 2X more capacity for a thread local values, than the
406 // number of threads in the pool to efficiently handle task stealing
407 // by threads that are not managed by the pool.
408 thread_local_capacity(2 * (parallelize_by_sharding_dim_only_ ? device_.numThreadsInPool() : 0)),
409 // We will use only one of the Lhs/Rhs thread local storage depending
410 // on the shard_by_col value and we parallelize by sharding dim ONLY.
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}) {
413 // These two options are mutually exclusive.
414 eigen_assert(!(parallel_pack && parallelize_by_sharding_dim_only));
415
416 for (Index x = 0; x < P; x++) {
417 // Normal number of notifications for k slice switch is
418 // nm_ + nn_ + nm_ * nn_. However, first P - 1 slices will receive only
419 // nm_ + nn_ notifications, because they will not receive notifications
420 // from preceding kernels.
421 state_switch_[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_];
427 // Kernels generally receive 3 notifications (previous kernel + 2
428 // packing), but the first slice won't get notifications from previous
429 // kernels.
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);
432 }
433 }
434
435 // Allocate memory for packed rhs/lhs matrices.
436 packed_mem_ = kernel_.allocateSlices( //
437 device_, //
438 /*num_lhs=*/nm0_, //
439 /*num_rhs=*/nn0_, //
440 /*num_slices=*/std::min<Index>(nk_, P - 1), //
441 packed_lhs_, packed_rhs_);
442
443 if (parallelize_by_sharding_dim_only_) {
444 const int num_worker_threads = device_.numThreadsInPool();
445
446 if (shard_by_col) {
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);
449
450 Index num_blocks = num_worker_threads * gn_;
451 thread_local_pre_allocated_mem_ = kernel_.allocateSlices( //
452 device_, //
453 /*num_lhs=*/0, //
454 /*num_rhs=*/num_blocks, //
455 /*num_slices=*/1, //
456 /*lhs_blocks=*/nullptr, &rhs_thread_local_pre_allocated_);
457
458 } else {
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);
461
462 Index num_blocks = num_worker_threads * gm_;
463 thread_local_pre_allocated_mem_ = kernel_.allocateSlices( //
464 device_, //
465 /*num_lhs=*/num_blocks, //
466 /*num_rhs=*/0, //
467 /*num_slices=*/1, &lhs_thread_local_pre_allocated_, //
468 /*rhs_blocks=*/nullptr);
469 }
470 }
471 }
472
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];
477 }
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_;
482 }
483 }
484
485 void run() {
486 // Kick off packing of the first slice.
487 signal_switch(0, 1);
488
489 // Wait for overall completion.
490 //
491 // If parallel evaluation is executed in async mode, this is a no-op, and
492 // Wait() will return immediately. In synchronous mode it will block the
493 // caller thread until it will receive notification from last task.
494 //
495 // In async mode, last task when completed will call done callback from
496 // the same thread, and will delete this context.
497 //
498 // TODO(dvyukov): This wait can lead to deadlock if contraction is
499 // evaluated in synchronous mode. If nthreads contractions are
500 // concurrently submitted from worker threads, this wait will block all
501 // worker threads and the system will deadlock.
502 done_.Wait();
503 }
504
505 private:
506 std::thread::id created_by_thread_id_;
507
508 // This notification is specialized on the type of DoneCallback and can be
509 // blocking or non-blocking.
510 EvalParallelNotification<DoneCallback, EvalParallelContext> done_;
511
512 const Device& device_;
513 LhsMapper lhs_;
514 RhsMapper rhs_;
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_;
523 // Matrix sizes.
524 const Index m_;
525 const Index n_;
526 const Index k_;
527 // Block sizes.
528 const Index bm_;
529 const Index bn_;
530 const Index bk_;
531 // Number of tasks.
532 const Index nm_;
533 const Index nn_;
534 const Index nk_;
535 // Task grain sizes (number of kernels executed per task).
536 const Index gm_;
537 const Index gn_;
538 // Number of blocks (this is different from nm_/nn_ because of task size
539 // coarsening).
540 const Index nm0_;
541 const Index nn0_;
542 // Tensor contraction kernel.
543 TensorContractionKernel kernel_;
544
545 // Parallelization strategy.
546 //
547 // Blocks related to the same k block can run in parallel because they write
548 // to different output blocks. So we parallelize within k slices, this
549 // gives us parallelism level of m x n. Before we can start any kernels
550 // related to k-th slice, we need to issue m lhs packing tasks and n rhs
551 // packing tasks.
552 //
553 // However, there is a bottleneck when we are finishing kernels for k-th
554 // slice (at the very end there is only 1 runnable kernel). To mitigate this
555 // bottleneck we allow kernels from k-th and k+1-th slices to run in
556 // parallel. Note that (m, n, k) and (m, n, k+1) kernels write to the same
557 // output block, so they must not run in parallel.
558 //
559 // This gives us the following dependency graph.
560 // On each k slice we have m x n kernel tasks, m lhs packing tasks and n rhs
561 // packing tasks.
562 // Kernel (m, n, k) can start when:
563 // - kernel (m, n, k-1) has finished
564 // - lhs packing (m, k) has finished
565 // - rhs packing (n, k) has finished
566 // Lhs/rhs packing can start when:
567 // - all k-1 packing has finished (artificially imposed to limit amount of
568 // parallel packing)
569 //
570 // On top of that we limit runnable tasks to two consecutive k slices.
571 // This is done to limit amount of memory we need for packed lhs/rhs
572 // (for each k slice we need m*bk + n*bk memory in packed_lhs_/packed_rhs_).
573 //
574 // state_switch_ tracks when we are ready to switch to the next k slice.
575 // state_kernel_[m][n] tracks when we are ready to kick off kernel (m, n).
576 // These variable are rolling over 3 consecutive k slices: first two we are
577 // actively executing + one to track completion of kernels in the second
578 // slice.
579 static constexpr Index P = 3;
580
581 // Handle to the allocated temporary storage for Lhs/Rhs blocks.
582 BlockMemHandle packed_mem_;
583 std::vector<LhsBlock> packed_lhs_[P - 1];
584 std::vector<RhsBlock> packed_rhs_[P - 1];
585
586 // If we choose to parallelize only by the sharding dimension, each thread
587 // will have its own "thread local" (not a C++ thread local storage) memory
588 // for packed_lhs or packed_rhs (shard_by_col = false of true). This memory
589 // can't be passed to a kernel that might execute on a different thread.
590 //
591 // In practice when we are ready to pack memory for the sharding dimension
592 // (rhs if shard_by_col==true) of the K-th slice, all kernels for K-1 slice
593 // already computed (99% of the time), and we can pack data into the thread
594 // local storage, and guarantee that all the kernels will be executed
595 // immediately in the same thread. This significantly increases L1 cache hit
596 // ratio and reduces pressure on the memory bus.
597 //
598 // It's still possible that kernel for the K-th slice will be ready before
599 // completion of the K-1 kernel, so we have to allocate "global" packed_lhs_
600 // and packed_rhs_ to allow kernels to be executed later on a thread
601 // different from the thread that was used for packing.
602
603 // Handle for pre-allocated thread local memory buffers.
604 BlockMemHandle thread_local_pre_allocated_mem_;
605
606 // Only one of these will be initialized depending on shard_by_col value
607 // (the size will be `num_worker_threads * num_grains_in_the_sharding_dim`).
608 std::vector<LhsBlock> lhs_thread_local_pre_allocated_;
609 std::vector<RhsBlock> rhs_thread_local_pre_allocated_;
610
611 // How many thread local blocks were already allocated.
612 std::atomic<int> num_thread_local_allocations_;
613 const int thread_local_capacity;
614
615 // We will use pre-allocated Lhs/Rhs blocks defined above, if the number of
616 // unique threads in a system is below or equal to the number of threads in
617 // a thread pool. We will fallback on dynamic memory allocation after that.
618
619 // ThreadLocalBlocks is a container for Lhs or Rhs thread local buffers. Its
620 // size is equal to the grain size in Lhs/Rhs sharding dimension.
621 template <typename BlockType>
622 class ThreadLocalBlocks {
623 public:
624 ThreadLocalBlocks() = default;
625
626 ThreadLocalBlocks(BlockType* base, size_t grain_size)
627 : is_pre_allocated_(true), thread_local_pre_allocated_base_(base), grain_size_(grain_size) {}
628
629 ThreadLocalBlocks(BlockMemHandle mem_handle, std::vector<BlockType> blocks)
630 : is_pre_allocated_(false), mem_handle_(std::move(mem_handle)), blocks_(std::move(blocks)) {}
631
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];
636 }
637
638 void Release(EvalParallelContext& ctx) const {
639 if (!is_pre_allocated_) {
640 ctx.kernel_.deallocate(ctx.device_, mem_handle_);
641 }
642 }
643
644 size_t size() const { return is_pre_allocated_ ? grain_size_ : blocks_.size(); }
645
646 private:
647 bool is_pre_allocated_;
648
649 // Reuse pre-allocated thread local buffers.
650 BlockType* thread_local_pre_allocated_base_ = nullptr;
651 size_t grain_size_ = 0;
652
653 // These will be initialized only if `is_pre_allocated == false`.
654 BlockMemHandle mem_handle_{};
655 std::vector<BlockType> blocks_;
656 };
657
658 // ThreadLocalBlocksInitialize callable does custom thread local blocks
659 // initialization, and will reuse pre-allocated buffers if possible, or will
660 // dynamically allocate new memory.
661 //
662 // Lhs/Rhs blocks might be of the same type, so we have to pass explicitly
663 // for what side do we plan to do block allocation.
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");
669
670 using Blocks = ThreadLocalBlocks<BlockType>;
671
672 public:
673 ThreadLocalBlocksInitialize(EvalParallelContext& ctx)
674 : ctx_(ctx), num_worker_threads_(ctx_.device_.numThreadsInPool()) {}
675
676 void operator()(Blocks& blocks) {
677 const int n = ctx_.num_thread_local_allocations_.fetch_add(1, std::memory_order_relaxed);
678
679 if (n >= num_worker_threads_) {
680 ThreadLocalBlocksAllocator<is_rhs>::allocate(ctx_, blocks);
681 } else {
682 ThreadLocalBlocksAllocator<is_rhs>::reuse(ctx_, n, blocks);
683 }
684 }
685
686 private:
687 // Explicit specializations are not allowed at class scope, so EvalCtx is
688 // a dummy template parameter to make these partial specializations.
689 template <bool pack_rhs, typename EvalCtx = EvalParallelContext>
690 struct ThreadLocalBlocksAllocator;
691
692 template <typename EvalCtx>
693 struct ThreadLocalBlocksAllocator</*pack_rhs=*/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_,
697 /*num_lhs=*/0,
698 /*num_rhs=*/ctx.gn_,
699 /*num_slices=*/1,
700 /*lhs_blocks=*/nullptr, /*rhs_blocks=*/&rhs_blocks);
701 blocks = ThreadLocalBlocks<RhsBlock>(std::move(mem_handle), std::move(rhs_blocks));
702 }
703
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_);
707 }
708 };
709
710 template <typename EvalCtx>
711 struct ThreadLocalBlocksAllocator</*pack_rhs=*/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_,
715 /*num_lhs=*/ctx.gm_,
716 /*num_rhs=*/0,
717 /*num_slices=*/1,
718 /*lhs_blocks=*/&lhs_blocks, /*rhs_blocks=*/nullptr);
719 blocks = ThreadLocalBlocks<LhsBlock>(std::move(mem_handle), std::move(lhs_blocks));
720 }
721
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_);
725 }
726 };
727
728 EvalParallelContext& ctx_;
729 const int num_worker_threads_;
730 };
731
732 template <typename BlockType>
733 class ThreadLocalBlocksRelease {
734 public:
735 using Blocks = ThreadLocalBlocks<BlockType>;
736 ThreadLocalBlocksRelease(EvalParallelContext& ctx) : ctx_(ctx) {}
737 void operator()(Blocks& blocks) { blocks.Release(ctx_); }
738
739 private:
740 EvalParallelContext& ctx_;
741 };
742
743 // ThreadLocalBlocks initialization callables.
744 using ThreadLocalLhsInit = ThreadLocalBlocksInitialize<LhsBlock, /*is_rhs=*/false>;
745 using ThreadLocalRhsInit = ThreadLocalBlocksInitialize<RhsBlock, /*is_rhs=*/true>;
746
747 // ThreadLocalBlocks release callables.
748 using ThreadLocalLhsRelease = ThreadLocalBlocksRelease<LhsBlock>;
749 using ThreadLocalRhsRelease = ThreadLocalBlocksRelease<RhsBlock>;
750
751 // Thread local containers for Lhs/Rhs block packs. In practice only one of
752 // them will be used, depending on the shard_by_col value.
753 Eigen::ThreadLocal<ThreadLocalBlocks<LhsBlock>, ThreadLocalLhsInit, ThreadLocalLhsRelease> lhs_thread_local_blocks_;
754 Eigen::ThreadLocal<ThreadLocalBlocks<RhsBlock>, ThreadLocalRhsInit, ThreadLocalRhsRelease> rhs_thread_local_blocks_;
755
756 // After a particular shard for Kth slice missed thread local execution
757 // opportunity (K-1 slice didn't complete kernels execution), we can no
758 // longer schedule K+1 and following slices in thread local mode, because
759 // there is no more guarantee that previous kernels were executed
760 // sequentially in the same thread (size is nn_ or nm_).
761 std::atomic<bool>* can_use_thread_local_packed_;
762
763 std::atomic<uint8_t>** state_kernel_[P];
764 // state_switch_ is frequently modified by worker threads, while other
765 // fields are read-only after constructor. Let's move it to a separate cache
766 // line to reduce cache-coherency traffic.
767 char pad_[128];
768 std::atomic<Index> state_packing_ready_[P];
769 std::atomic<Index> state_switch_[P];
770
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();
775
776 Index grain_index = m1 - m * gm_;
777 return blocks.block(
778 internal::convert_index<int>(grain_index)); // FIXME: Consider making ThreadLocalBlocks use Eigen::Index.
779 } else {
780 return packed_lhs_[k % (P - 1)][m1];
781 }
782 }
783
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();
788
789 Index grain_index = n1 - n * gn_;
790 return blocks.block(
791 internal::convert_index<int>(grain_index)); // FIXME: Consider making ThreadLocalBlocks use Eigen::Index.
792 } else {
793 return packed_rhs_[k % (P - 1)][n1];
794 }
795 }
796
797 // In following two methods (pack_lhs and pack_rhs), if we know for sure
798 // that we'll be able to immediately call a kernel with packed data, and do
799 // not submit it to the thread pool, we can use thread local memory for
800 // packed data.
801 //
802 // We can only reliably check it if we are running all kernels in sync mode
803 // (parallelize only by sharding dim). If kernel for m==0 (n==0) is ready to
804 // run, it's guaranteed that all kernels with larger values of m (n) are
805 // also ready, because we execute them in the same order for all K slices.
806
807 void pack_lhs(Index m, Index k) {
808 bool use_thread_local = false;
809
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;
814 } else {
815 // If we can't guarantee that all kernels in `k` slice will be
816 // executed sequentially in current thread, it's no longer safe to use
817 // thread local memory in following slices along the k dimensions.
818 eigen_assert(k > 0);
819 can_use_thread_local_packed_[m].store(false, std::memory_order_relaxed);
820 }
821 }
822
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));
826
827 if (!parallel_pack_ && shard_by_col_) {
828 eigen_assert(!use_thread_local);
829 signal_packing(k);
830 } else {
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);
835 }
836 }
837 }
838
839 void pack_rhs(Index n, Index k) {
840 bool use_thread_local = false;
841
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;
846 } else {
847 // If we can't guarantee that all kernels in `k` slice will be
848 // executed sequentially in current thread, it's no longer safe to use
849 // thread local memory in following slices along the k dimensions.
850 eigen_assert(k > 0);
851 can_use_thread_local_packed_[n].store(false, std::memory_order_relaxed);
852 }
853 }
854
855 const Index nend = n * gn_ + gn(n);
856 for (Index n1 = n * gn_; n1 < nend; n1++) {
857 EIGEN_IF_CONSTEXPR (!TensorContractionKernel::HasBeta) {
858 if (k == 0) {
859 // Zero the output memory in parallel, only if contraction kernel does
860 // not support `beta`. Otherwise we will pass beta 0.0 to the first
861 // call to the `TensorContractionKernel::invoke()`.
862 //
863 // On 10000x2x10000 mm zeroing can easily take half of time. Zero (bn
864 // x m) row. Safe to do here because all kernels that will write to
865 // this memory depend on completion of this task. Note: don't call
866 // device_.fill() here. device_.fill() blocks on thread pool
867 // worker thread, which can lead to underutilization and deadlocks.
868 std::fill_n(buffer_ + n1 * bn_ * m_, bn(n1) * m_, Scalar(0));
869 }
870 }
871 kernel_.packRhs(&packed_rhs(n, k, n1, use_thread_local), rhs_.getSubMapper(k * bk_, n1 * bn_), bk(k), bn(n1));
872 }
873
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);
879 }
880 } else {
881 eigen_assert(!use_thread_local);
882 signal_packing(k);
883 }
884 }
885
886 void kernel(Index m, Index n, Index k, bool use_thread_local) {
887 // Note: order of iteration matters here. Iteration over m is innermost
888 // because we want to reuse the same packed rhs in consecutive tasks
889 // (rhs fits into L2$ while lhs only into L3$).
890 const Index nend = n * gn_ + gn(n);
891 const Index mend = m * gm_ + gm(m);
892
893 // NOTE: output = alpha * LHS * RHS + beta * output.
894 const Scalar alpha = Scalar(1);
895 const Scalar beta = (TensorContractionKernel::HasBeta && k == 0) ? Scalar(0) : Scalar(1);
896
897 if (shard_by_col_) {
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);
903
904 // We are done with the last task for the [m1, n1] block.
905 if (k + 1 == nk_) {
906 output_kernel_(output_mapper, tensor_contraction_params_, m1 * bm_, n1 * bn_, bm(m1), bn(n1));
907 }
908 }
909 }
910 } else {
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);
916
917 // We are done with the last task for the [m1, n1] block.
918 if (k + 1 == nk_) {
919 output_kernel_(output_mapper, tensor_contraction_params_, m1 * bm_, n1 * bn_, bm(m1), bn(n1));
920 }
921 }
922 }
923 signal_kernel(m, n, k + 1, /*sync=*/false, /*use_thread_local=*/false);
924 signal_switch(k + 2);
925 }
926
927 void signal_packing(Index k) {
928 eigen_assert(!parallel_pack_);
929 Index s = state_packing_ready_[k % P].fetch_sub(1);
930 eigen_assert(s > 0);
931 if (s != 1) return;
932 state_packing_ready_[k % P] = shard_by_col_ ? nm_ : nn_;
933 enqueue_packing(k, shard_by_col_);
934 }
935
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();
939 eigen_assert(s > 0);
940 if (s != 1 && state->fetch_sub(1) != 1) {
941 eigen_assert(!use_thread_local);
942 return;
943 }
944 state->store(parallel_pack_ ? 3 : 2, std::memory_order_relaxed);
945 if (sync) {
946 kernel(m, n, k, use_thread_local);
947 } else {
948 eigen_assert(!use_thread_local);
949 device_.enqueue([this, m, n, k, use_thread_local]() { kernel(m, n, k, use_thread_local); });
950 }
951 }
952
953 void signal_switch(Index k, Index v = 1) {
954 Index s = state_switch_[k % P].fetch_sub(v);
955 eigen_assert(s >= v);
956 if (s != v) return;
957
958 // Ready to switch to the next k slice.
959 // Reset counter for the next iteration.
960 state_switch_[k % P] = (parallel_pack_ ? nm_ + nn_ : (shard_by_col_ ? nn_ : nm_)) + nm_ * nn_;
961 if (k < nk_) {
962 // Issue lhs/rhs packing. Their completion will in turn kick off
963 // kernels.
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);
969 } else {
970 enqueue_packing(k, true);
971 }
972
973 // Termination handling.
974 // Because kernel completion signals k + 2 switch, we need to finish nk
975 // + 2 slices without issuing any tasks on nk + 1 slice. So here we
976 // pretend that all nk + 1 packing tasks just finish instantly; so that
977 // nk + 2 switch only waits for completion of nk kernels.
978 } else if (k == nk_) {
979 signal_switch(k + 1, parallel_pack_ ? nm_ + nn_ : (shard_by_col_ ? nn_ : nm_));
980 } else {
981 done_.Notify();
982 }
983 }
984
985 // Enqueue all rhs/lhs packing for k-th slice.
986 void enqueue_packing(Index k, bool rhs) { enqueue_packing_helper(0, rhs ? nn_ : nm_, k, rhs); }
987
988 void enqueue_packing_helper(Index start, Index end, Index k, bool rhs) {
989 if (end - start == 1) {
990 if (rhs)
991 pack_rhs(start, k);
992 else
993 pack_lhs(start, k);
994 } else {
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); });
998 end = mid;
999 }
1000
1001 // Decide if we want to run first packing task (start == 0) in
1002 // async mode if we parallelize only by sharding dim:
1003 // (1) pack_lhs and pack_rhs call signal_switch before completing
1004 // all calls to signal_kernel, which in sync mode might lead
1005 // to the execution of the first kernel of the k+1 slice, before
1006 // completing a call to the last kernel of the k slice.
1007 // (2) all pack tasks for sharded dim must be executed in a thread
1008 // pool to get pre-allocated thread local buffers.
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_);
1011
1012 if (pack_async) {
1013 device_.enqueue([this, start, end, k, rhs]() { enqueue_packing_helper(start, end, k, rhs); });
1014 } else {
1015 enqueue_packing_helper(start, end, k, rhs);
1016 }
1017 }
1018 }
1019
1020 // Block sizes with accounting for potentially incomplete last block.
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_; }
1024 // Task grain sizes accounting for potentially incomplete last task.
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_; }
1027
1028 EvalParallelContext(const EvalParallelContext&) = delete;
1029 void operator=(const EvalParallelContext&) = delete;
1030 };
1031
1032 // ------------------------------------------------------------------------ //
1033
1034 // EvalShardedByInnerDimContext orchestrates sync/async contraction
1035 // evaluation, when we shard by inner dimension. When it is executed in
1036 // asynchronous mode, it owns all the shared state that might be accessible by
1037 // block processing tasks.
1038
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)
1043 : evaluator(self),
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),
1048 m(m_size),
1049 n(n_size),
1050 k(k_size),
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) {
1059 // Keep count of pending gemm tasks for each l0 range.
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));
1063 }
1064
1065 // Allocate temporary buffers for each block.
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);
1069 }
1070 }
1071
1072 ~EvalShardedByInnerDimContext() {
1073 for (Index i = 1; i < num_blocks; ++i) {
1074 evaluator->m_device.deallocate(block_buffers[i]);
1075 }
1076 }
1077
1078 template <int Alignment>
1079 void run() {
1080 Barrier barrier(internal::convert_index<int>(num_blocks));
1081 eval<Alignment>(barrier, 0, num_blocks);
1082 barrier.Wait();
1083
1084 // Aggregate partial sums from l0 ranges.
1085 aggregateL0Blocks<Alignment>();
1086
1087 // Apply output kernel.
1088 applyOutputKernel();
1089 }
1090
1091 template <int Alignment>
1092 void runAsync() {
1093 evalAsync<Alignment>(0, num_blocks);
1094 }
1095
1096 private:
1097 // The underlying GEMM kernel assumes that k is a multiple of
1098 // the packet size and subtle breakage occurs if this is violated.
1099 static constexpr Index packet_size = internal::packet_traits<RhsScalar>::size;
1100
1101 const Self* evaluator; // TensorContraction evaluator
1102
1103 // These fields cache values from the evaluator for use in processBlock dispatch.
1104 bool m_lhs_inner_dim_contiguous;
1105 bool m_rhs_inner_dim_contiguous;
1106 bool m_rhs_inner_dim_reordered;
1107
1108 Scalar* result;
1109
1110 Index m;
1111 Index n;
1112 Index k;
1113
1114 DoneCallback done;
1115
1116 // ----------------------------------------------------------------------//
1117 // Algorithm parameters.
1118
1119 // We will compute partial results into the buffers of this size.
1120 Index buffer_size_bytes;
1121
1122 Index block_size;
1123 Index num_blocks;
1124
1125 // Keep track of pending tasks when evaluate in async mode.
1126 std::atomic<int> num_pending_blocks;
1127
1128 // We compute partial gemm results in parallel, and to get the final result
1129 // we need to add them all together. For the large number of threads (>= 48)
1130 // this adds a very expensive sequential step at the end.
1131 //
1132 // We split the [0, num_blocks) into small ranges, and when a task for the
1133 // block finishes its partial gemm computation, it checks if it was the last
1134 // gemm in the range, and if so, it will add all blocks of the range.
1135 //
1136 // After all tasks done, we need to add only these pre-aggregated blocks.
1137
1138 // For now we use just a single level of ranges to compute pre-aggregated
1139 // partial sums, but in general we can use more layers to compute tree
1140 // aggregation in parallel and reduce the size of the sequential step.
1141 //
1142 // TODO(ezhulenev): Add multilevel tree aggregation? Probably will make
1143 // sense only if number of threads >= ~128?
1144 static constexpr Index l0_size = 4;
1145 Index l0_ranges;
1146
1147 // Keep count of pending gemm tasks for each l0 range.
1148 MaxSizeVector<std::atomic<int>> l0_state; // [0, l0_ranges)
1149
1150 // Buffers allocated for each temporary block computation.
1151 MaxSizeVector<Scalar*> block_buffers; // [0, num_blocks)
1152
1153 template <int Alignment>
1154 void processBlock(Index block_idx, Index begin, Index end) {
1155 Scalar* buf = block_buffers[block_idx];
1156
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, /*num_threads=*/internal::convert_index<int>(num_blocks));
1161 },
1162 m_lhs_inner_dim_contiguous, m_rhs_inner_dim_contiguous, m_rhs_inner_dim_reordered);
1163
1164 // Check if it was the last task in l0 range.
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);
1168
1169 // If we processed the last block of the range, we can aggregate all
1170 // partial results into the first block of the range.
1171 if (v == 1) {
1172 const Index rng_size = actualRangeSize(l0_ranges, l0_size, l0_index);
1173 const Index dst_block_idx = l0_index * l0_size;
1174
1175 if (rng_size == l0_size) {
1176 addAllToBuffer<Alignment>(m * n,
1177 /*src_buf0=*/block_buffers[dst_block_idx + 1],
1178 /*src_buf1=*/block_buffers[dst_block_idx + 2],
1179 /*src_buf2=*/block_buffers[dst_block_idx + 3],
1180 /*dst_buf= */ block_buffers[dst_block_idx]);
1181 } else {
1182 // Aggregate blocks of potentially incomplete last range.
1183 for (int i = 1; i < rng_size; ++i) {
1184 addToBuffer<Alignment>(m * n,
1185 /*src_buf=*/block_buffers[dst_block_idx + i],
1186 /*dst_buf=*/block_buffers[dst_block_idx]);
1187 }
1188 }
1189 }
1190 }
1191
1192 // Aggregate partial sums from l0 ranges.
1193 template <int Alignment>
1194 void aggregateL0Blocks() const {
1195 Index l0_index = 1;
1196
1197 for (; l0_index + 2 < l0_ranges; l0_index += 3) {
1198 addAllToBuffer<Alignment>(m * n,
1199 /*src_buf0=*/block_buffers[(l0_index + 0) * l0_size],
1200 /*src_buf1=*/block_buffers[(l0_index + 1) * l0_size],
1201 /*src_buf2=*/block_buffers[(l0_index + 2) * l0_size],
1202 /*dst_buf= */ block_buffers[0]);
1203 }
1204
1205 for (; l0_index < l0_ranges; ++l0_index) {
1206 addToBuffer<Alignment>(m * n, block_buffers[l0_index * l0_size], block_buffers[0]);
1207 }
1208 }
1209
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);
1214 }
1215
1216 // Compute block size with accounting for potentially incomplete last block.
1217 Index actualBlockSize(Index block_idx) const {
1218 return block_idx + 1 < num_blocks ? block_size : k + block_size - block_size * num_blocks;
1219 }
1220
1221 // Compute range size with accounting for potentially incomplete last range.
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;
1225 }
1226
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;
1230 size_t i = 0;
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);
1237 }
1238 for (; i < n; ++i) {
1239 tgt_buf[i] += src_buf[i];
1240 }
1241 }
1242
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;
1250
1251 const int output_packet_size = internal::unpacket_traits<PacketReturnType>::size;
1252
1253 size_t i = 0;
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);
1259
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));
1262
1263 pstoret<Scalar, PacketReturnType, Alignment>(dst_buf + i, sum);
1264 }
1265 for (; i < n; ++i) {
1266 dst_buf[i] += src_buf0[i] + src_buf1[i] + src_buf2[i];
1267 }
1268 }
1269
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);
1276 });
1277 end_block_idx = mid_block_idx;
1278 }
1279
1280 Index block_idx = start_block_idx;
1281 Index block_start = block_idx * block_size;
1282 Index block_end = block_start + actualBlockSize(block_idx);
1283
1284 processBlock<Alignment>(block_idx, block_start, block_end);
1285 barrier.Notify();
1286 }
1287
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;
1295 }
1296
1297 Index block_idx = start_block_idx;
1298
1299 Index block_start = block_idx * block_size;
1300 Index block_end = block_start + actualBlockSize(block_idx);
1301
1302 processBlock<Alignment>(block_idx, block_start, block_end);
1303
1304 int v = num_pending_blocks.fetch_sub(1);
1305 eigen_assert(v >= 1);
1306
1307 if (v == 1) {
1308 // Aggregate partial sums from l0 ranges.
1309 aggregateL0Blocks<Alignment>();
1310
1311 // Apply output kernel.
1312 applyOutputKernel();
1313
1314 // NOTE: If we call `done` callback before deleting this (context),
1315 // it might deallocate Self* pointer captured by context, and we'll
1316 // fail in destructor trying to deallocate temporary buffers.
1317
1318 // Move done call back from context before it will be destructed.
1319 DoneCallback done_copy = std::move(done);
1320
1321 // We are confident that we are the last one who touches context.
1322 delete this;
1323
1324 // Now safely call the done callback.
1325 done_copy();
1326 }
1327 }
1328
1329 // Cost model doesn't capture well the cost associated with constructing
1330 // tensor contraction mappers and computing loop bounds in gemm_pack_lhs
1331 // and gemm_pack_rhs, so we specify minimum desired block size.
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;
1336 };
1337
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;
1340
1341 return numext::mini<Index>(k, numext::maxi<Index>(desired_min_block_size, target_block_size));
1342 }
1343
1344 EvalShardedByInnerDimContext(const EvalShardedByInnerDimContext&) = delete;
1345 void operator=(const EvalShardedByInnerDimContext&) = delete;
1346 };
1347
1348 // ------------------------------------------------------------------------ //
1349
1350 // Below are the function used by evalProductImpl heuristics, trying to select
1351 // optimal parameters for parallelization algorithm.
1352
1353 // Decide whether we want to shard m x n contraction by columns or by rows.
1354 static bool shardByCol(Index m, Index n, Index num_threads) {
1355 // Note: we are comparing both n and m against Traits::nr, it is not
1356 // a mistake. We are trying to figure out how both n and m will fit into
1357 // the main sharding dimension.
1358
1359 // Sharding by column is the default
1360 // ... unless there is enough data for vectorization over rows
1361 if (m / num_threads >= Traits::nr &&
1362 // and not enough data for vectorization over columns
1363 (n / num_threads < Traits::nr ||
1364 // ... or barely enough data for vectorization over columns,
1365 // but it is not evenly dividable across threads
1366 (n / num_threads < 4 * Traits::nr && (n % (num_threads * Traits::nr)) != 0 &&
1367 // ... and it is evenly dividable across threads for rows
1368 ((m % (num_threads * Traits::nr)) == 0 ||
1369 // .. or it is not evenly dividable for both dimensions but
1370 // there is much more data over rows so that corner effects are
1371 // mitigated.
1372 (m / n >= 6)))))
1373 return false;
1374 // Wait, or if matrices are just substantially prolonged over the other
1375 // dimension.
1376 if (n / num_threads < 16 * Traits::nr && m > n * 32) return false;
1377 return true;
1378 }
1379
1380 Index coarsenM(Index m, Index n, Index bm, Index bn, Index bk, Index gn, int num_threads, bool shard_by_col) const {
1381 Index gm = 1;
1382 Index gm1 = 1;
1383 Index nm0 = numext::div_ceil(m, bm);
1384 Index nm1 = nm0;
1385 for (;;) {
1386 // Find the next candidate for m grain size. It needs to result in
1387 // different number of blocks. E.g. if we have 10 kernels, we want to try
1388 // 5 and 10, but not 6, 7, 8 and 9.
1389 while (gm1 <= nm0 && nm1 == numext::div_ceil(nm0, gm1)) gm1++;
1390 if (gm1 > nm0) break;
1391 // Check the candidate.
1392 int res = checkGrain(m, n, bm, bn, bk, gm1, gn, gm, gn, num_threads, shard_by_col);
1393 if (res < 0) break;
1394 nm1 = numext::div_ceil(nm0, gm1);
1395 if (res == 0) continue;
1396 // Commit new grain size.
1397 gm = gm1;
1398 }
1399 return gm;
1400 }
1401
1402 Index coarsenN(Index m, Index n, Index bm, Index bn, Index bk, Index gm, int num_threads, bool shard_by_col) const {
1403 Index gn = 1;
1404 Index gn1 = 1;
1405 Index nn0 = numext::div_ceil(n, bn);
1406 Index nn1 = nn0;
1407 for (;;) {
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);
1411 if (res < 0) break;
1412 nn1 = numext::div_ceil(nn0, gn1);
1413 if (res == 0) continue;
1414 gn = gn1;
1415 }
1416 return gn;
1417 }
1418
1419 // checkGrain checks whether grain (gm, gn) is suitable and is better than
1420 // (oldgm, oldgn).
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);
1425 // If the task is too small, then we agree on it regardless of anything
1426 // else. Otherwise synchronization overheads will dominate.
1427 if (taskSize < 1) return 1;
1428 // If it is too large, then we reject it and all larger tasks.
1429 if (taskSize > 2) return -1;
1430 // Now we are in presumably good task size range.
1431 // The main deciding factor here is parallelism. Consider that we have 12
1432 // kernels and 4 threads. Grains of 2, 3 and 4 all yield good task sizes.
1433 // But 2/4 yield 6/3 tasks, which gives us parallelism of 0.75 (at most 3/4
1434 // of cores will be busy). While grain size 3 gives us 4 tasks, which gives
1435 // us parallelism of 1 (we can load all cores).
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;
1445 return 0;
1446 }
1447
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);
1454 // Computations.
1455 TensorOpCost cost = TensorOpCost(0, 0, kd * compute_bandwidth, true, packed_size);
1456 // Output stores.
1457 cost += TensorOpCost(0, sizeof(CoeffReturnType), 0, true, output_packet_size);
1458 if (prepacked) {
1459 // Packing and kernels are executed in different tasks. When we calculate
1460 // task grain size we look only at kernel cost assuming that kernel
1461 // is more expensive than packing.
1462 return cost;
1463 }
1464 // Lhs/rhs loads + computations.
1465 TensorOpCost lhsCost = this->m_leftImpl.costPerCoeff(true) * (kd / n);
1466 TensorOpCost rhsCost = this->m_rightImpl.costPerCoeff(true) * (kd / m);
1467 // Lhs packing memory cost does not contribute considerably to overall
1468 // execution time because lhs is prefetched early and accessed sequentially.
1469 if (shard_by_col)
1470 lhsCost.dropMemoryCost();
1471 else
1472 rhsCost.dropMemoryCost();
1473 return cost + lhsCost + rhsCost;
1474 }
1475
1476 // Decide whether we want to shard m x k x n contraction over the inner
1477 // (contraction) dimension (k).
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;
1481 if (n == 1 || // If mat*vec or...
1482 num_threads_by_k < 2 || // running single threaded or...
1483 num_threads_by_k < num_threads || // sharding by k gives less parallelism or...
1484 bufsize > l3CacheSize() / num_threads_by_k || // need more buffer space
1485 // than L3 cache or...
1486 k / num_threads_by_k < 2 * Traits::nr) { // k per thread is tiny.
1487 shard_by_k = false;
1488 } else if (numext::maxi(m, n) / num_threads < Traits::nr || // both other dimensions are tiny or...
1489 // k per thread is not small and...
1490 (k / num_threads_by_k > 8 * Traits::nr &&
1491 // one of the outer dimensions is tiny or sharding by k offers
1492 // more parallelism.
1493 (numext::mini(m, n) < 2 * Traits::nr || num_threads_by_k > num_threads))) {
1494 shard_by_k = true;
1495 }
1496 return shard_by_k;
1497 }
1498
1499 TensorOpCost contractionCostPerInnerDim(Index m, Index n, Index k) const {
1500 // Compute cost.
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);
1503 // Output stores.
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;
1507 // Since the inner gemm kernel is always sharded by column, the lhs
1508 // load cost is negligible.
1509 lhsCost.dropMemoryCost();
1510 return cost + lhsCost + rhsCost;
1511 }
1512
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);
1517 // Cost of reduction step accumulating the m*n per-thread buffers into the
1518 // result.
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) {
1529 num_threads = nt;
1530 min_cost = parallel_cost;
1531 }
1532 }
1533 return num_threads;
1534 }
1535
1536 double computeBandwidth(bool shard_by_col, Index bm, Index bn, Index bk) const {
1537 // Peak VFMA bandwidth is 0.5. However if we have not enough data for
1538 // vectorization bandwidth drops. The 4.0 and 2.0 bandwidth is determined
1539 // experimentally.
1540 double computeBandwidth = bk == 1 ? 4.0
1541 : (shard_by_col ? bn : bm) < Traits::nr || (shard_by_col ? bm : bn) < Traits::mr ? 2.0
1542 : 0.5;
1543#ifndef EIGEN_VECTORIZE_FMA
1544 // Bandwidth of all of VFMA/MULPS/ADDPS is 0.5 on latest Intel processors.
1545 // However for MULPS/ADDPS we have dependent sequence of 2 such
1546 // instructions,
1547 // so overall bandwidth is 1.0.
1548 if (computeBandwidth == 0.5) computeBandwidth = 1.0;
1549#endif
1550 return computeBandwidth;
1551 }
1552};
1553
1554} // end namespace Eigen
1555
1556#endif // EIGEN_USE_THREADS
1557#endif // EIGEN_TENSOR_TENSOR_CONTRACTION_THREAD_POOL_H
Definition TensorContraction.h:335
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47