From 11af1e2a95012c4f7494c120d182d6f057a0d978 Mon Sep 17 00:00:00 2001 From: Alessio Quaglino Date: Tue, 15 Oct 2024 10:23:43 -0700 Subject: [PATCH] Fix a bug in mj_recompile. The qposadr_ and dofadr_ variables of the joints were not being reset to -1 when a joint is copied. This could lead to errors when calling mj_recompile after attaching a compiled spec. PiperOrigin-RevId: 686152004 Change-Id: I9e058089ce87ea0d8e51a7f43475bc48c6527a26 --- src/user/user_api.cc | 29 +++++++++++++---------- src/user/user_model.cc | 16 +++++++------ src/user/user_objects.cc | 2 ++ src/user/user_objects.h | 6 +++-- test/user/user_api_test.cc | 48 ++++++++++++++++++++++++++++++++++++++ 5 files changed, 80 insertions(+), 21 deletions(-) diff --git a/src/user/user_api.cc b/src/user/user_api.cc index 3745517c..9869141a 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -96,21 +96,26 @@ mjModel* mj_compile(mjSpec* s, const mjVFS* vfs) { mjCModel* modelC = static_cast(s->element); std::string state_name = "state"; mjtNum time = 0; - if (d) { - time = d->time; - modelC->SaveState(state_name, d->qpos, d->qvel, d->act, d->ctrl, d->mocap_pos, d->mocap_quat); - } - if (!modelC->Compile(vfs, &m)) { + try { if (d) { - mj_deleteData(d); + time = d->time; + modelC->SaveState(state_name, d->qpos, d->qvel, d->act, d->ctrl, d->mocap_pos, d->mocap_quat); } + if (!modelC->Compile(vfs, &m)) { + if (d) { + mj_deleteData(d); + } + return -1; + }; + if (d) { + modelC->MakeData(m, &d); + modelC->RestoreState(state_name, m->qpos0, m->body_pos, m->body_quat, d->qpos, d->qvel, + d->act, d->ctrl, d->mocap_pos, d->mocap_quat); + d->time = time; + } + } catch (mjCError& e) { + modelC->SetError(e); return -1; - }; - if (d) { - modelC->MakeData(m, &d); - modelC->RestoreState(state_name, m->qpos0, m->body_pos, m->body_quat, d->qpos, d->qvel, - d->act, d->ctrl, d->mocap_pos, d->mocap_quat); - d->time = time; } return 0; } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index 91ae83d5..60e06670 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -353,9 +353,7 @@ mjCModel& mjCModel::operator+=(const mjCModel& other) { // create global lists mjCBody *world = bodies_[0]; - if (compiled) { - ResetTreeLists(); - } + ResetTreeLists(); MakeLists(world); ProcessLists(/*checkrepeat=*/false); @@ -3034,11 +3032,15 @@ template void mjCModel::SaveState(const std::string& state_name, const T* qpos, const T* qvel, const T* act, const T* ctrl, const T* mpos, const T* mquat) { for (auto joint : joints_) { - if (joint->qposadr_ == -1 || joint->dofadr_ == -1) { - throw mjCError(nullptr, "SaveState: joint %s has no address", joint->name.c_str()); + if (joint->qposadr_ < -1 || joint->dofadr_ < -1) { + throw mjCError(nullptr, "SaveState: joint %s has invalid address", joint->name.c_str()); + } + if (qpos && joint->qposadr_ != -1) { + mjuu_copyvec(joint->qpos(state_name), qpos + joint->qposadr_, joint->nq()); + } + if (qvel && joint->dofadr_ != -1) { + mjuu_copyvec(joint->qvel(state_name), qvel + joint->dofadr_, joint->nv()); } - if (qpos) mjuu_copyvec(joint->qpos(state_name), qpos + joint->qposadr_, joint->nq()); - if (qvel) mjuu_copyvec(joint->qvel(state_name), qvel + joint->dofadr_, joint->nv()); } for (unsigned int i=0; ispec = other.spec; *static_cast(this) = static_cast(other); *static_cast(this) = static_cast(other); + qposadr_ = -1; + dofadr_ = -1; } PointToLocal(); return *this; diff --git a/src/user/user_objects.h b/src/user/user_objects.h index 08aca2f8..c2ee11b6 100644 --- a/src/user/user_objects.h +++ b/src/user/user_objects.h @@ -425,8 +425,6 @@ class mjCJoint_ : public mjCBase { mjCBody* body; // joint's body // variable used for temporarily storing the state of the joint - int qposadr_; // address of dof in data->qpos - int dofadr_; // address of dof in data->qvel std::map> qpos_; // qpos at the previous step std::map> qvel_; // qvel at the previous step @@ -473,6 +471,10 @@ class mjCJoint : public mjCJoint_, private mjsJoint { private: int Compile(void); // compiler; return dofnum void PointToLocal(void); + + // variables that should not be copied during copy assignment + int qposadr_; // address of dof in data->qpos + int dofadr_; // address of dof in data->qvel }; diff --git a/test/user/user_api_test.cc b/test/user/user_api_test.cc index bf5c6079..8404dc47 100644 --- a/test/user/user_api_test.cc +++ b/test/user/user_api_test.cc @@ -1368,6 +1368,54 @@ TEST_F(MujocoTest, PreserveState) { mj_deleteModel(m_expected); } +TEST_F(MujocoTest, RecompileAttach) { + std::array er; + std::string field = ""; + + static constexpr char xml[] = R"( + + + + + + + + )"; + + mjSpec* parent = mj_makeSpec(); + EXPECT_THAT(parent, NotNull()); + + mjSpec* child = mj_parseXMLString(xml, 0, er.data(), er.size()); + EXPECT_THAT(child, NotNull()); + + mjsFrame* frame1 = mjs_addFrame(mjs_findBody(parent, "world"), 0); + mjs_attachBody(frame1, mjs_findBody(child, "body"), "child-", "-1"); + + mjModel* model = mj_compile(parent, 0); + EXPECT_THAT(model, NotNull()); + EXPECT_THAT(model->nq, 1); + mjData* data = mj_makeData(model); + EXPECT_THAT(data, NotNull()); + + for (int i = 0; i < 100; i++) { + mj_step(model, data); + } + + mjsFrame* frame2 = mjs_addFrame(mjs_findBody(parent, "world"), 0); + mjs_attachBody(frame2, mjs_findBody(child, "body"), "child-", "-2"); + + EXPECT_EQ(mj_recompile(parent, 0, model, data), 0); + EXPECT_THAT(model, NotNull()); + EXPECT_THAT(model->nq, 2); + EXPECT_NE(data->qpos[0], data->qpos[1]); + EXPECT_EQ(data->qpos[1], 0); + + mj_deleteData(data); + mj_deleteModel(model); + mj_deleteSpec(child); + mj_deleteSpec(parent); +} + TEST_F(MujocoTest, AttachMocap) { std::array er; mjtNum tol = 0;