Support extra_geoms argument

PiperOrigin-RevId: 942283363
Change-Id: I7e0c3aa845c0d2ae1c4390fd5cc01d6ff2741720
This commit is contained in:
Matija Kecman
2026-07-03 16:38:39 -07:00
committed by Copybara-Service
parent 16f2276fe2
commit 7c0d527005
3 changed files with 36 additions and 11 deletions
@@ -164,7 +164,17 @@ class Viewer {
mujoco::python::MjvPerturbWrapper& perturb,
mujoco::python::MjvCameraWrapper& camera,
mujoco::python::MjvOptionWrapper& vis_options,
const std::vector<uint8_t>& render_flags) {
const std::vector<uint8_t>& render_flags,
const std::vector<mujoco::python::MjvGeomWrapper>& extra_geoms =
{}) {
std::vector<mjvGeom> 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<const std::string&, int, int, const std::string&>())
.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<mujoco::python::MjvGeomWrapper>())
.def("UploadImage", &Viewer::UploadImage)
.def("RenderToTexture", &Viewer::RenderToTexture)
.def("GetDropFile", &Viewer::GetDropFile)
@@ -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:
@@ -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:
...