Add find_all with string input.

PiperOrigin-RevId: 689487165
Change-Id: I48511aec1610b2112f9456e8f410d077c57c7a95
This commit is contained in:
Alessio Quaglino
2024-10-24 13:00:38 -07:00
committed by Copybara-Service
parent a3ea01e57e
commit bebec52869
2 changed files with 82 additions and 52 deletions
+66 -38
View File
@@ -155,6 +155,48 @@ void DefineArray(py::module& m, const std::string& typestr) {
}, py::keep_alive<0, 1>(), py::return_value_policy::reference_internal);
};
py::list FindAllImpl(raw::MjsBody& body, mjtObj objtype) {
py::list list;
raw::MjsElement* el = mjs_firstChild(&body, objtype, true);
std::string error = mjs_getError(mjs_getSpec(body.element));
if (!el && !error.empty()) {
throw pybind11::value_error(error);
}
while (el) {
switch (objtype) {
case mjOBJ_BODY:
list.append(mjs_asBody(el));
break;
case mjOBJ_CAMERA:
list.append(mjs_asCamera(el));
break;
case mjOBJ_FRAME:
list.append(mjs_asFrame(el));
break;
case mjOBJ_GEOM:
list.append(mjs_asGeom(el));
break;
case mjOBJ_JOINT:
list.append(mjs_asJoint(el));
break;
case mjOBJ_LIGHT:
list.append(mjs_asLight(el));
break;
case mjOBJ_SITE:
list.append(mjs_asSite(el));
break;
default:
// this should never happen
throw pybind11::value_error(
"body.find_all supports the types: body, frame, geom, site, "
"light, camera.");
break;
}
el = mjs_nextChild(&body, el, true);
}
return list; // list of pointers, so they can be copied
}
PYBIND11_MODULE(_specs, m) {
auto structs_m = py::module::import("mujoco._structs");
py::function mjmodel_from_spec_ptr =
@@ -419,45 +461,31 @@ PYBIND11_MODULE(_specs, m) {
mjsBody.def(
"find_all",
[](raw::MjsBody& self, mjtObj objtype) -> py::list {
py::list list;
raw::MjsElement* el = mjs_firstChild(&self, objtype, true);
std::string error = mjs_getError(mjs_getSpec(self.element));
if (!el && !error.empty()) {
throw pybind11::value_error(error);
return FindAllImpl(self, objtype);
},
py::return_value_policy::reference_internal);
mjsBody.def(
"find_all",
[](raw::MjsBody& self, std::string& name) -> py::list {
mjtObj objtype = mjOBJ_UNKNOWN;
if (name == "body") {
objtype = mjOBJ_BODY;
} else if (name == "frame") {
objtype = mjOBJ_FRAME;
} else if (name == "geom") {
objtype = mjOBJ_GEOM;
} else if (name == "site") {
objtype = mjOBJ_SITE;
} else if (name == "light") {
objtype = mjOBJ_LIGHT;
} else if (name == "camera") {
objtype = mjOBJ_CAMERA;
} else {
throw pybind11::value_error(
"body.find_all supports the types: body, frame, geom, site, "
"light, camera.");
}
while (el) {
switch (objtype) {
case mjOBJ_BODY:
list.append(mjs_asBody(el));
break;
case mjOBJ_CAMERA:
list.append(mjs_asCamera(el));
break;
case mjOBJ_FRAME:
list.append(mjs_asFrame(el));
break;
case mjOBJ_GEOM:
list.append(mjs_asGeom(el));
break;
case mjOBJ_JOINT:
list.append(mjs_asJoint(el));
break;
case mjOBJ_LIGHT:
list.append(mjs_asLight(el));
break;
case mjOBJ_SITE:
list.append(mjs_asSite(el));
break;
default:
// this should never happen
throw pybind11::value_error(
"body.find_all supports the types: body, frame, geom, site, "
"light, camera.");
break;
}
el = mjs_nextChild(&self, el, true);
}
return list;
return FindAllImpl(self, objtype);
},
py::return_value_policy::reference_internal);
mjsBody.def(
+16 -14
View File
@@ -611,7 +611,6 @@ class SpecsTest(absltest.TestCase):
"""
spec = mujoco.MjSpec.from_string(main_xml)
bodytype = mujoco.mjtObj.mjOBJ_BODY
sitetype = mujoco.mjtObj.mjOBJ_SITE
self.assertLen(spec.bodies, 5)
self.assertEqual(spec.bodies[1].name, 'body1')
self.assertEqual(spec.bodies[2].name, 'body2')
@@ -620,23 +619,26 @@ class SpecsTest(absltest.TestCase):
self.assertLen(spec.worldbody.find_all(bodytype), 4)
self.assertLen(spec.bodies[1].find_all(bodytype), 2)
self.assertLen(spec.bodies[3].find_all(bodytype), 1)
self.assertEqual(spec.worldbody.find_all(bodytype)[0].name, 'body1')
self.assertEqual(spec.worldbody.find_all(bodytype)[1].name, 'body2')
self.assertEqual(spec.worldbody.find_all(bodytype)[2].name, 'body3')
self.assertEqual(spec.worldbody.find_all(bodytype)[3].name, 'body4')
self.assertEqual(spec.bodies[1].find_all(bodytype)[0].name, 'body3')
self.assertEqual(spec.bodies[1].find_all(bodytype)[1].name, 'body4')
self.assertEqual(spec.bodies[3].find_all(bodytype)[0].name, 'body4')
self.assertEmpty(spec.bodies[2].find_all(bodytype))
self.assertEmpty(spec.bodies[4].find_all(bodytype))
self.assertEqual(spec.worldbody.find_all(sitetype)[0].name, 'site')
self.assertEqual(spec.worldbody.find_all('body')[0].name, 'body1')
self.assertEqual(spec.worldbody.find_all('body')[1].name, 'body2')
self.assertEqual(spec.worldbody.find_all('body')[2].name, 'body3')
self.assertEqual(spec.worldbody.find_all('body')[3].name, 'body4')
self.assertEqual(spec.bodies[1].find_all('body')[0].name, 'body3')
self.assertEqual(spec.bodies[1].find_all('body')[1].name, 'body4')
self.assertEqual(spec.bodies[3].find_all('body')[0].name, 'body4')
self.assertEmpty(spec.bodies[2].find_all('body'))
self.assertEmpty(spec.bodies[4].find_all('body'))
self.assertEqual(spec.worldbody.find_all('site')[0].name, 'site')
with self.assertRaises(ValueError) as cm:
spec.worldbody.find_all(mujoco.mjtObj.mjOBJ_ACTUATOR)
spec.worldbody.find_all('actuator')
self.assertEqual(
str(cm.exception),
'Error: Body.NextChild supports the types: body, frame, geom, site,'
' light, camera\nElement name \'world\', id 0',
'body.find_all supports the types: body, frame, geom, site,'
' light, camera.',
)
body4 = spec.worldbody.find_all('body')[3]
body4.name = 'body4_new'
self.assertEqual(spec.bodies[4].name, 'body4_new')
def test_iterators(self):
spec = mujoco.MjSpec()