diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 70bbc18f..c280911e 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -568,7 +568,6 @@ PYBIND11_MODULE(_specs, m) { // ============================= MJSBODY ===================================== mjsBody.def_property_readonly( "id", [](raw::MjsBody& self) -> int { return mjs_getId(self.element); }); - mjsBody.def("delete", [](raw::MjsBody& self) { mjs_delete(self.element); }); mjsBody.def( "add_body", [](raw::MjsBody& self, raw::MjsDefault* default_) -> raw::MjsBody* { diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 0459b013..ad068f91 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -98,6 +98,34 @@ class SpecsTest(absltest.TestCase): """),) + def test_load_xml(self): + filename = '../../test/testdata/model.xml' + state_type = mujoco.mjtState.mjSTATE_INTEGRATION + + # Load from file. + spec1 = mujoco.MjSpec() + spec1.from_file(filename) + model1 = spec1.compile() + data1 = mujoco.MjData(model1) + mujoco.mj_step(model1, data1) + size1 = mujoco.mj_stateSize(model1, state_type) + state1 = np.empty(size1, np.float64) + mujoco.mj_getState(model1, data1, state1, state_type) + + # Load from string. + spec2 = mujoco.MjSpec() + with open(filename, 'r') as file: + spec2.from_string(file.read().rstrip()) + model2 = spec2.compile() + data2 = mujoco.MjData(model2) + mujoco.mj_step(model2, data2) + size2 = mujoco.mj_stateSize(model2, state_type) + state2 = np.empty(size2, np.float64) + mujoco.mj_getState(model2, data2, state2, state_type) + + # Check that the state is the same. + np.testing.assert_array_equal(state1, state2) + def test_compile_errors_with_line_info(self): spec = mujoco.MjSpec() @@ -269,6 +297,30 @@ class SpecsTest(absltest.TestCase): model = spec.compile({'cube.obj': cube}) self.assertEqual(model.nmeshvert, 8) + def test_delete(self): + filename = '../../test/testdata/model.xml' + + spec = mujoco.MjSpec() + spec.from_file(filename) + + model = spec.compile() + self.assertIsNotNone(model) + self.assertEqual(model.nsite, 11) + self.assertEqual(model.nsensor, 11) + + head = spec.find_body('head') + self.assertIsNotNone(head) + site = head.first_site() + self.assertIsNotNone(site) + + site.delete() + spec.sensors[-1].delete() + spec.sensors[-1].delete() + + model = spec.compile() + self.assertIsNotNone(model) + self.assertEqual(model.nsite, 10) + self.assertEqual(model.nsensor, 9) if __name__ == '__main__': absltest.main() diff --git a/src/user/user_api.cc b/src/user/user_api.cc index 2c3b5aee..ba5890fe 100644 --- a/src/user/user_api.cc +++ b/src/user/user_api.cc @@ -148,7 +148,7 @@ int mjs_detachBody(mjSpec* s, mjsBody* b) { mjCModel* model = static_cast(s->element); mjCBody* body = static_cast(b->element); *model -= *body; - mjs_delete(b->element); + delete body; return 0; } @@ -181,7 +181,7 @@ void mjs_addSpec(mjSpec* s, mjSpec* child) { // delete object, it will call the appropriate destructor since ~mjCBase is virtual void mjs_delete(mjsElement* element) { mjCBase* object = static_cast(element); - delete object; + object->model->DeleteElement(element); } diff --git a/src/user/user_model.cc b/src/user/user_model.cc index fbd94a6b..c7fd5ccb 100644 --- a/src/user/user_model.cc +++ b/src/user/user_model.cc @@ -340,6 +340,68 @@ mjCModel_& mjCModel::operator+=(mjCDef& subtree) { +template +void deletefromlist(std::vector* list, mjsElement* element) { + if (!list) { + return; + } + for (int j = 0; j < list->size(); ++j) { + list->at(j)->id = -1; + if (list->at(j) == element) { + delete list->at(j); + list->erase(list->begin() + j); + j--; + } + } +} + + + +// discard all invalid elements from all lists +void mjCModel::DeleteElement(mjsElement* el) { + mjCBody *world = bodies_[0]; + if (compiled) { + ResetTreeLists(); + } + + switch (el->elemtype) { + case mjOBJ_BODY: + throw mjCError(NULL, "bodies cannot be deleted, use detach instead"); + break; + + case mjOBJ_GEOM: + deletefromlist(&(static_cast(el)->body->geoms), el); + break; + + case mjOBJ_SITE: + deletefromlist(&(static_cast(el)->body->sites), el); + break; + + case mjOBJ_JOINT: + deletefromlist(&(static_cast(el)->body->joints), el); + break; + + case mjOBJ_LIGHT: + deletefromlist(&(static_cast(el)->body->lights), el); + break; + + case mjOBJ_CAMERA: + deletefromlist(&(static_cast(el)->body->cameras), el); + break; + + default: + deletefromlist(object_lists_[el->elemtype], el); + break; + } + + if (compiled) { + MakeLists(world); + ProcessLists(/*checkrepeat=*/false); + } +} + + + // TODO: we should not use C-type casting with multiple C++ inheritance void mjCModel::CreateObjectLists() { for (int i = 0; i < mjNOBJECT; ++i) { diff --git a/src/user/user_model.h b/src/user/user_model.h index 52bef9c1..9f88bb5b 100644 --- a/src/user/user_model.h +++ b/src/user/user_model.h @@ -202,6 +202,9 @@ class mjCModel : public mjCModel_, private mjSpec { // delete all elements template void DeleteAll(std::vector& elements); + // delete object from the corresponding list + void DeleteElement(mjsElement* el); + // API for access to model elements (outside tree) int NumObjects(mjtObj type); // number of objects in specified list mjCBase* GetObject(mjtObj type, int id); // pointer to specified object