diff --git a/python/mujoco/experimental/studio/native_viewer.cc b/python/mujoco/experimental/studio/native_viewer.cc index de907eaf..1436e767 100644 --- a/python/mujoco/experimental/studio/native_viewer.cc +++ b/python/mujoco/experimental/studio/native_viewer.cc @@ -164,7 +164,17 @@ class Viewer { mujoco::python::MjvPerturbWrapper& perturb, mujoco::python::MjvCameraWrapper& camera, mujoco::python::MjvOptionWrapper& vis_options, - const std::vector& render_flags) { + const std::vector& render_flags, + const std::vector& extra_geoms = + {}) { + std::vector geoms; + geoms.reserve(extra_geoms.size()); + for (const auto& geom_wrapper : extra_geoms) { + if (geom_wrapper.get()) { + geoms.push_back(*geom_wrapper.get()); + } + } + py::gil_scoped_release no_gil; const float width = window_->GetWidth(); @@ -185,7 +195,7 @@ class Viewer { renderer_->Render(model.get(), data.get(), perturb.get(), camera.get(), vis_options.get(), width * scale, height * scale, - pixels_); + pixels_, geoms); window_->EndFrame(); window_->Present(pixels_); @@ -207,7 +217,12 @@ PYBIND11_MODULE(native_viewer_cc, m, pybind11::mod_gil_not_used()) { .def(pybind11::init()) .def("InitRenderer", &Viewer::InitRenderer) .def("NewFrame", &Viewer::NewFrame) - .def("Present", &Viewer::Present) + .def("Present", &Viewer::Present, pybind11::arg("model"), + pybind11::arg("data"), pybind11::arg("perturb"), + pybind11::arg("camera"), pybind11::arg("vis_options"), + pybind11::arg("render_flags"), + pybind11::arg("extra_geoms") = + std::vector()) .def("UploadImage", &Viewer::UploadImage) .def("RenderToTexture", &Viewer::RenderToTexture) .def("GetDropFile", &Viewer::GetDropFile) diff --git a/python/mujoco/experimental/studio/native_viewer.py b/python/mujoco/experimental/studio/native_viewer.py index cd50d8cf..5e1e341a 100644 --- a/python/mujoco/experimental/studio/native_viewer.py +++ b/python/mujoco/experimental/studio/native_viewer.py @@ -38,6 +38,7 @@ class NativeViewer(vp.Viewer): vis_options: mujoco.MjvOption | None = None, perturb: mujoco.MjvPerturb | None = None, render_flags: ux.RenderFlags | None = None, + extra_geoms: list[mujoco.MjvGeom] | None = None, ) -> None: """Initializes the NativeViewer. @@ -50,19 +51,15 @@ class NativeViewer(vp.Viewer): vis_options: Visualization options. Internal object is created if None. perturb: Perturbation parameters. Internal object is created if None. render_flags: Render flags. Internal object is created if None. + extra_geoms: List of extra geoms. Internal list is created if None. """ self.config = config + + # Set members of vp.Viewer. self.camera = camera or mujoco.MjvCamera() self.perturb = perturb or mujoco.MjvPerturb() self.vis_options = vis_options or mujoco.MjvOption() - self._viewer = _viewer.Viewer( - config.title, config.width, config.height, config.gfx or '' - ) - # This class does not own the model but we need to know if the model being - # rendered has changed, so we store the unique python object id here so we - # can use it to detect model changes. - self._renderer_model_id = id(None) - self._is_running = True + self.extra_geoms = extra_geoms or [] if render_flags is not None: self.render_flags = render_flags else: @@ -70,6 +67,17 @@ class NativeViewer(vp.Viewer): # Initted to match mujoco/src/engine/engine_vis_init.c self.render_flags.flags = [1, 0, 1, 0, 1, 0, 1, 0, 0, 0, 1] + # Create the renderer. + self._viewer = _viewer.Viewer( + config.title, config.width, config.height, config.gfx or '' + ) + + # This class does not own the model but we need to know if the model being + # rendered has changed, so we store the unique python object id here so we + # can use it to detect model changes. + self._renderer_model_id = id(None) + self._is_running = True + ctx = self._viewer.GetImGuiContext() imgui.SetCurrentContext(ctx) ux.set_imgui_context(ctx) @@ -106,6 +114,7 @@ class NativeViewer(vp.Viewer): self.camera, self.vis_options, self.render_flags.flags, + self.extra_geoms, ) def close(self) -> None: diff --git a/python/mujoco/experimental/studio/viewer_protocol.py b/python/mujoco/experimental/studio/viewer_protocol.py index fa1e7138..bff5841d 100644 --- a/python/mujoco/experimental/studio/viewer_protocol.py +++ b/python/mujoco/experimental/studio/viewer_protocol.py @@ -93,6 +93,7 @@ class Viewer(Protocol): perturb: mujoco.MjvPerturb vis_options: mujoco.MjvOption render_flags: ux.RenderFlags + extra_geoms: list[mujoco.MjvGeom] def is_running(self) -> bool: ...