Add optional kwarg inputs to mjSpec Python bindings for mjSpec add constructors.
PiperOrigin-RevId: 675929024 Change-Id: Idd8c584276d8f585b79585230f5167f79760eaef
This commit is contained in:
committed by
Copybara-Service
parent
4975c35657
commit
bfda9f4ace
@@ -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
@@ -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
@@ -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])
|
||||
|
||||
Reference in New Issue
Block a user