From 118b6810b368cd3ad0935e41438f826928a073b0 Mon Sep 17 00:00:00 2001 From: Yuval Tassa Date: Thu, 4 Jul 2024 08:34:35 -0700 Subject: [PATCH] Make "id" a property rather than a method of mjsElements in Python bindings. PiperOrigin-RevId: 649424035 Change-Id: I8226b3040aa2001ae9748399257e36ac263b7c8b --- python/mujoco/specs.cc | 165 ++++++++++++++++++------------------ python/mujoco/specs_test.py | 10 ++- 2 files changed, 90 insertions(+), 85 deletions(-) diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 227f8630..783a87b2 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -13,21 +13,21 @@ // limitations under the License. #include -#include +#include // IWYU pragma: keep #include #include #include -#include -#include +#include // IWYU pragma: keep +#include // IWYU pragma: keep #include #include -#include +#include // IWYU pragma: keep #include #include "errors.h" -#include "indexers.h" +#include "indexers.h" // IWYU pragma: keep #include "raw.h" -#include "structs.h" +#include "structs.h" // IWYU pragma: keep #include #include #include @@ -490,10 +490,10 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSBODY ==================================== - mjsBody.def("id", [](raw::MjsBody& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSBODY ===================================== + mjsBody.def_property_readonly( + "id", [](raw::MjsBody& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsBody.def("delete", [](raw::MjsBody& self) { mjs_delete(self.element); }); mjsBody.def( "add_body", @@ -665,9 +665,9 @@ PYBIND11_MODULE(_specs, m) { }); // ============================= MJSFRAME ==================================== - mjsFrame.def("id", [](raw::MjsFrame& self) -> int { - return mjs_getId(self.element); - }); + mjsFrame.def_property_readonly( + "id", [](raw::MjsFrame& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsFrame.def("delete", [](raw::MjsFrame& self) { mjs_delete(self.element); }); mjsFrame.def("set_frame", [](raw::MjsFrame& self, raw::MjsFrame& frame) { mjs_setFrame(self.element, &frame); @@ -677,10 +677,10 @@ PYBIND11_MODULE(_specs, m) { mjs_attachBody(&self, &body, prefix.c_str(), suffix.c_str()); }); - // ============================= MJSGEOM ==================================== - mjsGeom.def("id", [](raw::MjsGeom& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSGEOM ===================================== + mjsGeom.def_property_readonly( + "id", [](raw::MjsGeom& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsGeom.def("delete", [](raw::MjsGeom& self) { mjs_delete(self.element); }); mjsGeom.def("set_frame", [](raw::MjsGeom& self, raw::MjsFrame& frame) { mjs_setFrame(self.element, &frame); @@ -696,9 +696,9 @@ PYBIND11_MODULE(_specs, m) { py::return_value_policy::reference_internal); // ============================= MJSJOINT ==================================== - mjsJoint.def("id", [](raw::MjsJoint& self) -> int { - return mjs_getId(self.element); - }); + mjsJoint.def_property_readonly( + "id", [](raw::MjsJoint& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsJoint.def("delete", [](raw::MjsJoint& self) { mjs_delete(self.element); }); mjsJoint.def("set_frame", [](raw::MjsJoint& self, raw::MjsFrame& frame) { mjs_setFrame(self.element, &frame); @@ -713,10 +713,10 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSSITE ==================================== - mjsSite.def("id", [](raw::MjsSite& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSSITE ===================================== + mjsSite.def_property_readonly( + "id", [](raw::MjsSite& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsSite.def("delete", [](raw::MjsSite& self) { mjs_delete(self.element); }); mjsSite.def("set_frame", [](raw::MjsSite& self, raw::MjsFrame& frame) { mjs_setFrame(self.element, &frame); @@ -731,10 +731,10 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSCAMERA ================================== - mjsCamera.def("id", [](raw::MjsCamera& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSCAMERA =================================== + mjsCamera.def_property_readonly( + "id", [](raw::MjsCamera& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsCamera.def("delete", [](raw::MjsCamera& self) { mjs_delete(self.element); }); mjsCamera.def("set_frame", [](raw::MjsCamera& self, raw::MjsFrame& frame) { @@ -751,9 +751,9 @@ PYBIND11_MODULE(_specs, m) { py::return_value_policy::reference_internal); // ============================= MJSLIGHT ==================================== - mjsLight.def("id", [](raw::MjsLight& self) -> int { - return mjs_getId(self.element); - }); + mjsLight.def_property_readonly( + "id", [](raw::MjsLight& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsLight.def("delete", [](raw::MjsLight& self) { mjs_delete(self.element); }); mjsLight.def("set_frame", [](raw::MjsLight& self, raw::MjsFrame& frame) { mjs_setFrame(self.element, &frame); @@ -768,10 +768,11 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSMATERIAL ================================ - mjsMaterial.def("id", [](raw::MjsMaterial& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSMATERIAL ================================= + mjsMaterial.def_property_readonly( + "id", + [](raw::MjsMaterial& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsMaterial.def("delete", [](raw::MjsMaterial& self) { mjs_delete(self.element); }); mjsMaterial.def("set_default", @@ -785,10 +786,10 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSMESH ==================================== - mjsMesh.def("id", [](raw::MjsMesh& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSMESH ===================================== + mjsMesh.def_property_readonly( + "id", [](raw::MjsMesh& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsMesh.def("delete", [](raw::MjsMesh& self) { mjs_delete(self.element); }); mjsMesh.def("set_default", [](raw::MjsMesh& self, raw::MjsDefault& def) { mjs_setDefault(self.element, &def); @@ -800,10 +801,10 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSPAIR ==================================== - mjsPair.def("id", [](raw::MjsPair& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSPAIR ===================================== + mjsPair.def_property_readonly( + "id", [](raw::MjsPair& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsPair.def("delete", [](raw::MjsPair& self) { mjs_delete(self.element); }); mjsPair.def("set_default", [](raw::MjsPair& self, raw::MjsDefault& def) { mjs_setDefault(self.element, &def); @@ -816,9 +817,10 @@ PYBIND11_MODULE(_specs, m) { py::return_value_policy::reference_internal); // ============================= MJSEQUAL ==================================== - mjsEquality.def("id", [](raw::MjsEquality& self) -> int { - return mjs_getId(self.element); - }); + mjsEquality.def_property_readonly( + "id", + [](raw::MjsEquality& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsEquality.def("delete", [](raw::MjsEquality& self) { mjs_delete(self.element); }); mjsEquality.def("set_default", @@ -832,10 +834,11 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSACTUATOR ================================ - mjsActuator.def("id", [](raw::MjsActuator& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSACTUATOR ================================= + mjsActuator.def_property_readonly( + "id", + [](raw::MjsActuator& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsActuator.def("delete", [](raw::MjsActuator& self) { mjs_delete(self.element); }); mjsActuator.def("set_default", @@ -849,10 +852,10 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSTENDON ================================== - mjsTendon.def("id", [](raw::MjsTendon& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSTENDON =================================== + mjsTendon.def_property_readonly( + "id", [](raw::MjsTendon& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsTendon.def("delete", [](raw::MjsTendon& self) { mjs_delete(self.element); }); mjsTendon.def("set_default", [](raw::MjsTendon& self, raw::MjsDefault& def) { @@ -889,71 +892,72 @@ PYBIND11_MODULE(_specs, m) { }, py::return_value_policy::reference_internal); - // ============================= MJSSENSOR ================================== - mjsSensor.def("id", [](raw::MjsSensor& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSSENSOR =================================== + mjsSensor.def_property_readonly( + "id", [](raw::MjsSensor& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsSensor.def("delete", [](raw::MjsSensor& self) { mjs_delete(self.element); }); - // ============================= MJSFLEX ==================================== - mjsFlex.def("id", [](raw::MjsFlex& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSFLEX ===================================== + mjsFlex.def_property_readonly( + "id", [](raw::MjsFlex& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsFlex.def("delete", [](raw::MjsFlex& self) { mjs_delete(self.element); }); - // ============================= MJSHFIELD ================================== - mjsHField.def("id", [](raw::MjsHField& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSHFIELD =================================== + mjsHField.def_property_readonly( + "id", [](raw::MjsHField& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsHField.def("delete", [](raw::MjsHField& self) { mjs_delete(self.element); }); - // ============================= MJSSKIN ==================================== - mjsSkin.def("id", [](raw::MjsSkin& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSSKIN ===================================== + mjsSkin.def_property_readonly( + "id", [](raw::MjsSkin& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsSkin.def("delete", [](raw::MjsSkin& self) { mjs_delete(self.element); }); - // ============================= MJSTEXTURE ================================= - mjsTexture.def("id", [](raw::MjsTexture& self) -> int { - return mjs_getId(self.element); - }); + // ============================= MJSTEXTURE ================================== + mjsTexture.def_property_readonly( + "id", + [](raw::MjsTexture& self) -> int { return mjs_getId(self.element); }, + py::return_value_policy::reference_internal); mjsTexture.def("delete", [](raw::MjsTexture& self) { mjs_delete(self.element); }); - // ============================= MJSKEY ===================================== + // ============================= MJSKEY ====================================== mjsKey.def("id", [](raw::MjsKey& self) -> int { return mjs_getId(self.element); }); mjsKey.def("delete", [](raw::MjsKey& self) { mjs_delete(self.element); }); - // ============================= MJSTEXT ==================================== + // ============================= MJSTEXT ===================================== mjsText.def("id", [](raw::MjsText& self) -> int { return mjs_getId(self.element); }); mjsText.def("delete", [](raw::MjsText& self) { mjs_delete(self.element); }); - // ============================= MJSNUMERIC ================================= + // ============================= MJSNUMERIC ================================== mjsNumeric.def("id", [](raw::MjsNumeric& self) -> int { return mjs_getId(self.element); }); mjsNumeric.def("delete", [](raw::MjsNumeric& self) { mjs_delete(self.element); }); - // ============================= MJSEXCLUDE ================================ + // ============================= MJSEXCLUDE ================================== mjsExclude.def("id", [](raw::MjsExclude& self) -> int { return mjs_getId(self.element); }); mjsExclude.def("delete", [](raw::MjsExclude& self) { mjs_delete(self.element); }); - // ============================= MJSTUPLE =================================== + // ============================= MJSTUPLE ==================================== mjsTuple.def("id", [](raw::MjsTuple& self) -> int { return mjs_getId(self.element); }); mjsTuple.def("delete", [](raw::MjsTuple& self) { mjs_delete(self.element); }); - // ============================= MJSPLUGIN ================================== + // ============================= MJSPLUGIN =================================== mjsPlugin.def("id", [](raw::MjsPlugin& self) -> int { return mjs_getId(self.instance); }); @@ -961,6 +965,5 @@ PYBIND11_MODULE(_specs, m) { [](raw::MjsPlugin& self) { mjs_delete(self.instance); }); #include "specs.cc.inc" - } // PYBIND11_MODULE // NOLINT } // namespace mujoco::python diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index b8e2f2ac..9c36d728 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -66,13 +66,15 @@ class SpecsTest(absltest.TestCase): self.assertEqual(site.name, 'sitename') np.testing.assert_array_equal(site.userdata, [1, 2, 3, 4, 5, 6]) - # Check that the site has no id before compilation. - self.assertEqual(body.id(), -1) + # Check that the site and body have no id before compilation. + self.assertEqual(body.id, -1) + self.assertEqual(site.id, -1) # Compile the spec and check for expected values in the model. model = spec.compile() - self.assertEqual(spec.worldbody.id(), 0) - self.assertEqual(body.id(), 1) + self.assertEqual(spec.worldbody.id, 0) + self.assertEqual(body.id, 1) + self.assertEqual(site.id, 0) self.assertEqual(model.nbody, 2) # 2 bodies, including the world body np.testing.assert_array_equal(model.body_pos[1], [1, 2, 3]) np.testing.assert_array_equal(model.body_quat[1], [0, 1, 0, 0])