Implement threading for island constraint solving.
Humanoids22 No threads: Benchmark Execution Time: 9.475905812s Humanoids22 with 10 Threads: Benchmark Execution Time: 4.871352s PiperOrigin-RevId: 571214307 Change-Id: I1f4f2c761b4ae6bc8fac1f28c6c695f6c499f339
This commit is contained in:
committed by
Copybara-Service
parent
06670155c4
commit
071af3b015
@@ -14,6 +14,7 @@
|
||||
|
||||
#include "engine/engine_collision_driver.h"
|
||||
|
||||
#include <math.h>
|
||||
#include <stddef.h>
|
||||
#include <string.h>
|
||||
|
||||
@@ -61,8 +62,10 @@ static void collideGeoms(const mjModel* m, mjData* d,
|
||||
static inline void resetArena(mjData* d) {
|
||||
d->parena = d->ncon * sizeof(mjContact);
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
ASAN_POISON_MEMORY_REGION(
|
||||
(char*)d->arena + d->parena, d->narena - d->pstack - d->parena);
|
||||
if (!d->threadpool) {
|
||||
ASAN_POISON_MEMORY_REGION(
|
||||
(char*)d->arena + d->parena, d->narena - d->pstack - d->parena);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -40,6 +40,8 @@
|
||||
#include "engine/engine_util_misc.h"
|
||||
#include "engine/engine_util_solve.h"
|
||||
#include "engine/engine_util_sparse.h"
|
||||
#include "thread/thread_pool.h"
|
||||
#include "thread/thread_task.h"
|
||||
|
||||
|
||||
|
||||
@@ -492,6 +494,52 @@ static void warmstart(const mjModel* m, mjData* d) {
|
||||
|
||||
|
||||
|
||||
// struct encapsulating arguments to thread task
|
||||
struct mjSolIslandArgs_ {
|
||||
const mjModel* m;
|
||||
mjData* d;
|
||||
int island;
|
||||
};
|
||||
typedef struct mjSolIslandArgs_ mjSolIslandArgs;
|
||||
|
||||
// extract arguments, pass to solver
|
||||
void* mj_solCG_island_wrapper(void* args) {
|
||||
mjSolIslandArgs* solargs = (mjSolIslandArgs*) args;
|
||||
mj_solCG_island(solargs->m, solargs->d, solargs->island, solargs->m->opt.iterations);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
// CG solver, multi-threaded over islands
|
||||
void mj_solCG_island_multithreaded(const mjModel* m, mjData* d) {
|
||||
mj_markStack(d);
|
||||
// allocate array of arguments to be passed to threads
|
||||
mjSolIslandArgs* sol_cg_island_args =
|
||||
mj_stackAllocByte(d, sizeof(mjSolIslandArgs) * d->nisland, _Alignof(mjSolIslandArgs));
|
||||
mjTask* tasks = mj_stackAllocByte(d, sizeof(mjTask) * d->nisland, _Alignof(mjTask));
|
||||
|
||||
for (int island = 0; island < d->nisland; ++island) {
|
||||
sol_cg_island_args[island].m = m;
|
||||
sol_cg_island_args[island].d = d;
|
||||
sol_cg_island_args[island].island = island;
|
||||
|
||||
mju_defaultTask(&tasks[island]);
|
||||
tasks[island].func = mj_solCG_island_wrapper;
|
||||
tasks[island].args = &sol_cg_island_args[island];
|
||||
mju_threadPoolEnqueue((mjThreadPool*)d->threadpool, &tasks[island]);
|
||||
}
|
||||
|
||||
for (int island = 0; island < d->nisland; ++island) {
|
||||
mju_taskJoin(&tasks[island]);
|
||||
}
|
||||
|
||||
mj_freeStack(d);
|
||||
}
|
||||
|
||||
|
||||
|
||||
// compute efc_b, efc_force, qfrc_constraint; update qacc
|
||||
void mj_fwdConstraint(const mjModel* m, mjData* d) {
|
||||
TM_START;
|
||||
@@ -524,9 +572,15 @@ void mj_fwdConstraint(const mjModel* m, mjData* d) {
|
||||
|
||||
// run solver over constraint islands
|
||||
if (islands_supported) {
|
||||
// loop over islands
|
||||
for (int island=0; island < nisland; island++) {
|
||||
mj_solCG_island(m, d, island, m->opt.iterations);
|
||||
// no threadpool, loop over islands
|
||||
if (!d->threadpool) {
|
||||
for (int island=0; island < nisland; island++) {
|
||||
mj_solCG_island(m, d, island, m->opt.iterations);
|
||||
}
|
||||
}
|
||||
else {
|
||||
// solve using threads
|
||||
mj_solCG_island_multithreaded(m, d);
|
||||
}
|
||||
d->solver_nisland = nisland;
|
||||
}
|
||||
|
||||
@@ -1470,7 +1470,9 @@ static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
|
||||
//------------------------------ clear header
|
||||
|
||||
// clear stack pointer
|
||||
d->pstack = 0;
|
||||
if (!d->threadpool) {
|
||||
d->pstack = 0;
|
||||
}
|
||||
d->pbase = 0;
|
||||
|
||||
// clear arena pointers
|
||||
|
||||
@@ -206,7 +206,10 @@ mjStackInfo* mju_getStackInfoForThread(mjData* d, size_t thread_id) {
|
||||
// align the end of the shard to be mjStackInfo.
|
||||
misalignment = result % alignof(mjStackInfo);
|
||||
result -= misalignment;
|
||||
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
// Ensure StackInfo is always accessible
|
||||
ASAN_UNPOISON_MEMORY_REGION((void*)result, sizeof(mjStackInfo));
|
||||
#endif
|
||||
return (mjStackInfo*) result;
|
||||
}
|
||||
|
||||
@@ -253,7 +256,7 @@ static void ConfigureMultiThreadedStack(mjData* d) {
|
||||
}
|
||||
|
||||
// adds a thread pool to mjData and configures it for multi-threaded use.
|
||||
void mju_bindThreadPool(mjData* d, mjThreadPool* thread_pool) {
|
||||
void mju_bindThreadPool(mjData* d, void* thread_pool) {
|
||||
if (d->threadpool) {
|
||||
mju_error("Thread Pool already bound to mjData");
|
||||
}
|
||||
|
||||
@@ -54,7 +54,7 @@ MJAPI mjThreadPool* mju_threadPoolCreate(size_t number_of_threads);
|
||||
mjStackInfo* mju_getStackInfoForThread(mjData* d, size_t thread_id);
|
||||
|
||||
// Adds a thread pool to mjData and configures it for multi-threaded use.
|
||||
MJAPI void mju_bindThreadPool(mjData* d, mjThreadPool* thread_pool);
|
||||
MJAPI void mju_bindThreadPool(mjData* d, void* thread_pool);
|
||||
|
||||
// Gets the number of running threads in the thread pool.
|
||||
MJAPI size_t mju_threadPoolNumberOfThreads(mjThreadPool* thread_pool);
|
||||
|
||||
@@ -15,14 +15,13 @@
|
||||
#ifndef MUJOCO_SRC_THREAD_THREAD_TASK_H_
|
||||
#define MUJOCO_SRC_THREAD_THREAD_TASK_H_
|
||||
|
||||
#include <atomic>
|
||||
#include <new>
|
||||
#include <type_traits>
|
||||
|
||||
#include <mujoco/mjexport.h>
|
||||
#include <mujoco/mjthread.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
#include <atomic>
|
||||
#include <new>
|
||||
#include <type_traits>
|
||||
namespace mujoco {
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
Reference in New Issue
Block a user