From d39b2f4f9d31f940e7e18591c31e70a8db3afe95 Mon Sep 17 00:00:00 2001 From: Tarik Kelestemur Date: Wed, 15 Apr 2026 18:22:46 -0400 Subject: [PATCH] mjx segmentation --- doc/mjx.rst | 13 +- mjx/mujoco/mjx/__init__.py | 4 +- mjx/mujoco/mjx/_src/render.py | 24 +++- mjx/mujoco/mjx/_src/render_util.py | 90 +++++++++---- mjx/mujoco/mjx/_src/render_util_test.py | 86 ++++++++++++- mjx/mujoco/mjx/warp/io.py | 10 +- mjx/mujoco/mjx/warp/render.py | 162 ++++++++++++++++++++---- mjx/mujoco/mjx/warp/render_test.py | 91 ++++++++++++- 8 files changed, 414 insertions(+), 66 deletions(-) diff --git a/doc/mjx.rst b/doc/mjx.rst index 0d7f0b2d..247dafa8 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -233,6 +233,7 @@ pytree that should be passed into ``jit``/``vmap``-compiled functions: use_shadows=True, render_rgb=[True] * ncam, render_depth=[False] * ncam, + render_seg=[True] * ncam, enabled_geom_groups=[0, 1, 2], ) @@ -246,21 +247,23 @@ volume hierarchy (BVH) and executing the raycaster: .. code-block:: python from mujoco.mjx import get_rgb + from mujoco.mjx import get_segmentation @jax.jit def render_fn(mx, d, rc_pytree): # 1. Update the BVH for the current scene state d = mjx.refit_bvh(mx, d, rc_pytree) - # 2. Render all configured cameras - pixels, _ = mjx.render(mx, d, rc_pytree) + # 2. Render all configured cameras, including segmentation + pixels, _, segmentation = mjx.render_with_segmentation(mx, d, rc_pytree) - # 3. Extract the RGB tensor for the first camera (index 0) + # 3. Extract the RGB tensor and geom IDs for the first camera (index 0) rgb = get_rgb(rc_pytree, 0, pixels) + seg = get_segmentation(rc_pytree, 0, segmentation) - return rgb, d + return rgb, seg, d - rgb, d = render_fn(mx, d, rc.pytree()) + rgb, seg, d = render_fn(mx, d, rc.pytree()) .. WARNING:: The batch dimension ``nworld`` is fixed when the render context is created via diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index a0ee6dd4..4a7733c1 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -20,9 +20,9 @@ from mujoco.mjx._src.types import Model from mujoco.mjx._src.types import Data # isort: on +from mujoco.mjx._src.bvh import refit_bvh # pylint:disable=g-importing-member from mujoco.mjx._src.collision_driver import collision -from mujoco.mjx._src.bvh import refit_bvh from mujoco.mjx._src.constraint import make_constraint from mujoco.mjx._src.derivative import deriv_smooth_vel from mujoco.mjx._src.forward import euler @@ -46,8 +46,10 @@ from mujoco.mjx._src.io import state_size from mujoco.mjx._src.passive import passive from mujoco.mjx._src.ray import ray from mujoco.mjx._src.render import render +from mujoco.mjx._src.render import render_with_segmentation from mujoco.mjx._src.render_util import get_depth from mujoco.mjx._src.render_util import get_rgb +from mujoco.mjx._src.render_util import get_segmentation from mujoco.mjx._src.sensor import sensor_acc from mujoco.mjx._src.sensor import sensor_pos from mujoco.mjx._src.sensor import sensor_vel diff --git a/mjx/mujoco/mjx/_src/render.py b/mjx/mujoco/mjx/_src/render.py index 45f72aed..d6126529 100644 --- a/mjx/mujoco/mjx/_src/render.py +++ b/mjx/mujoco/mjx/_src/render.py @@ -15,6 +15,7 @@ """Render helpers for MJX.""" from typing import Any + # pylint: disable=g-importing-member from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import Impl @@ -26,8 +27,8 @@ import mujoco.mjx.warp as mjxw def render(m: Model, d: Data, ctx: Any) -> Data: """Render.""" if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: - import mujoco.mjx.warp.render_context as mjxw_rc # pylint: disable=g-import-not-at-top # pytype: disable=import-error - from mujoco.mjx.warp import render as mjxw_render # pylint: disable=g-import-not-at-top # pytype: disable=import-error + from mujoco.mjx.warp import render as mjxw_render + from mujoco.mjx.warp import render_context as mjxw_rc if not isinstance(ctx, mjxw_rc.RenderContextPytree): raise TypeError( @@ -38,3 +39,22 @@ def render(m: Model, d: Data, ctx: Any) -> Data: return mjxw_render.render(m, d, ctx) raise NotImplementedError('render only implemented for MuJoCo Warp.') + + +def render_with_segmentation(m: Model, d: Data, ctx: Any) -> Data: + """Render and return RGB, depth, and packed segmentation outputs.""" + if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: + from mujoco.mjx.warp import render as mjxw_render + from mujoco.mjx.warp import render_context as mjxw_rc + + if not isinstance(ctx, mjxw_rc.RenderContextPytree): + raise TypeError( + f'Expected RenderContextPytree, got {type(ctx).__name__}.' + ' Use rc.pytree() to get the JAX-compatible handle.' + ) + + return mjxw_render.render_with_segmentation(m, d, ctx) + + raise NotImplementedError( + 'render_with_segmentation only implemented for MuJoCo Warp.' + ) diff --git a/mjx/mujoco/mjx/_src/render_util.py b/mjx/mujoco/mjx/_src/render_util.py index 098291c3..56c7e6ad 100644 --- a/mjx/mujoco/mjx/_src/render_util.py +++ b/mjx/mujoco/mjx/_src/render_util.py @@ -25,6 +25,30 @@ if TYPE_CHECKING: from mujoco.mjx.warp.render_context import RenderContextPytree +def _get_warp_render_context(rc: 'RenderContextPytree'): + """Validates and returns the backing Warp render context.""" + if not mjxw.WARP_INSTALLED: + raise RuntimeError('Warp not installed.') + + from mujoco.mjx.warp import render_context as mjxw_rc + + if not isinstance(rc, mjxw_rc.RenderContextPytree): + raise TypeError( + f'Expected RenderContextPytree, got {type(rc).__name__}.' + ' Use rc.pytree() to get the JAX-compatible handle.' + ) + + # pylint: disable=protected-access + return mjxw_rc._MJX_RENDER_CONTEXT_BUFFERS[(rc.key, None)] + + +def _get_camera_resolution(warp_rc, cam_id: int) -> tuple[int, int]: + """Returns the render resolution for a given camera.""" + width = int(warp_rc.cam_res.numpy()[cam_id][0]) + height = int(warp_rc.cam_res.numpy()[cam_id][1]) + return width, height + + def get_rgb( rc: 'RenderContextPytree', cam_id: int, @@ -44,21 +68,9 @@ def get_rgb( Raises: RuntimeError: If Warp is not installed. """ - if not mjxw.WARP_INSTALLED: - raise RuntimeError('Warp not installed.') - - import mujoco.mjx.warp.render_context as mjxw_rc # pylint: disable=g-import-not-at-top # pytype: disable=import-error - - if not isinstance(rc, mjxw_rc.RenderContextPytree): - raise TypeError( - f'Expected RenderContextPytree, got {type(rc).__name__}.' - ' Use rc.pytree() to get the JAX-compatible handle.' - ) - - warp_rc = mjxw_rc._MJX_RENDER_CONTEXT_BUFFERS[(rc.key, None)] # pylint: disable=protected-access + warp_rc = _get_warp_render_context(rc) rgb_adr = int(warp_rc.rgb_adr.numpy()[cam_id]) - width = int(warp_rc.cam_res.numpy()[cam_id][0]) - height = int(warp_rc.cam_res.numpy()[cam_id][1]) + width, height = _get_camera_resolution(warp_rc, cam_id) packed = jax.lax.dynamic_slice_in_dim( rgb_data, rgb_adr, width * height, axis=rgb_data.ndim - 1 @@ -92,21 +104,9 @@ def get_depth( Raises: RuntimeError: If Warp is not installed. """ - if not mjxw.WARP_INSTALLED: - raise RuntimeError('Warp not installed.') - - import mujoco.mjx.warp.render_context as mjxw_rc # pylint: disable=g-import-not-at-top # pytype: disable=import-error - - if not isinstance(rc, mjxw_rc.RenderContextPytree): - raise TypeError( - f'Expected RenderContextPytree, got {type(rc).__name__}.' - ' Use rc.pytree() to get the JAX-compatible handle.' - ) - - warp_rc = mjxw_rc._MJX_RENDER_CONTEXT_BUFFERS[(rc.key, None)] # pylint: disable=protected-access + warp_rc = _get_warp_render_context(rc) depth_adr = int(warp_rc.depth_adr.numpy()[cam_id]) - width = int(warp_rc.cam_res.numpy()[cam_id][0]) - height = int(warp_rc.cam_res.numpy()[cam_id][1]) + width, height = _get_camera_resolution(warp_rc, cam_id) raw = jax.lax.dynamic_slice_in_dim( depth_data, depth_adr, width * height, axis=depth_data.ndim - 1 @@ -114,3 +114,37 @@ def get_depth( depth = jnp.clip(raw / depth_scale, 0.0, 1.0) return depth.reshape(raw.shape[:-1] + (height, width, 1)) + + +def get_segmentation( + rc: 'RenderContextPytree', + cam_id: int, + seg_data: jax.Array, +) -> jax.Array: + """Extract raw geom IDs for a camera. + + Args: + rc: RenderContextPytree. + cam_id: Camera index to extract. + seg_data: Packed segmentation output, shape (..., total_pixels) as integers. + + Returns: + Integer segmentation array with shape (..., H, W). + Any leading batch axes in `seg_data` are preserved. + + Raises: + RuntimeError: If Warp is not installed. + ValueError: If segmentation is not enabled for the selected camera. + """ + warp_rc = _get_warp_render_context(rc) + seg_adr = int(warp_rc.seg_adr.numpy()[cam_id]) + if seg_adr < 0: + raise ValueError( + f'Camera {cam_id} was not configured with segmentation rendering.' + ) + + width, height = _get_camera_resolution(warp_rc, cam_id) + packed = jax.lax.dynamic_slice_in_dim( + seg_data, seg_adr, width * height, axis=seg_data.ndim - 1 + ) + return packed.reshape(packed.shape[:-1] + (height, width)) diff --git a/mjx/mujoco/mjx/_src/render_util_test.py b/mjx/mujoco/mjx/_src/render_util_test.py index f97c1e0a..87065dcc 100644 --- a/mjx/mujoco/mjx/_src/render_util_test.py +++ b/mjx/mujoco/mjx/_src/render_util_test.py @@ -28,14 +28,23 @@ from mujoco.mjx.warp.render_context import RenderContextPytree _FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' -def _fake_render_context(ncam, width, height): +def _fake_render_context(ncam, width, height, render_seg=True): """Fake RenderContext for testing.""" rc = mock.MagicMock() rgb_adr = np.arange(ncam, dtype=np.int32) * width * height depth_adr = np.arange(ncam, dtype=np.int32) * width * height + if isinstance(render_seg, bool): + render_seg = [render_seg] * ncam + seg_adr = np.full(ncam, -1, dtype=np.int32) + seg_offset = 0 + for i, enabled in enumerate(render_seg): + if enabled: + seg_adr[i] = seg_offset + seg_offset += width * height cam_res = np.tile([width, height], (ncam, 1)).astype(np.int32) rc.rgb_adr.numpy.return_value = rgb_adr rc.depth_adr.numpy.return_value = depth_adr + rc.seg_adr.numpy.return_value = seg_adr rc.cam_res.numpy.return_value = cam_res return rc @@ -154,6 +163,81 @@ class RenderUtilTest(absltest.TestCase): self.assertEqual(depth.shape, (nworld, height, width, 1)) + def test_get_segmentation(self): + width, height = 4, 4 + warp_rc = _fake_render_context(1, width, height) + rc = mock.MagicMock(spec=RenderContextPytree, key=0) + seg_data = jnp.arange(width * height, dtype=jnp.int32) + + with mock.patch.dict( + 'mujoco.mjx.warp.render_context._MJX_RENDER_CONTEXT_BUFFERS', + {(0, None): warp_rc}, + ): + segmentation = jax.jit( + render_util.get_segmentation, static_argnums=(0, 1) + )(rc, 0, seg_data) + + self.assertEqual(segmentation.shape, (height, width)) + np.testing.assert_array_equal( + np.asarray(segmentation), + np.arange(width * height, dtype=np.int32).reshape(height, width), + ) + + def test_get_segmentation_preserves_leading_dims(self): + width, height = 4, 4 + warp_rc = _fake_render_context(1, width, height) + rc = mock.MagicMock(spec=RenderContextPytree, key=0) + + with mock.patch.dict( + 'mujoco.mjx.warp.render_context._MJX_RENDER_CONTEXT_BUFFERS', + {(0, None): warp_rc}, + ): + for leading_shape in ((1,), (3,), (2, 3)): + with self.subTest(leading_shape=leading_shape): + seg_data = jnp.arange( + np.prod(leading_shape) * width * height, dtype=jnp.int32 + ).reshape(leading_shape + (width * height,)) + segmentation = jax.jit( + render_util.get_segmentation, static_argnums=(0, 1) + )(rc, 0, seg_data) + + self.assertEqual(segmentation.shape, leading_shape + (height, width)) + + def test_get_segmentation_vmap(self): + nworld, width, height = 3, 4, 4 + warp_rc = _fake_render_context(1, width, height) + rc = mock.MagicMock(spec=RenderContextPytree, key=0) + seg_data = jnp.arange(nworld * width * height, dtype=jnp.int32).reshape( + nworld, width * height + ) + + with mock.patch.dict( + 'mujoco.mjx.warp.render_context._MJX_RENDER_CONTEXT_BUFFERS', + {(0, None): warp_rc}, + ): + segmentation = jax.jit( + jax.vmap(render_util.get_segmentation, in_axes=(None, None, 0)), + static_argnums=(0, 1), + )(rc, 0, seg_data) + + self.assertEqual(segmentation.shape, (nworld, height, width)) + + def test_get_segmentation_raises_for_disabled_camera(self): + width, height = 4, 4 + warp_rc = _fake_render_context(2, width, height, render_seg=[True, False]) + rc = mock.MagicMock(spec=RenderContextPytree, key=0) + seg_data = jnp.arange(width * height, dtype=jnp.int32) + + with mock.patch.dict( + 'mujoco.mjx.warp.render_context._MJX_RENDER_CONTEXT_BUFFERS', + {(0, None): warp_rc}, + ): + with self.assertRaisesWithLiteralMatch( + ValueError, + 'Camera 1 was not configured with segmentation rendering.', + ): + render_util.get_segmentation(rc, 1, seg_data) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/warp/io.py b/mjx/mujoco/mjx/warp/io.py index b32d7122..d3179990 100644 --- a/mjx/mujoco/mjx/warp/io.py +++ b/mjx/mujoco/mjx/warp/io.py @@ -14,11 +14,12 @@ # ============================================================================== """I/O functions for MJX Warp.""" -import mujoco -from mujoco.mjx.warp import render_context -import mujoco.mjx.third_party.mujoco_warp as mjw import warp as wp +import mujoco +import mujoco.mjx.third_party.mujoco_warp as mjw +from mujoco.mjx.warp import render_context + _MJX_RENDER_CONTEXT_COUNTER = 0 @@ -27,8 +28,11 @@ def _create_context(mjm, nworld, device, **kwargs): ctx = mjw.create_render_context(mjm=mjm, nworld=nworld, **kwargs) ctx.rgb_data_shape = ctx.rgb_data.shape ctx.depth_data_shape = ctx.depth_data.shape + ctx.seg_data_shape = ctx.seg_data.shape + ctx.seg_data_buffer = ctx.seg_data ctx.rgb_data = None ctx.depth_data = None + ctx.seg_data = None return ctx diff --git a/mjx/mujoco/mjx/warp/render.py b/mjx/mujoco/mjx/warp/render.py index c97003f0..1361aa11 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -14,17 +14,19 @@ # ============================================================================== """DO NOT EDIT. This file is auto-generated.""" + import dataclasses import functools + import jax +import warp as wp + from mujoco.mjx._src import types +import mujoco.mjx.third_party.mujoco_warp as mjwarp +from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types from mujoco.mjx.warp import ffi from mujoco.mjx.warp.render_context import _MJX_RENDER_CONTEXT_BUFFERS from mujoco.mjx.warp.render_context import RenderContextPytree -import mujoco.mjx.third_party.mujoco_warp as mjwarp -from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types -import warp as wp - _m = mjwarp.Model( **{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init} @@ -122,36 +124,129 @@ def _render_shim( render_context = _MJX_RENDER_CONTEXT_BUFFERS[(rc_id, wp.get_device().ordinal)] render_context.rgb_data = rgb render_context.depth_data = depth + render_context.seg_data = render_context.seg_data_buffer mjwarp.render(_m, _d, render_context) -def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree): +@ffi.format_args_for_warp +def _render_with_segmentation_shim( + # Model + nworld: int, + cam_fovy: wp.array2d[float], + cam_intrinsic: wp.array2d[wp.vec4], + cam_projection: wp.array[int], + cam_sensorsize: wp.array[wp.vec2], + flex_edge: wp.array[wp.vec2i], + flex_radius: wp.array[float], + flex_vertadr: wp.array[int], + geom_dataid: wp.array2d[int], + geom_matid: wp.array2d[int], + geom_rgba: wp.array2d[wp.vec4], + geom_size: wp.array2d[wp.vec3], + geom_type: wp.array[int], + light_active: wp.array2d[bool], + light_castshadow: wp.array2d[bool], + light_type: wp.array2d[int], + mat_rgba: wp.array2d[wp.vec4], + mat_texid: wp.array3d[int], + mat_texrepeat: wp.array2d[wp.vec2], + mesh_faceadr: wp.array[int], + nlight: int, + # Data + cam_xmat: wp.array2d[wp.mat33], + cam_xpos: wp.array2d[wp.vec3], + flexvert_xpos: wp.array2d[wp.vec3], + geom_xmat: wp.array2d[wp.mat33], + geom_xpos: wp.array2d[wp.vec3], + light_xdir: wp.array2d[wp.vec3], + light_xpos: wp.array2d[wp.vec3], + # Registry + rc_id: int, + rgb: wp.array2d[wp.uint32], + depth: wp.array2d[wp.float32], + seg: wp.array2d[int], +): + _m.stat = _s + _m.opt = _o + _m.callback = _cb + _d.efc = _e + _d.contact = _c + _m.cam_fovy = cam_fovy + _m.cam_intrinsic = cam_intrinsic + _m.cam_projection = cam_projection + _m.cam_sensorsize = cam_sensorsize + _m.flex_edge = flex_edge + _m.flex_radius = flex_radius + _m.flex_vertadr = flex_vertadr + _m.geom_dataid = geom_dataid + _m.geom_matid = geom_matid + _m.geom_rgba = geom_rgba + _m.geom_size = geom_size + _m.geom_type = geom_type + _m.light_active = light_active + _m.light_castshadow = light_castshadow + _m.light_type = light_type + _m.mat_rgba = mat_rgba + _m.mat_texid = mat_texid + _m.mat_texrepeat = mat_texrepeat + _m.mesh_faceadr = mesh_faceadr + _m.nlight = nlight + _d.cam_xmat = cam_xmat + _d.cam_xpos = cam_xpos + _d.flexvert_xpos = flexvert_xpos + _d.geom_xmat = geom_xmat + _d.geom_xpos = geom_xpos + _d.light_xdir = light_xdir + _d.light_xpos = light_xpos + _d.nworld = nworld + render_context = _MJX_RENDER_CONTEXT_BUFFERS[(rc_id, wp.get_device().ordinal)] + render_context.rgb_data = rgb + render_context.depth_data = depth + render_context.seg_data = seg + mjwarp.render(_m, _d, render_context) + + +_RENDER_STAGE_IN_ARGNAMES = set([ + 'cam_fovy', + 'cam_intrinsic', + 'cam_xmat', + 'cam_xpos', + 'geom_matid', + 'geom_rgba', + 'geom_size', + 'geom_xmat', + 'geom_xpos', + 'light_castshadow', + 'light_type', + 'mat_rgba', + 'mat_texid', +]) + + +def _render_jax_impl( + m: types.Model, + d: types.Data, + ctx: RenderContextPytree, + with_segmentation: bool = False, +): render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[(ctx.key, None)] output_dims = { 'rgb': render_ctx.rgb_data_shape, 'depth': render_ctx.depth_data_shape, } + render_shim = _render_shim + num_outputs = 2 + if with_segmentation: + output_dims['seg'] = render_ctx.seg_data_shape + render_shim = _render_with_segmentation_shim + num_outputs = 3 jf = ffi.jax_callable_variadic_tuple( - _render_shim, - num_outputs=2, + render_shim, + num_outputs=num_outputs, output_dims=output_dims, vmap_method=None, in_out_argnames=set([]), - stage_in_argnames=set([ - 'cam_fovy', - 'cam_intrinsic', - 'cam_xmat', - 'cam_xpos', - 'geom_matid', - 'geom_rgba', - 'geom_size', - 'geom_xmat', - 'geom_xpos', - 'light_castshadow', - 'light_type', - 'mat_rgba', - 'mat_texid', - ]), + stage_in_argnames=_RENDER_STAGE_IN_ARGNAMES, stage_out_argnames=set([]), graph_mode=m.opt._impl.graph_mode, has_side_effect=False, @@ -208,3 +303,26 @@ def render_vmap( ): out = render(m, d, ctx) return out, [True, True] + + +@jax.custom_batching.custom_vmap +@functools.partial(ffi.marshal_jax_warp_callable, tree_map_output=True) +def render_with_segmentation( + m: types.Model, + d: types.Data, + ctx: RenderContextPytree, +): + return _render_jax_impl(m, d, ctx, with_segmentation=True) + + +@render_with_segmentation.def_vmap +@functools.partial(ffi.marshal_custom_vmap, tree_map_output=True) +def render_with_segmentation_vmap( + unused_axis_size, + is_batched, + m: types.Model, + d: types.Data, + ctx: RenderContextPytree, +): + out = render_with_segmentation(m, d, ctx) + return out, [True, True, True] diff --git a/mjx/mujoco/mjx/warp/render_test.py b/mjx/mujoco/mjx/warp/render_test.py index d4864d9c..fc089617 100644 --- a/mjx/mujoco/mjx/warp/render_test.py +++ b/mjx/mujoco/mjx/warp/render_test.py @@ -19,22 +19,22 @@ from absl.testing import absltest from absl.testing import parameterized import jax from jax import numpy as jp +import numpy as np + import mujoco from mujoco import mjx from mujoco.mjx._src import bvh from mujoco.mjx._src import forward from mujoco.mjx._src import io from mujoco.mjx._src import render -import mujoco.mjx.warp as mjxw from mujoco.mjx.warp import test_util as tu from mujoco.mjx.warp import warp as wp # pylint: disable=g-importing-member -import numpy as np - +import mujoco.mjx.warp as mjxw _FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' -def _get_model_data_rc(xml, batch_size): +def _get_model_data_rc(xml, batch_size, render_seg=False): m = tu.load_test_file(xml) d = mujoco.MjData(m) mujoco.mj_forward(m, d) @@ -63,6 +63,7 @@ def _get_model_data_rc(xml, batch_size): use_shadows=True, render_rgb=True, render_depth=True, + render_seg=render_seg, enabled_geom_groups=[0, 1, 2], ) return mx, dx_batch, rc @@ -148,6 +149,88 @@ class RenderTest(parameterized.TestCase): self.assertGreater(np.count_nonzero(depth), 0) self.assertNotEqual(np.unique(depth).shape[0], 1) + @parameterized.product( + xml=('humanoid/humanoid.xml',), + batch_size=(1, 16), + ) + def test_render_with_segmentation(self, xml: str, batch_size: int): + """Tests MJX render pipeline with packed segmentation output.""" + self._maybe_skip() + mx, dx_batch, rc = _get_model_data_rc(xml, batch_size, render_seg=True) + + dx_batch = jax.jit(mjx.refit_bvh)(mx, dx_batch, rc.pytree()) + out_batch = jax.jit(mjx.render_with_segmentation)(mx, dx_batch, rc.pytree()) + + rgb = np.asarray(out_batch[0]) + depth = np.asarray(out_batch[1]) + seg = np.asarray(out_batch[2]) + + self.assertGreater(np.count_nonzero(rgb), 0) + self.assertGreater(np.count_nonzero(depth), 0) + self.assertTrue(np.any(seg != -1)) + self.assertGreater(np.unique(seg).shape[0], 1) + + unpacked_seg = jax.vmap(mjx.get_segmentation, in_axes=(None, None, 0))( + rc.pytree(), 0, out_batch[2] + ) + unpacked_seg = np.asarray(unpacked_seg) + width, height = rc._default.cam_res.numpy()[ + 0 + ] # pylint: disable=protected-access + seg_adr = int( + rc._default.seg_adr.numpy()[0] # pylint: disable=protected-access + ) + expected_seg = seg[:, seg_adr : seg_adr + width * height].reshape( + batch_size, height, width + ) + np.testing.assert_array_equal(unpacked_seg, expected_seg) + + @parameterized.product( + xml=('humanoid/humanoid.xml',), + batch_size=(4, 16), + ) + def test_render_with_segmentation_nested_vmap( + self, xml: str, batch_size: int + ): + """Tests MJX render_with_segmentation with nested vmap.""" + self._maybe_skip() + mx, dx_batch, rc = _get_model_data_rc(xml, batch_size, render_seg=True) + + def inner(mx, dx, rc): + dx = jax.vmap(bvh.refit_bvh, in_axes=(None, 0, None))(mx, dx, rc) + out = jax.vmap(render.render_with_segmentation, in_axes=(None, 0, None))( + mx, dx, rc + ) + return out + + dx_batch = jax.vmap(bvh.refit_bvh, in_axes=(None, 0, None))( + mx, dx_batch, rc.pytree() + ) + ref = jax.vmap(render.render_with_segmentation, in_axes=(None, 0, None))( + mx, dx_batch, rc.pytree() + ) + ref_rgb = np.asarray(ref[0]) + ref_depth = np.asarray(ref[1]) + ref_seg = np.asarray(ref[2]) + + def _reshape_batched(x): + if x.shape[0] == batch_size: + return x.reshape(2, batch_size // 2, *x.shape[1:]) + return x + + dx_2d = jax.tree.map(_reshape_batched, dx_batch) + + out_batch = jax.vmap(inner, in_axes=(None, 0, None))(mx, dx_2d, rc.pytree()) + out_batch = jax.tree.map(lambda x: x.reshape(-1, *x.shape[2:]), out_batch) + rgb = np.asarray(out_batch[0]) + depth = np.asarray(out_batch[1]) + seg = np.asarray(out_batch[2]) + + np.testing.assert_array_equal(rgb, ref_rgb) + np.testing.assert_array_equal(depth, ref_depth) + np.testing.assert_array_equal(seg, ref_seg) + self.assertTrue(np.any(seg != -1)) + class RenderContextGarbageCollectionTest(absltest.TestCase): """Tests that RenderContext cleans up buffers on deletion."""