Clean up MJX segmentation rendering
This commit is contained in:
+4
-1
@@ -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
|
||||
|
||||
@@ -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.')
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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.'
|
||||
)
|
||||
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user