diff --git a/doc/changelog.rst b/doc/changelog.rst index 9015ccce..25a57602 100644 --- a/doc/changelog.rst +++ b/doc/changelog.rst @@ -23,20 +23,22 @@ Python bindings 6. Fix the macOS ``mjpython`` launcher to work with the Python interpreter from Apple Command Line Tools. +7. Fixed a crash when copying instances of ``mujoco.MjData`` for models that use plugins. Introduced a ``model`` + attribute to ``MjData`` which is reference to the model that was used to create that ``MjData`` instance. Simulate ^^^^^^^^ -7. :ref:`simulate`: correct handling of "Pause update", "Fullscreen" and "VSync" buttons. +8. :ref:`simulate`: correct handling of "Pause update", "Fullscreen" and "VSync" buttons. Documentation ^^^^^^^^^^^^^ -8. Added documentation for the :ref:`UI` framework. -9. Fixed typos and supported fields in docs (fixes :github:issue:`1105` and :github:issue:`1106`). +9. Added documentation for the :ref:`UI` framework. +10. Fixed typos and supported fields in docs (fixes :github:issue:`1105` and :github:issue:`1106`). Bug fixes ^^^^^^^^^ -10. Fixed bug relating to welds modified with :ref:`torquescale`. +11. Fixed bug relating to welds modified with :ref:`torquescale`. Version 3.0.0 (October 18, 2023) -------------------------------- diff --git a/python/mujoco/CMakeLists.txt b/python/mujoco/CMakeLists.txt index 46f5f0ee..942397d5 100644 --- a/python/mujoco/CMakeLists.txt +++ b/python/mujoco/CMakeLists.txt @@ -265,7 +265,6 @@ target_sources( structs_header INTERFACE indexer_xmacro.h indexers.h - mjdata_meta.h structs.h ) set_target_properties(structs_header PROPERTIES PUBLIC_HEADER structs.h) diff --git a/python/mujoco/bindings_test.py b/python/mujoco/bindings_test.py index 104f8b66..e81bd3c2 100644 --- a/python/mujoco/bindings_test.py +++ b/python/mujoco/bindings_test.py @@ -74,10 +74,21 @@ TEST_XML_SENSOR = r""" """ TEST_XML_PLUGIN = r""" - + + """ @@ -1048,10 +1059,24 @@ Euler integrator, semi-implicit in velocity. while data_instances: d = data_instances.pop() self.assertEqual(sys.getrefcount(d), 2) + del d while model_instances: m = model_instances.pop() self.assertEqual(sys.getrefcount(m), 2) + # This test is disabled on PyPy as it uses sys.getrefcount + # However PyPy is not officially supported by MuJoCo + @absltest.skipIf(sys.implementation.name == 'pypy', + reason='requires sys.getrefcount') + def test_mjdata_holds_ref_to_model(self): + data = mujoco.MjData(mujoco.MjModel.from_xml_string('')) + model = data.model + # references: one in `data.model, one in `model`, one in the temporary ref + # passed to getrefcount. + self.assertEqual(sys.getrefcount(data.model), 3) + del data + self.assertEqual(sys.getrefcount(model), 2) + def test_can_initialize_mjv_structs(self): self.assertIsInstance(mujoco.MjvScene(), mujoco.MjvScene) self.assertIsInstance(mujoco.MjvCamera(), mujoco.MjvCamera) @@ -1287,6 +1312,34 @@ Euler integrator, semi-implicit in velocity. self.assertEqual(data.geom(3).xpos[2], 4) self.assertEqual(data.geom(4).xpos[2], 5) + def test_load_plugin(self): + model = mujoco.MjModel.from_xml_string(TEST_XML_PLUGIN) + data = mujoco.MjData(model) + mujoco.mj_forward(model, data) + + def test_copy_mjdata_with_plugin(self): + model = mujoco.MjModel.from_xml_string(TEST_XML_PLUGIN) + data1 = mujoco.MjData(model) + self.assertIs(data1.model, model) + mujoco.mj_step(model, data1) + data2 = copy.copy(data1) + mujoco.mj_step(model, data1) + mujoco.mj_step(model, data2) + np.testing.assert_array_equal(data1.qpos, data2.qpos) + self.assertIs(data1.model, data2.model) + + def test_deepcopy_mjdata_with_plugin(self): + model = mujoco.MjModel.from_xml_string(TEST_XML_PLUGIN) + data1 = mujoco.MjData(model) + self.assertIs(data1.model, model) + mujoco.mj_step(model, data1) + data2 = copy.deepcopy(data1) + mujoco.mj_step(model, data1) + mujoco.mj_step(model, data2) + np.testing.assert_array_equal(data1.qpos, data2.qpos) + self.assertIsNot(data1.model, data2.model) + self.assertNotEqual(data1.model._address, data2.model._address) + def _assert_attributes_equal(self, actual_obj, expected_obj, attr_to_compare): for name in attr_to_compare: actual_value = getattr(actual_obj, name) @@ -1300,9 +1353,6 @@ Euler integrator, semi-implicit in velocity. self.fail("Attribute '{}' differs from expected value: {}".format( name, str(e))) - def test_load_plugin(self): - mujoco.MjModel.from_xml_string(TEST_XML_PLUGIN) - if __name__ == '__main__': absltest.main() diff --git a/python/mujoco/functions.cc b/python/mujoco/functions.cc index 86402a89..cc8c774a 100644 --- a/python/mujoco/functions.cc +++ b/python/mujoco/functions.cc @@ -23,6 +23,7 @@ #include #include +#include #include "function_traits.h" #include "functions.h" #include "private.h" @@ -1367,7 +1368,7 @@ PYBIND11_MODULE(_functions, pymodule) { } #undef MJ_M -#define MJ_M(x) d.metadata().x +#define MJ_M(x) d.model().get()->x #undef MJ_D #define MJ_D(x) data->x #define X(type, name, nr, nc) \ @@ -1379,7 +1380,7 @@ PYBIND11_MODULE(_functions, pymodule) { } MJDATA_ARENA_POINTERS_PRIMAL - if (d.metadata().is_dual) { + if (mj_isDual(d.model().get())) { MJDATA_ARENA_POINTERS_DUAL } #undef X diff --git a/python/mujoco/indexers.cc b/python/mujoco/indexers.cc index e778272a..72311a93 100644 --- a/python/mujoco/indexers.cc +++ b/python/mujoco/indexers.cc @@ -13,16 +13,13 @@ // limitations under the License. #include -#include #include #include #include #include -#include #include #include "indexers.h" -#include "mjdata_meta.h" #include "raw.h" #include "util/crossplatform.h" @@ -76,7 +73,6 @@ IDToName MakeIDToName(int count, IntPtr name_offsets, CharPtr names) { // the MuJoCo category of the entity itself, e.g. nbody indicates that the // field belongs to a body. // T: Scalar data type of the field. -// M: Either raw::MjModel or MjDataMetadata. // // Args: // base_ptr: Pointer to the first entry in the entire field. @@ -87,35 +83,35 @@ IDToName MakeIDToName(int count, IntPtr name_offsets, CharPtr names) { // additional dimension of the size len(qvel) of the particular joint. // m: Used for dereferencing MjSize. // owner: The base object whose lifetime is tied to the returned array. -template +template py::array_t MakeArray(T* base_ptr, int index, std::vector&& shape, - const M& m, py::handle owner) { + const raw::MjModel& m, py::handle owner) { int offset; - if (MjSize == &M::nq) { + if (MjSize == &raw::MjModel::nq) { offset = m.jnt_qposadr[index]; shape.insert( shape.begin(), ((index < m.njnt-1) ? m.jnt_qposadr[index+1] : m.nq) - offset); - } else if (MjSize == &M::nv) { + } else if (MjSize == &raw::MjModel::nv) { offset = m.jnt_dofadr[index]; shape.insert( shape.begin(), ((index < m.njnt-1) ? m.jnt_dofadr[index+1] : m.nv) - offset); - } else if (MjSize == &M::nhfielddata) { + } else if (MjSize == &raw::MjModel::nhfielddata) { offset = m.hfield_adr[index]; shape.insert(shape.begin(), m.hfield_ncol[index]); shape.insert(shape.begin(), m.hfield_nrow[index]); - } else if (MjSize == &M::ntexdata) { + } else if (MjSize == &raw::MjModel::ntexdata) { offset = m.tex_adr[index]; shape.insert(shape.begin(), m.tex_width[index]); shape.insert(shape.begin(), m.tex_height[index]); - } else if (MjSize == &M::nsensordata) { + } else if (MjSize == &raw::MjModel::nsensordata) { offset = m.sensor_adr[index]; shape.insert(shape.begin(), m.sensor_dim[index]); - } else if (MjSize == &M::nnumericdata) { + } else if (MjSize == &raw::MjModel::nnumericdata) { offset = m.numeric_adr[index]; shape.insert(shape.begin(), m.numeric_size[index]); - } else if (MjSize == &M::ntupledata) { + } else if (MjSize == &raw::MjModel::ntupledata) { offset = m.tuple_adr[index]; shape.insert(shape.begin(), m.tuple_size[index]); } else { @@ -135,9 +131,7 @@ py::array_t MakeArray(T* base_ptr, int index, std::vector&& shape, } } // namespace -// M is either a raw::MjModel or MjDataMetadata. -template -NameToIDMappings::NameToIDMappings(const M& m) +NameToIDMappings::NameToIDMappings(const raw::MjModel& m) : body(MakeNameToID(m.nbody, m.name_bodyadr, m.names)), jnt(MakeNameToID(m.njnt, m.name_jntadr, m.names)), geom(MakeNameToID(m.ngeom, m.name_geomadr, m.names)), @@ -160,9 +154,7 @@ NameToIDMappings::NameToIDMappings(const M& m) tuple(MakeNameToID(m.ntuple, m.name_tupleadr, m.names)), key(MakeNameToID(m.nkey, m.name_keyadr, m.names)) {} -// M is either a raw::MjModel or MjDataMetadata. -template -IDToNameMappings::IDToNameMappings(const M& m) +IDToNameMappings::IDToNameMappings(const raw::MjModel& m) : body(MakeIDToName(m.nbody, m.name_bodyadr, m.names)), jnt(MakeIDToName(m.njnt, m.name_jntadr, m.names)), geom(MakeIDToName(m.ngeom, m.name_geomadr, m.names)), @@ -223,7 +215,7 @@ MJMODEL_VIEW_GROUPS MJMODEL_VIEW_GROUPS #undef XGROUP -MjDataIndexer::MjDataIndexer(raw::MjData* d, const MjDataMetadata* m, +MjDataIndexer::MjDataIndexer(raw::MjData* d, const raw::MjModel* m, py::handle owner) : d_(d), m_(m), @@ -369,7 +361,7 @@ MJMODEL_KEYFRAME #define X(type, prefix, var, dim0, dim1) \ py::array_t XGROUP::var() { \ if (!var##_.has_value()) { \ - var##_.emplace(MakeArray<&MjDataMetadata::dim0>( \ + var##_.emplace(MakeArray<&raw::MjModel::dim0>( \ d_->prefix##var, index_, MAKE_SHAPE(dim1), *m_, owner_)); \ } \ return *var##_; \ diff --git a/python/mujoco/indexers.h b/python/mujoco/indexers.h index 4b83185f..2b9defad 100644 --- a/python/mujoco/indexers.h +++ b/python/mujoco/indexers.h @@ -25,7 +25,6 @@ #include #include #include "indexer_xmacro.h" -#include "mjdata_meta.h" #include "raw.h" #include #include @@ -35,9 +34,7 @@ using NameToID = absl::flat_hash_map; using IDToName = std::vector; struct NameToIDMappings { - // M is either a raw::MjModel or MjDataMetadata. - template - explicit NameToIDMappings(const M& m); + explicit NameToIDMappings(const raw::MjModel& m); NameToID body; NameToID jnt; @@ -63,9 +60,7 @@ struct NameToIDMappings { }; struct IDToNameMappings { - // M is either a raw::MjModel or MjDataMetadata. - template - explicit IDToNameMappings(const M& m); + explicit IDToNameMappings(const raw::MjModel& m); IDToName body; IDToName jnt; @@ -163,7 +158,7 @@ class MjModelIndexer { class MjDataGroupedViewsBase { public: MjDataGroupedViewsBase(int index, std::string_view name, raw::MjData* d, - const MjDataMetadata* m, + const raw::MjModel* m, pybind11::handle owner) : index_(index), name_(name), d_(d), m_(m), owner_(owner) {} @@ -175,7 +170,7 @@ class MjDataGroupedViewsBase { int index_; std::string name_; raw::MjData* d_; - const MjDataMetadata* m_; + const raw::MjModel* m_; pybind11::handle owner_; }; @@ -201,7 +196,7 @@ MJDATA_VIEW_GROUPS // (e.g. a particular geom or joint) either by name or by ID. class MjDataIndexer { public: - MjDataIndexer(raw::MjData* d, const MjDataMetadata* m, + MjDataIndexer(raw::MjData* d, const raw::MjModel* m, pybind11::handle owner); #define XGROUP(MjDataGroupedViews, field, nfield, FIELD_XMACROS) \ @@ -213,7 +208,7 @@ class MjDataIndexer { private: raw::MjData* d_; - const MjDataMetadata* m_; + const raw::MjModel* m_; pybind11::handle owner_; NameToIDMappings name_to_id_; IDToNameMappings id_to_name_; diff --git a/python/mujoco/mjdata_meta.h b/python/mujoco/mjdata_meta.h deleted file mode 100644 index a9eb643c..00000000 --- a/python/mujoco/mjdata_meta.h +++ /dev/null @@ -1,108 +0,0 @@ -// Copyright 2021 DeepMind Technologies Limited -// -// Licensed under the Apache License, Version 2.0 (the "License"); -// you may not use this file except in compliance with the License. -// You may obtain a copy of the License at -// -// http://www.apache.org/licenses/LICENSE-2.0 -// -// Unless required by applicable law or agreed to in writing, software -// distributed under the License is distributed on an "AS IS" BASIS, -// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -// See the License for the specific language governing permissions and -// limitations under the License. - -#ifndef MUJOCO_PYTHON_MJDATA_META_H_ -#define MUJOCO_PYTHON_MJDATA_META_H_ - -#include -#include - -#include -#include -#include "raw.h" -#include "util/crossplatform.h" - -namespace mujoco::python { -namespace _impl { -template class MjWrapper; -} // namespace _impl - -#define MJDATA_METADATA \ - X( int, jnt_qposadr, njnt ) \ - X( int, jnt_dofadr, njnt ) \ - X( int, hfield_nrow, nhfield ) \ - X( int, hfield_ncol, nhfield ) \ - X( int, hfield_adr, nhfield ) \ - X( int, tex_height, ntex ) \ - X( int, tex_width, ntex ) \ - X( int, tex_adr, ntex ) \ - X( int, sensor_dim, nsensor ) \ - X( int, sensor_adr, nsensor ) \ - X( int, numeric_adr, nnumeric ) \ - X( int, numeric_size, nnumeric ) \ - X( int, tuple_adr, ntuple ) \ - X( int, tuple_size, ntuple ) \ - X( int, name_bodyadr, nbody ) \ - X( int, name_jntadr, njnt ) \ - X( int, name_geomadr, ngeom ) \ - X( int, name_siteadr, nsite ) \ - X( int, name_camadr, ncam ) \ - X( int, name_lightadr, nlight ) \ - X( int, name_meshadr, nmesh ) \ - X( int, name_skinadr, nskin ) \ - X( int, name_hfieldadr, nhfield ) \ - X( int, name_texadr, ntex ) \ - X( int, name_matadr, nmat ) \ - X( int, name_pairadr, npair ) \ - X( int, name_excludeadr, nexclude ) \ - X( int, name_eqadr, neq ) \ - X( int, name_tendonadr, ntendon ) \ - X( int, name_actuatoradr, nu ) \ - X( int, name_sensoradr, nsensor ) \ - X( int, name_numericadr, nnumeric ) \ - X( int, name_textadr, ntext ) \ - X( int, name_tupleadr, ntuple ) \ - X( int, name_keyadr, nkey ) \ - X( char, names, nnames ) - -// A subset of mjModel fields that are required to reconstruct an MjDataWrapper. -struct MjDataMetadata { - public: - friend class _impl::MjWrapper; - explicit MjDataMetadata(const raw::MjModel* m) - : -#define X(var) var(m->var), - MJMODEL_INTS -#undef X - -#define X(dtype, var, n) \ - var([](dtype* src, int len) { \ - dtype* dst = new dtype[len]; \ - std::memcpy(dst, src, len * sizeof(dtype)); \ - return dst; \ - }(m->var, m->n)), - - MJDATA_METADATA -#undef X - is_dual(mj_isDual(m)) { - } - -#define X(var) decltype(raw::MjModel::var) var; - MJMODEL_INTS -#undef X -#define X(type, var, n) std::shared_ptr var; - MJDATA_METADATA -#undef X - - bool is_dual; - - private: - MjDataMetadata() = default; - MjDataMetadata(const MjDataMetadata& other) = default; - MjDataMetadata(MjDataMetadata&& other) = default; -}; - -} // namespace mujoco::python - -#endif // MUJOCO_PYTHON_MJDATA_META_H_ diff --git a/python/mujoco/structs.cc b/python/mujoco/structs.cc index deb1ba88..1d7cc901 100644 --- a/python/mujoco/structs.cc +++ b/python/mujoco/structs.cc @@ -37,10 +37,11 @@ #include "errors.h" #include "function_traits.h" #include "indexers.h" -#include "mjdata_meta.h" #include "private.h" #include "raw.h" #include "serialization.h" +#include +#include #include #include #include @@ -453,7 +454,8 @@ void MjModelWrapper::Serialize(std::ostream& output) const { WriteBytes(output, buffer.data(), model_size); } -MjModelWrapper MjModelWrapper::Deserialize(std::istream& input) { +std::unique_ptr MjModelWrapper::Deserialize( + std::istream& input) { CheckInput(input, "mjModel"); char serializationVersion = ReadChar(input); @@ -479,7 +481,7 @@ MjModelWrapper MjModelWrapper::Deserialize(std::istream& input) { if (!model) { throw py::value_error("Invalid serialized mjModel."); } - return MjModelWrapper(model); + return std::unique_ptr(new MjModelWrapper(model)); } // ==================== MJCONTACT ============================================== @@ -554,13 +556,13 @@ MjDataWrapper* MjDataWrapper::FromRawPointer(raw::MjData* m) noexcept { } } -MjDataWrapper::MjWrapper(const MjModelWrapper& model) - : WrapperBase(InterceptMjErrors(mj_makeData)(model.get()), +MjDataWrapper::MjWrapper(MjModelWrapper* model) + : WrapperBase(InterceptMjErrors(mj_makeData)(model->get()), &MjDataCapsuleDestructor), #undef MJ_M -#define MJ_M(x) model.get()->x +#define MJ_M(x) model->get()->x #define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(model.get()->dim0, dim1), ptr_->var, owner_)), + var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)), MJDATA_POINTERS #undef MJ_M #define MJ_M(x) (x) @@ -571,8 +573,9 @@ MjDataWrapper::MjWrapper(const MjModelWrapper& model) #define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), MJDATA_VECTOR #undef X - metadata_(model.get()), - indexer_(ptr_, &metadata_, owner_) { + model_(model), + model_ref_(py::cast(model_)), + indexer_(ptr_, model_->get(), owner_) { bool is_newly_inserted = false; { py::gil_scoped_acquire gil; @@ -587,9 +590,9 @@ MjDataWrapper::MjWrapper(const MjModelWrapper& model) MjDataWrapper::MjWrapper(const MjDataWrapper& other) : WrapperBase(other.Copy(), &MjDataCapsuleDestructor), #undef MJ_M -#define MJ_M(x) other.metadata_.x +#define MJ_M(x) other.model_->get()->x #define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(other.metadata_.dim0, dim1), ptr_->var, \ + var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ owner_)), MJDATA_POINTERS #undef MJ_M @@ -601,8 +604,9 @@ MjDataWrapper::MjWrapper(const MjDataWrapper& other) #define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), MJDATA_VECTOR #undef X - metadata_(other.metadata_), - indexer_(ptr_, &metadata_, owner_) { + model_(other.model_), + model_ref_(other.model_ref_), + indexer_(ptr_, model_->get(), owner_) { bool is_newly_inserted = false; { py::gil_scoped_acquire gil; @@ -617,9 +621,9 @@ MjDataWrapper::MjWrapper(const MjDataWrapper& other) MjDataWrapper::MjWrapper(MjDataWrapper&& other) : WrapperBase(other.ptr_, other.owner_), #undef MJ_M -#define MJ_M(x) other.metadata_.x +#define MJ_M(x) other.model_->get()->x #define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(other.metadata_.dim0, dim1), ptr_->var, \ + var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ owner_)), MJDATA_POINTERS #undef MJ_M @@ -631,8 +635,9 @@ MjDataWrapper::MjWrapper(MjDataWrapper&& other) #define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), MJDATA_VECTOR #undef X - metadata_(other.metadata_), - indexer_(ptr_, &metadata_, owner_) { + model_(other.model_), + model_ref_(std::move(other.model_ref_)), + indexer_(ptr_, model_->get(), owner_) { bool is_newly_inserted = false; { py::gil_scoped_acquire gil; @@ -646,12 +651,13 @@ MjDataWrapper::MjWrapper(MjDataWrapper&& other) other.ptr_ = nullptr; } -MjDataWrapper::MjWrapper(MjDataMetadata&& metadata, raw::MjData* d) - : WrapperBase(d, &MjDataCapsuleDestructor), +MjDataWrapper::MjWrapper(const MjDataWrapper& other, MjModelWrapper* model) + : WrapperBase(other.Copy(), &MjDataCapsuleDestructor), #undef MJ_M -#define MJ_M(x) metadata.x -#define X(dtype, var, dim0, dim1) \ - var(InitPyArray(X_ARRAY_SHAPE(metadata.dim0, dim1), ptr_->var, owner_)), +#define MJ_M(x) other.model_->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(other.model_->get()->dim0, dim1), ptr_->var, \ + owner_)), MJDATA_POINTERS #undef MJ_M #define MJ_M(x) (x) @@ -662,8 +668,39 @@ MjDataWrapper::MjWrapper(MjDataMetadata&& metadata, raw::MjData* d) #define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), MJDATA_VECTOR #undef X - metadata_(std::move(metadata)), - indexer_(ptr_, &metadata_, owner_) { + model_(model), + model_ref_(py::cast(model_)), + indexer_(ptr_, model_->get(), owner_) { + bool is_newly_inserted = false; + { + py::gil_scoped_acquire gil; + is_newly_inserted = MjDataRawPointerMap().insert({ptr_, this}).second; + } + if (!is_newly_inserted) { + throw UnexpectedError( + "MjDataRawPointerMap already contains this raw mjData*"); + } +} + +MjDataWrapper::MjWrapper(MjModelWrapper* model, raw::MjData* d) + : WrapperBase(d, &MjDataCapsuleDestructor), +#undef MJ_M +#define MJ_M(x) model->get()->x +#define X(dtype, var, dim0, dim1) \ + var(InitPyArray(X_ARRAY_SHAPE(model->get()->dim0, dim1), ptr_->var, owner_)), + MJDATA_POINTERS +#undef MJ_M +#define MJ_M(x) (x) +#undef X + + contact(MjContactList(ptr_->contact, NConMax(ptr_), &ptr_->ncon, owner_)), + +#define X(dtype, var, dim0, dim1) var(InitPyArray(ptr_->var, owner_)), + MJDATA_VECTOR +#undef X + model_(model), + model_ref_(py::cast(model_)), + indexer_(ptr_, model_->get(), owner_) { bool is_newly_inserted = false; { py::gil_scoped_acquire gil; @@ -694,19 +731,7 @@ void MjDataWrapper::Serialize(std::ostream& output) const { // TODO: Replace this custom serialization with a protobuf WriteChar(output, kSerializationVersion); - // Write all size fields -#define X(var) WriteInt(output, this->metadata_.var); - MJMODEL_INTS -#undef X - - WriteInt(output, this->metadata_.is_dual); - -#define X(dtype, var, n) \ - WriteBytes(output, this->metadata_.var.get(), \ - this->metadata_.n * sizeof(dtype)); - - MJDATA_METADATA -#undef X + model_->Serialize(output); // Write struct and scalar fields #define X(var) WriteBytes(output, &ptr_->var, sizeof(ptr_->var)) @@ -729,15 +754,15 @@ void MjDataWrapper::Serialize(std::ostream& output) const { // Write buffer contents { - MJDATA_POINTERS_PREAMBLE((&this->metadata_)) + MJDATA_POINTERS_PREAMBLE((this->model_->get())) #define X(type, name, nr, nc) \ - WriteBytes(output, ptr_->name, sizeof(type)*(this->metadata_.nr)*(nc)); + WriteBytes(output, ptr_->name, sizeof(type)*(this->model_->get()->nr)*(nc)); MJDATA_POINTERS #undef X #undef MJ_M -#define MJ_M(x) this->metadata_.x +#define MJ_M(x) this->model_->get()->x #undef MJ_D #define MJ_D(x) this->ptr_->x #define X(type, name, nr, nc) \ @@ -747,7 +772,7 @@ void MjDataWrapper::Serialize(std::ostream& output) const { MJDATA_ARENA_POINTERS_CONTACT MJDATA_ARENA_POINTERS_PRIMAL - if (this->metadata_.is_dual) { + if (mj_isDual(this->model_->get())) { MJDATA_ARENA_POINTERS_DUAL } if (this->ptr_->nisland) { @@ -769,28 +794,12 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) { throw py::value_error("Incompatible serialization version."); } - // Read all size and address fields - MjDataMetadata metadata; - raw::MjModel m{0}; + // Read the model that was used to create the mjData. + std::unique_ptr m_wrapper = + MjModelWrapper::Deserialize(input); + raw::MjModel& m = *m_wrapper->get(); -#define X(var) \ - metadata.var = ReadInt(input); \ - CheckInput(input, "mjData"); \ - m.var = metadata.var; - - MJMODEL_INTS -#undef X - - metadata.is_dual = ReadInt(input); - -#define X(dtype, var, n) \ - metadata.var.reset(new dtype[metadata.n]); \ - ReadBytes(input, metadata.var.get(), metadata.n * sizeof(dtype)); \ - CheckInput(input, "mjData"); \ - m.var = metadata.var.get(); - - MJDATA_METADATA -#undef X + bool is_dual = mj_isDual(&m); raw::MjData* d = mj_makeData(&m); if (!d) { @@ -841,7 +850,7 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) { MJDATA_ARENA_POINTERS_CONTACT MJDATA_ARENA_POINTERS_PRIMAL - if (metadata.is_dual) { + if (is_dual) { MJDATA_ARENA_POINTERS_DUAL } if (d->nisland) { @@ -861,18 +870,12 @@ MjDataWrapper MjDataWrapper::Deserialize(std::istream& input) { throw py::value_error("Invalid serialized mjData."); } - return MjDataWrapper(std::move(metadata), d); + return MjDataWrapper(m_wrapper.release(), d); } raw::MjData* MjDataWrapper::Copy() const { - raw::MjModel m{0}; -#define X(var) m.var = this->metadata_.var; - MJMODEL_INTS -#undef X -#define X(dtype, var, n) m.var = this->metadata_.var.get(); - MJDATA_METADATA -#undef X - return InterceptMjErrors(mj_copyData)(NULL, &m, this->ptr_); + const raw::MjModel* m = model_->get(); + return InterceptMjErrors(mj_copyData)(NULL, m, this->ptr_); } // ==================== MJSTATISTIC ============================================ @@ -1890,15 +1893,19 @@ This is useful for example when the MJB is not available as a file on disk.)")); // ==================== MJDATA =============================================== py::class_ mjData(m, "MjData"); - mjData.def(py::init()); + mjData.def(py::init()); mjData.def_property_readonly("_address", [](const MjDataWrapper& d) { return reinterpret_cast(d.get()); }); + mjData.def_property_readonly("model", [](const MjDataWrapper& d) { + return &d.model(); + }); mjData.def("__copy__", [](const MjDataWrapper& other) { return MjDataWrapper(other); }); mjData.def("__deepcopy__", [](const MjDataWrapper& other, py::dict) { - return MjDataWrapper(other); + MjModelWrapper* model_copy = new MjModelWrapper(other.model()); + return MjDataWrapper(other, model_copy); }); mjData.def(py::pickle( [](const MjDataWrapper& d) { // __getstate__ @@ -1927,7 +1934,7 @@ This is useful for example when the MJB is not available as a file on disk.)")); #undef X #undef MJ_M -#define MJ_M(x) d.metadata().x +#define MJ_M(x) d.model().get()->x #undef MJ_D #define MJ_D(x) d.get()->x #define X(dtype, var, dim0, dim1) \ diff --git a/python/mujoco/structs.h b/python/mujoco/structs.h index 291dc96b..5fe3e1b3 100644 --- a/python/mujoco/structs.h +++ b/python/mujoco/structs.h @@ -29,10 +29,10 @@ #include #include #include "indexers.h" -#include "mjdata_meta.h" #include "raw.h" #include #include +#include namespace mujoco::python { namespace _impl { @@ -447,7 +447,8 @@ class MjWrapper : public WrapperBase { MjModelIndexer& indexer() { return indexer_; } void Serialize(std::ostream& output) const; - static MjWrapper Deserialize(std::istream& input); + static std::unique_ptr> Deserialize( + std::istream& input); static MjWrapper LoadXMLFile( const std::string& filename, @@ -564,12 +565,14 @@ struct is_mj_struct_list { template <> class MjWrapper: public WrapperBase { public: - explicit MjWrapper(const MjModelWrapper& model); + explicit MjWrapper(MjModelWrapper* model); MjWrapper(const MjWrapper& other); MjWrapper(MjWrapper&&); + // Used for deepcopy + MjWrapper(const MjWrapper& other, MjModelWrapper* model); ~MjWrapper(); - const MjDataMetadata& metadata() const { return metadata_; } + const MjModelWrapper& model() const { return *model_; } MjDataIndexer& indexer() { return indexer_; } void Serialize(std::ostream& output) const; @@ -598,10 +601,14 @@ class MjWrapper: public WrapperBase { protected: // Internal constructor which takes ownership of given mjData pointer. // Used for deserialization. - explicit MjWrapper(MjDataMetadata&& metadata, raw::MjData* d); + explicit MjWrapper(MjModelWrapper* model, raw::MjData* d); raw::MjData* Copy() const; - MjDataMetadata metadata_; + // A reference to the model that was used to create this mjData. + MjModelWrapper* model_; + // A py::object pointing to the same model as model_, to make sure Python + // doesn't doesn't garbage collect it until this mjData is released. + pybind11::object model_ref_; MjDataIndexer indexer_; };