Add bind() method to MjModel and MjData.

This method allows users to bind a spec to a model or data object, making it less verbose to access its arrays.

PiperOrigin-RevId: 699952211
Change-Id: I6b71a88d5a6968bc0f902bd49db7c568b46a4dfb
This commit is contained in:
Alessio Quaglino
2024-11-25 05:53:03 -08:00
committed by Copybara-Service
parent 13b6055098
commit c6ba8f27fd
6 changed files with 75 additions and 17 deletions
+3 -3
View File
@@ -457,9 +457,9 @@
"pos_x = []\n",
"pos_y = []\n",
"geoms = arena.worldbody.find_all(mj.mjtObj.mjOBJ_GEOM)\n",
"torsos_data = [data.geom(geom.name) for geom in geoms if 'torso' in geom.name]\n",
"torsos_model = [model.geom(geom.name) for geom in geoms if 'torso' in geom.name]\n",
"actuators = [data.actuator(actuator.name) for actuator in arena.actuators]\n",
"torsos_data = [data.bind(geom) for geom in geoms if 'torso' in geom.name]\n",
"torsos_model = [model.bind(geom) for geom in geoms if 'torso' in geom.name]\n",
"actuators = [data.bind(actuator) for actuator in arena.actuators]\n",
"\n",
"# Control signal frequency, phase, amplitude.\n",
"freq = 5\n",
+33
View File
@@ -307,6 +307,28 @@
XGROUP( MjModelTupleViews, tuple, ntuple, MJMODEL_TUPLE ) \
XGROUP( MjModelKeyframeViews, key, nkey, MJMODEL_KEYFRAME )
#define MJMODEL_BIND_GROUPS \
XGROUP( mjsActuator, actuator) \
XGROUP( mjsBody, body ) \
XGROUP( mjsCamera, cam ) \
XGROUP( mjsEquality, eq ) \
XGROUP( mjsExclude, exclude ) \
XGROUP( mjsGeom, geom ) \
XGROUP( mjsHField, hfield ) \
XGROUP( mjsJoint, jnt ) \
XGROUP( mjsLight, light ) \
XGROUP( mjsMaterial, mat ) \
XGROUP( mjsMesh, mesh ) \
XGROUP( mjsNumeric, numeric ) \
XGROUP( mjsPair, pair ) \
XGROUP( mjsSensor, sensor ) \
XGROUP( mjsSite, site ) \
XGROUP( mjsSkin, skin ) \
XGROUP( mjsTendon, tendon ) \
XGROUP( mjsTexture, tex ) \
XGROUP( mjsTuple, tuple ) \
XGROUP( mjsKey, key )
#define MJMODEL_VIEW_GROUPS_ALTNAMES \
XGROUP( cam, camera, MJMODEL_CAMERA ) \
XGROUP( eq, equality, MJMODEL_EQUALITY ) \
@@ -399,6 +421,17 @@
XGROUP( MjDataSiteViews, site, nsite, MJDATA_SITE ) \
XGROUP( MjDataTendonViews, tendon, ntendon, MJDATA_TENDON )
#define MJDATA_BIND_GROUPS \
XGROUP( mjsActuator, actuator) \
XGROUP( mjsBody, body ) \
XGROUP( mjsCamera, cam ) \
XGROUP( mjsGeom, geom ) \
XGROUP( mjsJoint, jnt ) \
XGROUP( mjsLight, light ) \
XGROUP( mjsSensor, sensor ) \
XGROUP( mjsSite, site ) \
XGROUP( mjsTendon, tendon )
#define MJDATA_VIEW_GROUPS_ALTNAMES \
XGROUP( cam, camera, MJDATA_CAMERA ) \
XGROUP( jnt, joint, MJDATA_JOINT ) \
+5 -2
View File
@@ -94,9 +94,12 @@ class SpecsTest(absltest.TestCase):
# Compile the spec and check for expected values in the model.
model = spec.compile()
data = mujoco.MjData(model)
mujoco.mj_forward(model, data)
self.assertEqual(model.nbody, 2) # 2 bodies, including the world body
np.testing.assert_array_equal(model.body('baz').pos, [1, 2, 3])
np.testing.assert_array_equal(model.body('baz').quat, [0, 1, 0, 0])
np.testing.assert_array_equal(model.bind(body).pos, [1, 2, 3])
np.testing.assert_array_equal(model.bind(body).quat, [0, 1, 0, 0])
np.testing.assert_array_equal(data.bind(body).xpos, [1, 2, 3])
self.assertEqual(model.nsite, 1)
self.assertEqual(model.nuser_site, 6)
np.testing.assert_array_equal(model.site_user[0], [1, 2, 3, 4, 5, 6])
+24
View File
@@ -1716,6 +1716,18 @@ This is useful for example when the MJB is not available as a file on disk.)"));
MJMODEL_VIEW_GROUPS
#undef XGROUP
#define XGROUP(spectype, field) \
mjModel.def( \
"bind", \
[](MjModelWrapper& m, spectype& spec) -> auto& { \
return m.indexer().field##_by_name(mjs_getString(spec.name)); \
}, \
py::return_value_policy::reference_internal, \
py::arg_v("spec", py::none()));
MJMODEL_BIND_GROUPS
#undef XGROUP
#define XGROUP(field, altname, FIELD_XMACROS) \
mjModel.def( \
#altname, \
@@ -2050,6 +2062,18 @@ This is useful for example when the MJB is not available as a file on disk.)"));
MJDATA_VIEW_GROUPS
#undef XGROUP
#define XGROUP(spectype, field) \
mjData.def( \
"bind", \
[](MjDataWrapper& d, spectype& spec) -> auto& { \
return d.indexer().field##_by_name(mjs_getString(spec.name)); \
}, \
py::return_value_policy::reference_internal, \
py::arg_v("spec", py::none()));
MJDATA_BIND_GROUPS
#undef XGROUP
#define XGROUP(field, altname, FIELD_XMACROS) \
mjData.def( \
#altname, \