Add new mju_threadpool API function, and delete old threading API.

PiperOrigin-RevId: 922838541
Change-Id: Id9f7e0fb298ffde61fcc49a802dc78971858ce51
This commit is contained in:
Kyle Bayes
2026-05-28 10:09:07 -07:00
committed by Copybara-Service
parent a22fc2423a
commit b935d4153c
47 changed files with 576 additions and 1755 deletions
+14 -48
View File
@@ -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
+3 -8
View File
@@ -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