Merge pull request #3235 from tkelestemur:mjx-warp-segmentation

PiperOrigin-RevId: 910401768
Change-Id: I015bd05d9d823db5efb92f660c707a67ddca6591
This commit is contained in:
Copybara-Service
2026-05-04 20:42:32 -07:00
13 changed files with 518 additions and 51 deletions
+3 -1
View File
@@ -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
+46 -12
View File
@@ -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.'
)
+132
View File
@@ -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()
+56 -21
View File
@@ -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))
+75 -1
View File
@@ -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()
+15 -5
View File
@@ -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"""
+3 -2
View File
@@ -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,
+4 -1
View File
@@ -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
+4
View File
@@ -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
+7 -3
View File
@@ -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]
+10
View File
@@ -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)]
+101 -2
View File
@@ -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."""
+62 -3
View File
@@ -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