File size: 4,307 Bytes
d20d01e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
// Fork-safe data parallelism for the runtime encode/decode kernels.
//
// These kernels run inside forked processes (for example PyTorch DataLoader
// workers).  GNU OpenMP keeps a process-wide thread pool that does not survive
// fork(): once the parent has entered a parallel region -- ours, or any other
// library's sharing libgomp -- the child hangs or crashes when it enters one.
// This helper instead starts its threads per call and joins them before
// returning, so no thread state outlives a call and fork() is always safe.
//
// Iterations are split into contiguous static blocks, like OpenMP's
// ``schedule(static)``, and every iteration writes only its own outputs, so
// results do not depend on the thread count.
//
// Starting a thread costs tens of microseconds, and one iteration costs from
// well under a microsecond to about a hundred depending on the kernel and the
// profile's primitive counts.  So each call times its first iteration on the
// calling thread and sizes the split of the remaining work ``W`` from it: the
// calling thread starts threads one after another at ``s`` seconds each, so a
// split over ``w`` threads takes about ``W / w + s * w``, least at
// ``w = sqrt(W / s)``.
#pragma once

#include <algorithm>
#include <atomic>
#include <chrono>
#include <cmath>
#include <cstdint>
#include <exception>
#include <mutex>
#include <thread>
#include <vector>

#ifdef __linux__
#include <sched.h>
#endif

namespace ac2 {

// Process-wide cap on threads per call; 0 means every CPU the process may use.
inline std::atomic<int>& thread_budget() {
  static std::atomic<int> value{0};
  return value;
}

inline int available_cpus() {
#ifdef __linux__
  cpu_set_t allowed;
  CPU_ZERO(&allowed);
  if (sched_getaffinity(0, sizeof(allowed), &allowed) == 0) {
    return std::max(1, CPU_COUNT(&allowed));
  }
#endif
  return std::max(1, static_cast<int>(std::thread::hardware_concurrency()));
}

// Measured cost of starting and joining one thread (seconds).
constexpr double kThreadStartSeconds = 30e-6;

// Threads for ``count`` iterations of ``seconds_each``: at most ``requested``
// when positive, at most the process budget, and at most ``sqrt(W / s)``.
inline int64_t plan_threads(int64_t count, int requested, double seconds_each) {
  const int budget = thread_budget().load();
  int64_t threads = budget > 0 ? budget : available_cpus();
  if (requested > 0) threads = std::min<int64_t>(threads, requested);
  const double by_work = std::sqrt(count * seconds_each / kThreadStartSeconds);
  if (by_work < static_cast<double>(threads)) threads = static_cast<int64_t>(by_work);
  return std::max<int64_t>(1, std::min(threads, count));
}

// Call ``body(index)`` for every index in [0, count).  Iteration 0 runs first
// on the calling thread and sizes the split of the rest; the calling thread
// also runs the first block of the rest.  The first exception thrown by any
// block is rethrown after all threads have joined.
template <class Body>
void parallel_for(int64_t count, int threads, Body&& body) {
  if (count <= 0) return;
  const auto start = std::chrono::steady_clock::now();
  body(0);
  const double seconds_each =
      std::chrono::duration<double>(std::chrono::steady_clock::now() - start).count();
  const int64_t rest = count - 1;
  if (rest == 0) return;
  const int64_t workers = plan_threads(rest, threads, seconds_each);
  auto run = [&](int64_t worker) {
    const int64_t begin = 1 + rest * worker / workers;
    const int64_t end = 1 + rest * (worker + 1) / workers;
    for (int64_t index = begin; index < end; ++index) body(index);
  };
  if (workers == 1) {
    run(0);
    return;
  }
  std::exception_ptr failure;
  std::mutex failure_lock;
  auto guarded = [&](int64_t worker) {
    try {
      run(worker);
    } catch (...) {
      std::lock_guard<std::mutex> hold(failure_lock);
      if (!failure) failure = std::current_exception();
    }
  };
  std::vector<std::thread> pool;
  pool.reserve(static_cast<size_t>(workers - 1));
  try {
    for (int64_t worker = 1; worker < workers; ++worker) pool.emplace_back(guarded, worker);
  } catch (...) {
    for (auto& thread : pool) thread.join();
    throw;
  }
  guarded(0);
  for (auto& thread : pool) thread.join();
  if (failure) std::rethrow_exception(failure);
}

}  // namespace ac2