Threading api refactor (#1955)
Refactor the multi-threading api to support
using custom user-provided thread factory
instead of always spawning POSIX Threads.
diff --git a/docs/user_guide.md b/docs/user_guide.md
index b3c1cce..0bcfe15 100644
--- a/docs/user_guide.md
+++ b/docs/user_guide.md
@@ -863,6 +863,46 @@
Without `UseRealTime`, CPU time is used by default.
+### Manual Multithreaded Benchmarks
+
+Google/benchmark uses `std::thread` as multithreading environment per default.
+If you want to use another multithreading environment (e.g. OpenMP), you can provide
+a factory function to your benchmark using the `ThreadRunner` function.
+The factory function takes the number of threads as argument and creates a custom class
+derived from `benchmark::ThreadRunnerBase`.
+This custom class must override the function
+`void RunThreads(const std::function<void(int)>& fn)`.
+`RunThreads` is called by the main thread and spawns the requested number of threads.
+Each spawned thread must call `fn(thread_index)`, where `thread_index` is its own
+thread index. Before `RunThreads` returns, all spawned threads must be joined.
+```c++
+class OpenMPThreadRunner : public benchmark::ThreadRunnerBase
+{
+ OpenMPThreadRunner(int num_threads)
+ : num_threads_(num_threads)
+ {}
+
+ void RunThreads(const std::function<void(int)>& fn) final
+ {
+#pragma omp parallel num_threads(num_threads_)
+ fn(omp_get_thread_num());
+ }
+
+private:
+ int num_threads_;
+};
+
+BENCHMARK(BM_MultiThreaded)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<OpenMPThreadRunner>(num_threads);
+ })
+ ->Threads(1)->Threads(2)->Threads(4);
+```
+The above example creates a parallel OpenMP region before it enters `BM_MultiThreaded`.
+The actual benchmark code can remain the same and is therefore not tied to a specific
+thread runner. The measurement does not include the time for creating and joining the
+threads.
+
<a name="cpu-timers" />
## CPU Timers
diff --git a/include/benchmark/benchmark.h b/include/benchmark/benchmark.h
index 624ab29..5e40975 100644
--- a/include/benchmark/benchmark.h
+++ b/include/benchmark/benchmark.h
@@ -1093,8 +1093,18 @@
return StateIterator();
}
+// Base class for user-defined multi-threading
+struct ThreadRunnerBase {
+ virtual ~ThreadRunnerBase() {}
+ virtual void RunThreads(const std::function<void(int)>& fn) = 0;
+};
+
namespace internal {
+// Define alias of ThreadRunner factory function type
+using threadrunner_factory =
+ std::function<std::unique_ptr<ThreadRunnerBase>(int)>;
+
typedef void(Function)(State&);
// ------------------------------------------------------
@@ -1299,6 +1309,9 @@
// Equivalent to ThreadRange(NumCPUs(), NumCPUs())
Benchmark* ThreadPerCpu();
+ // Sets a user-defined threadrunner (see ThreadRunnerBase)
+ Benchmark* ThreadRunner(threadrunner_factory&& factory);
+
virtual void Run(State& state) = 0;
TimeUnit GetTimeUnit() const;
@@ -1340,6 +1353,8 @@
callback_function setup_;
callback_function teardown_;
+ threadrunner_factory threadrunner_;
+
BENCHMARK_DISALLOW_COPY_AND_ASSIGN(Benchmark);
};
diff --git a/src/benchmark_api_internal.h b/src/benchmark_api_internal.h
index 82ab71f..efa0602 100644
--- a/src/benchmark_api_internal.h
+++ b/src/benchmark_api_internal.h
@@ -41,6 +41,9 @@
int threads() const { return threads_; }
void Setup() const;
void Teardown() const;
+ const auto& GetUserThreadRunnerFactory() const {
+ return benchmark_.threadrunner_;
+ }
State Run(IterationCount iters, int thread_id, internal::ThreadTimer* timer,
internal::ThreadManager* manager,
diff --git a/src/benchmark_register.cc b/src/benchmark_register.cc
index 8b94540..d8cefe4 100644
--- a/src/benchmark_register.cc
+++ b/src/benchmark_register.cc
@@ -484,6 +484,11 @@
return this;
}
+Benchmark* Benchmark::ThreadRunner(threadrunner_factory&& factory) {
+ threadrunner_ = std::move(factory);
+ return this;
+}
+
void Benchmark::SetName(const std::string& name) { name_ = name; }
const char* Benchmark::GetName() const { return name_.c_str(); }
diff --git a/src/benchmark_runner.cc b/src/benchmark_runner.cc
index 55ad69c..427bd85 100644
--- a/src/benchmark_runner.cc
+++ b/src/benchmark_runner.cc
@@ -34,6 +34,7 @@
#include <cstdio>
#include <cstdlib>
#include <fstream>
+#include <functional>
#include <iostream>
#include <limits>
#include <memory>
@@ -182,6 +183,38 @@
return iters_or_time.iters;
}
+class ThreadRunnerDefault : public ThreadRunnerBase {
+ public:
+ explicit ThreadRunnerDefault(int num_threads)
+ : pool(static_cast<size_t>(num_threads - 1)) {}
+
+ void RunThreads(const std::function<void(int)>& fn) final {
+ // Run all but one thread in separate threads
+ for (std::size_t ti = 0; ti < pool.size(); ++ti) {
+ pool[ti] = std::thread(fn, static_cast<int>(ti + 1));
+ }
+ // And run one thread here directly.
+ // (If we were asked to run just one thread, we don't create new threads.)
+ // Yes, we need to do this here *after* we start the separate threads.
+ fn(0);
+
+ // The main thread has finished. Now let's wait for the other threads.
+ for (std::thread& thread : pool) {
+ thread.join();
+ }
+ }
+
+ private:
+ std::vector<std::thread> pool;
+};
+
+std::unique_ptr<ThreadRunnerBase> GetThreadRunner(
+ const threadrunner_factory& userThreadRunnerFactory, int num_threads) {
+ return userThreadRunnerFactory
+ ? userThreadRunnerFactory(num_threads)
+ : std::make_unique<ThreadRunnerDefault>(num_threads);
+}
+
} // end namespace
BenchTimeType ParseBenchMinTime(const std::string& value) {
@@ -258,7 +291,8 @@
has_explicit_iteration_count(b.iterations() != 0 ||
parsed_benchtime_flag.tag ==
BenchTimeType::ITERS),
- pool(static_cast<size_t>(b.threads() - 1)),
+ thread_runner(
+ GetThreadRunner(b.GetUserThreadRunnerFactory(), b.threads())),
iters(FLAGS_benchmark_dry_run
? 1
: (has_explicit_iteration_count
@@ -289,22 +323,10 @@
std::unique_ptr<internal::ThreadManager> manager;
manager.reset(new internal::ThreadManager(b.threads()));
- // Run all but one thread in separate threads
- for (std::size_t ti = 0; ti < pool.size(); ++ti) {
- pool[ti] = std::thread(&RunInThread, &b, iters, static_cast<int>(ti + 1),
- manager.get(), perf_counters_measurement_ptr,
- /*profiler_manager=*/nullptr);
- }
- // And run one thread here directly.
- // (If we were asked to run just one thread, we don't create new threads.)
- // Yes, we need to do this here *after* we start the separate threads.
- RunInThread(&b, iters, 0, manager.get(), perf_counters_measurement_ptr,
- /*profiler_manager=*/nullptr);
-
- // The main thread has finished. Now let's wait for the other threads.
- for (std::thread& thread : pool) {
- thread.join();
- }
+ thread_runner->RunThreads([&](int thread_idx) {
+ RunInThread(&b, iters, thread_idx, manager.get(),
+ perf_counters_measurement_ptr, /*profiler_manager=*/nullptr);
+ });
IterationResults i;
// Acquire the measurements/counters from the manager, UNDER THE LOCK!
diff --git a/src/benchmark_runner.h b/src/benchmark_runner.h
index bc76c81..9a2231a 100644
--- a/src/benchmark_runner.h
+++ b/src/benchmark_runner.h
@@ -15,6 +15,7 @@
#ifndef BENCHMARK_RUNNER_H_
#define BENCHMARK_RUNNER_H_
+#include <memory>
#include <thread>
#include <vector>
@@ -89,7 +90,7 @@
int num_repetitions_done = 0;
- std::vector<std::thread> pool;
+ std::unique_ptr<ThreadRunnerBase> thread_runner;
IterationCount iters; // preserved between repetitions!
// So only the first repetition has to find/calculate it,
diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt
index df575d9..bde248f 100644
--- a/test/CMakeLists.txt
+++ b/test/CMakeLists.txt
@@ -189,6 +189,9 @@
compile_output_test(internal_threading_test)
benchmark_add_test(NAME internal_threading_test COMMAND internal_threading_test --benchmark_min_time=0.01s)
+compile_output_test(manual_threading_test)
+benchmark_add_test(NAME manual_threading_test COMMAND manual_threading_test --benchmark_min_time=0.01s)
+
compile_output_test(report_aggregates_only_test)
benchmark_add_test(NAME report_aggregates_only_test COMMAND report_aggregates_only_test --benchmark_min_time=0.01s)
diff --git a/test/manual_threading_test.cc b/test/manual_threading_test.cc
new file mode 100644
index 0000000..e85d495
--- /dev/null
+++ b/test/manual_threading_test.cc
@@ -0,0 +1,174 @@
+
+#include <memory>
+#undef NDEBUG
+
+#include <chrono>
+#include <thread>
+
+#include "../src/timers.h"
+#include "benchmark/benchmark.h"
+
+namespace {
+
+const std::chrono::duration<double, std::milli> time_frame(50);
+const double time_frame_in_sec(
+ std::chrono::duration_cast<std::chrono::duration<double, std::ratio<1, 1>>>(
+ time_frame)
+ .count());
+
+void MyBusySpinwait() {
+ const auto start = benchmark::ChronoClockNow();
+
+ while (true) {
+ const auto now = benchmark::ChronoClockNow();
+ const auto elapsed = now - start;
+
+ if (std::chrono::duration<double, std::chrono::seconds::period>(elapsed) >=
+ time_frame) {
+ return;
+ }
+ }
+}
+
+int numRunThreadsCalled_ = 0;
+
+class ManualThreadRunner : public benchmark::ThreadRunnerBase {
+ public:
+ explicit ManualThreadRunner(int num_threads)
+ : pool(static_cast<size_t>(num_threads - 1)) {}
+
+ void RunThreads(const std::function<void(int)>& fn) final {
+ for (std::size_t ti = 0; ti < pool.size(); ++ti) {
+ pool[ti] = std::thread(fn, static_cast<int>(ti + 1));
+ }
+
+ fn(0);
+
+ for (std::thread& thread : pool) {
+ thread.join();
+ }
+
+ ++numRunThreadsCalled_;
+ }
+
+ private:
+ std::vector<std::thread> pool;
+};
+
+// ========================================================================= //
+// --------------------------- TEST CASES BEGIN ---------------------------- //
+// ========================================================================= //
+
+// ========================================================================= //
+// BM_ManualThreading
+// Creation of threads is done before the start of the measurement,
+// joining after the finish of the measurement.
+void BM_ManualThreading(benchmark::State& state) {
+ for (auto _ : state) {
+ MyBusySpinwait();
+ state.SetIterationTime(time_frame_in_sec);
+ }
+ state.counters["invtime"] =
+ benchmark::Counter{1, benchmark::Counter::kIsRate};
+}
+
+} // end namespace
+
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(1);
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(1)
+ ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(1)
+ ->UseManualTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(1)
+ ->MeasureProcessCPUTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(1)
+ ->MeasureProcessCPUTime()
+ ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(1)
+ ->MeasureProcessCPUTime()
+ ->UseManualTime();
+
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(2);
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(2)
+ ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(2)
+ ->UseManualTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(2)
+ ->MeasureProcessCPUTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(2)
+ ->MeasureProcessCPUTime()
+ ->UseRealTime();
+BENCHMARK(BM_ManualThreading)
+ ->Iterations(1)
+ ->ThreadRunner([](int num_threads) {
+ return std::make_unique<ManualThreadRunner>(num_threads);
+ })
+ ->Threads(2)
+ ->MeasureProcessCPUTime()
+ ->UseManualTime();
+
+// ========================================================================= //
+// ---------------------------- TEST CASES END ----------------------------- //
+// ========================================================================= //
+
+int main(int argc, char* argv[]) {
+ benchmark::Initialize(&argc, argv);
+ benchmark::RunSpecifiedBenchmarks();
+ benchmark::Shutdown();
+ assert(numRunThreadsCalled_ > 0);
+}