diff --git a/python/mujoco/codegen/generate_spec_bindings.py b/python/mujoco/codegen/generate_spec_bindings.py index 97854082..06d217a5 100644 --- a/python/mujoco/codegen/generate_spec_bindings.py +++ b/python/mujoco/codegen/generate_spec_bindings.py @@ -232,15 +232,33 @@ def generate() -> None: print(code) -def generate_body_add() -> None: - """Generate add functions for bodies.""" - for key in [ - 'mjsSite', - 'mjsGeom', - 'mjsJoint', - 'mjsLight', - 'mjsCamera', - 'mjsBody', +def generate_add() -> None: + """Generate add constructors with optional keyword arguments.""" + for key, parent, default in [ + ('mjsSite', 'Body', True), + ('mjsGeom', 'Body', True), + ('mjsJoint', 'Body', True), + ('mjsLight', 'Body', True), + ('mjsCamera', 'Body', True), + ('mjsBody', 'Body', True), + ('mjsFrame', 'Body', True), + ('mjsMaterial', 'Spec', True), + ('mjsMesh', 'Spec', True), + ('mjsPair', 'Spec', True), + ('mjsEquality', 'Spec', True), + ('mjsTendon', 'Spec', True), + ('mjsActuator', 'Spec', True), + ('mjsSkin', 'Spec', False), + ('mjsTexture', 'Spec', False), + ('mjsText', 'Spec', False), + ('mjsTuple', 'Spec', False), + ('mjsFlex', 'Spec', False), + ('mjsHField', 'Spec', False), + ('mjsKey', 'Spec', False), + ('mjsNumeric', 'Spec', False), + ('mjsExclude', 'Spec', False), + ('mjsSensor', 'Spec', False), + ('mjsPlugin', 'Spec', False), ]: def _field(f: ast_nodes.StructFieldDecl): @@ -262,11 +280,9 @@ def generate_body_add() -> None: return f'set_vec("{f.name}", out->{f.name});', 'vec', f.name elif isinstance(f.type, ast_nodes.ArrayType): return ( - ( - f'set_array("{f.name}", out->{f.name},' - f' {f.type.extents[0]});' - ), - 'array', f.name + f'set_array("{f.name}", out->{f.name}, {f.type.extents[0]});', + 'array', + f.name, ) elif isinstance(f.type, ast_nodes.ValueType): return f'set_value("{f.name}", out->{f.name});', 'value', f.name @@ -289,10 +305,33 @@ def generate_body_add() -> None: titlecase = 'Mjs' + elem # function definition and call to mjs_add_ - code = f""" - mjsBody.def("add_{elemlower}", [](raw::MjsBody& self, raw::MjsDefault* default_, py::kwargs kwargs) -> raw::{titlecase}* {{ - auto out = mjs_add{elem}(&self, default_); - """ + if parent == 'Spec': + if default: + code = f""" + {'mj' + parent}.def("add_{elemlower}", []({'Mj' + parent}& self, + raw::MjsDefault* default_, py::kwargs kwargs) -> raw::{titlecase}* {{ + auto out = mjs_add{elem}(self.ptr, default_); + """ + else: + code = f""" + {'mj' + parent}.def("add_{elemlower}", []({'Mj' + parent}& self, py::kwargs kwargs) -> raw::{titlecase}* {{ + auto out = mjs_add{elem}(self.ptr); + """ + elif parent == 'Body': + if key == 'mjsFrame': + code = f""" + {'mjs' + parent}.def("add_{elemlower}", []({'raw::Mjs' + parent}& self, + raw::MjsFrame* parentframe_, py::kwargs kwargs) -> raw::{titlecase}* {{ + auto out = mjs_add{elem}(&self, parentframe_); + """ + else: + code = f""" + {'mjs' + parent}.def("add_{elemlower}", []({'raw::Mjs' + parent}& self, + raw::MjsDefault* default_, py::kwargs kwargs) -> raw::{titlecase}* {{ + auto out = mjs_add{elem}(&self, default_); + """ + else: + raise NotImplementedError(f'{parent} parent is not implement.') # check for valid kwargs code += '\n std::set valid_kwargs = {' @@ -381,10 +420,10 @@ def generate_body_add() -> None: """ code += code_field - code += """\n + code += f"""\n return out; - }, - py::arg_v("default", nullptr), + }}, + {'py::arg_v("default", nullptr),' if default else ''} py::return_value_policy::reference_internal); """ @@ -395,7 +434,7 @@ def main(argv: Sequence[str]) -> None: if len(argv) > 1: raise app.UsageError('Too many command-line arguments.') generate() - generate_body_add() + generate_add() if __name__ == '__main__': diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index 4fd9e97f..db35d826 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -271,92 +271,6 @@ PYBIND11_MODULE(_specs, m) { return mjs_getSpecDefault(self.ptr); }, py::return_value_policy::reference_internal); - mjSpec.def( - "add_material", - [](MjSpec& self, raw::MjsDefault* default_) -> raw::MjsMaterial* { - return mjs_addMaterial(self.ptr, default_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); - mjSpec.def( - "add_mesh", - [](MjSpec& self, raw::MjsDefault* default_) -> raw::MjsMesh* { - return mjs_addMesh(self.ptr, default_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); - mjSpec.def( - "add_skin", - [](MjSpec& self) -> raw::MjsSkin* { return mjs_addSkin(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_texture", - [](MjSpec& self) -> raw::MjsTexture* { return mjs_addTexture(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_text", - [](MjSpec& self) -> raw::MjsText* { return mjs_addText(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_tuple", - [](MjSpec& self) -> raw::MjsTuple* { return mjs_addTuple(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_flex", - [](MjSpec& self) -> raw::MjsFlex* { return mjs_addFlex(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_hfield", - [](MjSpec& self) -> raw::MjsHField* { return mjs_addHField(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_key", - [](MjSpec& self) -> raw::MjsKey* { return mjs_addKey(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_numeric", - [](MjSpec& self) -> raw::MjsNumeric* { return mjs_addNumeric(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_pair", - [](MjSpec& self, raw::MjsDefault* default_) -> raw::MjsPair* { - return mjs_addPair(self.ptr, default_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); - mjSpec.def( - "add_exclude", - [](MjSpec& self) -> raw::MjsExclude* { return mjs_addExclude(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_equality", - [](MjSpec& self, raw::MjsDefault* default_) -> raw::MjsEquality* { - return mjs_addEquality(self.ptr, default_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); - mjSpec.def( - "add_tendon", - [](MjSpec& self, raw::MjsDefault* default_) -> raw::MjsTendon* { - return mjs_addTendon(self.ptr, default_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); - mjSpec.def( - "add_sensor", - [](MjSpec& self) -> raw::MjsSensor* { return mjs_addSensor(self.ptr); }, - py::return_value_policy::reference_internal); - mjSpec.def( - "add_actuator", - [](MjSpec& self, raw::MjsDefault* default_) -> raw::MjsActuator* { - return mjs_addActuator(self.ptr, default_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); - mjSpec.def( - "add_plugin", - [](MjSpec& self) -> raw::MjsPlugin* { return mjs_addPlugin(self.ptr); }, - py::return_value_policy::reference_internal); mjSpec.def("detach_body", [](MjSpec& self, raw::MjsBody& body) { mjs_detachBody(self.ptr, &body); }); @@ -556,17 +470,38 @@ PYBIND11_MODULE(_specs, m) { // ============================= MJSBODY ===================================== mjsBody.def_property_readonly( "id", [](raw::MjsBody& self) -> int { return mjs_getId(self.element); }); - mjsBody.def( - "add_frame", - [](raw::MjsBody& self, raw::MjsFrame* parentframe_) -> raw::MjsFrame* { - return mjs_addFrame(&self, parentframe_); - }, - py::arg_v("default", nullptr), - py::return_value_policy::reference_internal); mjsBody.def( "add_freejoint", - [](raw::MjsBody& self) -> raw::MjsJoint* { - return mjs_addFreeJoint(&self); + [](raw::MjsBody& self, py::kwargs kwargs) -> raw::MjsJoint* { + auto out = mjs_addFreeJoint(&self); + py::dict kwarg_dict = kwargs; + for (auto item : kwarg_dict) { + std::string key = py::str(item.first); + if (key == "align") { + try { + out->align = kwargs["align"].cast(); + } catch (const py::cast_error& e) { + throw pybind11::value_error("align is the wrong type."); + } + } else if (key == "name") { + try { + *out->name = kwargs["name"].cast(); + } catch (const py::cast_error& e) { + throw pybind11::value_error("name is the wrong type."); + } + } else if (key == "group") { + try { + out->group = kwargs["group"].cast(); + } catch (const py::cast_error& e) { + throw pybind11::value_error("group is the wrong type."); + } + } else { + throw pybind11::type_error( + "Invalid '" + key + + "' keyword argument. Valid options are: align, group, name."); + } + } + return out; }, py::return_value_policy::reference_internal); mjsBody.def("set_frame", diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index 15434c00..33363b99 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -102,6 +102,89 @@ class SpecsTest(absltest.TestCase): # Create a spec. spec = mujoco.MjSpec() + # Add material. + material = spec.add_material(texrepeat=[1, 2], emission=-1) + np.testing.assert_array_equal(material.texrepeat, [1, 2]) + self.assertEqual(material.emission, -1) + + # Add mesh. + mesh = spec.add_mesh(refpos=[1, 2, 3]) + np.testing.assert_array_equal(mesh.refpos, [1, 2, 3]) + + # Add pair. + pair = spec.add_pair(gap=0.1) + self.assertEqual(pair.gap, 0.1) + + # Add equality. + equality = spec.add_equality(objtype=mujoco.mjtObj.mjOBJ_SITE) + self.assertEqual(equality.objtype, mujoco.mjtObj.mjOBJ_SITE) + + # Add tendon. + tendon = spec.add_tendon(stiffness=2, springlength=[0.1, 0.2]) + self.assertEqual(tendon.stiffness, 2) + np.testing.assert_array_equal(tendon.springlength, [0.1, 0.2]) + + # Add actuator. + actuator = spec.add_actuator(actdim=10, ctrlrange=[-1, 10]) + self.assertEqual(actuator.actdim, 10) + np.testing.assert_array_equal(actuator.ctrlrange, [-1, 10]) + + # Add skin. + skin = spec.add_skin(inflate=2.0, vertid=[[1, 1], [2, 2]]) + self.assertEqual(skin.inflate, 2.0) + np.testing.assert_array_equal(skin.vertid, [[1, 1], [2, 2]]) + + # Add texture. + texture = spec.add_texture(builtin=0, nchannel=3) + self.assertEqual(texture.builtin, 0) + self.assertEqual(texture.nchannel, 3) + + # Add text. + text = spec.add_text(data='data', info='info') + self.assertEqual(text.data, 'data') + self.assertEqual(text.info, 'info') + + # Add tuple. + tuple_ = spec.add_tuple(objprm=[2.0, 3.0, 5.0], objname=['obj']) + np.testing.assert_array_equal(tuple_.objprm, [2.0, 3.0, 5.0]) + self.assertEqual(tuple_.objname, ['obj']) + + # Add flex. + flex = spec.add_flex(friction=[1, 2, 3], texcoord=[1.0, 2.0, 3.0]) + np.testing.assert_array_equal(flex.friction, [1, 2, 3]) + np.testing.assert_array_equal(flex.texcoord, [1.0, 2.0, 3.0]) + + # Add hfield. + hfield = spec.add_hfield(nrow=2, content_type='type') + self.assertEqual(hfield.nrow, 2) + self.assertEqual(hfield.content_type, 'type') + + # Add key. + key = spec.add_key(time=1.2, qpos=[1.0, 2.0]) + self.assertEqual(key.time, 1.2) + np.testing.assert_array_equal(key.qpos, [1.0, 2.0]) + + # Add numeric. + numeric = spec.add_numeric(data=[1.0, 1.1, 1.2], size=2) + np.testing.assert_array_equal(numeric.data, [1.0, 1.1, 1.2]) + self.assertEqual(numeric.size, 2) + + # Add exclude. + exclude = spec.add_exclude(bodyname2='body2') + self.assertEqual(exclude.bodyname2, 'body2') + + # Add sensor. + sensor = spec.add_sensor( + needstage=mujoco.mjtStage.mjSTAGE_ACC, objtype=mujoco.mjtObj.mjOBJ_SITE + ) + self.assertEqual(sensor.needstage, mujoco.mjtStage.mjSTAGE_ACC) + self.assertEqual(sensor.objtype, mujoco.mjtObj.mjOBJ_SITE) + + # Add plugin. + plugin = spec.add_plugin(plugin_slot=7, instance_name='plugin') + self.assertEqual(plugin.plugin_slot, 7) + self.assertEqual(plugin.instance_name, 'plugin') + # Add a body. body = spec.worldbody.add_body( name='body', pos=[1, 2, 3], quat=[0, 0, 0, 1] @@ -154,10 +237,38 @@ class SpecsTest(absltest.TestCase): self.assertEqual(cam.orthographic, 1) np.testing.assert_array_equal(cam.resolution, [10, 20]) + # Add frame. + framea0 = body.add_frame(name='framea', pos=[1, 2, 3], quat=[0, 0, 0, 1]) + np.testing.assert_array_equal(framea0.pos, [1, 2, 3]) + np.testing.assert_array_equal(framea0.quat, [0, 0, 0, 1]) + + frameb0 = body.add_frame( + framea0, name='frameb', pos=[4, 5, 6], quat=[0, 1, 0, 0] + ) + + framea1 = body.first_frame() + frameb1 = body.next_frame(framea1) + + self.assertEqual(framea1.name, framea0.name) + self.assertEqual(frameb1.name, frameb0.name) + np.testing.assert_array_equal(framea1.pos, framea0.pos) + np.testing.assert_array_equal(framea1.quat, framea0.quat) + np.testing.assert_array_equal(frameb1.pos, frameb0.pos) + np.testing.assert_array_equal(frameb1.quat, frameb0.quat) + # Add joint. - jnt = body.add_joint(type=mujoco.mjtJoint.mjJNT_HINGE, axis=[0, 1, 0]) - self.assertEqual(jnt.type, mujoco.mjtJoint.mjJNT_HINGE) - np.testing.assert_array_equal(jnt.axis, [0, 1, 0]) + joint = body.add_joint(type=mujoco.mjtJoint.mjJNT_HINGE, axis=[0, 1, 0]) + self.assertEqual(joint.type, mujoco.mjtJoint.mjJNT_HINGE) + np.testing.assert_array_equal(joint.axis, [0, 1, 0]) + + # Add freejoint. + freejoint = body.add_freejoint() + self.assertEqual(freejoint.type, mujoco.mjtJoint.mjJNT_FREE) + freejoint_align = body.add_freejoint(align=True) + self.assertEqual(freejoint_align.align, True) + + with self.assertRaises(TypeError): + body.add_freejoint(axis=[1, 2, 3]) # invalid keyword argument # Add light. light = body.add_light(attenuation=[1, 2, 3])