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
+2
View File
@@ -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
+73 -34
View File
@@ -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
View File
@@ -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);
+4 -3
View File
@@ -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
View File
@@ -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;
}
-3
View File
@@ -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
View File
@@ -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
+190
View File
@@ -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;
}
+41
View File
@@ -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_
-23
View File
@@ -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})
-312
View File
@@ -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
-85
View File
@@ -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_
-152
View File
@@ -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_
-33
View File
@@ -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
-49
View File
@@ -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_