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
+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;
}