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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user