mjx segmentation

This commit is contained in:
Tarik Kelestemur
2026-04-15 18:22:46 -04:00
parent eacbce95f4
commit d39b2f4f9d
8 changed files with 414 additions and 66 deletions
+8 -5
View File
@@ -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
+3 -1
View File
@@ -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
+22 -2
View File
@@ -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.'
)
+62 -28
View File
@@ -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))
+85 -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,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()
+7 -3
View File
@@ -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
View File
@@ -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]
+87 -4
View File
@@ -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."""