Refactor stack allocations to have a clean byte allocation interface.

PiperOrigin-RevId: 558881470
Change-Id: I08c45cb5fa25a1a9c39560e598e805cfa5141c3e
This commit is contained in:
Matthew Bennice
2023-08-21 13:17:48 -07:00
committed by Copybara-Service
parent 22fd0586b1
commit 59138af10b
10 changed files with 67 additions and 64 deletions
+1 -1
View File
@@ -948,7 +948,7 @@ Euler integrator, semi-implicit in velocity.
def test_can_raise_error(self):
self.data.pstack = self.data.nstack
with self.assertRaisesRegex(mujoco.FatalError,
r'\Amj_stackAlloc: stack overflow'):
r'\Amj_stackAllocBytes: stack overflow'):
mujoco.mj_forward(self.model, self.data)
def test_mjcb_time(self):
+3 -9
View File
@@ -208,11 +208,7 @@ int mj_collideOBB(const mjtNum aabb1[6], const mjtNum aabb2[6],
}
static mjCollisionTree* mj_stackAllocTree(mjData* d, int max_stack) {
// check that the quotient is an integer
_Static_assert(sizeof(mjCollisionTree*) % sizeof(mjtNum) == 0,
"mjCollisionTree has a different size from mjtNum");
return (mjCollisionTree*)mj_stackAlloc(
d, max_stack * sizeof(mjCollisionTree*) / sizeof(mjtNum));
return (mjCollisionTree*)mj_stackAllocBytes(d, max_stack * sizeof(mjCollisionTree*));
}
// binary search between two body trees
@@ -754,10 +750,8 @@ int mj_broadphase(const mjModel* m, mjData* d, int* pair, int maxpair) {
}
// allocate sort buffer
int quot = sizeof(mjtBroadphase)/sizeof(mjtNum);
int rem = sizeof(mjtBroadphase)%sizeof(mjtNum);
sortbuf = (mjtBroadphase*)mj_stackAlloc(d, 2*bufcnt*(quot + (rem ? 1 : 0)));
activebuf = (mjtBroadphase*)mj_stackAlloc(d, 2*bufcnt*(quot + (rem ? 1 : 0)));
sortbuf = (mjtBroadphase*)mj_stackAllocBytes(d, 2 * bufcnt * sizeof(mjtBroadphase));
activebuf = (mjtBroadphase*)mj_stackAllocBytes(d, 2 *bufcnt * sizeof(mjtBroadphase));
// init sortbuf with axis0
int k = 0;
+2 -2
View File
@@ -456,8 +456,8 @@ static void collideBVH(const mjModel* m, mjData* d, int g,
int node;
};
typedef struct CollideTreeArgs_ CollideTreeArgs;
CollideTreeArgs* stack = (CollideTreeArgs*)mj_stackAlloc(
d, max_stack * sizeof(CollideTreeArgs*) / sizeof(mjtNum));
CollideTreeArgs* stack = (CollideTreeArgs*)mj_stackAllocBytes(
d, max_stack * sizeof(CollideTreeArgs));
int nstack = 0;
stack[nstack].node = 0;
+1 -1
View File
@@ -251,7 +251,7 @@ void mjd_smooth_velFD(const mjModel* m, mjData* d, mjtNum eps) {
mjtNum* plus = mj_stackAlloc(d, nv);
mjtNum* minus = mj_stackAlloc(d, nv);
mjtNum* fd = mj_stackAlloc(d, nv);
int* cnt = (int*) mj_stackAlloc(d, nv);
int* cnt = mj_stackAllocInt(d, nv);
// clear row counters
memset(cnt, 0, nv*sizeof(int));
+46 -40
View File
@@ -1186,8 +1186,8 @@ void* mj_arenaAlloc(mjData* d, int bytes, int alignment) {
// allocate size mjtNums on the mjData stack
mjtNum* mj_stackAlloc(mjData* d, int size) {
// allocate size bytes on the mjData stack
void* mj_stackAllocBytes(mjData* d, size_t size) {
// return NULL if empty
if (!size) {
return NULL;
@@ -1195,76 +1195,82 @@ mjtNum* mj_stackAlloc(mjData* d, int size) {
// add red zone padding when built with asan, to detect out-of-bound accesses
#ifdef ADDRESS_SANITIZER
#define mjREDZONE 4
#define mjREDZONE 32
#else
#define mjREDZONE 0
#endif
// size of entire arena/stack in bytes
size_t stack_size_bytes = d->nstack * sizeof(mjtNum);
// end of the arena
uintptr_t end_of_arena_ptr = (uintptr_t)d->arena + stack_size_bytes;
// current top of the stack
uintptr_t end_ptr = end_of_arena_ptr - (d->pstack * sizeof(mjtNum));
// start of the memory to be allocated to the buffer
uintptr_t start_ptr = end_ptr - (size + mjREDZONE);
// move start_ptr back to align to max_align_t
start_ptr -= start_ptr % _Alignof(max_align_t);
// new top of the stack
uintptr_t new_pstack_ptr = start_ptr - mjREDZONE;
size_t new_pstack = (end_of_arena_ptr - new_pstack_ptr) / sizeof(mjtNum);
// exclude red zone from stack usage statistics
size_t current_alloc_usage = (end_ptr - new_pstack_ptr - 2 * mjREDZONE) / sizeof(mjtNum);
size_t usage = current_alloc_usage + d->pstack;
// check size
size_t stack_available_bytes = d->nstack * sizeof(mjtNum) - d->parena;
size_t stack_required_bytes = (d->pstack + size + 2*mjREDZONE) * sizeof(mjtNum);
size_t stack_available_bytes = end_ptr - ((uintptr_t)d->arena + d->parena);
size_t stack_required_bytes = end_ptr - new_pstack_ptr;
if (stack_required_bytes > stack_available_bytes) {
mjERROR("stack overflow: max = %zu, available = %zu, requested = %zu "
"(ne = %d, nf = %d, nefc = %d, ncon = %d)",
d->nstack * sizeof(mjtNum), stack_available_bytes, stack_required_bytes,
stack_size_bytes, stack_available_bytes, stack_required_bytes,
d->ne, d->nf, d->nefc, d->ncon);
}
// allocate at end of arena
char* end_ptr = (char*)d->arena + d->nstack * sizeof(mjtNum);
char* result = end_ptr - (d->pstack + size + mjREDZONE) * sizeof(mjtNum);
size_t new_pstack = d->pstack + size + 2*mjREDZONE;
#undef mjREDZONE
// new stack usage level
size_t usage;
#ifdef ADDRESS_SANITIZER
if ((uintptr_t)result % sizeof(mjtNum)) {
mjERROR("mj_stackAlloc fails to align to sizeof(mjtNum)");
if ((uintptr_t)start_ptr % sizeof(mjtNum)) {
mjERROR("mj_stackAlloc failed to align to sizeof(mjtNum)");
}
// actual stack usage (without red zone bytes) is stored in the red zone
if (d->pstack) {
size_t* prev_ptr = (size_t*)(end_ptr - d->pstack*sizeof(mjtNum));
ASAN_UNPOISON_MEMORY_REGION(prev_ptr, sizeof(size_t));
usage = *prev_ptr + size;
ASAN_POISON_MEMORY_REGION(prev_ptr, sizeof(size_t));
} else {
usage = size;
size_t* prev_usage_ptr = (size_t*)(end_of_arena_ptr - d->pstack*sizeof(mjtNum));
ASAN_UNPOISON_MEMORY_REGION(prev_usage_ptr, sizeof(size_t));
usage = current_alloc_usage + *prev_usage_ptr;
ASAN_POISON_MEMORY_REGION(prev_usage_ptr, sizeof(size_t));
}
// store new stack usage in the red zone
size_t* cur_ptr = (size_t*)(end_ptr - new_pstack*sizeof(mjtNum));
ASAN_UNPOISON_MEMORY_REGION(cur_ptr, sizeof(size_t));
*cur_ptr = usage;
ASAN_POISON_MEMORY_REGION(cur_ptr, sizeof(size_t));
ASAN_UNPOISON_MEMORY_REGION(new_pstack_ptr, sizeof(size_t));
*(size_t*)new_pstack_ptr = usage;
ASAN_POISON_MEMORY_REGION(new_pstack_ptr, sizeof(size_t));
// unpoison the actual usable allocation
ASAN_UNPOISON_MEMORY_REGION(result, size*sizeof(mjtNum));
#else
usage = d->pstack + size;
ASAN_UNPOISON_MEMORY_REGION(start_ptr, size);
#endif
#undef mjREDZONE
// update pstack and max usage statistics
d->pstack = new_pstack;
d->maxuse_stack = mjMAX(d->maxuse_stack, usage);
d->maxuse_arena = mjMAX(d->maxuse_arena, usage*sizeof(mjtNum) + d->parena);
return (mjtNum*)result;
return (void*)start_ptr;
}
mjtNum* mj_stackAlloc(mjData* d, int size) {
return (mjtNum*)mj_stackAllocBytes(d, size * sizeof(mjtNum));
}
int* mj_stackAllocInt(mjData* d, int size) {
// optimize for mjtNum being twice the size of int
if (2*sizeof(int) == sizeof(mjtNum)) {
return (int*)mj_stackAlloc(d, (size + 1) >> 1);
}
// arbitrary bytes sizes
int new_size = (sizeof(int)*size + sizeof(mjtNum) - 1) / sizeof(mjtNum);
return (int*)mj_stackAlloc(d, new_size);
return (int*)mj_stackAllocBytes(d, size * sizeof(int));
}
+3
View File
@@ -108,6 +108,9 @@ MJAPI mjtNum* mj_stackAlloc(mjData* d, int size);
// mjData stack allocate for array of ints
MJAPI int* mj_stackAllocInt(mjData* d, int size);
// mjData stack allocate for a specific size of bytes
MJAPI void* mj_stackAllocBytes(mjData* d, size_t size);
// de-allocate data
MJAPI void mj_deleteData(mjData* d);
+2 -2
View File
@@ -69,14 +69,14 @@ static void set0(mjModel* m, mjData* d) {
// save camera and light mode, set to fixed
if (m->ncam) {
cammode = (int*) mj_stackAlloc(d, m->ncam);
cammode = mj_stackAllocInt(d, m->ncam);
for (int i=0; i < m->ncam; i++) {
cammode[i] = m->cam_mode[i];
m->cam_mode[i] = mjCAMLIGHT_FIXED;
}
}
if (m->nlight) {
lightmode = (int*) mj_stackAlloc(d, m->nlight);
lightmode = mj_stackAllocInt(d, m->nlight);
for (int i=0; i < m->nlight; i++) {
lightmode[i] = m->light_mode[i];
m->light_mode[i] = mjCAMLIGHT_FIXED;
+4 -4
View File
@@ -1277,7 +1277,7 @@ static void HessianCone(const mjModel* m, mjData* d, mjCGContext* ctx) {
// storage for L'*J
mjtNum* LTJ = mj_stackAlloc(d, 6*nv);
mjtNum* LTJ_row = mj_stackAlloc(d, nv);
int* LTJ_ind = (int*) mj_stackAlloc(d, nv);
int* LTJ_ind = mj_stackAllocInt(d, nv);
// start with Hcone = H
mju_copy(ctx->Hcone, ctx->H, ctx->nnz);
@@ -1366,8 +1366,8 @@ static void HessianDirect(const mjModel* m, mjData* d, mjCGContext* ctx) {
if (mj_isSparse(m)) {
// create sparse inertia matrix M
int nnz = m->nD; // use sparse dof-dof matrix
int* M_rownnz = (int*) mj_stackAlloc(d, nv); // actual nnz count
int* M_colind = (int*) mj_stackAlloc(d, nnz);
int* M_rownnz = mj_stackAllocInt(d, nv); // actual nnz count
int* M_colind = mj_stackAllocInt(d, nnz);
mjtNum* M = mj_stackAlloc(d, nnz);
mj_makeMSparse(m, d, M, M_rownnz, NULL, M_colind);
@@ -1443,7 +1443,7 @@ static void HessianIncremental(const mjModel* m, mjData* d,
// local space
mjtNum* vec = mj_stackAlloc(d, nv);
int* vec_ind = (int*) mj_stackAlloc(d, nv);
int* vec_ind = mj_stackAllocInt(d, nv);
// clear update counter
ctx->nupdate = 0;
+3 -3
View File
@@ -953,8 +953,8 @@ void mj_addM(const mjModel* m, mjData* d, mjtNum* dst,
mjMARKSTACK;
// create sparse inertia matrix M
int nnz = m->nD; // use sparse dof-dof matrix
int* M_rownnz = (int*) mj_stackAlloc(d, nv); // actual nnz count
int* M_colind = (int*) mj_stackAlloc(d, nnz);
int* M_rownnz = mj_stackAllocInt(d, nv); // actual nnz count
int* M_colind = mj_stackAllocInt(d, nnz);
mjtNum* M = mj_stackAlloc(d, nnz);
mj_makeMSparse(m, d, M, M_rownnz, NULL, M_colind);
@@ -1048,7 +1048,7 @@ void mj_addMSparse(const mjModel* m, mjData* d, mjtNum* dst,
}
mjMARKSTACK;
int* buf_ind = (int*) mj_stackAlloc(d, nv);
int* buf_ind = mj_stackAllocInt(d, nv);
mjtNum* sparse_buf = mj_stackAlloc(d, nv);
// add to destination
+2 -2
View File
@@ -149,7 +149,7 @@ int mju_cholFactorSparse(mjtNum* mat, int n, mjtNum mindiag,
int rank = n;
mjMARKSTACK;
int* buf_ind = (int*) mj_stackAlloc(d, n);
int* buf_ind = mj_stackAllocInt(d, n);
mjtNum* sparse_buf = mj_stackAlloc(d, n);
// shrink rows so that rownnz ends at diagonal
@@ -255,7 +255,7 @@ int mju_cholUpdateSparse(mjtNum* mat, mjtNum* x, int n, int flg_plus,
int* rownnz, int* rowadr, int* colind, int x_nnz, int* x_ind,
mjData* d) {
mjMARKSTACK;
int* buf_ind = (int*) mj_stackAlloc(d, n);
int* buf_ind = mj_stackAllocInt(d, n);
mjtNum* sparse_buf = mj_stackAlloc(d, n);
// backpass over rows corresponding to non-zero x(r)