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
@@ -75,6 +75,8 @@ set(MUJOCO_ENGINE_SRCS
|
||||
engine_sort.h
|
||||
engine_support.c
|
||||
engine_support.h
|
||||
engine_thread.cc
|
||||
engine_thread.h
|
||||
engine_util_blas.c
|
||||
engine_util_blas.h
|
||||
engine_util_blas_avx.h
|
||||
|
||||
@@ -26,6 +26,7 @@
|
||||
#include "engine/engine_collision_gjk.h"
|
||||
#include "engine/engine_collision_primitive.h"
|
||||
#include "engine/engine_collision_sdf.h"
|
||||
#include "engine/engine_thread.h"
|
||||
#include "engine/engine_core_constraint.h"
|
||||
#include "engine/engine_core_util.h"
|
||||
#include "engine/engine_inline.h"
|
||||
@@ -178,10 +179,8 @@ static inline mjtNum getGap(const mjModel* m, int g1, int g2, int ipair) {
|
||||
static inline void resetArena(mjData* d) {
|
||||
d->parena = d->ncon * sizeof(mjContact);
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
if (!d->threadpool) {
|
||||
ASAN_POISON_MEMORY_REGION(
|
||||
(char*)d->arena + d->parena, d->narena - d->pstack - d->parena);
|
||||
}
|
||||
#endif
|
||||
}
|
||||
|
||||
@@ -937,6 +936,7 @@ int mj_collideOBB(const mjtNum aabb1[6], const mjtNum aabb2[6],
|
||||
}
|
||||
|
||||
|
||||
// binary search between two bodyflex trees
|
||||
void mj_collideTree(const mjModel* m, mjData* d, int bf1, int bf2,
|
||||
int merged, int startadr, int pairadr) {
|
||||
int nbody = m->nbody, nbvhstatic = m->nbvhstatic;
|
||||
@@ -1814,13 +1814,67 @@ static void mj_makeCapsule(const mjModel* m, mjData* d, int f, const int vid[2],
|
||||
}
|
||||
|
||||
|
||||
// struct for collision task
|
||||
typedef struct {
|
||||
mjPreContact* conbuffer; // pre-contact buffer returned by collision functions
|
||||
int* nconbuffer; // contact count for each collision pair
|
||||
char* epabuffer; // buffer for nativeccd
|
||||
int ccd_size; // size of nativeccd buffer
|
||||
const int* pairbuffer; // collision pairs (g1, g2, ipair, index into conbuffer)
|
||||
int npair; // number of collision pairs
|
||||
int chunksize; // number of pairs to process per task
|
||||
int maxcon; // maximum number of contacts (size of conbuffer)
|
||||
} mjContactArg;
|
||||
|
||||
|
||||
static void collisionTask(const mjModel* m, mjData* d, void* arg, int thread_id, int idx) {
|
||||
mjContactArg* conargs = (mjContactArg*)arg;
|
||||
mjPreContact* conbuffer = conargs->conbuffer;
|
||||
char* epabuffer = conargs->epabuffer;
|
||||
int chunksize = conargs->chunksize;
|
||||
int globalidx = chunksize * idx;
|
||||
const int* pair = conargs->pairbuffer + 4 * globalidx;
|
||||
int* ncon = conargs->nconbuffer + chunksize * idx;
|
||||
|
||||
int npair = conargs->npair;
|
||||
int n = mjMIN(chunksize, npair - globalidx);
|
||||
|
||||
mjc_setCCDBuffer(epabuffer + thread_id * conargs->ccd_size);
|
||||
for (int i = 0; i < n; i++) {
|
||||
int g1 = pair[4*i + 0];
|
||||
int g2 = pair[4*i + 1];
|
||||
int ipair = pair[4*i + 2];
|
||||
int conpos = pair[4*i + 3];
|
||||
|
||||
mjfCollision collision_func = mjCOLLISIONFUNC[m->geom_type[g1]][m->geom_type[g2]];
|
||||
mjtNum margin = getMargin(m, g1, g2, ipair);
|
||||
mjtNum gap = getGap(m, g1, g2, ipair);
|
||||
ncon[i] = collision_func(m, d, conbuffer + conpos, g1, g2, margin + gap);
|
||||
|
||||
// SHOULD NOT OCCUR
|
||||
int expected_max = (globalidx + i + 1 < npair ? pair[4*(i+1) + 3] : conargs->maxcon) - conpos;
|
||||
if (ncon[i] > expected_max) {
|
||||
mjERROR("collision function returned %d contacts for geom pair (%d, %d), "
|
||||
"expected at most %d from mj_maxContact", ncon[i], g1, g2, expected_max);
|
||||
}
|
||||
}
|
||||
mjc_setCCDBuffer(NULL);
|
||||
}
|
||||
|
||||
|
||||
// compute contacts for a batch of collision pairs contained in a buffer of
|
||||
// stride 3 ints (g1, g2, ipair)
|
||||
// if buffer is NULL, results are read from arena starting at parena
|
||||
void mj_narrowphase(const mjModel* m, mjData* d, const int* buffer, int npair, size_t parena) {
|
||||
int nthread = mju_numThread(d);
|
||||
int ccd_size = mjc_ccdSize(m->opt.ccd_iterations);
|
||||
mjtNum margin, gap;
|
||||
|
||||
// try to balance load of 5 chunks per thread (chunksize should be divisible by 16)
|
||||
int chunksize = npair / mjMAX(1, 5 * nthread);
|
||||
chunksize = mjMAX(16, (chunksize + 15) & ~15); // round up to next 16
|
||||
int nchunk = (npair + chunksize - 1) / chunksize;
|
||||
|
||||
// set buffer and arena pointer
|
||||
if (!buffer) {
|
||||
buffer = (const int*) ((char*) d->arena + parena);
|
||||
@@ -1830,9 +1884,6 @@ void mj_narrowphase(const mjModel* m, mjData* d, const int* buffer, int npair, s
|
||||
|
||||
mj_markStack(d);
|
||||
|
||||
// buffer store how many contacts are generated for each pair
|
||||
int* nconbuffer = mj_stackAllocInt(d, npair);
|
||||
|
||||
// buffer for pair data (g1, g2, ipair, index into conbuffer)
|
||||
int* pairbuffer = mj_stackAllocInt(d, 4 * npair);
|
||||
int maxcon = 0;
|
||||
@@ -1850,42 +1901,30 @@ void mj_narrowphase(const mjModel* m, mjData* d, const int* buffer, int npair, s
|
||||
maxcon += mj_maxContact(m, g1, g2, margin + gap > 0);
|
||||
}
|
||||
|
||||
// buffer for precontact data
|
||||
mjPreContact* conbuffer = mjSTACKALLOC(d, maxcon, mjPreContact);
|
||||
|
||||
// buffer data has been copied to metadata on the stack;
|
||||
// reclaim arena space so contacts can overwrite the buffer region
|
||||
d->parena = parena;
|
||||
|
||||
// set buffer for nativeccd
|
||||
mj_markStack(d);
|
||||
mjc_setCCDBuffer(mj_stackAllocByte(d, ccd_size, sizeof(mjtNum)));
|
||||
mjContactArg arg;
|
||||
arg.ccd_size = ccd_size;
|
||||
arg.pairbuffer = pairbuffer;
|
||||
arg.nconbuffer = mjSTACKALLOC(d, npair, int);
|
||||
arg.conbuffer = mjSTACKALLOC(d, maxcon, mjPreContact);
|
||||
arg.npair = npair;
|
||||
arg.chunksize = chunksize;
|
||||
arg.maxcon = maxcon;
|
||||
|
||||
for (int i = 0; i < npair; i++) {
|
||||
int g1 = pairbuffer[4*i + 0];
|
||||
int g2 = pairbuffer[4*i + 1];
|
||||
int ipair = pairbuffer[4*i + 2];
|
||||
int idx = pairbuffer[4*i + 3];
|
||||
mjfCollision collision_func = mjCOLLISIONFUNC[m->geom_type[g1]][m->geom_type[g2]];
|
||||
margin = getMargin(m, g1, g2, ipair);
|
||||
gap = getGap(m, g1, g2, ipair);
|
||||
nconbuffer[i] = collision_func(m, d, conbuffer + idx, g1, g2, margin + gap);
|
||||
|
||||
// SHOULD NOT OCCUR
|
||||
int expected_max = (i + 1 < npair ? pairbuffer[4*(i+1) + 3] : maxcon) - idx;
|
||||
if (nconbuffer[i] > expected_max) {
|
||||
mjERROR("collision function returned %d contacts for geom pair (%d, %d), "
|
||||
"expected at most %d from mj_maxContact", nconbuffer[i], g1, g2, expected_max);
|
||||
}
|
||||
// dispatch narrowphase to threads with local stack allocation for EPA
|
||||
{
|
||||
mj_markStack(d);
|
||||
arg.epabuffer = mj_stackAllocByte(d, ccd_size * nthread, sizeof(mjtNum));
|
||||
mju_dispatch(m, d, collisionTask, &arg, nchunk);
|
||||
mj_freeStack(d);
|
||||
}
|
||||
|
||||
// set nativeccd buffer back to NULL
|
||||
mjc_setCCDBuffer(NULL);
|
||||
mj_freeStack(d);
|
||||
|
||||
int ncon = 0;
|
||||
for (int i = 0; i < npair; i++) {
|
||||
ncon += nconbuffer[i];
|
||||
ncon += arg.nconbuffer[i];
|
||||
}
|
||||
|
||||
if (ncon == 0) {
|
||||
@@ -1906,7 +1945,7 @@ void mj_narrowphase(const mjModel* m, mjData* d, const int* buffer, int npair, s
|
||||
// fill in contact data
|
||||
int conpos = 0;
|
||||
for (int i = 0; i < npair; i++) {
|
||||
if (!(ncon = nconbuffer[i]))
|
||||
if (!(ncon = arg.nconbuffer[i]))
|
||||
continue;
|
||||
|
||||
int condim;
|
||||
@@ -1928,7 +1967,7 @@ void mj_narrowphase(const mjModel* m, mjData* d, const int* buffer, int npair, s
|
||||
mj_contactParam(m, &condim, solref, solimp, friction, g1, g2, -1, -1);
|
||||
}
|
||||
|
||||
mjPreContact* bc = conbuffer + pairbuffer[4*i + 3];
|
||||
mjPreContact* bc = arg.conbuffer + pairbuffer[4*i+3];
|
||||
margin = getMargin(m, g1, g2, ipair);
|
||||
for (int j=0; j < ncon; j++) {
|
||||
mjContact* c = con + conpos + j;
|
||||
|
||||
+13
-111
@@ -44,8 +44,7 @@
|
||||
#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"
|
||||
#include "engine/engine_thread.h"
|
||||
|
||||
|
||||
|
||||
@@ -116,28 +115,6 @@ void mj_checkAcc(const mjModel* m, mjData* d) {
|
||||
|
||||
//-------------------------- solver components -----------------------------------------------------
|
||||
|
||||
// args for internal functions in mj_fwdPosition
|
||||
struct mjFwdPositionArgs_ {
|
||||
const mjModel* m;
|
||||
mjData* d;
|
||||
};
|
||||
typedef struct mjFwdPositionArgs_ mjFwdPositionArgs;
|
||||
|
||||
// wrapper for mj_crb and mj_factorM
|
||||
void* mj_inertialThreaded(void* args) {
|
||||
mjFwdPositionArgs* forward_args = (mjFwdPositionArgs*) args;
|
||||
mj_makeM(forward_args->m, forward_args->d);
|
||||
mj_factorM(forward_args->m, forward_args->d);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
// wrapper for mj_collision
|
||||
void* mj_collisionThreaded(void* args) {
|
||||
mjFwdPositionArgs* forward_args = (mjFwdPositionArgs*) args;
|
||||
mj_collision(forward_args->m, forward_args->d);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
// kinematics-related computations
|
||||
void mj_fwdKinematics(const mjModel* m, mjData* d) {
|
||||
mj_kinematics(m, d);
|
||||
@@ -162,36 +139,12 @@ void mj_fwdPosition(const mjModel* m, mjData* d) {
|
||||
|
||||
TM_END(mjTIMER_POS_KINEMATICS);
|
||||
|
||||
// no threadpool: inertia and collision on main thread
|
||||
if (!d->threadpool) {
|
||||
// inertia, timed internally (POS_INERTIA)
|
||||
mj_makeM(m, d);
|
||||
mj_factorM(m, d);
|
||||
// inertia, timed internally (POS_INERTIA)
|
||||
mj_makeM(m, d);
|
||||
mj_factorM(m, d);
|
||||
|
||||
// collision, timed internally (POS_COLLISION)
|
||||
mj_collision(m, d);
|
||||
}
|
||||
|
||||
// have threadpool: inertia and collision on separate threads
|
||||
else {
|
||||
mjTask tasks[2];
|
||||
mjFwdPositionArgs forward_args;
|
||||
forward_args.m = m;
|
||||
forward_args.d = d;
|
||||
|
||||
mju_defaultTask(&tasks[0]);
|
||||
tasks[0].func = mj_inertialThreaded;
|
||||
tasks[0].args = &forward_args;
|
||||
mju_threadPoolEnqueue((mjThreadPool*)d->threadpool, &tasks[0]);
|
||||
|
||||
mju_defaultTask(&tasks[1]);
|
||||
tasks[1].func = mj_collisionThreaded;
|
||||
tasks[1].args = &forward_args;
|
||||
mju_threadPoolEnqueue((mjThreadPool*)d->threadpool, &tasks[1]);
|
||||
|
||||
mju_taskJoin(&tasks[0]);
|
||||
mju_taskJoin(&tasks[1]);
|
||||
}
|
||||
// collision, timed internally (POS_COLLISION)
|
||||
mj_collision(m, d);
|
||||
|
||||
if (mj_wakeCollision(m, d)) {
|
||||
mj_updateSleep(m, d);
|
||||
@@ -909,51 +862,13 @@ 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 CG solver
|
||||
static void* CG_wrapper(void* args) {
|
||||
mjSolIslandArgs* solargs = (mjSolIslandArgs*) args;
|
||||
mj_solCG_island(solargs->m, solargs->d, solargs->island, solargs->m->opt.iterations);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
// extract arguments, pass to Newton solver
|
||||
static void* Newton_wrapper(void* args) {
|
||||
mjSolIslandArgs* solargs = (mjSolIslandArgs*) args;
|
||||
mj_solNewton_island(solargs->m, solargs->d, solargs->island, solargs->m->opt.iterations);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
// CG solver, multi-threaded over islands
|
||||
static void solve_threaded(const mjModel* m, mjData* d, int flg_Newton) {
|
||||
mj_markStack(d);
|
||||
// allocate array of arguments to be passed to threads
|
||||
mjSolIslandArgs* sol_island_args = mjSTACKALLOC(d, d->nisland, mjSolIslandArgs);
|
||||
mjTask* tasks = mjSTACKALLOC(d, d->nisland, mjTask);
|
||||
|
||||
for (int island = 0; island < d->nisland; ++island) {
|
||||
sol_island_args[island].m = m;
|
||||
sol_island_args[island].d = d;
|
||||
sol_island_args[island].island = island;
|
||||
|
||||
mju_defaultTask(&tasks[island]);
|
||||
tasks[island].func = flg_Newton ? Newton_wrapper : CG_wrapper;
|
||||
tasks[island].args = &sol_island_args[island];
|
||||
mju_threadPoolEnqueue((mjThreadPool*)d->threadpool, &tasks[island]);
|
||||
// mju_dispatch callback: solve one island
|
||||
static void solveIslandTask(const mjModel* m, mjData* d, void* arg, int thread_id, int island) {
|
||||
if (m->opt.solver == mjSOL_NEWTON) {
|
||||
mj_solNewton_island(m, d, island, m->opt.iterations);
|
||||
} else {
|
||||
mj_solCG_island(m, d, island, m->opt.iterations);
|
||||
}
|
||||
|
||||
for (int island = 0; island < d->nisland; ++island) {
|
||||
mju_taskJoin(&tasks[island]);
|
||||
}
|
||||
|
||||
mj_freeStack(d);
|
||||
}
|
||||
|
||||
|
||||
@@ -1009,20 +924,7 @@ void mj_fwdConstraint(const mjModel* m, mjData* d) {
|
||||
mju_gather(d->iefc_force, d->efc_force, d->map_iefc2efc, nefc);
|
||||
mju_gather(d->iefc_aref, d->efc_aref, d->map_iefc2efc, nefc);
|
||||
|
||||
// solve per island, with or without threads
|
||||
if (!d->threadpool) {
|
||||
// no threadpool, loop over islands
|
||||
for (int island=0; island < nisland; island++) {
|
||||
if (m->opt.solver == mjSOL_NEWTON) {
|
||||
mj_solNewton_island(m, d, island, m->opt.iterations);
|
||||
} else {
|
||||
mj_solCG_island(m, d, island, m->opt.iterations);
|
||||
}
|
||||
}
|
||||
} else {
|
||||
// have threadpool, solve using threads
|
||||
solve_threaded(m, d, m->opt.solver == mjSOL_NEWTON);
|
||||
}
|
||||
mju_dispatch(m, d, solveIslandTask, NULL, nisland);
|
||||
|
||||
// copy back solver outputs (scatter dofs since ni <= nv)
|
||||
mju_scatter(d->qacc, d->iacc, d->map_idof2dof, nidof);
|
||||
|
||||
@@ -33,6 +33,7 @@
|
||||
#include "engine/engine_memory.h"
|
||||
#include "engine/engine_plugin.h"
|
||||
#include "engine/engine_sleep.h"
|
||||
#include "engine/engine_thread.h"
|
||||
#include "engine/engine_util_blas.h"
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "engine/engine_util_misc.h"
|
||||
@@ -1081,6 +1082,7 @@ void mj_makeRawData(mjData** dest, const mjModel* m) {
|
||||
|
||||
// clear threadpool
|
||||
d->threadpool = 0;
|
||||
d->threadlock = 0;
|
||||
|
||||
// clear nplugin (overwritten by _initPlugin)
|
||||
d->nplugin = 0;
|
||||
@@ -1140,6 +1142,7 @@ mjData* mj_copyDataVisual(mjData* dest, const mjModel* m, const mjData* src, int
|
||||
*dest = *src;
|
||||
dest->buffer = save_buffer;
|
||||
dest->arena = save_arena;
|
||||
dest->threadpool = 0;
|
||||
mj_setPtrData(m, dest);
|
||||
|
||||
// save plugin_data, since the X macro copying block below will override it
|
||||
@@ -1239,8 +1242,6 @@ mjData* mj_copyDataVisual(mjData* dest, const mjModel* m, const mjData* src, int
|
||||
}
|
||||
}
|
||||
|
||||
dest->threadpool = src->threadpool;
|
||||
|
||||
return dest;
|
||||
}
|
||||
|
||||
@@ -1300,7 +1301,6 @@ static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
|
||||
|
||||
// clear memory utilization stats
|
||||
d->maxuse_stack = 0;
|
||||
memset(d->maxuse_threadstack, 0, mjMAXTHREAD*sizeof(mjtSize));
|
||||
d->maxuse_arena = 0;
|
||||
d->maxuse_con = 0;
|
||||
d->maxuse_efc = 0;
|
||||
@@ -1572,6 +1572,7 @@ void mj_resetDataKeyframe(const mjModel* m, mjData* d, int key) {
|
||||
// de-allocate mjData
|
||||
void mj_deleteData(mjData* d) {
|
||||
if (d) {
|
||||
mju_threadpool(d, 0);
|
||||
freeDataBuffers(d);
|
||||
mju_free(d);
|
||||
}
|
||||
|
||||
+63
-73
@@ -16,6 +16,7 @@
|
||||
|
||||
#include <inttypes.h> // IWYU pragma: keep
|
||||
#include <limits.h>
|
||||
#include <stdatomic.h>
|
||||
#include <stddef.h>
|
||||
#include <stdint.h>
|
||||
#include <stdio.h>
|
||||
@@ -26,7 +27,7 @@
|
||||
#include <mujoco/mjsan.h> // IWYU pragma: keep
|
||||
#include "engine/engine_crossplatform.h"
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "thread/thread_pool.h"
|
||||
|
||||
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
#include <sanitizer/asan_interface.h>
|
||||
@@ -57,25 +58,20 @@ static inline size_t fastmod(size_t a, size_t b) {
|
||||
return a % b;
|
||||
}
|
||||
|
||||
typedef struct {
|
||||
uintptr_t bottom; // first memory address available to the stack
|
||||
uintptr_t top; // current memory address used by the stack
|
||||
uintptr_t limit; // top limit of the stack (stack grows down)
|
||||
uintptr_t stack_base; // current stack base for mark and free stack
|
||||
} mjStackInfo;
|
||||
|
||||
|
||||
typedef struct {
|
||||
size_t pbase; // value of d->pbase immediately before mj_markStack
|
||||
size_t pstack; // value of d->pstack immediately before mj_markStack
|
||||
void* pc; // program counter of the call site of mj_markStack (only set when under asan)
|
||||
} mjStackFrame;
|
||||
|
||||
static void maybe_lock_alloc_mutex(mjData* d) {
|
||||
if (d->threadpool != 0) {
|
||||
mju_threadPoolLockAllocMutex((mjThreadPool*)d->threadpool);
|
||||
}
|
||||
}
|
||||
|
||||
static void maybe_unlock_alloc_mutex(mjData* d) {
|
||||
if (d->threadpool != 0) {
|
||||
mju_threadPoolUnlockAllocMutex((mjThreadPool*)d->threadpool);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
static inline mjStackInfo get_stack_info_from_data(const mjData* d) {
|
||||
mjStackInfo stack_info;
|
||||
stack_info.bottom = (uintptr_t)d->arena + (uintptr_t)d->narena;
|
||||
@@ -110,14 +106,12 @@ static size_t stack_usage_redzone(const mjStackInfo* stack_info) {
|
||||
|
||||
// allocate memory from the mjData arena
|
||||
void* mj_arenaAllocByte(mjData* d, size_t bytes, size_t alignment) {
|
||||
maybe_lock_alloc_mutex(d);
|
||||
size_t misalignment = fastmod(d->parena, alignment);
|
||||
size_t padding = misalignment ? alignment - misalignment : 0;
|
||||
|
||||
// check size
|
||||
size_t bytes_available = d->narena - d->pstack;
|
||||
if (mjUNLIKELY(d->parena + padding + bytes > bytes_available)) {
|
||||
maybe_unlock_alloc_mutex(d);
|
||||
return NULL;
|
||||
}
|
||||
|
||||
@@ -125,16 +119,8 @@ void* mj_arenaAllocByte(mjData* d, size_t bytes, size_t alignment) {
|
||||
|
||||
// under ASAN, get stack usage from red zone
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
mjStackInfo stack_info;
|
||||
mjStackInfo* stack_info_ptr;
|
||||
if (!d->threadpool) {
|
||||
stack_info = get_stack_info_from_data(d);
|
||||
stack_info_ptr = &stack_info;
|
||||
} else {
|
||||
size_t thread_id = mju_threadPoolCurrentWorkerId((mjThreadPool*)d->threadpool);
|
||||
stack_info_ptr = mju_getStackInfoForThread(d, thread_id);
|
||||
}
|
||||
stack_usage = stack_usage_redzone(stack_info_ptr);
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
stack_usage = stack_usage_redzone(&stack_info);
|
||||
#endif
|
||||
|
||||
// allocate, update max, return pointer to buffer
|
||||
@@ -150,7 +136,6 @@ void* mj_arenaAllocByte(mjData* d, size_t bytes, size_t alignment) {
|
||||
__msan_allocated_memory(result, bytes);
|
||||
#endif
|
||||
|
||||
maybe_unlock_alloc_mutex(d);
|
||||
return result;
|
||||
}
|
||||
|
||||
@@ -212,13 +197,8 @@ static inline void* stackallocinternal(mjData* d, mjStackInfo* stack_info, size_
|
||||
|
||||
// update max usage statistics
|
||||
stack_info->top = new_top_ptr;
|
||||
if (!d->threadpool) {
|
||||
d->maxuse_stack = mjMAX(d->maxuse_stack, usage);
|
||||
d->maxuse_arena = mjMAX(d->maxuse_arena, usage + d->parena);
|
||||
} else {
|
||||
size_t thread_id = mju_threadPoolCurrentWorkerId((mjThreadPool*)d->threadpool);
|
||||
d->maxuse_threadstack[thread_id] = mjMAX(d->maxuse_threadstack[thread_id], usage);
|
||||
}
|
||||
d->maxuse_stack = mjMAX(d->maxuse_stack, usage);
|
||||
d->maxuse_arena = mjMAX(d->maxuse_arena, usage + d->parena);
|
||||
|
||||
return (void*)start_ptr;
|
||||
}
|
||||
@@ -228,18 +208,46 @@ static inline void* stackallocinternal(mjData* d, mjStackInfo* stack_info, size_
|
||||
// declared inline so that modular arithmetic with specific alignments can be optimized out
|
||||
static inline void* stackalloc(mjData* d, size_t size, size_t alignment,
|
||||
const char* caller, int line) {
|
||||
// single threaded allocation
|
||||
if (!d->threadpool) {
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
void* result = stackallocinternal(d, &stack_info, size, alignment, caller, line);
|
||||
d->pstack = stack_info.bottom - stack_info.top;
|
||||
return result;
|
||||
// size zero: no-op
|
||||
if (!size) {
|
||||
return NULL;
|
||||
}
|
||||
|
||||
// multi threaded allocation
|
||||
size_t thread_id = mju_threadPoolCurrentWorkerId((mjThreadPool*)d->threadpool);
|
||||
mjStackInfo* stack_info = mju_getStackInfoForThread(d, thread_id);
|
||||
return stackallocinternal(d, stack_info, size, alignment, caller, line);
|
||||
// call in mju_dispatch: atomically reserve space on the stack
|
||||
if (d->threadlock) {
|
||||
size_t alloc_size = size + alignment - 1 + 2 * mjREDZONE;
|
||||
size_t old_pstack = atomic_fetch_add_explicit(
|
||||
(_Atomic size_t*)&d->pstack, alloc_size, memory_order_relaxed);
|
||||
|
||||
// check for stack overflow
|
||||
size_t stack_available_bytes = (size_t)d->narena - d->parena;
|
||||
if (mjUNLIKELY(old_pstack + alloc_size > stack_available_bytes)) {
|
||||
char info[1024];
|
||||
if (caller) {
|
||||
snprintf(info, sizeof(info), " at %s, line %d", caller, line);
|
||||
} else {
|
||||
info[0] = '\0';
|
||||
}
|
||||
mju_error(
|
||||
"mj_stackAlloc: out of memory, stack overflow%s (threadlock)\n"
|
||||
" max = %" PRIuPTR ", available = %" PRIuPTR ", requested = %" PRIuPTR
|
||||
"\n nefc = %d, ncon = %d",
|
||||
info, (uintptr_t)stack_available_bytes,
|
||||
(uintptr_t)(stack_available_bytes - old_pstack),
|
||||
(uintptr_t)alloc_size, d->nefc, d->ncon);
|
||||
}
|
||||
|
||||
uintptr_t bottom = (uintptr_t)d->arena + (uintptr_t)d->narena;
|
||||
uintptr_t start_ptr = bottom - old_pstack - size - mjREDZONE;
|
||||
start_ptr -= fastmod(start_ptr, alignment);
|
||||
ASAN_UNPOISON_MEMORY_REGION((void*)start_ptr, size);
|
||||
return (void*)start_ptr;
|
||||
}
|
||||
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
void* result = stackallocinternal(d, &stack_info, size, alignment, caller, line);
|
||||
d->pstack = stack_info.bottom - stack_info.top;
|
||||
return result;
|
||||
}
|
||||
|
||||
|
||||
@@ -268,17 +276,15 @@ void mj_markStack(mjData* d)
|
||||
void mj__markStack(mjData* d)
|
||||
#endif
|
||||
{
|
||||
if (!d->threadpool) {
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
markstackinternal(d, &stack_info);
|
||||
d->pstack = stack_info.bottom - stack_info.top;
|
||||
d->pbase = stack_info.stack_base;
|
||||
// no-op if called from mju_dispatch
|
||||
if (d->threadlock) {
|
||||
return;
|
||||
}
|
||||
|
||||
size_t thread_id = mju_threadPoolCurrentWorkerId((mjThreadPool*)d->threadpool);
|
||||
mjStackInfo* stack_info = mju_getStackInfoForThread(d, thread_id);
|
||||
markstackinternal(d, stack_info);
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
markstackinternal(d, &stack_info);
|
||||
d->pstack = stack_info.bottom - stack_info.top;
|
||||
d->pbase = stack_info.stack_base;
|
||||
}
|
||||
|
||||
|
||||
@@ -319,30 +325,14 @@ void mj_freeStack(mjData* d)
|
||||
void mj__freeStack(mjData* d)
|
||||
#endif
|
||||
{
|
||||
if (!d->threadpool) {
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
freestackinternal(&stack_info);
|
||||
d->pstack = stack_info.bottom - stack_info.top;
|
||||
d->pbase = stack_info.stack_base;
|
||||
if (d->threadlock) {
|
||||
return;
|
||||
}
|
||||
|
||||
size_t thread_id = mju_threadPoolCurrentWorkerId((mjThreadPool*)d->threadpool);
|
||||
mjStackInfo* stack_info = mju_getStackInfoForThread(d, thread_id);
|
||||
freestackinternal(stack_info);
|
||||
}
|
||||
|
||||
|
||||
// returns the number of bytes available on the stack
|
||||
size_t mj_stackBytesAvailable(mjData* d) {
|
||||
if (!d->threadpool) {
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
return stack_info.top - stack_info.limit;
|
||||
} else {
|
||||
size_t thread_id = mju_threadPoolCurrentWorkerId((mjThreadPool*)d->threadpool);
|
||||
mjStackInfo* stack_info = mju_getStackInfoForThread(d, thread_id);
|
||||
return stack_info->top - stack_info->limit;
|
||||
}
|
||||
mjStackInfo stack_info = get_stack_info_from_data(d);
|
||||
freestackinternal(&stack_info);
|
||||
d->pstack = stack_info.bottom - stack_info.top;
|
||||
d->pbase = stack_info.stack_base;
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -49,9 +49,6 @@ void mj__freeStack(mjData* d) __attribute__((noinline));
|
||||
|
||||
#endif // ADDRESS_SANITIZER
|
||||
|
||||
// returns the number of bytes available on the stack
|
||||
MJAPI size_t mj_stackBytesAvailable(mjData* d);
|
||||
|
||||
// allocate bytes on the stack
|
||||
MJAPI void* mj_stackAllocByte(mjData* d, size_t bytes, size_t alignment);
|
||||
|
||||
|
||||
+17
-28
@@ -35,8 +35,7 @@
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "engine/engine_util_misc.h"
|
||||
#include "engine/engine_util_spatial.h"
|
||||
#include "thread/thread_pool.h"
|
||||
#include "thread/thread_task.h"
|
||||
#include "engine/engine_thread.h"
|
||||
|
||||
|
||||
|
||||
@@ -64,8 +63,6 @@ mjPARTIAL_SORT(ContactSelect, ContactInfo, ContactInfoCompare);
|
||||
|
||||
// arguments for parallel tactile sensor computation
|
||||
typedef struct mjTactileTaskArgs_ {
|
||||
const mjModel* m;
|
||||
mjData* d;
|
||||
int sensor_id;
|
||||
int mesh_id;
|
||||
int geom_id;
|
||||
@@ -80,10 +77,8 @@ typedef struct mjTactileTaskArgs_ {
|
||||
|
||||
|
||||
// worker function for parallel tactile computation over taxel batches
|
||||
static void* tactile_taxel_batch(void* args) {
|
||||
static void* tactile_taxel_batch(const mjModel* m, mjData* d, void* args) {
|
||||
mjTactileTaskArgs* t = (mjTactileTaskArgs*)args;
|
||||
const mjModel* m = t->m;
|
||||
mjData* d = t->d;
|
||||
int mesh_id = t->mesh_id;
|
||||
int geom_id = t->geom_id;
|
||||
int parent_weld = t->parent_weld;
|
||||
@@ -193,6 +188,12 @@ static void* tactile_taxel_batch(void* args) {
|
||||
}
|
||||
|
||||
|
||||
static void tactileTask(const mjModel* m, mjData* d, void* arg, int thread_id, int task_id) {
|
||||
mjTactileTaskArgs* args_array = (mjTactileTaskArgs*)arg;
|
||||
tactile_taxel_batch(m, d, &args_array[task_id]);
|
||||
}
|
||||
|
||||
|
||||
// apply cutoff to sensor i, clamping values in data buffer
|
||||
static void apply_cutoff(const mjModel* m, int i, mjtNum* data) {
|
||||
mjtNum cutoff = m->sensor_cutoff[i];
|
||||
@@ -1261,18 +1262,15 @@ static void mj_computeSensorAcc(const mjModel* m, mjData* d, int i, mjtNum* sens
|
||||
// threshold for parallelization (taxel count below which sequential is faster)
|
||||
const int kTactileParallelThreshold = 1000;
|
||||
|
||||
// parallel path: use threadpool to process taxel batches
|
||||
if (d->threadpool && ncon >= kTactileParallelThreshold) {
|
||||
int nthreads = mju_threadPoolNumberOfThreads((mjThreadPool*)d->threadpool);
|
||||
int batch_size = (ncon + nthreads - 1) / nthreads;
|
||||
int ntasks = (ncon + batch_size - 1) / batch_size;
|
||||
// parallel path: use mj_batch to process taxel batches
|
||||
int nthread = mju_numThread(d);
|
||||
if (nthread > 0 && ncon >= kTactileParallelThreshold) {
|
||||
int batch_size = (ncon + nthread - 1) / nthread;
|
||||
int ntask = (ncon + batch_size - 1) / batch_size;
|
||||
|
||||
mjTask* tasks = mjSTACKALLOC(d, ntasks, mjTask);
|
||||
mjTactileTaskArgs* task_args = mjSTACKALLOC(d, ntasks, mjTactileTaskArgs);
|
||||
mjTactileTaskArgs* task_args = mjSTACKALLOC(d, ntask, mjTactileTaskArgs);
|
||||
|
||||
for (int t = 0; t < ntasks; t++) {
|
||||
task_args[t].m = m;
|
||||
task_args[t].d = d;
|
||||
for (int t = 0; t < ntask; t++) {
|
||||
task_args[t].sensor_id = i;
|
||||
task_args[t].mesh_id = mesh_id;
|
||||
task_args[t].geom_id = geom_id;
|
||||
@@ -1283,22 +1281,13 @@ static void mj_computeSensorAcc(const mjModel* m, mjData* d, int i, mjtNum* sens
|
||||
task_args[t].start_taxel = t * batch_size;
|
||||
task_args[t].end_taxel = mju_min((t+1) * batch_size, ncon);
|
||||
task_args[t].forcesT = forcesT;
|
||||
|
||||
mju_defaultTask(&tasks[t]);
|
||||
tasks[t].func = tactile_taxel_batch;
|
||||
tasks[t].args = &task_args[t];
|
||||
mju_threadPoolEnqueue((mjThreadPool*)d->threadpool, &tasks[t]);
|
||||
}
|
||||
|
||||
for (int t = 0; t < ntasks; t++) {
|
||||
mju_taskJoin(&tasks[t]);
|
||||
}
|
||||
mju_dispatch(m, d, tactileTask, task_args, ntask);
|
||||
}
|
||||
// sequential path: call tactile_taxel_batch with full range
|
||||
else {
|
||||
mjTactileTaskArgs args;
|
||||
args.m = m;
|
||||
args.d = d;
|
||||
args.sensor_id = i;
|
||||
args.mesh_id = mesh_id;
|
||||
args.geom_id = geom_id;
|
||||
@@ -1309,7 +1298,7 @@ static void mj_computeSensorAcc(const mjModel* m, mjData* d, int i, mjtNum* sens
|
||||
args.start_taxel = 0;
|
||||
args.end_taxel = ncon;
|
||||
args.forcesT = forcesT;
|
||||
tactile_taxel_batch(&args);
|
||||
tactile_taxel_batch(m, d, &args);
|
||||
}
|
||||
|
||||
// compute sensor output
|
||||
|
||||
@@ -0,0 +1,190 @@
|
||||
// Copyright 2026 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 "engine/engine_thread.h"
|
||||
|
||||
#include <atomic>
|
||||
#include <cstdint>
|
||||
#include <thread>
|
||||
#include <vector>
|
||||
|
||||
#include <mujoco/mjdata.h>
|
||||
#include <mujoco/mjmacro.h>
|
||||
#include <mujoco/mjmodel.h>
|
||||
#include "engine/engine_memory.h"
|
||||
|
||||
// context for thread pool stored on mjData
|
||||
class ThreadPoolContext {
|
||||
public:
|
||||
explicit ThreadPoolContext(int nthread) : threads_(nthread) {
|
||||
for (int i = 0; i < nthread; i++) {
|
||||
threads_[i] = std::thread(&ThreadPoolContext::Worker, this, i + 1);
|
||||
}
|
||||
}
|
||||
|
||||
// non-copyable, non-movable
|
||||
ThreadPoolContext(const ThreadPoolContext&) = delete;
|
||||
ThreadPoolContext& operator=(const ThreadPoolContext&) = delete;
|
||||
|
||||
~ThreadPoolContext() {
|
||||
signal_.store(0, std::memory_order_release);
|
||||
signal_.notify_all();
|
||||
for (auto& thread : threads_) {
|
||||
if (thread.joinable()) {
|
||||
thread.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// dispatch tasks to the thread pool and work on them on the main thread
|
||||
void Dispatch(const mjModel* model, mjData* data, mjTaskFunc func, void* arg,
|
||||
int ntask) {
|
||||
func_ = func;
|
||||
model_ = model;
|
||||
data_ = data;
|
||||
arg_ = arg;
|
||||
ntask_ = ntask;
|
||||
next_.store(0, std::memory_order_relaxed);
|
||||
ndone_.store(0, std::memory_order_relaxed);
|
||||
signal_.store(-signal_.load(std::memory_order_relaxed),
|
||||
std::memory_order_release);
|
||||
signal_.notify_all();
|
||||
|
||||
// process tasks on main thread
|
||||
while (true) {
|
||||
int taskId = next_.fetch_add(1, std::memory_order_relaxed);
|
||||
if (taskId >= ntask_) {
|
||||
break;
|
||||
}
|
||||
func_(model_, data_, arg_, 0, taskId);
|
||||
}
|
||||
|
||||
// busy wait for rest of workers to finish
|
||||
int nthread = threads_.size();
|
||||
while (ndone_.load(std::memory_order_acquire) < nthread) {
|
||||
}
|
||||
}
|
||||
|
||||
int ThreadCount() const { return threads_.size(); }
|
||||
|
||||
private:
|
||||
// worker loop for each worker thread
|
||||
void Worker(int threadId) {
|
||||
int status = 1;
|
||||
|
||||
// main loop waiting for next batch of tasks
|
||||
while (true) {
|
||||
// wait until signal atomic is notified and sign flips
|
||||
signal_.wait(status, std::memory_order_acquire);
|
||||
|
||||
// if signal was set to zero, halt
|
||||
status = signal_.load(std::memory_order_acquire);
|
||||
if (status == 0) {
|
||||
return;
|
||||
}
|
||||
|
||||
// subloop to process tasks for the current batch
|
||||
while (true) {
|
||||
int taskId = next_.fetch_add(1, std::memory_order_relaxed);
|
||||
if (taskId >= ntask_) {
|
||||
break;
|
||||
}
|
||||
func_(model_, data_, arg_, threadId, taskId);
|
||||
}
|
||||
|
||||
// let main thread know this worker is done
|
||||
ndone_.fetch_add(1, std::memory_order_release);
|
||||
}
|
||||
}
|
||||
|
||||
// arguments for the current batch set by Dispatch
|
||||
const mjModel* model_;
|
||||
mjData* data_;
|
||||
mjTaskFunc func_;
|
||||
void* arg_;
|
||||
int ntask_; // total number of tasks for workers to do
|
||||
|
||||
// atomic for each worker to grab the next task
|
||||
std::atomic<int> next_{0};
|
||||
|
||||
// atomic counter for number of workers who completed their tasks
|
||||
alignas(64) std::atomic<int> ndone_{0};
|
||||
|
||||
// alternating signal from -1, 1 to start / halt the worker threads,
|
||||
// set to 0 to force all workers to exit
|
||||
std::atomic<int> signal_{1};
|
||||
|
||||
std::vector<std::thread> threads_;
|
||||
};
|
||||
|
||||
|
||||
|
||||
// create a thread pool with nthread threads
|
||||
void mju_threadpool(mjData* d, int nthread) {
|
||||
if (d->threadpool) {
|
||||
ThreadPoolContext* ctx =
|
||||
reinterpret_cast<ThreadPoolContext*>(d->threadpool);
|
||||
// same size, nothing to do
|
||||
if (nthread == ctx->ThreadCount()) {
|
||||
return;
|
||||
}
|
||||
delete ctx;
|
||||
d->threadpool = 0; // null out in case nthread == 0
|
||||
}
|
||||
|
||||
if (nthread >= 1) {
|
||||
d->threadpool = reinterpret_cast<uintptr_t>(new ThreadPoolContext(nthread));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// dispatch ntask tasks to the thread pool; passes arg into func along with
|
||||
// thread_id and task_id
|
||||
void mju_dispatch(const mjModel* m, mjData* d, mjTaskFunc func, void* arg,
|
||||
int ntask) {
|
||||
// no thread pool or trivial number of tasks: run on main thread
|
||||
if (!d->threadpool || ntask < 2) {
|
||||
for (int i = 0; i < ntask; i++) {
|
||||
func(m, d, arg, 0, i);
|
||||
}
|
||||
return;
|
||||
}
|
||||
|
||||
ThreadPoolContext& ctx = *reinterpret_cast<ThreadPoolContext*>(d->threadpool);
|
||||
|
||||
// lock mjData and mark stack frame, memory will be freed after thread completion
|
||||
if (!d->threadlock) {
|
||||
mj_markStack(d);
|
||||
d->threadlock = true;
|
||||
}
|
||||
|
||||
ctx.Dispatch(m, d, func, arg, ntask);
|
||||
|
||||
if (d->threadlock) {
|
||||
// update max usage statistics
|
||||
d->maxuse_stack = mjMAX(d->maxuse_stack, d->pstack);
|
||||
d->maxuse_arena = mjMAX(d->maxuse_arena, d->pstack + d->parena);
|
||||
|
||||
// unlock mjData and free stack used during worker execution
|
||||
d->threadlock = false;
|
||||
mj_freeStack(d);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// return total number of threads in the pool (including main thread)
|
||||
int mju_numThread(const mjData* d) {
|
||||
ThreadPoolContext* ctx = reinterpret_cast<ThreadPoolContext*>(d->threadpool);
|
||||
return ctx ? ctx->ThreadCount() + 1 : 1;
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
// Copyright 2026 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.
|
||||
|
||||
#ifndef MUJOCO_SRC_ENGINE_ENGINE_THREAD_H_
|
||||
#define MUJOCO_SRC_ENGINE_ENGINE_THREAD_H_
|
||||
|
||||
#include <mujoco/mjdata.h>
|
||||
#include <mujoco/mjexport.h>
|
||||
#include <mujoco/mjmodel.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
// dispatch function for mju_dispatch
|
||||
typedef void (*mjTaskFunc)(const mjModel* m, mjData* d, void* arg, int thread_id, int task_id);
|
||||
|
||||
// create a thread pool with nthread worker threads.
|
||||
MJAPI void mju_threadpool(mjData* d, int nthread);
|
||||
|
||||
// return total number of threads in the pool (including main thread)
|
||||
MJAPI int mju_numThread(const mjData* d);
|
||||
|
||||
// dispatch ntask tasks to the thread pool; passes arg into func along with thread_id and task_id
|
||||
MJAPI void mju_dispatch(const mjModel* m, mjData* d, mjTaskFunc func, void* arg, int ntask);
|
||||
|
||||
#ifdef __cplusplus
|
||||
}
|
||||
#endif
|
||||
#endif // MUJOCO_SRC_ENGINE_ENGINE_THREAD_H_
|
||||
@@ -1,23 +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.
|
||||
|
||||
set(MUJOCO_THREAD_SRCS
|
||||
thread_pool.cc
|
||||
thread_pool.h
|
||||
thread_queue.h
|
||||
thread_task.cc
|
||||
thread_task.h
|
||||
)
|
||||
|
||||
target_sources(mujoco PRIVATE ${MUJOCO_THREAD_SRCS})
|
||||
@@ -1,312 +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 "thread/thread_pool.h"
|
||||
|
||||
#include <stdint.h>
|
||||
|
||||
#include <algorithm>
|
||||
#include <atomic>
|
||||
#include <cstddef>
|
||||
#include <memory>
|
||||
#include <mutex>
|
||||
#include <thread>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include <mujoco/mjsan.h> // IWYU pragma: keep
|
||||
#include <mujoco/mjthread.h>
|
||||
#include <mujoco/mujoco.h>
|
||||
#include "engine/engine_crossplatform.h"
|
||||
#include "engine/engine_util_errmem.h"
|
||||
#include "thread/thread_queue.h"
|
||||
#include "thread/thread_task.h"
|
||||
|
||||
namespace mujoco {
|
||||
namespace {
|
||||
constexpr size_t kThreadPoolQueueSize = 640;
|
||||
|
||||
// Each thread being run will be assigned a worker_id.
|
||||
// 0: main thread
|
||||
// 1->n: workers
|
||||
thread_local size_t worker_id = 0;
|
||||
|
||||
struct WorkerThread {
|
||||
// Shutdown function passed to running threads to ensure clean shutdown.
|
||||
static void* ShutdownFunction(void* args) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
// Thread for the worker.
|
||||
std::unique_ptr<std::thread> thread_;
|
||||
|
||||
// An mjTask for shutting down this worker.
|
||||
mjTask shutdown_task_ {
|
||||
&ShutdownFunction,
|
||||
nullptr,
|
||||
mjTASK_NEW
|
||||
};
|
||||
};
|
||||
} // namespace
|
||||
|
||||
// Concrete C++ class definition for mjThreadPool.
|
||||
// (The public mjThreadPool C struct is an opaque one.)
|
||||
class ThreadPoolImpl : public mjThreadPool {
|
||||
public:
|
||||
ThreadPoolImpl(int num_worker) : mjThreadPool{num_worker} {
|
||||
// initialize worker threads
|
||||
for (int i = 0; i < std::min(num_worker, mjMAXTHREAD); ++i) {
|
||||
WorkerThread worker{
|
||||
std::make_unique<std::thread>(ThreadPoolWorker, this, i)};
|
||||
workers_.push_back(std::move(worker));
|
||||
}
|
||||
}
|
||||
|
||||
size_t NumberOfThreads() {
|
||||
return workers_.size();
|
||||
}
|
||||
|
||||
// start a task in the threadpool
|
||||
void Enqueue(mjTask* task) {
|
||||
if (mjUNLIKELY(GetAtomicTaskStatus(task).exchange(mjTASK_QUEUED) !=
|
||||
mjTASK_NEW)) {
|
||||
mjERROR("task->status is not mjTASK_NEW");
|
||||
}
|
||||
lockless_queue_.push(task);
|
||||
}
|
||||
|
||||
// shutdown the threadpool
|
||||
void Shutdown() {
|
||||
if (shutdown_) {
|
||||
return;
|
||||
}
|
||||
|
||||
shutdown_ = true;
|
||||
std::vector<mjTask> shutdown_tasks(workers_.size());
|
||||
for (auto& worker : workers_) {
|
||||
Enqueue(&worker.shutdown_task_);
|
||||
}
|
||||
|
||||
for (auto& worker : workers_) {
|
||||
worker.thread_->join();
|
||||
}
|
||||
}
|
||||
|
||||
// registers a worker ID for a given thread
|
||||
void RegisterWorker(const size_t input_worker_id) {
|
||||
worker_id = input_worker_id;
|
||||
}
|
||||
|
||||
// gets the worker id of the current thread
|
||||
size_t GetWorkerId() {
|
||||
return worker_id;
|
||||
}
|
||||
|
||||
void LockAlloc() {
|
||||
alloc_mutex_.lock();
|
||||
}
|
||||
|
||||
void UnlockAlloc() {
|
||||
alloc_mutex_.unlock();
|
||||
}
|
||||
|
||||
bool IsThreadPoolBound() {
|
||||
return thread_pool_bound_;
|
||||
}
|
||||
|
||||
void BindThreadPool() {
|
||||
thread_pool_bound_ = true;
|
||||
}
|
||||
|
||||
~ThreadPoolImpl() { Shutdown(); }
|
||||
|
||||
private:
|
||||
// method executed by running threads
|
||||
static void ThreadPoolWorker(
|
||||
ThreadPoolImpl* thread_pool, const size_t thread_index) {
|
||||
worker_id = thread_index + 1;
|
||||
while (!thread_pool->shutdown_) {
|
||||
auto task = static_cast<mjTask*>(thread_pool->lockless_queue_.pop());
|
||||
task->args = task->func(task->args);
|
||||
GetAtomicTaskStatus(task).store(mjTASK_COMPLETED);
|
||||
}
|
||||
}
|
||||
|
||||
// indicates whether the thread pool is being shut down
|
||||
std::atomic<bool> shutdown_ = false;
|
||||
|
||||
// OS threads that are running in this pool
|
||||
std::vector<WorkerThread> workers_;
|
||||
|
||||
// queue of tasks to execute
|
||||
mujoco::LocklessQueue<void*, kThreadPoolQueueSize> lockless_queue_;
|
||||
|
||||
// Mutex to protect arena allocations.
|
||||
std::mutex alloc_mutex_;
|
||||
|
||||
// Whether or not a ThreadPool was bound using mju_bindThreadPool.
|
||||
bool thread_pool_bound_ = false;
|
||||
};
|
||||
|
||||
// create a thread pool
|
||||
mjThreadPool* mju_threadPoolCreate(size_t number_of_threads) {
|
||||
return reinterpret_cast<mjThreadPool*>(new ThreadPoolImpl(number_of_threads));
|
||||
}
|
||||
|
||||
// gets the number of shards the stack is currently broken into
|
||||
static size_t GetNumberOfShards(mjData* d) {
|
||||
if (!d->threadpool) {
|
||||
return 1;
|
||||
}
|
||||
return mju_threadPoolNumberOfThreads((mjThreadPool*)d->threadpool) + 1;
|
||||
}
|
||||
|
||||
// returns the stack information for the specified thread's shard
|
||||
mjStackInfo* mju_getStackInfoForThread(mjData* d, size_t thread_id) {
|
||||
auto thread_pool = (ThreadPoolImpl*)d->threadpool;
|
||||
if (!thread_pool || !thread_pool->IsThreadPoolBound()) {
|
||||
mju_error("Thread Pool not bound, use mju_bindThreadPool to add an mjThreadPool to mjData");
|
||||
}
|
||||
|
||||
// number of threads running in the threadpool plus the main thread
|
||||
size_t number_of_shards = GetNumberOfShards(d);
|
||||
|
||||
// size of entire arena/stack in bytes
|
||||
size_t total_arena_size_bytes = d->narena;
|
||||
|
||||
// set the shard cursor to the end of the arena
|
||||
uintptr_t end_of_arena_ptr = (uintptr_t)d->arena + total_arena_size_bytes;
|
||||
|
||||
// each thread including the main one will get an equal shard of the stack
|
||||
size_t bytes_per_shard = total_arena_size_bytes / (2 * (number_of_shards));
|
||||
|
||||
// ensure the shard is larger than the cache line
|
||||
size_t misalignment = bytes_per_shard % mju_getDestructiveInterferenceSize();
|
||||
|
||||
if (misalignment != 0) {
|
||||
bytes_per_shard += mju_getDestructiveInterferenceSize() - misalignment;
|
||||
}
|
||||
|
||||
if (bytes_per_shard * number_of_shards > total_arena_size_bytes) {
|
||||
mju_error("Arena is not large enough for %zu shards", number_of_shards);
|
||||
}
|
||||
|
||||
uintptr_t result = (end_of_arena_ptr - (thread_id + 1) * bytes_per_shard);
|
||||
|
||||
// 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;
|
||||
}
|
||||
|
||||
// shards the stack for each thread
|
||||
static void ConfigureMultiThreadedStack(mjData* d) {
|
||||
if (!d->threadpool) {
|
||||
mju_error("No thread pool specified for multithreaded operation");
|
||||
}
|
||||
|
||||
size_t number_of_shards = GetNumberOfShards(d);
|
||||
|
||||
// current top of the stack
|
||||
uintptr_t current_limit = (uintptr_t)d->arena + d->narena - d->pstack;
|
||||
|
||||
// set the shard cursor to the end of the arena
|
||||
uintptr_t begin_shard_cursor_ptr = (uintptr_t)d->arena + d->narena;
|
||||
|
||||
for (size_t shard_index = 0; shard_index < number_of_shards; ++shard_index) {
|
||||
mjStackInfo* end_shard_cursor_ptr = mju_getStackInfoForThread(d, shard_index);
|
||||
#ifdef ADDRESS_SANITIZER
|
||||
// unpoison stack info
|
||||
ASAN_UNPOISON_MEMORY_REGION((void*)end_shard_cursor_ptr, sizeof(mjStackInfo));
|
||||
#endif
|
||||
// handle the main thread's stack which may already have data in it
|
||||
if (shard_index == 0) {
|
||||
// abort if the current stack is already larger than the portion of the stack
|
||||
// that would be reserved for the main thread
|
||||
if ((uintptr_t)end_shard_cursor_ptr > current_limit) {
|
||||
mju_error("mj_bindThreadPool: sharding stack - existing stack larger than shard size: current_size = %zu, "
|
||||
"max_size = %zu", current_limit, (uintptr_t) end_shard_cursor_ptr);
|
||||
}
|
||||
end_shard_cursor_ptr->top = current_limit;
|
||||
end_shard_cursor_ptr->stack_base = d->pbase;
|
||||
} else {
|
||||
// all other stacks are empty because threads have not been used yet
|
||||
end_shard_cursor_ptr->top = begin_shard_cursor_ptr;
|
||||
end_shard_cursor_ptr->stack_base = 0;
|
||||
}
|
||||
|
||||
end_shard_cursor_ptr->bottom = begin_shard_cursor_ptr;
|
||||
end_shard_cursor_ptr->limit = (uintptr_t)end_shard_cursor_ptr + sizeof(mjStackInfo);
|
||||
begin_shard_cursor_ptr = (uintptr_t)end_shard_cursor_ptr - 1;
|
||||
}
|
||||
}
|
||||
|
||||
// adds a thread pool to mjData and configures it for multi-threaded use.
|
||||
void mju_bindThreadPool(mjData* d, void* thread_pool) {
|
||||
if (d->threadpool) {
|
||||
mju_error("Thread Pool already bound to mjData");
|
||||
}
|
||||
|
||||
d->threadpool = (uintptr_t) thread_pool;
|
||||
((ThreadPoolImpl*)thread_pool)->BindThreadPool();
|
||||
ConfigureMultiThreadedStack(d);
|
||||
}
|
||||
|
||||
// gets the number of running threads in the thread pool.
|
||||
size_t mju_threadPoolNumberOfThreads(mjThreadPool* thread_pool) {
|
||||
auto thread_pool_impl = static_cast<ThreadPoolImpl*>(thread_pool);
|
||||
return thread_pool_impl->NumberOfThreads();
|
||||
}
|
||||
|
||||
size_t mju_threadPoolCurrentWorkerId(mjThreadPool* thread_pool) {
|
||||
auto thread_pool_impl = static_cast<ThreadPoolImpl*>(thread_pool);
|
||||
return thread_pool_impl->GetWorkerId();
|
||||
}
|
||||
|
||||
// start a task in the threadpool
|
||||
void mju_threadPoolEnqueue(mjThreadPool* thread_pool, mjTask* task) {
|
||||
auto thread_pool_impl = static_cast<ThreadPoolImpl*>(thread_pool);
|
||||
thread_pool_impl->Enqueue(task);
|
||||
}
|
||||
|
||||
// shutdown the threadpool and free the memory
|
||||
void mju_threadPoolDestroy(mjThreadPool* thread_pool) {
|
||||
auto thread_pool_impl = static_cast<ThreadPoolImpl*>(thread_pool);
|
||||
thread_pool_impl->Shutdown();
|
||||
delete thread_pool_impl;
|
||||
}
|
||||
|
||||
// locks the allocation mutex to protect Stack and Arena allocations
|
||||
void mju_threadPoolLockAllocMutex(mjThreadPool* thread_pool) {
|
||||
auto thread_pool_impl = static_cast<ThreadPoolImpl*>(thread_pool);
|
||||
thread_pool_impl->LockAlloc();
|
||||
}
|
||||
|
||||
// unlocks the allocation mutex to protect Stack and Arena allocations
|
||||
void mju_threadPoolUnlockAllocMutex(mjThreadPool* thread_pool) {
|
||||
auto thread_pool_impl = static_cast<ThreadPoolImpl*>(thread_pool);
|
||||
thread_pool_impl->UnlockAlloc();
|
||||
}
|
||||
|
||||
// Get the destructive interference size for the architecture.
|
||||
size_t mju_getDestructiveInterferenceSize(void) {
|
||||
// return std::hardware_destructive_interference_size;
|
||||
return 128;
|
||||
}
|
||||
|
||||
} // namespace mujoco
|
||||
@@ -1,85 +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.
|
||||
|
||||
#ifndef MUJOCO_SRC_THREAD_THREAD_POOL_H_
|
||||
#define MUJOCO_SRC_THREAD_THREAD_POOL_H_
|
||||
|
||||
#include <stddef.h>
|
||||
|
||||
#include <mujoco/mjexport.h>
|
||||
#include <mujoco/mjthread.h>
|
||||
#include <mujoco/mujoco.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
namespace mujoco {
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
// MultiThreaded Stack will be an approximately 50/50 split of the entire buffer, with a little
|
||||
// wiggle for alignment and caching concerns. The basic layout is to reuse the existing single
|
||||
// threaded markers, and then create shards for each thread to use as its stack.
|
||||
// Not to scale.
|
||||
// |----------|-----------|-----------|-----------|-----------|----------|-----------|-----------|
|
||||
// |Used Arena|Free Arena |Shard1 |Shard1 |Shard1 |Shard0 |Shard0 |Shard0 |
|
||||
// |%%%%%%%%%%| |StackInfo |Free Stack |Used Stack |StackInfo |Free Stack |Used Stack |
|
||||
// |%%%%%%%%%%| | | |%%%%%%%%%%%| | |%%%%%%%%%%%|
|
||||
// |%%%%%%%%%%| | | |%%%%%%%%%%%| | |%%%%%%%%%%%|
|
||||
// |----------|-----------|-----------|-----------|-----------|----------|-----------|-----------|
|
||||
// d->arena d->parena d->pstack shard1->stack_info shard1->bottom_of_stack shard1->bottom_of_stack
|
||||
// shard1->stack_info shard0->stack_info shard0->current_stack
|
||||
// shard1->top_of_stack shard1->top_of_stack
|
||||
// shard1->current_stack
|
||||
typedef struct {
|
||||
uintptr_t bottom; // First memory address available to the stack
|
||||
uintptr_t top; // Current memory address used by the stack
|
||||
uintptr_t limit; // Top limit of the stack (note this is smaller than bottom, stack grows down)
|
||||
uintptr_t stack_base; // Current stack base for mark and free stack
|
||||
} mjStackInfo;
|
||||
|
||||
// Create a thread pool with the specified number of threads running.
|
||||
MJAPI mjThreadPool* mju_threadPoolCreate(size_t number_of_threads);
|
||||
|
||||
// Returns the stack information for the specified thread's shard.
|
||||
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, void* thread_pool);
|
||||
|
||||
// Gets the number of running threads in the thread pool.
|
||||
MJAPI size_t mju_threadPoolNumberOfThreads(mjThreadPool* thread_pool);
|
||||
|
||||
// Gets the ID of the current thread being executed
|
||||
MJAPI size_t mju_threadPoolCurrentWorkerId(mjThreadPool* thread_pool);
|
||||
|
||||
// Enqueue a task in a thread pool.
|
||||
MJAPI void mju_threadPoolEnqueue(mjThreadPool* thread_pool, mjTask* task);
|
||||
|
||||
// Locks the allocation mutex to protect Arena allocations.
|
||||
MJAPI void mju_threadPoolLockAllocMutex(mjThreadPool* thread_pool);
|
||||
|
||||
// Unlocks the allocation mutex to protect Arena allocations.
|
||||
MJAPI void mju_threadPoolUnlockAllocMutex(mjThreadPool* thread_pool);
|
||||
|
||||
// Destroy a thread pool.
|
||||
MJAPI void mju_threadPoolDestroy(mjThreadPool* thread_pool);
|
||||
|
||||
// Get the destructive interference size for the architecture.
|
||||
MJAPI size_t mju_getDestructiveInterferenceSize(void);
|
||||
|
||||
#ifdef __cplusplus
|
||||
} // extern "C"
|
||||
} // namespace mujoco
|
||||
#endif // __cplusplus
|
||||
|
||||
#endif // MUJOCO_SRC_THREAD_THREAD_POOL_H_
|
||||
@@ -1,152 +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.
|
||||
|
||||
#ifndef MUJOCO_SRC_THREAD_LOCKLESS_QUEUE_H_
|
||||
#define MUJOCO_SRC_THREAD_LOCKLESS_QUEUE_H_
|
||||
|
||||
#include <atomic>
|
||||
#include <climits>
|
||||
#include <cstddef>
|
||||
#include <thread>
|
||||
|
||||
namespace mujoco {
|
||||
|
||||
// A Lockless Queue allows for sending information quickly between different
|
||||
// threads. This is a Multi-Producer Multi-Consumer Lockless Queue allowing for
|
||||
// multiple threads to be adding items to the queue while multiple threads are
|
||||
// consuming items from the queue. Internally it uses a Ring Buffer for storage
|
||||
// so it will not grow as items are added. Push will block if the Queue is full
|
||||
// and Pop will block if it is empty.
|
||||
//
|
||||
// For a basic overview of this category of structures:
|
||||
// https://www.linuxjournal.com/content/lock-free-multi-producer-multi-consumer-queue-ring-buffer
|
||||
template <typename T, size_t buffer_capacity>
|
||||
class LocklessQueue {
|
||||
public:
|
||||
bool full() const {
|
||||
return full_internal(
|
||||
convert_to_index(read_cursor_), convert_to_index(write_cursor_));
|
||||
}
|
||||
|
||||
bool empty() const {
|
||||
return maximum_read_cursor_ == read_cursor_;
|
||||
}
|
||||
|
||||
// Push an element into the queue.
|
||||
void push(const T& input) {
|
||||
// Reserve a slot in the queue
|
||||
size_t current_write_cursor;
|
||||
size_t dummy_current_write_cursor;
|
||||
size_t next_write_cursor;
|
||||
size_t current_write_index;
|
||||
size_t current_read_index;
|
||||
do {
|
||||
// Check if the queue is full.
|
||||
do {
|
||||
current_write_cursor = write_cursor_.load();
|
||||
current_write_index = convert_to_index(current_write_cursor);
|
||||
next_write_cursor = get_next_cursor(current_write_cursor);
|
||||
|
||||
current_read_index = convert_to_index(read_cursor_.load());
|
||||
} while (full_internal(current_read_index, current_write_index));
|
||||
|
||||
// Once it's not full, attempt to grab a slot to write.
|
||||
dummy_current_write_cursor = current_write_cursor;
|
||||
} while (!write_cursor_.compare_exchange_weak(
|
||||
dummy_current_write_cursor, next_write_cursor));
|
||||
|
||||
// Write the entry.
|
||||
buffer_[current_write_index].store(input);
|
||||
|
||||
// Increment maximum read cursor. Note here it has to wait if the compare
|
||||
// and exchange fails as another thread might not have completed its write.
|
||||
do {
|
||||
dummy_current_write_cursor = current_write_cursor;
|
||||
} while (!maximum_read_cursor_.compare_exchange_weak(
|
||||
dummy_current_write_cursor, next_write_cursor));
|
||||
}
|
||||
|
||||
// Pop an element from the queue.
|
||||
T pop() {
|
||||
size_t current_read_cursor;
|
||||
size_t dummy_current_read_cursor;
|
||||
size_t current_read_index;
|
||||
size_t next_read_cursor;
|
||||
size_t current_maximum_read_cursor;
|
||||
size_t current_maximum_read_index;
|
||||
bool empty = false;
|
||||
T result;
|
||||
do {
|
||||
// Wait until the queue has an element
|
||||
do {
|
||||
if (empty) {
|
||||
std::this_thread::yield();
|
||||
}
|
||||
current_read_cursor = read_cursor_.load();
|
||||
current_maximum_read_cursor = maximum_read_cursor_.load();
|
||||
|
||||
current_read_index = convert_to_index(current_read_cursor);
|
||||
current_maximum_read_index = convert_to_index(
|
||||
current_maximum_read_cursor);
|
||||
|
||||
empty = empty_internal(
|
||||
current_read_index, current_maximum_read_index);
|
||||
} while (empty);
|
||||
|
||||
next_read_cursor = get_next_cursor(current_read_cursor);
|
||||
|
||||
// Attempt to grab the element, if unsuccessful then wait for the next
|
||||
// element to arrive.
|
||||
result = buffer_[current_read_index].load();
|
||||
dummy_current_read_cursor = current_read_cursor;
|
||||
} while (!read_cursor_.compare_exchange_weak(
|
||||
dummy_current_read_cursor, next_read_cursor));
|
||||
|
||||
return result;
|
||||
}
|
||||
|
||||
private:
|
||||
size_t convert_to_index(size_t input) const {
|
||||
return input % internal_buffer_capacity_;
|
||||
}
|
||||
|
||||
size_t get_next_cursor(size_t input) const {
|
||||
return (input + 1) % cursor_max_;
|
||||
}
|
||||
|
||||
size_t get_next_index(size_t input) const {
|
||||
return convert_to_index(get_next_cursor(input));
|
||||
}
|
||||
|
||||
bool full_internal(size_t read_index, size_t write_index) const {
|
||||
return get_next_index(write_index) == read_index;
|
||||
}
|
||||
|
||||
bool empty_internal(size_t read_index, size_t write_index) const {
|
||||
return read_index == write_index;
|
||||
}
|
||||
|
||||
const size_t internal_buffer_capacity_ = buffer_capacity + 1;
|
||||
const size_t cursor_max_ = UINT_MAX - (UINT_MAX % internal_buffer_capacity_);
|
||||
|
||||
std::atomic<size_t> read_cursor_ = 0;
|
||||
std::atomic<size_t> write_cursor_ = 0;
|
||||
std::atomic<size_t> maximum_read_cursor_ = 0;
|
||||
|
||||
std::atomic<T> buffer_[(buffer_capacity + 1)];
|
||||
};
|
||||
|
||||
} // namespace mujoco
|
||||
|
||||
#endif // MUJOCO_SRC_THREAD_LOCKLESS_QUEUE_H_
|
||||
@@ -1,33 +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 "thread/thread_task.h"
|
||||
|
||||
#include <thread>
|
||||
|
||||
#include <mujoco/mjthread.h>
|
||||
|
||||
namespace mujoco {
|
||||
void mju_defaultTask(mjTask* task) {
|
||||
task->func = nullptr;
|
||||
task->args = nullptr;
|
||||
task->status = mjTASK_NEW;
|
||||
}
|
||||
|
||||
void mju_taskJoin(mjTask* task) {
|
||||
while (GetAtomicTaskStatus(task) != mjTASK_COMPLETED) {
|
||||
std::this_thread::yield();
|
||||
}
|
||||
}
|
||||
} // namespace mujoco
|
||||
@@ -1,49 +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.
|
||||
|
||||
#ifndef MUJOCO_SRC_THREAD_THREAD_TASK_H_
|
||||
#define MUJOCO_SRC_THREAD_THREAD_TASK_H_
|
||||
|
||||
#include <mujoco/mjexport.h>
|
||||
#include <mujoco/mjthread.h>
|
||||
|
||||
#ifdef __cplusplus
|
||||
#include <atomic>
|
||||
#include <new>
|
||||
#include <type_traits>
|
||||
namespace mujoco {
|
||||
extern "C" {
|
||||
#endif
|
||||
|
||||
// Initialize an mjTask.
|
||||
MJAPI void mju_defaultTask(mjTask* task);
|
||||
|
||||
// Wait for a task to complete.
|
||||
MJAPI void mju_taskJoin(mjTask* task);
|
||||
|
||||
#ifdef __cplusplus
|
||||
} // extern "C"
|
||||
|
||||
using TaskStatus = std::remove_volatile_t<decltype(mjTask::status)>;
|
||||
inline std::atomic<TaskStatus>& GetAtomicTaskStatus(mjTask* task) {
|
||||
static_assert(sizeof(std::atomic<TaskStatus>) == sizeof(TaskStatus));
|
||||
static_assert(alignof(std::atomic<TaskStatus>) == alignof(TaskStatus));
|
||||
static_assert(std::atomic<TaskStatus>::is_always_lock_free);
|
||||
return *std::launder(reinterpret_cast<std::atomic<TaskStatus>*>(
|
||||
const_cast<TaskStatus*>(&task->status)));
|
||||
}
|
||||
} // namespace mujoco
|
||||
#endif // __cplusplus
|
||||
|
||||
#endif // MUJOCO_SRC_THREAD_THREAD_TASK_H_
|
||||
Reference in New Issue
Block a user