mjx segmentation
This commit is contained in:
+8
-5
@@ -233,6 +233,7 @@ pytree that should be passed into ``jit``/``vmap``-compiled functions:
|
||||
use_shadows=True,
|
||||
render_rgb=[True] * ncam,
|
||||
render_depth=[False] * ncam,
|
||||
render_seg=[True] * ncam,
|
||||
enabled_geom_groups=[0, 1, 2],
|
||||
)
|
||||
|
||||
@@ -246,21 +247,23 @@ volume hierarchy (BVH) and executing the raycaster:
|
||||
.. code-block:: python
|
||||
|
||||
from mujoco.mjx import get_rgb
|
||||
from mujoco.mjx import get_segmentation
|
||||
|
||||
@jax.jit
|
||||
def render_fn(mx, d, rc_pytree):
|
||||
# 1. Update the BVH for the current scene state
|
||||
d = mjx.refit_bvh(mx, d, rc_pytree)
|
||||
|
||||
# 2. Render all configured cameras
|
||||
pixels, _ = mjx.render(mx, d, rc_pytree)
|
||||
# 2. Render all configured cameras, including segmentation
|
||||
pixels, _, segmentation = mjx.render_with_segmentation(mx, d, rc_pytree)
|
||||
|
||||
# 3. Extract the RGB tensor for the first camera (index 0)
|
||||
# 3. Extract the RGB tensor and geom IDs for the first camera (index 0)
|
||||
rgb = get_rgb(rc_pytree, 0, pixels)
|
||||
seg = get_segmentation(rc_pytree, 0, segmentation)
|
||||
|
||||
return rgb, d
|
||||
return rgb, seg, d
|
||||
|
||||
rgb, d = render_fn(mx, d, rc.pytree())
|
||||
rgb, seg, d = render_fn(mx, d, rc.pytree())
|
||||
|
||||
.. WARNING::
|
||||
The batch dimension ``nworld`` is fixed when the render context is created via
|
||||
|
||||
@@ -20,9 +20,9 @@ from mujoco.mjx._src.types import Model
|
||||
from mujoco.mjx._src.types import Data
|
||||
# isort: on
|
||||
|
||||
from mujoco.mjx._src.bvh import refit_bvh
|
||||
# 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.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,6 +15,7 @@
|
||||
"""Render helpers for MJX."""
|
||||
|
||||
from typing import Any
|
||||
|
||||
# pylint: disable=g-importing-member
|
||||
from mujoco.mjx._src.types import Data
|
||||
from mujoco.mjx._src.types import Impl
|
||||
@@ -26,8 +27,8 @@ import mujoco.mjx.warp as mjxw
|
||||
def render(m: Model, d: Data, ctx: Any) -> Data:
|
||||
"""Render."""
|
||||
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
|
||||
from mujoco.mjx.warp import render_context as mjxw_rc
|
||||
|
||||
if not isinstance(ctx, mjxw_rc.RenderContextPytree):
|
||||
raise TypeError(
|
||||
@@ -38,3 +39,22 @@ def render(m: Model, d: Data, ctx: Any) -> Data:
|
||||
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:
|
||||
"""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.'
|
||||
)
|
||||
|
||||
return mjxw_render.render_with_segmentation(m, d, ctx)
|
||||
|
||||
raise NotImplementedError(
|
||||
'render_with_segmentation only implemented for MuJoCo Warp.'
|
||||
)
|
||||
|
||||
@@ -25,6 +25,30 @@ 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,
|
||||
@@ -44,21 +68,9 @@ def get_rgb(
|
||||
Raises:
|
||||
RuntimeError: If Warp is not installed.
|
||||
"""
|
||||
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
|
||||
|
||||
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 = _get_warp_render_context(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
|
||||
@@ -92,21 +104,9 @@ def get_depth(
|
||||
Raises:
|
||||
RuntimeError: If Warp is not installed.
|
||||
"""
|
||||
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
|
||||
|
||||
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 = _get_warp_render_context(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 +114,37 @@ 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 raw geom IDs for a camera.
|
||||
|
||||
Args:
|
||||
rc: RenderContextPytree.
|
||||
cam_id: Camera index to extract.
|
||||
seg_data: Packed segmentation output, shape (..., total_pixels) as integers.
|
||||
|
||||
Returns:
|
||||
Integer segmentation array with shape (..., H, W).
|
||||
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)
|
||||
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)
|
||||
packed = jax.lax.dynamic_slice_in_dim(
|
||||
seg_data, seg_adr, width * height, axis=seg_data.ndim - 1
|
||||
)
|
||||
return packed.reshape(packed.shape[:-1] + (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,81 @@ 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)
|
||||
seg_data = jnp.arange(width * height, dtype=jnp.int32)
|
||||
|
||||
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):
|
||||
seg_data = jnp.arange(
|
||||
np.prod(leading_shape) * width * height, dtype=jnp.int32
|
||||
).reshape(leading_shape + (width * height,))
|
||||
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)
|
||||
seg_data = jnp.arange(nworld * width * height, dtype=jnp.int32).reshape(
|
||||
nworld, width * height
|
||||
)
|
||||
|
||||
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))
|
||||
|
||||
def test_get_segmentation_raises_for_disabled_camera(self):
|
||||
width, height = 4, 4
|
||||
warp_rc = _fake_render_context(2, width, height, render_seg=[True, False])
|
||||
rc = mock.MagicMock(spec=RenderContextPytree, key=0)
|
||||
seg_data = jnp.arange(width * height, dtype=jnp.int32)
|
||||
|
||||
with mock.patch.dict(
|
||||
'mujoco.mjx.warp.render_context._MJX_RENDER_CONTEXT_BUFFERS',
|
||||
{(0, None): warp_rc},
|
||||
):
|
||||
with self.assertRaisesWithLiteralMatch(
|
||||
ValueError,
|
||||
'Camera 1 was not configured with segmentation rendering.',
|
||||
):
|
||||
render_util.get_segmentation(rc, 1, seg_data)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
absltest.main()
|
||||
|
||||
@@ -14,11 +14,12 @@
|
||||
# ==============================================================================
|
||||
"""I/O functions for MJX Warp."""
|
||||
|
||||
import mujoco
|
||||
from mujoco.mjx.warp import render_context
|
||||
import mujoco.mjx.third_party.mujoco_warp as mjw
|
||||
import warp as wp
|
||||
|
||||
import mujoco
|
||||
import mujoco.mjx.third_party.mujoco_warp as mjw
|
||||
from mujoco.mjx.warp import render_context
|
||||
|
||||
_MJX_RENDER_CONTEXT_COUNTER = 0
|
||||
|
||||
|
||||
@@ -27,8 +28,11 @@ 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
|
||||
ctx.seg_data_buffer = ctx.seg_data
|
||||
ctx.rgb_data = None
|
||||
ctx.depth_data = None
|
||||
ctx.seg_data = None
|
||||
return ctx
|
||||
|
||||
|
||||
|
||||
+140
-22
@@ -14,17 +14,19 @@
|
||||
# ==============================================================================
|
||||
|
||||
"""DO NOT EDIT. This file is auto-generated."""
|
||||
|
||||
import dataclasses
|
||||
import functools
|
||||
|
||||
import jax
|
||||
import warp as wp
|
||||
|
||||
from mujoco.mjx._src import types
|
||||
import mujoco.mjx.third_party.mujoco_warp as mjwarp
|
||||
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
|
||||
from mujoco.mjx.warp import ffi
|
||||
from mujoco.mjx.warp.render_context import _MJX_RENDER_CONTEXT_BUFFERS
|
||||
from mujoco.mjx.warp.render_context import RenderContextPytree
|
||||
import mujoco.mjx.third_party.mujoco_warp as mjwarp
|
||||
from mujoco.mjx.third_party.mujoco_warp._src import types as mjwp_types
|
||||
import warp as wp
|
||||
|
||||
|
||||
_m = mjwarp.Model(
|
||||
**{f.name: None for f in dataclasses.fields(mjwarp.Model) if f.init}
|
||||
@@ -122,36 +124,129 @@ 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 = render_context.seg_data_buffer
|
||||
mjwarp.render(_m, _d, render_context)
|
||||
|
||||
|
||||
def _render_jax_impl(m: types.Model, d: types.Data, ctx: RenderContextPytree):
|
||||
@ffi.format_args_for_warp
|
||||
def _render_with_segmentation_shim(
|
||||
# Model
|
||||
nworld: int,
|
||||
cam_fovy: wp.array2d[float],
|
||||
cam_intrinsic: wp.array2d[wp.vec4],
|
||||
cam_projection: wp.array[int],
|
||||
cam_sensorsize: wp.array[wp.vec2],
|
||||
flex_edge: wp.array[wp.vec2i],
|
||||
flex_radius: wp.array[float],
|
||||
flex_vertadr: wp.array[int],
|
||||
geom_dataid: wp.array2d[int],
|
||||
geom_matid: wp.array2d[int],
|
||||
geom_rgba: wp.array2d[wp.vec4],
|
||||
geom_size: wp.array2d[wp.vec3],
|
||||
geom_type: wp.array[int],
|
||||
light_active: wp.array2d[bool],
|
||||
light_castshadow: wp.array2d[bool],
|
||||
light_type: wp.array2d[int],
|
||||
mat_rgba: wp.array2d[wp.vec4],
|
||||
mat_texid: wp.array3d[int],
|
||||
mat_texrepeat: wp.array2d[wp.vec2],
|
||||
mesh_faceadr: wp.array[int],
|
||||
nlight: int,
|
||||
# Data
|
||||
cam_xmat: wp.array2d[wp.mat33],
|
||||
cam_xpos: wp.array2d[wp.vec3],
|
||||
flexvert_xpos: wp.array2d[wp.vec3],
|
||||
geom_xmat: wp.array2d[wp.mat33],
|
||||
geom_xpos: wp.array2d[wp.vec3],
|
||||
light_xdir: wp.array2d[wp.vec3],
|
||||
light_xpos: wp.array2d[wp.vec3],
|
||||
# Registry
|
||||
rc_id: int,
|
||||
rgb: wp.array2d[wp.uint32],
|
||||
depth: wp.array2d[wp.float32],
|
||||
seg: wp.array2d[int],
|
||||
):
|
||||
_m.stat = _s
|
||||
_m.opt = _o
|
||||
_m.callback = _cb
|
||||
_d.efc = _e
|
||||
_d.contact = _c
|
||||
_m.cam_fovy = cam_fovy
|
||||
_m.cam_intrinsic = cam_intrinsic
|
||||
_m.cam_projection = cam_projection
|
||||
_m.cam_sensorsize = cam_sensorsize
|
||||
_m.flex_edge = flex_edge
|
||||
_m.flex_radius = flex_radius
|
||||
_m.flex_vertadr = flex_vertadr
|
||||
_m.geom_dataid = geom_dataid
|
||||
_m.geom_matid = geom_matid
|
||||
_m.geom_rgba = geom_rgba
|
||||
_m.geom_size = geom_size
|
||||
_m.geom_type = geom_type
|
||||
_m.light_active = light_active
|
||||
_m.light_castshadow = light_castshadow
|
||||
_m.light_type = light_type
|
||||
_m.mat_rgba = mat_rgba
|
||||
_m.mat_texid = mat_texid
|
||||
_m.mat_texrepeat = mat_texrepeat
|
||||
_m.mesh_faceadr = mesh_faceadr
|
||||
_m.nlight = nlight
|
||||
_d.cam_xmat = cam_xmat
|
||||
_d.cam_xpos = cam_xpos
|
||||
_d.flexvert_xpos = flexvert_xpos
|
||||
_d.geom_xmat = geom_xmat
|
||||
_d.geom_xpos = geom_xpos
|
||||
_d.light_xdir = light_xdir
|
||||
_d.light_xpos = light_xpos
|
||||
_d.nworld = nworld
|
||||
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)
|
||||
|
||||
|
||||
_RENDER_STAGE_IN_ARGNAMES = set([
|
||||
'cam_fovy',
|
||||
'cam_intrinsic',
|
||||
'cam_xmat',
|
||||
'cam_xpos',
|
||||
'geom_matid',
|
||||
'geom_rgba',
|
||||
'geom_size',
|
||||
'geom_xmat',
|
||||
'geom_xpos',
|
||||
'light_castshadow',
|
||||
'light_type',
|
||||
'mat_rgba',
|
||||
'mat_texid',
|
||||
])
|
||||
|
||||
|
||||
def _render_jax_impl(
|
||||
m: types.Model,
|
||||
d: types.Data,
|
||||
ctx: RenderContextPytree,
|
||||
with_segmentation: bool = False,
|
||||
):
|
||||
render_ctx = _MJX_RENDER_CONTEXT_BUFFERS[(ctx.key, None)]
|
||||
output_dims = {
|
||||
'rgb': render_ctx.rgb_data_shape,
|
||||
'depth': render_ctx.depth_data_shape,
|
||||
}
|
||||
render_shim = _render_shim
|
||||
num_outputs = 2
|
||||
if with_segmentation:
|
||||
output_dims['seg'] = render_ctx.seg_data_shape
|
||||
render_shim = _render_with_segmentation_shim
|
||||
num_outputs = 3
|
||||
jf = ffi.jax_callable_variadic_tuple(
|
||||
_render_shim,
|
||||
num_outputs=2,
|
||||
render_shim,
|
||||
num_outputs=num_outputs,
|
||||
output_dims=output_dims,
|
||||
vmap_method=None,
|
||||
in_out_argnames=set([]),
|
||||
stage_in_argnames=set([
|
||||
'cam_fovy',
|
||||
'cam_intrinsic',
|
||||
'cam_xmat',
|
||||
'cam_xpos',
|
||||
'geom_matid',
|
||||
'geom_rgba',
|
||||
'geom_size',
|
||||
'geom_xmat',
|
||||
'geom_xpos',
|
||||
'light_castshadow',
|
||||
'light_type',
|
||||
'mat_rgba',
|
||||
'mat_texid',
|
||||
]),
|
||||
stage_in_argnames=_RENDER_STAGE_IN_ARGNAMES,
|
||||
stage_out_argnames=set([]),
|
||||
graph_mode=m.opt._impl.graph_mode,
|
||||
has_side_effect=False,
|
||||
@@ -208,3 +303,26 @@ def render_vmap(
|
||||
):
|
||||
out = render(m, d, ctx)
|
||||
return out, [True, True]
|
||||
|
||||
|
||||
@jax.custom_batching.custom_vmap
|
||||
@functools.partial(ffi.marshal_jax_warp_callable, tree_map_output=True)
|
||||
def render_with_segmentation(
|
||||
m: types.Model,
|
||||
d: types.Data,
|
||||
ctx: RenderContextPytree,
|
||||
):
|
||||
return _render_jax_impl(m, d, ctx, with_segmentation=True)
|
||||
|
||||
|
||||
@render_with_segmentation.def_vmap
|
||||
@functools.partial(ffi.marshal_custom_vmap, tree_map_output=True)
|
||||
def render_with_segmentation_vmap(
|
||||
unused_axis_size,
|
||||
is_batched,
|
||||
m: types.Model,
|
||||
d: types.Data,
|
||||
ctx: RenderContextPytree,
|
||||
):
|
||||
out = render_with_segmentation(m, d, ctx)
|
||||
return out, [True, True, True]
|
||||
|
||||
@@ -19,22 +19,22 @@ from absl.testing import absltest
|
||||
from absl.testing import parameterized
|
||||
import jax
|
||||
from jax import numpy as jp
|
||||
import numpy as np
|
||||
|
||||
import mujoco
|
||||
from mujoco import mjx
|
||||
from mujoco.mjx._src import bvh
|
||||
from mujoco.mjx._src import forward
|
||||
from mujoco.mjx._src import io
|
||||
from mujoco.mjx._src import render
|
||||
import mujoco.mjx.warp as mjxw
|
||||
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
|
||||
|
||||
import mujoco.mjx.warp as mjxw
|
||||
|
||||
_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 +63,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 +149,88 @@ 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 != -1))
|
||||
self.assertGreater(np.unique(seg).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
|
||||
)
|
||||
expected_seg = seg[:, seg_adr : seg_adr + width * height].reshape(
|
||||
batch_size, height, width
|
||||
)
|
||||
np.testing.assert_array_equal(unpacked_seg, expected_seg)
|
||||
|
||||
@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 != -1))
|
||||
|
||||
|
||||
class RenderContextGarbageCollectionTest(absltest.TestCase):
|
||||
"""Tests that RenderContext cleans up buffers on deletion."""
|
||||
|
||||
Reference in New Issue
Block a user