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);
|
||||
|
||||
+12
-3
@@ -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
|
||||
|
||||
@@ -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/>";
|
||||
|
||||
|
||||
Reference in New Issue
Block a user