Add optional kwarg inputs to mjSpec Python bindings for mjSpec add constructors.

PiperOrigin-RevId: 675929024
Change-Id: Idd8c584276d8f585b79585230f5167f79760eaef
This commit is contained in:
Taylor Howell
2024-09-18 04:00:52 -07:00
committed by Copybara-Service
parent 4975c35657
commit bfda9f4ace
3 changed files with 205 additions and 120 deletions
+61 -22
View File
@@ -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<std::string> 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__':
+30 -95
View File
@@ -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<int>();
} 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<std::string>();
} 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<int>();
} 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",
+114 -3
View File
@@ -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])