diff --git a/python/mujoco/simulate.cc b/python/mujoco/simulate.cc index 868bec28..326879cc 100644 --- a/python/mujoco/simulate.cc +++ b/python/mujoco/simulate.cc @@ -149,7 +149,7 @@ class SimulateWrapper { void ClearFigures() { simulate_->user_figures_.clear(); } - void SetOverlayText( + void SetText( const std::vector>& overlay_texts) { // Collection of [font, gridpos, text1, text2] tuples for overlay text std::vector> user_overlay_text; @@ -161,21 +161,24 @@ class SimulateWrapper { simulate_->user_text_ = user_overlay_text; } - void ClearOverlayText() { simulate_->user_text_.clear(); } + void ClearText() { simulate_->user_text_.clear(); } void SetImages( - const std::vector>>& viewport_images + const std::vector> viewports_images ) { // Clear previous images to prevent memory leaks ClearImages(); - for (const auto& [viewport, image] : viewport_images) { + for (const auto& [viewport, image] : viewports_images) { auto buf = image.request(); + if (buf.ndim != 3) { + throw std::invalid_argument("image must have 3 dimensions (H, W, C)"); + } 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)"); + if (buf.itemsize != sizeof(unsigned char)) { + throw std::invalid_argument("image must be uint8 format"); } // Calculate size of the image data @@ -300,9 +303,9 @@ PYBIND11_MODULE(_simulate, pymodule) { .def("set_figures", &SimulateWrapper::SetFigures, py::arg("viewports_figures")) .def("clear_figures", &SimulateWrapper::ClearFigures) - .def("overlay_text", &SimulateWrapper::SetOverlayText, + .def("set_text", &SimulateWrapper::SetText, py::arg("overlay_texts")) - .def("clear_overlay_text", &SimulateWrapper::ClearOverlayText) + .def("clear_text", &SimulateWrapper::ClearText) .def("set_images", &SimulateWrapper::SetImages, py::arg("viewports_images")) .def("clear_images", &SimulateWrapper::ClearImages) diff --git a/python/mujoco/viewer.py b/python/mujoco/viewer.py index 21a1eeea..527fda50 100644 --- a/python/mujoco/viewer.py +++ b/python/mujoco/viewer.py @@ -117,10 +117,21 @@ class Handle: return None def set_figures( - self, viewports_figures: List[Tuple[mujoco.MjrRect, mujoco.MjvFigure]] + self, viewports_figures: Union[Tuple[mujoco.MjrRect, mujoco.MjvFigure], + List[Tuple[mujoco.MjrRect, mujoco.MjvFigure]]] ): + """Overlay figures on the viewer. + + Args: + viewports_figures: Single tuple or list of tuples of (viewport, figure) + viewport: Rectangle defining position and size of the figure + figure: MjvFigure object containing the figure data to display + """ sim = self._sim() if sim is not None: + # Convert single tuple to list if needed + if isinstance(viewports_figures, tuple): + viewports_figures = [viewports_figures] sim.set_figures(viewports_figures) def clear_figures(self): @@ -128,43 +139,67 @@ class Handle: if sim is not None: sim.clear_figures() - def overlay_text(self, overlay_texts: List[Tuple[int, int, str, str]]): + def set_text(self, overlay_texts: Union[Tuple[Optional[int], Optional[int], Optional[str], Optional[str]], + List[Tuple[Optional[int], Optional[int], Optional[str], Optional[str]]]]): """Overlay text on the viewer. Args: - overlay_texts: List of tuples of (font, gridpos, text1, text2) - let: + overlay_texts: Single tuple or list of tuples of (font, gridpos, text1, text2) font: Font style from mujoco.mjtFontScale gridpos: Position of text box from mujoco.mjtGridPos - text1: Left text column - text2: Right text column + text1: Left text column, defaults to empty string if None + text2: Right text column, defaults to empty string if None """ sim = self._sim() if sim is not None: - sim.overlay_text(overlay_texts) + # Convert single tuple to list if needed + if isinstance(overlay_texts, tuple): + overlay_texts = [overlay_texts] + + # Convert None values to empty strings + default_font = mujoco.mjtFontScale.mjFONTSCALE_150 + default_gridpos = mujoco.mjtGridPos.mjGRID_TOPLEFT + processed_texts = [( + default_font if font is None else font, + default_gridpos if gridpos is None else gridpos, + "" if text1 is None else text1, + "" if text2 is None else text2) + for font, gridpos, text1, text2 in overlay_texts] + + sim.set_text(processed_texts) - def clear_overlay_text(self): + def clear_text(self): sim = self._sim() if sim is not None: - sim.clear_overlay_text() + sim.clear_text() def set_images( - self, viewports_images: List[Tuple[mujoco.MjrRect, np.ndarray]] + self, viewports_images: Union[Tuple[mujoco.MjrRect, np.ndarray], + List[Tuple[mujoco.MjrRect, np.ndarray]]] ): + """Overlay images on the viewer. + + Args: + viewports_images: Single tuple or list of tuples of (viewport, image) + viewport: Rectangle defining position and size of the image + image: RGB image with shape (height, width, 3) + """ 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 = [] + # Convert single tuple to list if needed + if isinstance(viewports_images, tuple): + viewports_images = [viewports_images] + + processed_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) + # Check if image is already the correct shape + if image.shape[:2] != targ_shape: + raise ValueError(f"Image shape {image.shape[:2]} does not match target shape {targ_shape}") + flipped = np.flip(image, axis=0) + contiguous = np.ascontiguousarray(flipped) + processed_images.append((viewport, contiguous)) + sim.set_images(processed_images) def clear_images(self): sim = self._sim() diff --git a/simulate/simulate.h b/simulate/simulate.h index 0bf6ad25..00dc5fe5 100644 --- a/simulate/simulate.h +++ b/simulate/simulate.h @@ -249,7 +249,7 @@ class Simulate { mjvFigure figsize = {}; mjvFigure figsensor = {}; - // additional user-defined visualization geoms (used in passive mode) + // additional user-defined visualization mjvScene* user_scn = nullptr; mjtByte user_scn_flags_prev_[mjNRNDFLAG]; std::vector> user_figures_;