From f02fdfe2d1d3149579b29650d75117190c189cdb Mon Sep 17 00:00:00 2001 From: andrew Date: Fri, 14 Mar 2025 21:41:28 -0400 Subject: [PATCH] add handles for text and image overlays --- python/mujoco/simulate.cc | 50 +++++++++++++++++++++++++++++++++++++++ python/mujoco/viewer.py | 44 +++++++++++++++++++++++++++++++--- simulate/simulate.cc | 18 ++++++++++++++ simulate/simulate.h | 2 ++ 4 files changed, 111 insertions(+), 3 deletions(-) diff --git a/python/mujoco/simulate.cc b/python/mujoco/simulate.cc index 56fe4437..8fdb668c 100644 --- a/python/mujoco/simulate.cc +++ b/python/mujoco/simulate.cc @@ -148,6 +148,50 @@ class SimulateWrapper { void ClearFigures() { simulate_->user_figures_.clear(); } + void SetOverlayText( + const std::vector>& overlay_texts) { + // Collection of [font, gridpos, text1, text2] tuples for overlay text + std::vector> user_overlay_text; + for (const auto& [font, gridpos, text1, text2] : overlay_texts) { + user_overlay_text.push_back(std::make_tuple(font, gridpos, text1, text2)); + } + + // Set them all at once to prevent overlay text flickering. + simulate_->user_text_ = user_overlay_text; + } + + void ClearOverlayText() { simulate_->user_text_.clear(); } + + void SetImages( + const std::vector>>& viewport_images + ) { + // Clear previous images to prevent memory leaks + simulate_->user_images_.clear(); + + for (const auto& [viewport, image] : viewport_images) { + auto buf = image.request(); + if (static_cast(buf.shape[2]) != 3) { + throw std::invalid_argument("image must have 3 channels"); + } + if (buf.ndim != 3) { + throw std::invalid_argument("image must have 3 dimensions (H, W, C)"); + } + + // Calculate size of the image data + size_t height = buf.shape[0]; + size_t width = buf.shape[1]; + size_t size = height * width * 3; + + // Make a copy of the image data to prevent flickering + unsigned char* image_copy = new unsigned char[size]; + std::memcpy(image_copy, buf.ptr, size); + + simulate_->user_images_.push_back(std::make_tuple(viewport, image_copy)); + } + } + + void ClearImages() { simulate_->user_images_.clear(); } + private: mujoco::Simulate* simulate_; std::atomic_int destroyed_ = 0; @@ -249,6 +293,12 @@ PYBIND11_MODULE(_simulate, pymodule) { .def("set_figures", &SimulateWrapper::SetFigures, py::arg("viewports_figures")) .def("clear_figures", &SimulateWrapper::ClearFigures) + .def("overlay_text", &SimulateWrapper::SetOverlayText, + py::arg("overlay_texts")) + .def("clear_overlay_text", &SimulateWrapper::ClearOverlayText) + .def("set_images", &SimulateWrapper::SetImages, + py::arg("viewports_images")) + .def("clear_images", &SimulateWrapper::ClearImages) .def_property_readonly("m", &SimulateWrapper::GetModel) .def_property_readonly("d", &SimulateWrapper::GetData) .def_property_readonly("viewport", &SimulateWrapper::GetViewport) diff --git a/python/mujoco/viewer.py b/python/mujoco/viewer.py index 65852c87..b37355cc 100644 --- a/python/mujoco/viewer.py +++ b/python/mujoco/viewer.py @@ -23,9 +23,8 @@ import queue import sys import threading import time -from typing import Callable, Optional, Tuple, Union +from typing import Callable, List, Optional, Tuple, Union import weakref - import glfw import mujoco from mujoco import _simulate @@ -115,7 +114,7 @@ class Handle: return sim.viewport return None - def set_figures(self, viewports_figures): + def set_figures(self, viewports_figures: List[Tuple[mujoco.MjrRect, mujoco.MjvFigure]]): sim = self._sim() if sim is not None: sim.set_figures(viewports_figures) @@ -125,6 +124,45 @@ class Handle: if sim is not None: sim.clear_figures() + def overlay_text(self, overlay_texts: List[Tuple[int, int, str, str]]): + """ Overlay text on the viewer. + + Args: + overlay_texts: List of tuples of (font, gridpos, text1, text2) + let: + font: Font style from mujoco.mjtFontScale + gridpos: Position of text box from mujoco.mjtGridPos + text1: Left text column + text2: Right text column + """ + sim = self._sim() + if sim is not None: + sim.overlay_text(overlay_texts) + + def clear_overlay_text(self): + sim = self._sim() + if sim is not None: + sim.clear_overlay_text() + + def set_images(self, viewports_images: List[Tuple[mujoco.MjrRect, np.ndarray]]): + sim = self._sim() + if sim is not None: + # Nearest neighbor resize + resize = lambda a, s: a[(np.arange(s[0]) * a.shape[0]) // s[0]][:, (np.arange(s[1]) * a.shape[1]) // s[1]] + resized_viewports_images = [] + for viewport, image in viewports_images: + targ_shape = (viewport.height, viewport.width) + resized = resize(image, targ_shape) + resized = np.flip(resized, axis=0) + resized = np.ascontiguousarray(resized) + resized_viewports_images.append((viewport, resized)) + sim.set_images(resized_viewports_images) + + def clear_images(self): + sim = self._sim() + if sim is not None: + sim.clear_images() + def close(self): sim = self._sim() if sim is not None: diff --git a/simulate/simulate.cc b/simulate/simulate.cc index bf4b4a13..57e3eaaa 100644 --- a/simulate/simulate.cc +++ b/simulate/simulate.cc @@ -546,6 +546,14 @@ void ShowFigure(mj::Simulate* sim, mjrRect viewport, mjvFigure* fig){ mjr_figure(viewport, fig, &sim->platform_ui->mjr_context()); } +void ShowOverlayText(mj::Simulate* sim, mjrRect viewport, int font, int gridpos, std::string text1, std::string text2){ + mjr_overlay(font, gridpos, viewport, text1.c_str(), text2.c_str(), &sim->platform_ui->mjr_context()); +} + +void ShowImage(mj::Simulate* sim, mjrRect viewport, const unsigned char* image) { + mjr_drawPixels(image, nullptr, viewport, &sim->platform_ui->mjr_context()); +} + // load state from history buffer static void LoadScrubState(mj::Simulate* sim) { // get index into circular buffer @@ -2597,6 +2605,16 @@ void Simulate::Render() { ShowFigure(this, viewport, &figure); } + // overlay text + for (auto& [font, gridpos, text1, text2] : this->user_text_) { + ShowOverlayText(this, rect, font, gridpos, text1, text2); + } + + // user images + for (auto& [viewport, image] : this->user_images_) { + ShowImage(this, viewport, image); + } + // finalize this->platform_ui->SwapBuffers(); } diff --git a/simulate/simulate.h b/simulate/simulate.h index cd654192..0bf6ad25 100644 --- a/simulate/simulate.h +++ b/simulate/simulate.h @@ -253,6 +253,8 @@ class Simulate { mjvScene* user_scn = nullptr; mjtByte user_scn_flags_prev_[mjNRNDFLAG]; std::vector> user_figures_; + std::vector> user_text_; + std::vector> user_images_; // OpenGL rendering and UI int refresh_rate = 60;