From 0cf5500b6ad2c1d6893bb5f4b4ddbdc67bf7ae07 Mon Sep 17 00:00:00 2001 From: AaronYoung5 Date: Sat, 14 Dec 2024 07:47:24 -0500 Subject: [PATCH] Fixed MjSpec introspection with visual.rgba and visual.headlight. Added access to MjSpec.visual.global_. --- python/mujoco/codegen/generate_spec_bindings.py | 7 ++++++- python/mujoco/specs.cc | 9 +++++++++ python/mujoco/specs_test.py | 9 +++++++++ 3 files changed, 24 insertions(+), 1 deletion(-) diff --git a/python/mujoco/codegen/generate_spec_bindings.py b/python/mujoco/codegen/generate_spec_bindings.py index 2db53870..c6764a3d 100644 --- a/python/mujoco/codegen/generate_spec_bindings.py +++ b/python/mujoco/codegen/generate_spec_bindings.py @@ -227,8 +227,13 @@ def _binding_code(field: ast_nodes.StructFieldDecl, key: str) -> str: if isinstance(field.type, ast_nodes.ValueType): return _value_binding_code(field.type, key, field.name) elif isinstance(field.type, ast_nodes.AnonymousStructDecl): + code = "" + if field.name in ['headlight', 'rgba']: + for subfield in field.type.fields: + code += _binding_code(subfield, 'mjVisual'+field.name.title()) field.type = ast_nodes.ValueType(name='mjVisual'+field.name.title()) - return _value_binding_code(field.type, key, field.name) + code += _value_binding_code(field.type, key, field.name) + return code elif isinstance(field.type, ast_nodes.PointerType): return _ptr_binding_code(field.type, key, field.name) elif isinstance(field.type, ast_nodes.ArrayType): diff --git a/python/mujoco/specs.cc b/python/mujoco/specs.cc index c376ea35..60a4fa69 100644 --- a/python/mujoco/specs.cc +++ b/python/mujoco/specs.cc @@ -236,6 +236,8 @@ PYBIND11_MODULE(_specs, m) { py::class_ mjOption(m, "MjOption"); py::class_ mjStatistic(m, "MjStatistic"); py::class_ mjVisual(m, "MjVisual"); + py::class_ mjVisualHeadlight(m, "MjVisualHeadlight"); + py::class_ mjVisualRgba(m, "MjVisualRgba"); py::class_ mjsCompiler(m, "MjsCompiler"); DefineArray(m, "MjCharVec"); DefineArray(m, "MjStringVec"); @@ -979,6 +981,13 @@ PYBIND11_MODULE(_specs, m) { }); mjsPlugin.def("delete", [](raw::MjsPlugin& self) { mjs_delete(self.element); }); + // ============================= MJVISUAL ==================================== + mjVisual.def_property( + "global_", + [](raw::MjVisual& self) -> raw::MjVisualGlobal& { return self.global; }, + [](raw::MjVisual& self, raw::MjVisualGlobal& value) { + self.global = value; + }); #include "specs.cc.inc" } // PYBIND11_MODULE // NOLINT diff --git a/python/mujoco/specs_test.py b/python/mujoco/specs_test.py index e0d3e836..329f08c9 100644 --- a/python/mujoco/specs_test.py +++ b/python/mujoco/specs_test.py @@ -848,22 +848,31 @@ class SpecsTest(absltest.TestCase): + + """) self.assertEqual(spec.option.timestep, 0.001) self.assertEqual(spec.stat.meansize, 0.05) self.assertEqual(spec.visual.quality.shadowsize, 4096) + self.assertEqual(spec.visual.headlight.active, 0) + self.assertEqual(spec.visual.global_, getattr(spec.visual, 'global')) + np.testing.assert_array_equal(spec.visual.rgba.camera, [0, 0, 0, 0]) spec.option.timestep = 0.002 spec.stat.meansize = 0.06 spec.visual.quality.shadowsize = 8192 + spec.visual.headlight.active = 1 + spec.visual.rgba.camera = [1, 1, 1, 1] model = spec.compile() self.assertEqual(model.opt.timestep, 0.002) self.assertEqual(model.stat.meansize, 0.06) self.assertEqual(model.vis.quality.shadowsize, 8192) + self.assertEqual(model.vis.headlight.active, 1) + np.testing.assert_array_equal(model.vis.rgba.camera, [1, 1, 1, 1]) def test_assign_list_element(self): spec = mujoco.MjSpec()