Don't initialize plugins during model compilation.

PiperOrigin-RevId: 585701389
Change-Id: I4e5e34276f404abadf6f5d4d9764ca717304ee34
This commit is contained in:
Yuval Tassa
2023-11-27 11:00:28 -08:00
committed by Copybara-Service
parent 6412c95a37
commit 51d80564fa
4 changed files with 95 additions and 139 deletions
+78 -36
View File
@@ -415,6 +415,7 @@ static void mj_setPtrModel(mjModel* m) {
}
// increases buffer size without causing integer overflow, returns 0 if
// operation would cause overflow
// performs the following operations:
@@ -595,6 +596,8 @@ mjModel* mj_makeModel(
return m;
}
// copy mjModel, if dest==NULL create new model
mjModel* mj_copyModel(mjModel* dest, const mjModel* src) {
void* save_bufptr;
@@ -1094,8 +1097,25 @@ static void mj_setPtrData(const mjModel* m, mjData* d) {
// allocate and initialize mjData structure
static mjData* _makeData(const mjModel* m) {
// initialize plugins, copy into d (required for deletion)
static void _initPlugin(const mjModel* m, mjData* d) {
d->nplugin = m->nplugin;
for (int i = 0; i < m->nplugin; ++i) {
d->plugin[i] = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (plugin->init && plugin->init(m, d, i) < 0) {
mju_free(d->buffer);
mju_free(d->arena);
mju_free(d);
mjERROR("plugin->init failed for plugin id %d", i);
}
}
}
// allocate and initialize raw mjData structure
mjData* mj_makeRawData(const mjModel* m) {
intptr_t offset = 0;
// allocate mjData
@@ -1138,38 +1158,32 @@ static mjData* _makeData(const mjModel* m) {
mjERROR("could not allocate mjData arena");
}
// set pointers into buffer, reset data
// set pointers into buffer
mj_setPtrData(m, d);
// copy plugins into d, required for deletion
d->nplugin = m->nplugin;
for (int i = 0; i < m->nplugin; ++i) {
d->plugin[i] = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (plugin->init) {
if (plugin->init(m, d, i) < 0) {
mju_free(d->buffer);
mju_free(d->arena);
mju_free(d);
mjERROR("plugin->init failed for plugin id %d", i);
}
}
}
// clear threadpool
d->threadpool = 0;
// clear nplugin (overwritten by _initPlugin)
d->nplugin = 0;
return d;
}
// allocate and initialize mjData structure
mjData* mj_makeData(const mjModel* m) {
mjData* d = _makeData(m);
mjData* d = mj_makeRawData(m);
if (d) {
_initPlugin(m, d);
mj_resetData(m, d);
}
return d;
}
// copy mjData, if dest==NULL create new data
mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src) {
void* save_buffer;
@@ -1177,7 +1191,8 @@ mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src) {
// allocate new data if needed
if (!dest) {
dest = _makeData(m);
dest = mj_makeRawData(m);
_initPlugin(m, dest);
}
// check sizes
@@ -1252,6 +1267,7 @@ mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src) {
}
static void maybe_lock_alloc_mutex(mjData* d) {
if (d->threadpool != 0) {
mju_threadPoolLockAllocMutex((mjThreadPool*)d->threadpool);
@@ -1264,7 +1280,6 @@ static void maybe_unlock_alloc_mutex(mjData* d) {
}
}
// allocate memory from the mjData arena
void* mj_arenaAllocByte(mjData* d, size_t bytes, size_t alignment) {
maybe_lock_alloc_mutex(d);
@@ -1366,6 +1381,7 @@ static inline void* stackallocinternal(mjData* d, mjStackInfo* stack_info, size_
}
static inline mjStackInfo get_stack_info_from_data(mjData* d) {
mjStackInfo stack_info;
stack_info.bottom = (uintptr_t)d->arena + (uintptr_t)d->narena;
@@ -1377,6 +1393,7 @@ static inline mjStackInfo get_stack_info_from_data(mjData* d) {
}
// internal: allocate size bytes in mjData
// 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) {
@@ -1396,6 +1413,7 @@ static inline void* stackalloc(mjData* d, size_t size, size_t alignment) {
}
// mjStackInfo mark stack frame, inline so ASAN errors point to correct code unit
#ifdef ADDRESS_SANITIZER
__attribute__((always_inline))
@@ -1414,6 +1432,7 @@ static inline void markstackinternal(mjData* d, mjStackInfo* stack_info) {
}
// mjData mark stack frame
#ifdef ADDRESS_SANITIZER
__attribute__((noinline))
@@ -1432,6 +1451,8 @@ void mj_markStack(mjData* d) {
markstackinternal(d, stack_info);
}
#ifdef ADDRESS_SANITIZER
__attribute__((always_inline))
#endif
@@ -1466,6 +1487,7 @@ static inline void freestackinternal(mjStackInfo* stack_info) {
}
// mjData free stack frame
#ifdef ADDRESS_SANITIZER
__attribute__((noinline))
@@ -1484,14 +1506,23 @@ void mj_freeStack(mjData* d) {
freestackinternal(stack_info);
}
// allocate bytes on the stack
void* mj_stackAllocByte(mjData* d, size_t bytes, size_t alignment) {
return stackalloc(d, bytes, alignment);
}
// allocate mjtNums on the stack
mjtNum* mj_stackAllocNum(mjData* d, int size) {
return (mjtNum*) stackalloc(d, size * sizeof(mjtNum), _Alignof(mjtNum));
}
// allocate ints on the stack
int* mj_stackAllocInt(mjData* d, int size) {
return (int*) stackalloc(d, size * sizeof(int), _Alignof(int));
}
@@ -1501,10 +1532,14 @@ int* mj_stackAllocInt(mjData* d, int size) {
// clear data, set defaults
static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
//------------------------------ save plugin state and data
mjtNum* plugin_state = mju_malloc(sizeof(mjtNum) * m->npluginstate);
memcpy(plugin_state, d->plugin_state, sizeof(mjtNum) * m->npluginstate);
uintptr_t* plugindata = mju_malloc(sizeof(uintptr_t) * m->nplugin);
memcpy(plugindata, d->plugin_data, sizeof(uintptr_t) * m->nplugin);
mjtNum* plugin_state;
uintptr_t* plugindata;
if (d->nplugin) {
plugin_state = mju_malloc(sizeof(mjtNum) * m->npluginstate);
memcpy(plugin_state, d->plugin_state, sizeof(mjtNum) * m->npluginstate);
plugindata = mju_malloc(sizeof(uintptr_t) * m->nplugin);
memcpy(plugindata, d->plugin_data, sizeof(uintptr_t) * m->nplugin);
}
//------------------------------ clear header
@@ -1627,18 +1662,20 @@ static void _resetData(const mjModel* m, mjData* d, unsigned char debug_value) {
}
// restore pluginstate and plugindata
memcpy(d->plugin_state, plugin_state, sizeof(mjtNum) * m->npluginstate);
mju_free(plugin_state);
memcpy(d->plugin_data, plugindata, sizeof(uintptr_t) * m->nplugin);
mju_free(plugindata);
if (d->nplugin) {
memcpy(d->plugin_state, plugin_state, sizeof(mjtNum) * m->npluginstate);
mju_free(plugin_state);
memcpy(d->plugin_data, plugindata, sizeof(uintptr_t) * m->nplugin);
mju_free(plugindata);
// restore the plugin array back into d and reset the instances
for (int i = 0; i < m->nplugin; ++i) {
d->plugin[i] = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (plugin->reset) {
plugin->reset(m, &d->plugin_state[m->plugin_stateadr[i]],
(void*)(d->plugin_data[i]), i);
// restore the plugin array back into d and reset the instances
for (int i = 0; i < m->nplugin; ++i) {
d->plugin[i] = m->plugin[i];
const mjpPlugin* plugin = mjp_getPluginAtSlot(m->plugin[i]);
if (plugin->reset) {
plugin->reset(m, &d->plugin_state[m->plugin_stateadr[i]],
(void*)(d->plugin_data[i]), i);
}
}
}
}
@@ -1699,6 +1736,7 @@ void mj_deleteData(mjData* d) {
}
// number of position and velocity coordinates for each joint type
const int nPOS[4] = {7, 4, 1, 1};
const int nVEL[4] = {6, 3, 1, 1};
@@ -1762,6 +1800,8 @@ static int sensorSize(mjtSensor sensor_type, int sensor_dim) {
return -1;
}
// returns the number of objects of the given type
// -1: mjOBJ_UNKNOWN
// -2: invalid objtype
@@ -1822,6 +1862,8 @@ static int numObjects(const mjModel* m, mjtObj objtype) {
return -2;
}
// validate reference fields in a model; return null if valid, error message otherwise
const char* mj_validateReferences(const mjModel* m) {
// for each field in mjModel that refers to another field, call X with:
+5 -2
View File
@@ -87,10 +87,13 @@ MJAPI const char* mj_validateReferences(const mjModel* m);
//------------------------------- mjData -----------------------------------------------------------
// Allocate mjData corresponding to given model.
// If the model buffer is unallocated the initial configuration will not be set.
// allocate mjData corresponding to given model, initialize plugins, reset the state
// if the model buffer is unallocated the initial configuration will not be set
MJAPI mjData* mj_makeData(const mjModel* m);
// allocate mjData corresponding to given model, used internally
MJAPI mjData* mj_makeRawData(const mjModel* m);
// Copy mjData.
// m is only required to contain the size fields from MJMODEL_INTS.
MJAPI mjData* mj_copyData(mjData* dest, const mjModel* m, const mjData* src);