Add new mju_threadpool API function, and delete old threading API.
PiperOrigin-RevId: 922838541 Change-Id: Id9f7e0fb298ffde61fcc49a802dc78971858ce51
This commit is contained in:
committed by
Copybara-Service
parent
a22fc2423a
commit
b935d4153c
@@ -15,8 +15,8 @@
|
||||
// A benchmark for parsing and compiling models from XML.
|
||||
|
||||
#include <cstddef>
|
||||
#include <vector>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include <benchmark/benchmark.h>
|
||||
#include <absl/base/attributes.h>
|
||||
@@ -36,6 +36,9 @@ static const int kNumWarmupSteps = 1000;
|
||||
// number of steps to benchmark (before resetting state)
|
||||
static const int kBatchSize = 50;
|
||||
|
||||
// number of threads to test
|
||||
static const int kNumThreads = 9;
|
||||
|
||||
static const char kBoxMeshPath[] =
|
||||
"../test/engine/testdata/collision_convex/perf/boxmesh.xml";
|
||||
static const char kBoxBoxPath[] =
|
||||
@@ -47,11 +50,10 @@ static const char kMixedPath[] =
|
||||
|
||||
class TestHarness {
|
||||
public:
|
||||
TestHarness(const char* xml_path, std::string label, int disable_flags = 0) {
|
||||
TestHarness(const char* xml_path, std::string label) {
|
||||
// Fail test if there are any mujoco errors
|
||||
MujocoErrorTestGuard guard;
|
||||
model_ = LoadModelFromPath(xml_path);
|
||||
model_->opt.disableflags |= disable_flags;
|
||||
data_ = mj_makeData(model_);
|
||||
for (int i=0; i < kNumWarmupSteps; i++) {
|
||||
mj_step(model_, data_);
|
||||
@@ -62,13 +64,18 @@ class TestHarness {
|
||||
int size = mj_stateSize(model_, spec_);
|
||||
initial_state_.resize(size);
|
||||
mj_getState(model_, data_, initial_state_.data(), spec_);
|
||||
label_ = label;
|
||||
name_ = label;
|
||||
}
|
||||
|
||||
void Reset() {
|
||||
mj_setState(model_, data_, initial_state_.data(), spec_);
|
||||
}
|
||||
|
||||
void SetThreads(int nthread) {
|
||||
mju_threadpool(data_, nthread);
|
||||
nthread_ = nthread;
|
||||
}
|
||||
|
||||
void RunBenchmark(benchmark::State& state) {
|
||||
std::size_t ncon = 0;
|
||||
while (state.KeepRunningBatch(kBatchSize)) {
|
||||
@@ -81,7 +88,9 @@ class TestHarness {
|
||||
}
|
||||
}
|
||||
|
||||
state.SetLabel(label_);
|
||||
std::string label = name_ + " " + std::to_string(nthread_ + 1) +
|
||||
" thread(s)";
|
||||
state.SetLabel(label);
|
||||
state.SetItemsProcessed(ncon); // report number of contacts per second
|
||||
}
|
||||
|
||||
@@ -91,7 +100,8 @@ class TestHarness {
|
||||
}
|
||||
|
||||
private:
|
||||
std::string label_;
|
||||
std::string name_;
|
||||
int nthread_ = 0;
|
||||
int spec_;
|
||||
mjModel* model_;
|
||||
mjData* data_;
|
||||
@@ -103,62 +113,47 @@ class TestHarness {
|
||||
// separately in CPU profiles (and don't get replaced with raw calls to
|
||||
// run_parse_benchmark).
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL
|
||||
BM_BoxMesh_NativeCCD(benchmark::State& state) {
|
||||
static TestHarness harness(kBoxMeshPath, "boxmesh.xml (nativeccd)");
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_BoxMesh(benchmark::State& state) {
|
||||
int nthread = state.range(0);
|
||||
static TestHarness harness(kBoxMeshPath, "boxmesh.xml");
|
||||
harness.SetThreads(nthread);
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_BoxMesh_NativeCCD);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL
|
||||
BM_BoxMesh_LibCCD(benchmark::State& state) {
|
||||
static TestHarness harness(kBoxMeshPath, "boxmesh.xml (libccd)",
|
||||
mjDSBL_NATIVECCD);
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_BoxMesh_LibCCD);
|
||||
BENCHMARK(BM_BoxMesh)->Arg(0)->Arg(kNumThreads);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_BoxBox(benchmark::State& state) {
|
||||
static TestHarness harness(kBoxBoxPath, "box.xml (BoxBox)", mjDSBL_NATIVECCD);
|
||||
int nthread = state.range(0);
|
||||
static TestHarness harness(kBoxBoxPath, "box.xml (BoxBox)");
|
||||
harness.SetThreads(nthread);
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_BoxBox);
|
||||
BENCHMARK(BM_BoxBox)->Arg(0)->Arg(kNumThreads);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_BoxBox_NativeCCD(benchmark::State& state) {
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_BoxBoxConvex(benchmark::State& state) {
|
||||
int nthread = state.range(0);
|
||||
mjCOLLISIONFUNC[mjGEOM_BOX][mjGEOM_BOX] = mjc_Convex;
|
||||
static TestHarness harness(kBoxBoxPath, "box.xml (NativeCCD)");
|
||||
harness.SetThreads(nthread);
|
||||
harness.RunBenchmark(state);
|
||||
mjCOLLISIONFUNC[mjGEOM_BOX][mjGEOM_BOX] = mjc_BoxBox;
|
||||
}
|
||||
BENCHMARK(BM_BoxBox_NativeCCD);
|
||||
BENCHMARK(BM_BoxBoxConvex)->Arg(0)->Arg(kNumThreads);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL
|
||||
BM_Ellipsoid_NativeCCD(benchmark::State& state) {
|
||||
static TestHarness harness(kEllipsoidPath, "ellipsoid.xml (nativeccd)");
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_Ellipsoid(benchmark::State& state) {
|
||||
int nthread = state.range(0);
|
||||
static TestHarness harness(kEllipsoidPath, "ellipsoid.xml");
|
||||
harness.SetThreads(nthread);
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_Ellipsoid_NativeCCD);
|
||||
BENCHMARK(BM_Ellipsoid)->Arg(0)->Arg(kNumThreads);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL
|
||||
BM_Ellipsoid_LibCCD(benchmark::State& state) {
|
||||
static TestHarness harness(kEllipsoidPath, "ellipsoid.xml (libccd)",
|
||||
mjDSBL_NATIVECCD);
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_Mixed(benchmark::State& state) {
|
||||
int nthread = state.range(0);
|
||||
static TestHarness harness(kMixedPath, "mixed.xml");
|
||||
harness.SetThreads(nthread);
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_Ellipsoid_LibCCD);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_Mixed_NativeCCD(benchmark::State& state) {
|
||||
static TestHarness harness(kMixedPath, "mixed.xml (nativeccd)");
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_Mixed_NativeCCD);
|
||||
|
||||
void ABSL_ATTRIBUTE_NO_TAIL_CALL BM_Mixed_LibCCD(benchmark::State& state) {
|
||||
static TestHarness harness(kMixedPath, "mixed.xml (libccd)",
|
||||
mjDSBL_NATIVECCD);
|
||||
harness.RunBenchmark(state);
|
||||
}
|
||||
BENCHMARK(BM_Mixed_LibCCD);
|
||||
BENCHMARK(BM_Mixed)->Arg(0)->Arg(kNumThreads);
|
||||
|
||||
} // namespace
|
||||
} // namespace mujoco
|
||||
|
||||
@@ -17,7 +17,6 @@
|
||||
|
||||
#include <benchmark/benchmark.h>
|
||||
#include <mujoco/mjmodel.h>
|
||||
#include <mujoco/mjthread.h>
|
||||
#include <mujoco/mujoco.h>
|
||||
#include "test/fixture.h"
|
||||
|
||||
@@ -30,6 +29,9 @@ static const int kNumWarmupSteps = 500;
|
||||
// number of steps to benchmark (before resetting state)
|
||||
static const int kBatchSize = 50;
|
||||
|
||||
// number of threads to test
|
||||
static const int kNumThreads = 6;
|
||||
|
||||
void BM_StepHumanoid200(benchmark::State& state) {
|
||||
int nthread = state.range(0);
|
||||
std::string label = std::to_string(nthread) + " thread(s)";
|
||||
@@ -43,10 +45,8 @@ void BM_StepHumanoid200(benchmark::State& state) {
|
||||
model->opt.disableflags &= ~mjDSBL_ISLAND; // enable islands
|
||||
|
||||
mjData* data = mj_makeData(model);
|
||||
mjThreadPool* threadpool = nullptr;
|
||||
if (nthread > 1) {
|
||||
threadpool = mju_threadPoolCreate(nthread);
|
||||
mju_bindThreadPool(data, threadpool);
|
||||
if (nthread) {
|
||||
mju_threadpool(data, nthread);
|
||||
}
|
||||
|
||||
// warm-up rollout to get a steady state
|
||||
@@ -73,9 +73,6 @@ void BM_StepHumanoid200(benchmark::State& state) {
|
||||
state.SetLabel(label);
|
||||
state.SetItemsProcessed(state.iterations());
|
||||
mj_deleteData(data);
|
||||
if (threadpool) {
|
||||
mju_threadPoolDestroy(threadpool);
|
||||
}
|
||||
}
|
||||
|
||||
void BM_Step22Humanoids(benchmark::State& state) {
|
||||
@@ -90,10 +87,8 @@ void BM_Step22Humanoids(benchmark::State& state) {
|
||||
model->opt.disableflags &= ~mjDSBL_ISLAND; // enable islands
|
||||
|
||||
mjData* data = mj_makeData(model);
|
||||
mjThreadPool* threadpool = nullptr;
|
||||
if (nthread > 1) {
|
||||
threadpool = mju_threadPoolCreate(nthread);
|
||||
mju_bindThreadPool(data, threadpool);
|
||||
if (nthread) {
|
||||
mju_threadpool(data, nthread);
|
||||
}
|
||||
|
||||
// warm-up rollout to get a steady state
|
||||
@@ -123,12 +118,9 @@ void BM_Step22Humanoids(benchmark::State& state) {
|
||||
state.SetLabel(label);
|
||||
state.SetItemsProcessed(state.iterations());
|
||||
mj_deleteData(data);
|
||||
if (threadpool) {
|
||||
mju_threadPoolDestroy(threadpool);
|
||||
}
|
||||
}
|
||||
|
||||
BENCHMARK(BM_StepHumanoid200)->Arg(1)->Arg(6);
|
||||
BENCHMARK(BM_Step22Humanoids)->Arg(1)->Arg(6);
|
||||
BENCHMARK(BM_StepHumanoid200)->Arg(0)->Arg(kNumThreads);
|
||||
BENCHMARK(BM_Step22Humanoids)->Arg(0)->Arg(kNumThreads);
|
||||
} // namespace
|
||||
} // namespace mujoco
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
#include <mujoco/mjxmacro.h>
|
||||
#include <mujoco/mujoco.h>
|
||||
#include "src/engine/engine_util_errmem.h"
|
||||
#include "src/thread/thread_pool.h"
|
||||
#include "src/engine/engine_thread.h"
|
||||
#include "test/fixture.h"
|
||||
|
||||
namespace mujoco {
|
||||
@@ -804,32 +804,18 @@ TEST_F(EngineIoTest, VeryLargeMemory) {
|
||||
}
|
||||
}
|
||||
|
||||
struct TestFunctionArgs_ {
|
||||
mjData* d;
|
||||
int input;
|
||||
int stack_output;
|
||||
int arena_output;
|
||||
size_t output_thread_worker;
|
||||
struct TestFunctionArgs {
|
||||
int stack_output[1000];
|
||||
int arena_output[1000];
|
||||
};
|
||||
typedef TestFunctionArgs_ TestFunctionArgs;
|
||||
|
||||
void* TestFunction(void* args) {
|
||||
void TestFunction(const mjModel* m, mjData* d, void* args, int i, int j) {
|
||||
TestFunctionArgs* test_args = static_cast<TestFunctionArgs*>(args);
|
||||
test_args->output_thread_worker =
|
||||
mju_threadPoolCurrentWorkerId((mjThreadPool*)test_args->d->threadpool);
|
||||
mj_markStack(test_args->d);
|
||||
int* test_ints = mj_stackAllocInt(test_args->d, 10);
|
||||
test_ints[0] = test_args->input;
|
||||
test_args->stack_output = test_ints[0];
|
||||
|
||||
int* test_arena_ints =
|
||||
(int*)mj_arenaAllocByte(test_args->d, sizeof(int) * 10, alignof(int));
|
||||
test_arena_ints[0] = test_args->input;
|
||||
test_args->arena_output = test_arena_ints[0];
|
||||
|
||||
|
||||
mj_freeStack(test_args->d);
|
||||
return nullptr;
|
||||
mj_markStack(d);
|
||||
int* test_ints = mj_stackAllocInt(d, 10);
|
||||
test_ints[0] = j;
|
||||
test_args->stack_output[j] = test_ints[0];
|
||||
mj_freeStack(d);
|
||||
}
|
||||
|
||||
TEST_F(EngineIoTest, TestStackShardingForThreads) {
|
||||
@@ -846,38 +832,18 @@ TEST_F(EngineIoTest, TestStackShardingForThreads) {
|
||||
|
||||
mjData* data = mj_makeData(model);
|
||||
ASSERT_THAT(data, NotNull());
|
||||
mjThreadPool* thread_pool = mju_threadPoolCreate(10);
|
||||
mju_bindThreadPool(data, thread_pool);
|
||||
mju_threadpool(data, 10);
|
||||
|
||||
constexpr int kTasks = 1000;
|
||||
TestFunctionArgs test_function_args[kTasks];
|
||||
mjTask tasks[kTasks];
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
test_function_args[i].d = data;
|
||||
test_function_args[i].input = i;
|
||||
mju_defaultTask(&tasks[i]);
|
||||
tasks[i].func = TestFunction;
|
||||
tasks[i].args = &test_function_args[i];
|
||||
mju_threadPoolEnqueue(thread_pool, &tasks[i]);
|
||||
}
|
||||
|
||||
mj_markStack(data);
|
||||
int* test_ints = mj_stackAllocInt(data, 10);
|
||||
test_ints[0] = 1;
|
||||
mj_freeStack(data);
|
||||
TestFunctionArgs test_function_args;
|
||||
mju_dispatch(model, data, TestFunction, &test_function_args, kTasks);
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
mju_taskJoin(&tasks[i]);
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
EXPECT_EQ(test_function_args[i].input, test_function_args[i].stack_output);
|
||||
EXPECT_EQ(test_function_args[i].input, test_function_args[i].arena_output);
|
||||
EXPECT_EQ(i, test_function_args.stack_output[i]);
|
||||
}
|
||||
|
||||
mj_deleteData(data);
|
||||
mj_deleteModel(model);
|
||||
mju_threadPoolDestroy(thread_pool);
|
||||
}
|
||||
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
|
||||
@@ -22,10 +22,9 @@
|
||||
#include <gmock/gmock.h>
|
||||
#include <gtest/gtest.h>
|
||||
#include <mujoco/mjmodel.h>
|
||||
#include <mujoco/mjthread.h>
|
||||
#include <mujoco/mjxmacro.h>
|
||||
#include <mujoco/mujoco.h>
|
||||
#include "src/thread/thread_pool.h"
|
||||
#include "src/engine/engine_thread.h"
|
||||
#include "test/fixture.h"
|
||||
|
||||
namespace mujoco {
|
||||
@@ -61,8 +60,7 @@ TEST_F(ThreadTest, SingleAndMultiThreadedMatch) {
|
||||
mj_setState(model_threaded, data_threaded, initial_state.data(), spec);
|
||||
|
||||
// bind a threadpool to the data_threaded
|
||||
mjThreadPool* threadpool = mju_threadPoolCreate(10);
|
||||
mju_bindThreadPool(data_threaded, threadpool);
|
||||
mju_threadpool(data_threaded, 10);
|
||||
|
||||
for (int i = 0; i < 10; ++i) {
|
||||
mj_step(model, data);
|
||||
@@ -83,7 +81,6 @@ TEST_F(ThreadTest, SingleAndMultiThreadedMatch) {
|
||||
mj_deleteModel(model);
|
||||
mj_deleteData(data_threaded);
|
||||
mj_deleteModel(model_threaded);
|
||||
mju_threadPoolDestroy(threadpool);
|
||||
}
|
||||
|
||||
TEST_F(ThreadTest, IslandSingleAndMultiThreadedMatch) {
|
||||
@@ -116,8 +113,7 @@ TEST_F(ThreadTest, IslandSingleAndMultiThreadedMatch) {
|
||||
mj_setState(model_threaded, data_threaded, initial_state.data(), spec);
|
||||
|
||||
// bind a threadpool to the data_threaded
|
||||
mjThreadPool* threadpool = mju_threadPoolCreate(10);
|
||||
mju_bindThreadPool(data_threaded, threadpool);
|
||||
mju_threadpool(data_threaded, 10);
|
||||
|
||||
for (int i = 0; i < 10; ++i) {
|
||||
mj_step(model, data);
|
||||
@@ -138,7 +134,6 @@ TEST_F(ThreadTest, IslandSingleAndMultiThreadedMatch) {
|
||||
mj_deleteModel(model);
|
||||
mj_deleteData(data_threaded);
|
||||
mj_deleteModel(model_threaded);
|
||||
mju_threadPoolDestroy(threadpool);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
# Copyright 2023 DeepMind Technologies Limited
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# https://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
mujoco_test(thread_pool_test)
|
||||
|
||||
mujoco_test(thread_queue_test)
|
||||
@@ -1,142 +0,0 @@
|
||||
// Copyright 2023 DeepMind Technologies Limited
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#include <atomic>
|
||||
#include <condition_variable>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <thread>
|
||||
|
||||
#include <gtest/gtest.h>
|
||||
#include <mujoco/mujoco.h>
|
||||
|
||||
namespace mujoco {
|
||||
namespace {
|
||||
|
||||
struct TestFunctionArgs_ {
|
||||
int input;
|
||||
// make this atomic to avoid red-herring tsan failures.
|
||||
std::atomic<int> output;
|
||||
};
|
||||
typedef struct TestFunctionArgs_ TestFunctionArgs;
|
||||
|
||||
void* test_function(void* args) {
|
||||
TestFunctionArgs* test_function_args = static_cast<TestFunctionArgs*>(args);
|
||||
if (!test_function_args) {
|
||||
return nullptr;
|
||||
}
|
||||
test_function_args->output = test_function_args->input;
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
TEST(TestMjThreadPool, TestMjThreadPool10Threads) {
|
||||
mjThreadPool* thread_pool = mju_threadPoolCreate(10);
|
||||
|
||||
constexpr int kTasks = 1000;
|
||||
TestFunctionArgs test_function_args[kTasks];
|
||||
mjTask tasks[kTasks];
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
test_function_args[i].input = i;
|
||||
mju_defaultTask(&tasks[i]);
|
||||
tasks[i].func = test_function;
|
||||
tasks[i].args = &test_function_args[i];
|
||||
mju_threadPoolEnqueue(thread_pool, &tasks[i]);
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
mju_taskJoin(&tasks[i]);
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
EXPECT_EQ(test_function_args[i].input, test_function_args[i].output);
|
||||
}
|
||||
mju_threadPoolDestroy(thread_pool);
|
||||
}
|
||||
|
||||
TEST(TestMjThreadPool, TestMjThreadPool100Threads) {
|
||||
mjThreadPool* thread_pool = mju_threadPoolCreate(100);
|
||||
|
||||
constexpr int kTasks = 1000;
|
||||
TestFunctionArgs test_function_args[kTasks];
|
||||
mjTask tasks[kTasks];
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
test_function_args[i].input = i;
|
||||
mju_defaultTask(&tasks[i]);
|
||||
tasks[i].func = test_function;
|
||||
tasks[i].args = &test_function_args[i];
|
||||
mju_threadPoolEnqueue(thread_pool, &tasks[i]);
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
mju_taskJoin(&tasks[i]);
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
EXPECT_EQ(test_function_args[i].input, test_function_args[i].output);
|
||||
}
|
||||
|
||||
mju_threadPoolDestroy(thread_pool);
|
||||
}
|
||||
|
||||
TEST(TestMjThreadPool, TestMjThreadPoolManyWriters) {
|
||||
mjThreadPool* thread_pool = mju_threadPoolCreate(10);
|
||||
|
||||
constexpr int kTasks = 20;
|
||||
TestFunctionArgs test_function_args[kTasks];
|
||||
mjTask tasks[kTasks];
|
||||
std::unique_ptr<std::thread> enqueue_threads[kTasks];
|
||||
|
||||
// add tasks to the thread pool from many threads
|
||||
std::condition_variable start_cv;
|
||||
std::mutex start_mutex;
|
||||
bool start = false;
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
test_function_args[i].input = i;
|
||||
mju_defaultTask(&tasks[i]);
|
||||
tasks[i].func = &test_function;
|
||||
tasks[i].args = &test_function_args[i];
|
||||
|
||||
enqueue_threads[i] = std::make_unique<std::thread>([&, i] {
|
||||
// synchronize all threads adding to the thread_pool at the same time
|
||||
{
|
||||
std::unique_lock<std::mutex> lock(start_mutex);
|
||||
start_cv.wait(lock, [&] { return start; });
|
||||
}
|
||||
// enqueue outside the lock, to get some concurrency
|
||||
mju_threadPoolEnqueue(thread_pool, &tasks[i]);
|
||||
});
|
||||
}
|
||||
{
|
||||
std::unique_lock<std::mutex> lock(start_mutex);
|
||||
start = true;
|
||||
}
|
||||
start_cv.notify_all();
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
enqueue_threads[i]->join();
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
mju_taskJoin(&tasks[i]);
|
||||
}
|
||||
|
||||
for (int i = 0; i < kTasks; ++i) {
|
||||
EXPECT_EQ(test_function_args[i].input, test_function_args[i].output);
|
||||
}
|
||||
|
||||
mju_threadPoolDestroy(thread_pool);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
} // namespace mujoco
|
||||
@@ -1,46 +0,0 @@
|
||||
// Copyright 2023 DeepMind Technologies Limited
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#include "src/thread/thread_queue.h"
|
||||
|
||||
#include <cstddef>
|
||||
|
||||
#include <gtest/gtest.h>
|
||||
|
||||
namespace mujoco {
|
||||
namespace {
|
||||
|
||||
constexpr size_t kBufferCapacity = 640;
|
||||
|
||||
TEST(TestMujocoLocklessQueue, TestMujocoLocklessQueue) {
|
||||
LocklessQueue<void*, 640> queue;
|
||||
EXPECT_TRUE(queue.empty());
|
||||
int test_integers[kBufferCapacity];
|
||||
for (int h = 0; h < 10; ++h) {
|
||||
for (int i = 0; i < kBufferCapacity; ++i) {
|
||||
test_integers[i] = i;
|
||||
queue.push(&test_integers[i]);
|
||||
}
|
||||
EXPECT_TRUE(queue.full());
|
||||
|
||||
for (int i = 0; i < kBufferCapacity; ++i) {
|
||||
void* output_ptr = queue.pop();
|
||||
ASSERT_EQ(output_ptr, &test_integers[i]);
|
||||
}
|
||||
EXPECT_TRUE(queue.empty());
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
} // namespace mujoco
|
||||
Reference in New Issue
Block a user