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_