diff --git a/doc/mjx.rst b/doc/mjx.rst index 247dafa8..42ad03e8 100644 --- a/doc/mjx.rst +++ b/doc/mjx.rst @@ -257,7 +257,7 @@ volume hierarchy (BVH) and executing the raycaster: # 2. Render all configured cameras, including segmentation pixels, _, segmentation = mjx.render_with_segmentation(mx, d, rc_pytree) - # 3. Extract the RGB tensor and geom IDs for the first camera (index 0) + # 3. Extract the RGB tensor and segmentation IDs for the first camera rgb = get_rgb(rc_pytree, 0, pixels) seg = get_segmentation(rc_pytree, 0, segmentation) @@ -265,6 +265,9 @@ volume hierarchy (BVH) and executing the raycaster: rgb, seg, d = render_fn(mx, d, rc.pytree()) +The segmentation image contains MuJoCo geom IDs per pixel, ``-1`` for +background, and ``-2`` for flex bodies. + .. WARNING:: The batch dimension ``nworld`` is fixed when the render context is created via :func:`~mujoco.mjx.create_render_context` since the underlying Warp render context allocates diff --git a/mjx/mujoco/mjx/_src/bvh.py b/mjx/mujoco/mjx/_src/bvh.py index 1a75a8a0..87e75ae9 100644 --- a/mjx/mujoco/mjx/_src/bvh.py +++ b/mjx/mujoco/mjx/_src/bvh.py @@ -15,26 +15,23 @@ """BVH helpers for MJX.""" from typing import Any + +import mujoco.mjx.warp as mjxw + +from mujoco.mjx._src.warp_context import get_warp_render_context # 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 refit_bvh(m: Model, d: Data, ctx: Any): """Refit the scene BVH for the current pose.""" 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 bvh as mjxw_bvh # pylint: disable=g-import-not-at-top # 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.' - ) - + get_warp_render_context(ctx) return mjxw_bvh.refit_bvh(m, d, ctx) raise NotImplementedError('refit_bvh only implemented for MuJoCo Warp.') diff --git a/mjx/mujoco/mjx/_src/render.py b/mjx/mujoco/mjx/_src/render.py index d6126529..56666dc9 100644 --- a/mjx/mujoco/mjx/_src/render.py +++ b/mjx/mujoco/mjx/_src/render.py @@ -16,42 +16,38 @@ from typing import Any +import jax +import mujoco.mjx.warp as mjxw + +from mujoco.mjx._src.warp_context import get_warp_render_context +from mujoco.mjx._src.warp_context import require_segmentation_enabled # 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 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: 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.' - ) + get_warp_render_context(ctx) 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: +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.""" 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.' - ) + warp_rc = get_warp_render_context(ctx) + require_segmentation_enabled(warp_rc) return mjxw_render.render_with_segmentation(m, d, ctx) diff --git a/mjx/mujoco/mjx/_src/render_util.py b/mjx/mujoco/mjx/_src/render_util.py index 56c7e6ad..78f537b1 100644 --- a/mjx/mujoco/mjx/_src/render_util.py +++ b/mjx/mujoco/mjx/_src/render_util.py @@ -19,36 +19,13 @@ from typing import TYPE_CHECKING import jax import jax.numpy as jnp -import mujoco.mjx.warp as mjxw +from mujoco.mjx._src.warp_context import get_camera_resolution +from mujoco.mjx._src.warp_context import get_warp_render_context 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, @@ -68,9 +45,9 @@ def get_rgb( Raises: RuntimeError: If Warp is not installed. """ - warp_rc = _get_warp_render_context(rc) + warp_rc = get_warp_render_context(rc) rgb_adr = int(warp_rc.rgb_adr.numpy()[cam_id]) - width, height = _get_camera_resolution(warp_rc, cam_id) + 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 @@ -104,9 +81,9 @@ def get_depth( Raises: RuntimeError: If Warp is not installed. """ - warp_rc = _get_warp_render_context(rc) + warp_rc = get_warp_render_context(rc) depth_adr = int(warp_rc.depth_adr.numpy()[cam_id]) - width, height = _get_camera_resolution(warp_rc, cam_id) + 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 @@ -121,7 +98,7 @@ def get_segmentation( cam_id: int, seg_data: jax.Array, ) -> jax.Array: - """Extract raw geom IDs for a camera. + """Extract raw segmentation IDs for a camera. Args: rc: RenderContextPytree. @@ -129,21 +106,23 @@ def get_segmentation( seg_data: Packed segmentation output, shape (..., total_pixels) as integers. Returns: - Integer segmentation array with shape (..., H, W). + Integer segmentation array with shape (..., H, W). Each pixel contains the + MuJoCo geom ID of the hit geometry, ``-1`` for background, or ``-2`` for a + flex body. 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) + 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) + 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 ) diff --git a/mjx/mujoco/mjx/_src/warp_context.py b/mjx/mujoco/mjx/_src/warp_context.py new file mode 100644 index 00000000..3479ccfd --- /dev/null +++ b/mjx/mujoco/mjx/_src/warp_context.py @@ -0,0 +1,56 @@ +# 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. +# ============================================================================== +"""Shared helpers for MJX-Warp render contexts.""" + +from typing import TYPE_CHECKING + +import mujoco.mjx.warp as mjxw + +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 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.' + ) diff --git a/mjx/mujoco/mjx/warp/render_test.py b/mjx/mujoco/mjx/warp/render_test.py index fc089617..6745c2e6 100644 --- a/mjx/mujoco/mjx/warp/render_test.py +++ b/mjx/mujoco/mjx/warp/render_test.py @@ -185,6 +185,22 @@ class RenderTest(parameterized.TestCase): ) 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),