Don't initialize plugins during model compilation.
PiperOrigin-RevId: 585701389 Change-Id: I4e5e34276f404abadf6f5d4d9764ca717304ee34
This commit is contained in:
committed by
Copybara-Service
parent
6412c95a37
commit
51d80564fa
+78
-36
@@ -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:
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user