resolve PR comments

This commit is contained in:
andrew
2025-03-27 19:35:36 -04:00
parent faff1224ba
commit 97092facd6
3 changed files with 67 additions and 29 deletions
+11 -8
View File
@@ -149,7 +149,7 @@ class SimulateWrapper {
void ClearFigures() { simulate_->user_figures_.clear(); }
void SetOverlayText(
void SetText(
const std::vector<std::tuple<int, int, std::string, std::string>>& overlay_texts) {
// Collection of [font, gridpos, text1, text2] tuples for overlay text
std::vector<std::tuple<int, int, std::string, std::string>> 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<std::tuple<mjrRect, pybind11::array_t<unsigned char>>>& viewport_images
const std::vector<std::tuple<mjrRect, pybind11::array&>> 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<int>(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)
+55 -20
View File
@@ -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()