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);
+12 -3
View File
@@ -32,7 +32,6 @@
#include "engine/engine_io.h"
#include "engine/engine_plugin.h"
#include "engine/engine_setconst.h"
#include "engine/engine_resource.h"
#include "engine/engine_support.h"
#include "engine/engine_util_blas.h"
#include "engine/engine_util_errmem.h"
@@ -3100,11 +3099,12 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) {
// create data
int disableflags = m->opt.disableflags;
m->opt.disableflags |= mjDSBL_CONTACT;
d = mj_makeData(m);
d = mj_makeRawData(m);
if (!d) {
mj_deleteModel(m);
throw mjCError(0, "could not create mjData");
}
mj_resetData(m, d);
// normalize keyframe quaternions
for (int i=0; i<m->nkey; i++) {
@@ -3134,8 +3134,17 @@ void mjCModel::TryCompile(mjModel*& m, mjData*& d, const mjVFS* vfs) {
mj_deleteModel(m);
throw mjCError(0, validationerr);
}
// delete partial mjData (no plugins), make a complete one
mj_deleteData(d);
d = nullptr;
d = mj_makeData(m);
if (!d) {
mj_deleteModel(m);
throw mjCError(0, "could not create mjData");
}
// test forward simulation
mj_resetData(m, d);
mj_step(m, d);
// delete data
-98
View File
@@ -44,16 +44,6 @@ using ::testing::NotNull;
using EngineIoTest = MujocoTest;
// Return an mjModel with just the ints set.
mjModel PartialModel(const mjModel* m) {
mjModel partial_model = {0};
#define X(var) partial_model.var = m->var;
MJMODEL_INTS;
#undef X
partial_model.nbuffer = 0;
return partial_model;
}
TEST_F(EngineIoTest, VerifySizeModel) {
constexpr char xml[] = R"(
<mujoco>
@@ -84,49 +74,6 @@ TEST_F(EngineIoTest, VerifySizeModel) {
EXPECT_EQ(file_size, model_size);
}
TEST_F(EngineIoTest, MakeDataFromPartialModel) {
constexpr char xml[] = R"(
<mujoco>
<worldbody>
<body>
<joint/>
<geom size="1"/>
</body>
</worldbody>
</mujoco>
)";
std::array<char, 1024> error;
mjModel* model = LoadModelFromString(xml, error.data(), error.size());
ASSERT_THAT(model, NotNull()) << "Failed to load model: " << error.data();
mjData* data_from_model = mj_makeData(model);
ASSERT_THAT(data_from_model, NotNull());
mjModel partial_model = PartialModel(model);
mj_deleteModel(model);
mjData* data_from_partial = mj_makeData(&partial_model);
ASSERT_THAT(data_from_partial, NotNull());
EXPECT_EQ(data_from_partial->nbuffer, data_from_model->nbuffer);
// If there are no mocap bodies and qpos0 is all zero, mjData should be the
// same whether it was made from the full model or the partial model.
{
MJDATA_POINTERS_PREAMBLE((&partial_model))
#define X(type, name, nr, nc) \
if (strcmp(#name, "D_rownnz") && strcmp(#name, "D_rowadr") && \
strcmp(#name, "B_rownnz") && strcmp(#name, "B_rowadr")) \
EXPECT_EQ(std::memcmp(data_from_partial->name, data_from_model->name, \
sizeof(type)*(partial_model.nr)*(nc)), \
0) << "mjData::" #name " differs";
MJDATA_POINTERS
#undef X
}
mj_deleteData(data_from_model);
mj_deleteData(data_from_partial);
}
TEST_F(EngineIoTest, MakeDataLoadsQpos0) {
constexpr char xml[] = R"(
<mujoco>
@@ -173,51 +120,6 @@ TEST_F(EngineIoTest, MakeDataLoadsMocapBodies) {
mj_deleteModel(model);
}
TEST_F(EngineIoTest, CopyDataWithPartialModel) {
constexpr char xml[] = R"(
<mujoco>
<worldbody>
<body>
<joint/>
<geom size="1"/>
</body>
<body mocap="true" pos="42 0 42">
<geom type="sphere" size="0.1"/>
</body>
</worldbody>
</mujoco>
)";
std::array<char, 1024> error;
mjModel* model = LoadModelFromString(xml, error.data(), error.size());
ASSERT_THAT(model, NotNull()) << "Failed to load model: " << error.data();
mjData* data = mj_makeData(model);
ASSERT_THAT(data, NotNull());
mjModel partial_model = PartialModel(model);
mj_deleteModel(model);
mjData* copy = mj_makeData(&partial_model);
ASSERT_THAT(copy, NotNull());
data->qpos[0] = 1;
mj_copyData(copy, &partial_model, data);
EXPECT_EQ(copy->nbuffer, data->nbuffer);
EXPECT_EQ(copy->qpos[0], 1);
{
MJDATA_POINTERS_PREAMBLE((&partial_model))
#define X(type, name, nr, nc) \
EXPECT_EQ(std::memcmp(copy->name, data->name, \
sizeof(type)*(partial_model.nr)*(nc)), \
0) << "mjData::" #name " differs";
MJDATA_POINTERS
#undef X
}
mj_deleteData(data);
mj_deleteData(copy);
}
TEST_F(EngineIoTest, MakeDataReturnsNullOnFailure) {
constexpr char xml[] = "<mujoco/>";