Merge pull request #3235 from tkelestemur:mjx-warp-segmentation
PiperOrigin-RevId: 910401768 Change-Id: I015bd05d9d823db5efb92f660c707a67ddca6591
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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.'
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
@@ -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))
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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"""
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)]
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user