Make "id" a property rather than a method of mjsElements in Python bindings.

PiperOrigin-RevId: 649424035
Change-Id: I8226b3040aa2001ae9748399257e36ac263b7c8b
This commit is contained in:
Yuval Tassa
2024-07-04 08:34:35 -07:00
committed by Copybara-Service
parent f8b944bf7f
commit 118b6810b3
2 changed files with 90 additions and 85 deletions
+84 -81
View File
@@ -13,21 +13,21 @@
// limitations under the License.
#include <array>
#include <cstddef>
#include <cstddef> // IWYU pragma: keep
#include <cstdint>
#include <memory>
#include <string>
#include <string_view>
#include <vector>
#include <string_view> // IWYU pragma: keep
#include <vector> // IWYU pragma: keep
#include <Eigen/Core>
#include <Eigen/Eigen>
#include <mujoco/mjspec.h>
#include <mujoco/mjspec.h> // IWYU pragma: keep
#include <mujoco/mujoco.h>
#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 <pybind11/cast.h>
#include <pybind11/eigen.h>
#include <pybind11/eigen/matrix.h>
@@ -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
+6 -4
View File
@@ -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])