diff --git a/src/user/user_api.cc b/src/user/user_api.cc index 79c20224..8b7d449a 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -86,9 +86,14 @@ mjModel* mj_compile(mjSpec* s, const mjVFS* vfs) { // recompile spec into existing model and data while preserving the state void mj_recompile(mjSpec* s, const mjVFS* vfs, mjModel* m, mjData* d) { mjCModel* modelC = static_cast(s->element); - modelC->SaveState(d); + if (d) { + modelC->SaveState(d->qpos, d->qvel, d->act); + } modelC->Compile(vfs, &m); - modelC->RestoreState(m, &d); + if (d) { + modelC->MakeData(m, &d); + modelC->RestoreState(d->qpos, d->qvel, d->act); + } } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index b63f2c12..5f0631b8 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -2776,68 +2776,73 @@ void mjCModel::CopyObjects(mjModel* m) { // save the current state -void mjCModel::SaveState(const mjData* d) { +void mjCModel::SaveState(const mjtNum* qpos, const mjtNum* qvel, const mjtNum* act) { for (auto joint : joints_) { switch (joint->type) { case mjJNT_FREE: - mjuu_copyvec(joint->qpos, d->qpos + joint->qposadr_, 7); - mjuu_copyvec(joint->qvel, d->qvel + joint->dofadr_, 6); + if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 7); + if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 6); break; case mjJNT_BALL: - mjuu_copyvec(joint->qpos, d->qpos + joint->qposadr_, 4); - mjuu_copyvec(joint->qvel, d->qvel + joint->dofadr_, 3); + if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 4); + if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 3); break; case mjJNT_HINGE: case mjJNT_SLIDE: - mjuu_copyvec(joint->qpos, d->qpos + joint->qposadr_, 1); - mjuu_copyvec(joint->qvel, d->qvel + joint->dofadr_, 1); + if (qpos) mjuu_copyvec(joint->qpos, qpos + joint->qposadr_, 1); + if (qvel) mjuu_copyvec(joint->qvel, qvel + joint->dofadr_, 1); break; } } for (auto actuator : actuators_) { - if (actuator->actadr_ != -1) { + if (actuator->actadr_ != -1 && act) { actuator->act.assign(actuator->actnum_, 0); - mjuu_copyvec(actuator->act.data(), d->act + actuator->actadr_, actuator->actnum_); + mjuu_copyvec(actuator->act.data(), act + actuator->actadr_, actuator->actnum_); } } } -// restore the previous state -void mjCModel::RestoreState(const mjModel* m, mjData** dest) { +// clear existing data +void mjCModel::MakeData(const mjModel* m, mjData** dest) { mj_makeRawData(dest, m); mjData* d = *dest; if (d) { mj_initPlugin(m, d); mj_resetData(m, d); } +} + + +// restore the previous state +void mjCModel::RestoreState(mjtNum* qpos, mjtNum* qvel, mjtNum* act) { for (auto joint : joints_) { if (!mjuu_defined(joint->qpos[0]) || !mjuu_defined(joint->qvel[0])) { continue; } switch (joint->type) { case mjJNT_FREE: - mjuu_copyvec(d->qpos + joint->qposadr_, joint->qpos, 7); - mjuu_copyvec(d->qvel + joint->dofadr_, joint->qvel, 6); + if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 7); + if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 6); break; case mjJNT_BALL: - mjuu_copyvec(d->qpos + joint->qposadr_, joint->qpos, 4); - mjuu_copyvec(d->qvel + joint->dofadr_, joint->qvel, 3); + if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 4); + if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 3); break; case mjJNT_HINGE: case mjJNT_SLIDE: - mjuu_copyvec(d->qpos + joint->qposadr_, joint->qpos, 1); - mjuu_copyvec(d->qvel + joint->dofadr_, joint->qvel, 1); + if (qpos) mjuu_copyvec(qpos + joint->qposadr_, joint->qpos, 1); + if (qvel) mjuu_copyvec(qvel + joint->dofadr_, joint->qvel, 1); break; } } for (auto actuator : actuators_) { - if (mjuu_defined(actuator->act[0])) { - mjuu_copyvec(d->act + actuator->actadr_, actuator->act.data(), actuator->actnum_); + if (mjuu_defined(actuator->act[0]) && act) { + mjuu_copyvec(act + actuator->actadr_, actuator->act.data(), actuator->actnum_); } } } diff --git a/src/user/user_model.h b/src/user/user_model.h index a4bd11ed..7e36d6ec 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -273,8 +273,9 @@ class mjCModel : public mjCModel_, private mjSpec { std::string_view name = ""); // save/restore the current state - void SaveState(const mjData* d); - void RestoreState(const mjModel* m, mjData** dest); + void SaveState(const mjtNum* qpos, const mjtNum* qvel, const mjtNum* act); + void MakeData(const mjModel* m, mjData** dest); + void RestoreState(mjtNum* qpos, mjtNum* qvel, mjtNum* act); // map from default class name to default class pointer std::unordered_map def_map; diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index 64d87650..bebdc010 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -863,8 +863,11 @@ TEST_F(MujocoTest, PreserveState) { EXPECT_EQ(data->act[i], d_expected->act[i]) << i; } - // destroy everything + // check that the function is callable with no data mj_deleteData(data); + mj_recompile(spec, 0, model, nullptr); + + // destroy everything mj_deleteData(d_expected); mj_deleteSpec(spec); mj_deleteModel(model);