diff --git a/mjx/mujoco/mjx/__init__.py b/mjx/mujoco/mjx/__init__.py index a0ee6dd4..b97abb4b 100644 --- a/mjx/mujoco/mjx/__init__.py +++ b/mjx/mujoco/mjx/__init__.py @@ -21,8 +21,8 @@ from mujoco.mjx._src.types import Data # isort: on # 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.collision_driver import collision 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..68b843cf 100644 --- a/mjx/mujoco/mjx/_src/render.py +++ b/mjx/mujoco/mjx/_src/render.py @@ -15,26 +15,60 @@ """Render helpers for MJX.""" from typing import Any + +import jax +import mujoco.mjx.warp as mjxw + # pylint: disable=g-importing-member from mujoco.mjx._src.types import Data from mujoco.mjx._src.types import Impl from mujoco.mjx._src.types import Model # pylint: enable=g-importing-member -import mujoco.mjx.warp as mjxw -def render(m: Model, d: Data, ctx: Any) -> Data: - """Render.""" +def _require_segmentation_enabled(warp_rc) -> None: + """Raises if the render context has no segmentation-enabled cameras.""" + if not (warp_rc.seg_adr.numpy() >= 0).any(): + raise ValueError( + 'Render context was not configured with segmentation rendering. ' + 'Pass render_seg=True or enable it for at least one camera in ' + 'create_render_context.' + ) + + +def render(m: Model, d: Data, ctx: Any) -> tuple[jax.Array, jax.Array]: + """Render packed RGB and depth buffers.""" 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 # pytype: disable=import-error + from mujoco.mjx.warp import render_context # pytype: disable=import-error - 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(m, d, ctx) + render_context.get(ctx) + out = mjxw_render.render(m, d, ctx) + return out[0], out[1] raise NotImplementedError('render only implemented for MuJoCo Warp.') + + +def render_with_segmentation( + m: Model, d: Data, ctx: Any +) -> tuple[jax.Array, jax.Array, jax.Array]: + """Render and return RGB, depth, and packed segmentation outputs. + + Returns: + A tuple ``(rgb, depth, seg)`` of packed buffers. The segmentation buffer + stores per-pixel ``(object_id, object_type)`` pairs matching the + ``mujoco_warp`` convention. + """ + if m.impl == Impl.WARP and d.impl == Impl.WARP and mjxw.WARP_INSTALLED: + from mujoco.mjx.warp import render as mjxw_render # pytype: disable=import-error + from mujoco.mjx.warp import render_context # pytype: disable=import-error + + warp_rc = render_context.get(ctx) + _require_segmentation_enabled(warp_rc) + + out = mjxw_render.render(m, d, ctx) + return out[0], out[1], out[2] + + raise NotImplementedError( + 'render_with_segmentation only implemented for MuJoCo Warp.' + ) diff --git a/mjx/mujoco/mjx/_src/render_test.py b/mjx/mujoco/mjx/_src/render_test.py new file mode 100644 index 00000000..df09837a --- /dev/null +++ b/mjx/mujoco/mjx/_src/render_test.py @@ -0,0 +1,132 @@ +# Copyright 2026 DeepMind Technologies Limited +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Integration tests for render + get_rgb / get_depth / get_segmentation.""" + +import functools +import os + +from absl.testing import absltest +from absl.testing import parameterized +import jax +from jax import numpy as jp +import mujoco +from mujoco import mjx +from mujoco.mjx._src import forward +from mujoco.mjx._src import io +import mujoco.mjx.warp as mjxw +from mujoco.mjx.warp import test_util as tu +import numpy as np + +_FORCE_TEST = os.environ.get('MJX_WARP_FORCE_TEST', '0') == '1' + +_WIDTH, _HEIGHT = 32, 32 + + +def _setup(batch_size): + """Returns (mx, dx_batch, rc) for humanoid with segmentation enabled.""" + m = tu.load_test_file('humanoid/humanoid.xml') + d = mujoco.MjData(m) + mujoco.mj_forward(m, d) + + mx = mjx.put_model(m, impl='warp') + worldids = jp.arange(batch_size) + dx_batch = jax.vmap(functools.partial(tu.make_data, m))(worldids) + dx_batch = jax.jit(jax.vmap(forward.forward, in_axes=(None, 0)))( + mx, dx_batch + ) + + rc = mjx.create_render_context( + mjm=m, + nworld=batch_size, + cam_res=(_WIDTH, _HEIGHT), + render_rgb=True, + render_depth=True, + render_seg=True, + enabled_geom_groups=[0, 1, 2], + ) + dx_batch = jax.jit(mjx.refit_bvh)(mx, dx_batch, rc.pytree()) + return mx, dx_batch, rc + + +class RenderIntegrationTest(parameterized.TestCase): + """Tests the full render → unpack pipeline.""" + + def setUp(self): + super().setUp() + if mjxw.WARP_INSTALLED: + import warp # pylint: disable=g-import-not-at-top + warp.config.kernel_cache_dir = '/tmp/wp_kernel_cache_dir_RenderIntTest' + np.random.seed(0) + + def _maybe_skip(self): + if not _FORCE_TEST: + if not mjxw.WARP_INSTALLED: + self.skipTest('Warp not installed.') + if not io.has_cuda_gpu_device(): + self.skipTest('No CUDA GPU device available.') + + @parameterized.parameters(1, 4) + def test_render_unpack(self, batch_size): + """render_with_segmentation → get_rgb / get_depth / get_segmentation.""" + self._maybe_skip() + mx, dx_batch, rc = _setup(batch_size) + + rgb_packed, depth_packed, seg_packed = jax.jit( + mjx.render_with_segmentation + )(mx, dx_batch, rc.pytree()) + + rc_pytree = rc.pytree() + rgb = mjx.get_rgb(rc_pytree, 0, rgb_packed) + depth = mjx.get_depth(rc_pytree, 0, depth_packed, 5.0) + seg = mjx.get_segmentation(rc_pytree, 0, seg_packed) + + self.assertEqual(rgb.shape, (batch_size, _HEIGHT, _WIDTH, 3)) + self.assertEqual(depth.shape, (batch_size, _HEIGHT, _WIDTH, 1)) + self.assertEqual(seg.shape, (batch_size, _HEIGHT, _WIDTH)) + self.assertGreater(np.count_nonzero(np.asarray(rgb)), 0) + self.assertGreater(np.count_nonzero(np.asarray(depth)), 0) + self.assertTrue(np.any(np.asarray(seg) != -1)) + + @parameterized.parameters((4,),) + def test_render_unpack_vmap(self, batch_size): + """render_with_segmentation → vmap(get_rgb / get_depth / get_seg).""" + self._maybe_skip() + mx, dx_batch, rc = _setup(batch_size) + + rgb_packed, depth_packed, seg_packed = jax.jit( + mjx.render_with_segmentation + )(mx, dx_batch, rc.pytree()) + + rc_pytree = rc.pytree() + rgb = jax.vmap(mjx.get_rgb, in_axes=(None, None, 0))( + rc_pytree, 0, rgb_packed + ) + depth = jax.vmap(mjx.get_depth, in_axes=(None, None, 0, None))( + rc_pytree, 0, depth_packed, 5.0 + ) + seg = jax.vmap(mjx.get_segmentation, in_axes=(None, None, 0))( + rc_pytree, 0, seg_packed + ) + + self.assertEqual(rgb.shape, (batch_size, _HEIGHT, _WIDTH, 3)) + self.assertEqual(depth.shape, (batch_size, _HEIGHT, _WIDTH, 1)) + self.assertEqual(seg.shape, (batch_size, _HEIGHT, _WIDTH)) + self.assertGreater(np.count_nonzero(np.asarray(rgb)), 0) + self.assertGreater(np.count_nonzero(np.asarray(depth)), 0) + self.assertTrue(np.any(np.asarray(seg) != -1)) + + +if __name__ == '__main__': + absltest.main() diff --git a/mjx/mujoco/mjx/_src/render_util.py b/mjx/mujoco/mjx/_src/render_util.py index 098291c3..69c31675 100644 --- a/mjx/mujoco/mjx/_src/render_util.py +++ b/mjx/mujoco/mjx/_src/render_util.py @@ -18,13 +18,19 @@ from typing import TYPE_CHECKING import jax import jax.numpy as jnp - import mujoco.mjx.warp as mjxw if TYPE_CHECKING: from mujoco.mjx.warp.render_context import RenderContextPytree +def _get_camera_resolution(warp_rc, cam_id: int) -> tuple[int, int]: + """Returns (width, height) 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, @@ -47,18 +53,11 @@ def get_rgb( 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 + from mujoco.mjx.warp import render_context # pylint: disable=g-import-not-at-top - 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 = render_context.get(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 @@ -95,18 +94,11 @@ def get_depth( 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 + from mujoco.mjx.warp import render_context # pylint: disable=g-import-not-at-top - 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 = render_context.get(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 +106,46 @@ 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 segmentation object IDs for a camera. + + Args: + rc: RenderContextPytree. + cam_id: Camera index to extract. + seg_data: Packed segmentation output, shape (..., total_pixels, 2). Each + pixel stores a ``(object_id, object_type)`` pair matching the + ``mujoco_warp`` convention. + + Returns: + Integer segmentation array with shape (..., H, W). Each pixel contains the + object ID (geom or mesh index, ``-1`` for background). + + Raises: + RuntimeError: If Warp is not installed. + ValueError: If segmentation is not enabled for the selected camera. + """ + if not mjxw.WARP_INSTALLED: + raise RuntimeError('Warp not installed.') + + from mujoco.mjx.warp import render_context # pylint: disable=g-import-not-at-top + + warp_rc = render_context.get(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) + # seg_data shape: (..., total_pixels, 2); slice along pixel axis. + packed = jax.lax.dynamic_slice_in_dim( + seg_data, seg_adr, width * height, axis=seg_data.ndim - 2 + ) + # Extract object_id (index 0), discard object_type (index 1). + return packed[..., 0].reshape(packed.shape[:-2] + (height, width)) diff --git a/mjx/mujoco/mjx/_src/render_util_test.py b/mjx/mujoco/mjx/_src/render_util_test.py index f97c1e0a..963745d9 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,71 @@ 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) + objids = jnp.arange(width * height, dtype=jnp.int32) + # Shape: (total_pixels, 2) — (object_id, object_type) per pixel. + seg_data = jnp.stack([objids, jnp.ones_like(objids)], axis=-1) + + 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): + n = int(np.prod(leading_shape)) * width * height + objids = jnp.arange(n, dtype=jnp.int32) + seg_data = jnp.stack( + [objids, jnp.ones_like(objids)], axis=-1 + ).reshape(leading_shape + (width * height, 2)) + 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) + n = nworld * width * height + objids = jnp.arange(n, dtype=jnp.int32) + seg_data = jnp.stack( + [objids, jnp.ones_like(objids)], axis=-1 + ).reshape(nworld, width * height, 2) + + 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)) + if __name__ == '__main__': absltest.main() diff --git a/mjx/mujoco/mjx/codegen/generate_warp_shim.py b/mjx/mujoco/mjx/codegen/generate_warp_shim.py index 84c63925..4c163153 100644 --- a/mjx/mujoco/mjx/codegen/generate_warp_shim.py +++ b/mjx/mujoco/mjx/codegen/generate_warp_shim.py @@ -251,8 +251,10 @@ def _warp_function( if fn_name == 'render': render_context_args.append('rgb: wp.array2d[wp.uint32],') render_context_args.append('depth: wp.array2d[wp.float32],') + render_context_args.append('seg: wp.array2d[wp.vec2i],') fn_assignments.append(' render_context.rgb_data = rgb') fn_assignments.append(' render_context.depth_data = depth') + fn_assignments.append(' render_context.seg_data = seg') else: fn_assignments.append(' dummy.zero_()') @@ -287,13 +289,16 @@ def _jax_shim_fn( for arg in warp_fn_args: if 'nworld' in arg: - jax_args.append('d.qpos.shape[0]') + if field_usage.render_context_in_caller: + jax_args.append('render_ctx.nworld') + else: + jax_args.append('d.qpos.shape[0]') continue if arg in ('rc_id', 'dummy'): continue - if arg in ('rgb', 'depth') and fn_name == 'render': + if arg in ('rgb', 'depth', 'seg') and fn_name == 'render': num_outputs += 1 continue @@ -342,13 +347,17 @@ def _jax_shim_fn( needs_dummy_output = not field_usage.data_out_fields if needs_dummy_output and fn_name != 'render': num_outputs = 1 - output_dims = ["'dummy': (d.qpos.shape[0],)"] + if field_usage.render_context_in_caller: + output_dims = ["'dummy': (render_ctx.nworld,)"] + else: + output_dims = ["'dummy': (d.qpos.shape[0],)"] has_side_effect = True if fn_name == 'render': output_dims = [ "'rgb': render_ctx.rgb_data_shape", "'depth': render_ctx.depth_data_shape", + "'seg': render_ctx.seg_data_shape", ] tree_replace = [] @@ -438,8 +447,9 @@ def _{fn_name}_shim( ) = _jax_shim_fn(fn_name, field_usage, warp_fn_args, mjwarp_field_info) render_ctx_line = '' return_stmt = 'return d' - if fn_name == 'render': + if field_usage.render_context_in_caller: render_ctx_line = f' render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[(ctx.key, None)]\n' + if fn_name == 'render': return_stmt = 'return out' output_dims_str = '{' + ','.join(output_dims) + '}' data_tree_replace = f"d = d.tree_replace({{ {','.join(tree_replace)} }})" @@ -478,7 +488,7 @@ def _{fn_name}_jax_impl({','.join(fn_args)}): '@functools.partial(ffi.marshal_custom_vmap, tree_map_output=True)' ) vmap_return_stmt = ( - f'out = {fn_name}({fn_call_str})\n return out, [True, True]' + f'out = {fn_name}({fn_call_str})\n return out, [True, True, True]' ) src += f""" diff --git a/mjx/mujoco/mjx/warp/bvh.py b/mjx/mujoco/mjx/warp/bvh.py index d80b12fd..7f648694 100644 --- a/mjx/mujoco/mjx/warp/bvh.py +++ b/mjx/mujoco/mjx/warp/bvh.py @@ -110,7 +110,8 @@ def _refit_bvh_shim( def _refit_bvh_jax_impl( m: types.Model, d: types.Data, ctx: RenderContextPytree ): - output_dims = {'dummy': (d.qpos.shape[0],)} + render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[(ctx.key, None)] + output_dims = {'dummy': (render_ctx.nworld,)} jf = ffi.jax_callable_variadic_tuple( _refit_bvh_shim, num_outputs=1, @@ -123,7 +124,7 @@ def _refit_bvh_jax_impl( has_side_effect=True, ) out = jf( - d.qpos.shape[0], + render_ctx.nworld, m._impl.flex_dim, m._impl.flex_edge, m._impl.flex_elem, diff --git a/mjx/mujoco/mjx/warp/ffi.py b/mjx/mujoco/mjx/warp/ffi.py index a485cbbb..23dbbb9a 100644 --- a/mjx/mujoco/mjx/warp/ffi.py +++ b/mjx/mujoco/mjx/warp/ffi.py @@ -404,7 +404,10 @@ def marshal_custom_vmap( ) if tree_map_output: out = jax.tree.map( - lambda x: x.reshape(axis_size, -1), d_broadcast_flat_result + lambda x: x + if x.shape[0] == axis_size + else x.reshape(axis_size, -1, *x.shape[1:]), + d_broadcast_flat_result, ) return out, out_batched diff --git a/mjx/mujoco/mjx/warp/io.py b/mjx/mujoco/mjx/warp/io.py index b32d7122..6dbb1f83 100644 --- a/mjx/mujoco/mjx/warp/io.py +++ b/mjx/mujoco/mjx/warp/io.py @@ -27,8 +27,12 @@ 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, 2) # vec2i → trailing dim + ctx.seg_data_buffer = ctx.seg_data + ctx.nworld = nworld 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 57de37f9..719e8df9 100644 --- a/mjx/mujoco/mjx/warp/render.py +++ b/mjx/mujoco/mjx/warp/render.py @@ -48,6 +48,7 @@ _cb = mjwp_types.Callback( **{f.name: None for f in dataclasses.fields(mjwp_types.Callback) if f.init} ) + @ffi.format_args_for_warp def _render_shim( # Model @@ -84,6 +85,7 @@ def _render_shim( rc_id: int, rgb: wp.array2d[wp.uint32], depth: wp.array2d[wp.float32], + seg: wp.array2d[wp.vec2i], ): _m.stat = _s _m.opt = _o @@ -121,6 +123,7 @@ 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 = seg mjwarp.render(_m, _d, render_context) @@ -129,10 +132,11 @@ def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree): output_dims = { 'rgb': render_ctx.rgb_data_shape, 'depth': render_ctx.depth_data_shape, + 'seg': render_ctx.seg_data_shape, } jf = ffi.jax_callable_variadic_tuple( _render_shim, - num_outputs=2, + num_outputs=3, output_dims=output_dims, vmap_method=None, in_out_argnames=set([]), @@ -156,7 +160,7 @@ def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree): has_side_effect=False, ) out = jf( - d.qpos.shape[0], + render_ctx.nworld, m.cam_fovy, m.cam_intrinsic, m._impl.cam_projection, @@ -206,4 +210,4 @@ def render_vmap( ctx: RenderContextPytree, ): out = render(m, d, ctx) - return out, [True, True] + return out, [True, True, True] diff --git a/mjx/mujoco/mjx/warp/render_context.py b/mjx/mujoco/mjx/warp/render_context.py index 774e6030..e6d574a6 100644 --- a/mjx/mujoco/mjx/warp/render_context.py +++ b/mjx/mujoco/mjx/warp/render_context.py @@ -68,3 +68,13 @@ class RenderContextPytree(mjx_dataclasses.PyTreeNode): """ key: int + + +def get(rc: RenderContextPytree): + """Validates and returns the backing Warp render context.""" + if not isinstance(rc, RenderContextPytree): + raise TypeError( + f'Expected RenderContextPytree, got {type(rc).__name__}.' + ' Use rc.pytree() to get the JAX-compatible handle.' + ) + return _MJX_RENDER_CONTEXT_BUFFERS[(rc.key, None)] diff --git a/mjx/mujoco/mjx/warp/render_test.py b/mjx/mujoco/mjx/warp/render_test.py index d4864d9c..280e846d 100644 --- a/mjx/mujoco/mjx/warp/render_test.py +++ b/mjx/mujoco/mjx/warp/render_test.py @@ -30,11 +30,10 @@ 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 - _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 +62,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 +148,105 @@ 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[..., 0] != -1)) + self.assertGreater(np.unique(seg[..., 0]).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 + ) + # seg shape: (batch, total_pixels, 2); extract objid channel + expected_seg = seg[:, seg_adr : seg_adr + width * height, 0].reshape( + batch_size, height, width + ) + np.testing.assert_array_equal(unpacked_seg, expected_seg) + + def test_render_with_segmentation_raises_when_disabled(self): + """Tests render_with_segmentation rejects contexts without seg output.""" + self._maybe_skip() + mx, dx_batch, rc = _get_model_data_rc( + 'humanoid/humanoid.xml', 1, render_seg=False + ) + + dx_batch = jax.jit(mjx.refit_bvh)(mx, dx_batch, rc.pytree()) + with self.assertRaisesWithLiteralMatch( + ValueError, + 'Render context was not configured with segmentation rendering. ' + 'Pass render_seg=True or enable it for at least one camera in ' + 'create_render_context.', + ): + jax.jit(mjx.render_with_segmentation)(mx, dx_batch, rc.pytree()) + + @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[..., 0] != -1)) + class RenderContextGarbageCollectionTest(absltest.TestCase): """Tests that RenderContext cleans up buffers on deletion.""" diff --git a/mjx/mujoco/mjx/warp/visualize_render.py b/mjx/mujoco/mjx/warp/visualize_render.py index 7e470856..7ebd58b1 100644 --- a/mjx/mujoco/mjx/warp/visualize_render.py +++ b/mjx/mujoco/mjx/warp/visualize_render.py @@ -52,6 +52,9 @@ _RANDOMIZE_QPOS = flags.DEFINE_boolean( ) _USE_TEXTURES = flags.DEFINE_boolean('use_textures', True, 'enable textures') _USE_SHADOWS = flags.DEFINE_boolean('use_shadows', True, 'enable shadows') +_RENDER_SEGMENTATION = flags.DEFINE_boolean( + 'render_segmentation', False, 'enable segmentation rendering' +) _WP_KERNEL_CACHE_DIR = flags.DEFINE_string( 'wp_kernel_cache_dir', '/tmp/wp_kernel_cache_dir_visualize_render', @@ -91,6 +94,31 @@ def _save_tiled(rgb, out_path): print(f' tiled image: {out_path}') +def _colorize_segmentation(seg_ids: np.ndarray) -> np.ndarray: + """Map integer geom IDs to deterministic RGB colors. + + Background (-1) and flex (-2) pixels are mapped to black. + """ + seg = np.asarray(seg_ids) + h, w = seg.shape[-2], seg.shape[-1] + flat = seg.reshape(*seg.shape[:-2], -1) + + # Deterministic pastel palette via golden-ratio hue spacing. + r = np.zeros_like(flat, dtype=np.uint8) + g = np.zeros_like(flat, dtype=np.uint8) + b = np.zeros_like(flat, dtype=np.uint8) + + mask = flat >= 0 + ids = flat[mask] + # Simple hash-based colouring. + r[mask] = ((ids * 67 + 11) % 256).astype(np.uint8) + g[mask] = ((ids * 113 + 59) % 256).astype(np.uint8) + b[mask] = ((ids * 197 + 37) % 256).astype(np.uint8) + + rgb = np.stack([r, g, b], axis=-1) + return rgb.reshape(*seg.shape[:-2], h, w, 3) + + def _main(_: Sequence[str]): os.environ['MJX_WARP_ENABLED'] = 'true' @@ -110,6 +138,7 @@ def _main(_: Sequence[str]): print(f' camera_id : {_CAMERA_ID.value}') print(f' use_textures: {_USE_TEXTURES.value}') print(f' use_shadows : {_USE_SHADOWS.value}') + print(f' render_seg : {_RENDER_SEGMENTATION.value}') print(f' pmap : {_PMAP.value}') print(f' output_dir : {_OUTPUT_DIR.value}\n') @@ -144,6 +173,7 @@ def _main(_: Sequence[str]): use_shadows=_USE_SHADOWS.value, render_rgb=True, render_depth=True, + render_seg=_RENDER_SEGMENTATION.value, enabled_geom_groups=[0, 1, 2], ) @@ -151,14 +181,23 @@ def _main(_: Sequence[str]): mx, dx_batch, rc.pytree() ) - out_batch = jax_jit(jax.vmap(render.render, in_axes=(None, 0, None)))( + if _RENDER_SEGMENTATION.value: + render_fn = render.render_with_segmentation + else: + render_fn = render.render + + out_batch = jax_jit(jax.vmap(render_fn, in_axes=(None, 0, None)))( mx, dx_batch, rc.pytree() ) rgb_packed = out_batch[0] depth_packed = out_batch[1] + seg_packed = out_batch[2] if _RENDER_SEGMENTATION.value else None print(f' rgb shape: {rgb_packed.shape}') - print(f' depth shape: {depth_packed.shape}\n') + print(f' depth shape: {depth_packed.shape}') + if seg_packed is not None: + print(f' seg shape: {seg_packed.shape}') + print() rgb = jax.vmap(render_util.get_rgb, in_axes=(None, None, 0))( rc.pytree(), _CAMERA_ID.value, rgb_packed @@ -173,12 +212,26 @@ def _main(_: Sequence[str]): ) _save_single(rgb, single_path) - depth_rgb = np.repeat(np.asarray(depth)[..., None], 3, axis=-1) + depth_np = np.asarray(depth).squeeze(-1) # (nworld, H, W) + depth_rgb = np.repeat(depth_np[..., None], 3, axis=-1) depth_single_path = os.path.join( _OUTPUT_DIR.value, f'depth_{_CAMERA_ID.value}.png' ) _save_single(depth_rgb, depth_single_path) + if _RENDER_SEGMENTATION.value: + seg = jax.vmap(render_util.get_segmentation, in_axes=(None, None, 0))( + rc.pytree(), _CAMERA_ID.value, seg_packed + ) + seg_rgb = _colorize_segmentation(np.asarray(seg)) + # Convert to float [0, 1] so _save_single / _save_tiled work. + seg_rgb_f = seg_rgb.astype(np.float32) / 255.0 + + seg_single_path = os.path.join( + _OUTPUT_DIR.value, f'seg_{_CAMERA_ID.value}.png' + ) + _save_single(seg_rgb_f, seg_single_path) + if _NWORLD.value > 1: tiled_path = os.path.join( _OUTPUT_DIR.value, f'tiled_{_CAMERA_ID.value}.png' @@ -190,6 +243,12 @@ def _main(_: Sequence[str]): ) _save_tiled(depth_rgb, depth_tiled_path) + if _RENDER_SEGMENTATION.value: + seg_tiled_path = os.path.join( + _OUTPUT_DIR.value, f'seg_tiled_{_CAMERA_ID.value}.png' + ) + _save_tiled(seg_rgb_f, seg_tiled_path) + if _PMAP.value: ndevices = jax.local_device_count() nworld = _NWORLD.value