11#ifndef EIGEN_PARALLELIZER_H
12#define EIGEN_PARALLELIZER_H
15#include "../InternalHeaderCheck.h"
38#if defined(EIGEN_HAS_OPENMP) && defined(EIGEN_GEMM_THREADPOOL)
39#error "EIGEN_HAS_OPENMP and EIGEN_GEMM_THREADPOOL may not both be defined."
45inline void manage_multi_threading(Action action,
int* v);
51EIGEN_DEPRECATED_WITH_REASON(
"Initialization is no longer needed.") inline
void initParallel() {}
55inline int nbThreads() {
57 internal::manage_multi_threading(GetAction, &ret);
63inline void setNbThreads(
int v) { internal::manage_multi_threading(SetAction, &v); }
65#ifdef EIGEN_VECTORIZE_SME
66#ifndef EIGEN_SME_MIN_TASK_SIZE
67#define EIGEN_SME_MIN_TASK_SIZE (double(1 << 20))
72template <
typename Scalar>
74 static constexpr double value = (std::is_same<typename NumTraits<Scalar>::Real,
double>::value ? 4.0 : 1.0) *
75 (NumTraits<Scalar>::IsComplex ? 4.0 : 1.0);
80inline int detect_sme_units() {
81#if defined(EIGEN_SME_UNITS)
82 return EIGEN_SME_UNITS;
86 int32_t cores = 0, per_l2 = 0;
87 size_t sz =
sizeof(cores);
88 if (sysctlbyname(
"hw.perflevel0.physicalcpu", &cores, &sz,
nullptr, 0) != 0 || cores <= 0)
return 0;
90 if (sysctlbyname(
"hw.perflevel0.cpusperl2", &per_l2, &sz,
nullptr, 0) != 0 || per_l2 <= 0)
return 0;
91 return numext::maxi<int32_t>(1, cores / per_l2);
96inline void manage_sme_units(Action action,
int* v) {
97 static int m_units = detect_sme_units();
98 if (action == SetAction)
109inline int nbSmeUnits() {
110#ifdef EIGEN_VECTORIZE_SME
112 internal::manage_sme_units(GetAction, &ret);
123inline void setNbSmeUnits(
int v) {
124#ifdef EIGEN_VECTORIZE_SME
125 internal::manage_sme_units(SetAction, &v);
127 EIGEN_UNUSED_VARIABLE(v);
131#ifdef EIGEN_GEMM_THREADPOOL
140inline ThreadPool* setGemmThreadPool(ThreadPool* new_pool) {
141 static ThreadPool* pool =
nullptr;
142 if (new_pool !=
nullptr) {
149 setNbThreads(pool->NumThreads());
155inline ThreadPool* getGemmThreadPool() {
return setGemmThreadPool(
nullptr); }
166constexpr int kGemmColGrain = 4;
174template <
typename Index>
175EIGEN_ALWAYS_INLINE
void balanced_gemm_range(Index total, Index parts, Index grain, Index part, Index& start,
177 const Index chunks = numext::div_ceil(total, grain);
178 const Index base = chunks / parts;
179 const Index extra = chunks % parts;
180 const Index first_chunk = part * base + numext::mini(part, extra);
181 const Index chunk_count = base + (part < extra ? 1 : 0);
182 start = numext::mini(first_chunk * grain, total);
183 length = numext::mini(chunk_count * grain, total - start);
186#if defined(EIGEN_USE_BLAS) || (!defined(EIGEN_HAS_OPENMP) && !defined(EIGEN_GEMM_THREADPOOL))
188inline void manage_multi_threading(Action action,
int* v) {
189 if (action == SetAction) {
190 eigen_internal_assert(v !=
nullptr);
191 }
else if (action == GetAction) {
192 eigen_internal_assert(v !=
nullptr);
195 eigen_internal_assert(
false);
198template <
typename Index>
199struct GemmParallelInfo {};
200template <
bool Condition,
typename Functor,
typename Index>
201EIGEN_STRONG_INLINE
void parallelize_gemm(
const Functor& func, Index rows, Index cols, Index ,
203 func(0, rows, 0, cols);
208template <
typename Index>
209struct GemmParallelTaskInfo {
210 GemmParallelTaskInfo() {}
211 std::atomic<Index> sync{Index(-1)};
212 std::atomic<int> users{0};
214 Index lhs_length = 0;
217template <
typename Index>
218struct GemmParallelInfo {
219 const int logical_thread_id;
220 const int num_threads;
221 GemmParallelTaskInfo<Index>* task_info;
223 GemmParallelInfo(
int logical_thread_id_,
int num_threads_, GemmParallelTaskInfo<Index>* task_info_)
224 : logical_thread_id(logical_thread_id_), num_threads(num_threads_), task_info(task_info_) {}
227inline void manage_multi_threading(Action action,
int* v) {
228 static int m_maxThreads = -1;
229 if (action == SetAction) {
230 eigen_internal_assert(v !=
nullptr);
231#if defined(EIGEN_HAS_OPENMP)
235 eigen_internal_assert(*v >= 0);
236 int omp_threads = omp_get_max_threads();
237 m_maxThreads = (*v == 0 ? omp_threads : std::min<int>(*v, omp_threads));
238#elif defined(EIGEN_GEMM_THREADPOOL)
242 eigen_internal_assert(*v >= 0);
243 ThreadPool* pool = getGemmThreadPool();
244 int pool_threads = pool !=
nullptr ? pool->NumThreads() : 1;
245 m_maxThreads = (*v == 0 ? pool_threads : numext::mini(pool_threads, *v));
247 }
else if (action == GetAction) {
248 eigen_internal_assert(v !=
nullptr);
249#if defined(EIGEN_HAS_OPENMP)
250 if (m_maxThreads > 0)
253 *v = omp_get_max_threads();
258 eigen_internal_assert(
false);
262template <
bool Condition,
typename Functor,
typename Index>
263EIGEN_STRONG_INLINE
void parallelize_gemm(
const Functor& func, Index rows, Index cols, Index depth,
bool transpose) {
273 Index size = transpose ? rows : cols;
274 Index pb_max_threads = std::max<Index>(1, size / Functor::Traits::nr);
277 double work =
static_cast<double>(rows) *
static_cast<double>(cols) *
static_cast<double>(depth);
278 double kMinTaskSize = 50000;
279#ifdef EIGEN_VECTORIZE_SME
283 constexpr bool kSme =
284 sme_has_gebp_kernel<typename Functor::Traits::LhsScalar, typename Functor::Traits::RhsScalar>::value;
285 const bool sme_units_known = kSme && nbSmeUnits() > 0;
286 const bool sme_disjoint = sme_units_known && func.ownsNoBuffers();
288 kMinTaskSize = EIGEN_SME_MIN_TASK_SIZE / sme_madd_cost<typename Functor::Traits::LhsScalar>::value;
290 const Index kernel_rows = transpose ? cols : rows, kernel_cols = transpose ? rows : cols;
292 std::max<Index>(1, (std::max)(kernel_rows / Functor::Traits::mr, kernel_cols / Functor::Traits::nr));
295 pb_max_threads = std::max<Index>(1, std::min<Index>(pb_max_threads,
static_cast<Index
>(work / kMinTaskSize)));
298 int threads = std::min<int>(nbThreads(),
static_cast<int>(pb_max_threads));
299#ifdef EIGEN_VECTORIZE_SME
301 EIGEN_IF_CONSTEXPR (kSme) {
302 const int units = nbSmeUnits();
303 if (units > 0) threads = std::min<int>(threads, units);
309 bool dont_parallelize = (!Condition) || (threads <= 1);
310#if defined(EIGEN_HAS_OPENMP)
312 dont_parallelize |= omp_get_num_threads() > 1;
313#elif defined(EIGEN_GEMM_THREADPOOL)
318 ThreadPool* pool = getGemmThreadPool();
319 dont_parallelize |= (pool ==
nullptr || pool->CurrentThreadId() != -1);
321 if (dont_parallelize)
return func(0, rows, 0, cols);
323#ifdef EIGEN_VECTORIZE_SME
327 EIGEN_IF_CONSTEXPR (kSme) {
332 const Index kernel_rows = transpose ? cols : rows, kernel_cols = transpose ? rows : cols;
333 const Index row_parts = kernel_rows / (8 * Functor::Traits::mr);
334 const bool split_kernel_rows = kernel_rows >= 2 * kernel_cols && row_parts >= 2;
335 const bool split_rows = transpose ? !split_kernel_rows : split_kernel_rows;
336 if (split_kernel_rows) threads =
static_cast<int>((std::min)(Index(threads), row_parts));
338 const Index row_grain = transpose ? Functor::Traits::nr : Functor::Traits::mr;
339 const Index col_grain = transpose ? Functor::Traits::mr : Functor::Traits::nr;
340 const Index chunks = split_rows ? (rows + row_grain - 1) / row_grain : (cols + col_grain - 1) / col_grain;
341 threads =
static_cast<int>((std::min)(Index(threads), chunks));
342 if (threads <= 1)
return func(0, rows, 0, cols);
343 auto part = [&func, rows, cols, threads, split_rows, row_grain, col_grain](
int i) {
346 balanced_gemm_range<Index>(rows, threads, row_grain, i, start, length);
347 if (length > 0) func(start, length, 0, cols);
349 balanced_gemm_range<Index>(cols, threads, col_grain, i, start, length);
350 if (length > 0) func(0, rows, start, length);
353#if defined(EIGEN_HAS_OPENMP)
354#pragma omp parallel for num_threads(threads) schedule(static, 1)
355 for (
int i = 0; i < threads; ++i) part(i);
356#elif defined(EIGEN_GEMM_THREADPOOL)
360 std::atomic<int> pending(threads - 1);
361 Barrier done(threads - 1);
362 for (
int i = 0; i < threads - 1; ++i)
363 pool->Schedule([&part, &pending, &done, i] {
365 pending.fetch_sub(1, std::memory_order_release);
368 const auto wait = [&pending, &done] {
369 for (
int spin = 0; spin < 65536 && pending.load(std::memory_order_acquire) != 0; ++spin) {
373 EIGEN_TRY { part(threads - 1); }
385 func.initParallelSession(threads);
387 if (transpose) std::swap(rows, cols);
389 ei_declare_aligned_stack_constructed_variable(GemmParallelTaskInfo<Index>, task_info, threads, 0);
391#if defined(EIGEN_HAS_OPENMP)
392#pragma omp parallel num_threads(threads)
394 Index i = omp_get_thread_num();
397 Index actual_threads = omp_get_num_threads();
398 GemmParallelInfo<Index> info(
static_cast<int>(i),
static_cast<int>(actual_threads), task_info);
400 Index r0, actualBlockRows;
401 balanced_gemm_range<Index>(rows, actual_threads, Index(Functor::Traits::mr), i, r0, actualBlockRows);
403 Index c0, actualBlockCols;
404 balanced_gemm_range<Index>(cols, actual_threads, Index(kGemmColGrain), i, c0, actualBlockCols);
406 info.task_info[i].lhs_start = r0;
407 info.task_info[i].lhs_length = actualBlockRows;
410 func(c0, actualBlockCols, 0, rows, &info);
412 func(0, rows, c0, actualBlockCols, &info);
415#elif defined(EIGEN_GEMM_THREADPOOL)
416 Barrier barrier(threads);
417 auto task = [=, &func, &barrier, &task_info](
int i) {
418 Index actual_threads = threads;
419 GemmParallelInfo<Index> info(i,
static_cast<int>(actual_threads), task_info);
420 Index r0, actualBlockRows;
421 balanced_gemm_range<Index>(rows, actual_threads, Index(Functor::Traits::mr), i, r0, actualBlockRows);
423 Index c0, actualBlockCols;
424 balanced_gemm_range<Index>(cols, actual_threads, Index(kGemmColGrain), i, c0, actualBlockCols);
426 info.task_info[i].lhs_start = r0;
427 info.task_info[i].lhs_length = actualBlockRows;
430 func(c0, actualBlockCols, 0, rows, &info);
432 func(0, rows, c0, actualBlockCols, &info);
439 for (
int i = 0; i < threads - 1; ++i) {
440 pool->Schedule([=, task = std::move(task)] { task(i); });