Clean up MJX segmentation rendering

This commit is contained in:
Tarik Kelestemur
2026-04-19 13:12:52 -04:00
parent fa8b6311a2
commit 8262280f5f
6 changed files with 106 additions and 59 deletions
+4 -1
View File
@@ -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
+5 -8
View File
@@ -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.')
+13 -17
View File
@@ -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)
+12 -33
View File
@@ -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
)
+56
View File
@@ -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.'
)
+16
View File
@@ -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),