Keep render_util arg order consistent between mujoco_warp and mjx-warp.

PiperOrigin-RevId: 872055868
Change-Id: Ic5b4bdedbdb353dc0b1385cf555261a86407cfdb
This commit is contained in:
Baruch Tabanpour
2026-02-18 14:56:04 -08:00
committed by Copybara-Service
parent 57c2316b9e
commit 3c858f871b
4 changed files with 41 additions and 67 deletions
+4 -4
View File
@@ -23,15 +23,15 @@ import mujoco.mjx.warp as mjxw
def get_rgb(
rc: Any,
rgb_data: jax.Array,
cam_id: int,
rgb_data: jax.Array,
) -> jax.Array:
"""Unpack uint32 ABGR pixel data into float32 RGB.
Args:
rc: The RenderContext handle.
rgb_data: Packed render output, shape (total_pixels,) as uint32.
cam_id: Camera index to extract.
rgb_data: Packed render output, shape (total_pixels,) as uint32.
Returns:
Float32 RGB array with shape (H, W, 3), values in [0, 1].
@@ -62,16 +62,16 @@ def get_rgb(
def get_depth(
rc: Any,
depth_data: jax.Array,
cam_id: int,
depth_data: jax.Array,
depth_scale: float,
) -> jax.Array:
"""Extract and normalize depth data for a camera.
Args:
rc: The RenderContext handle.
depth_data: Raw depth output, shape (total_pixels,) as float32.
cam_id: Camera index to extract.
depth_data: Raw depth output, shape (total_pixels,) as float32.
depth_scale: Scale factor for normalizing depth values.
Returns:
+12 -18
View File
@@ -58,9 +58,7 @@ class RenderUtilTest(absltest.TestCase):
'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS',
{0: warp_rc},
):
rgb = jax.jit(render_util.get_rgb, static_argnums=(0, 2))(
rc, rgb_data, 0
)
rgb = jax.jit(render_util.get_rgb, static_argnums=(0, 1))(rc, 0, rgb_data)
self.assertEqual(rgb.shape, (height, width, 3))
@@ -68,18 +66,16 @@ class RenderUtilTest(absltest.TestCase):
nworld, width, height = 3, 4, 4
warp_rc = _fake_render_context(1, width, height)
rc = mock.MagicMock(key=0)
rgb_data = jnp.zeros(
(nworld, width * height), dtype=jnp.uint32
)
rgb_data = jnp.zeros((nworld, width * height), dtype=jnp.uint32)
with mock.patch.dict(
'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS',
{0: warp_rc},
):
rgb = jax.jit(
jax.vmap(render_util.get_rgb, in_axes=(None, 0, None)),
static_argnums=(0, 2),
)(rc, rgb_data, 0)
jax.vmap(render_util.get_rgb, in_axes=(None, None, 0)),
static_argnums=(0, 1),
)(rc, 0, rgb_data)
self.assertEqual(rgb.shape, (nworld, height, width, 3))
@@ -93,9 +89,9 @@ class RenderUtilTest(absltest.TestCase):
'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS',
{0: warp_rc},
):
depth = jax.jit(
render_util.get_depth, static_argnums=(0, 2, 3)
)(rc, depth_data, 0, 5.0)
depth = jax.jit(render_util.get_depth, static_argnums=(0, 1, 3))(
rc, 0, depth_data, 5.0
)
self.assertEqual(depth.shape, (height, width))
@@ -103,18 +99,16 @@ class RenderUtilTest(absltest.TestCase):
nworld, width, height = 3, 4, 4
warp_rc = _fake_render_context(1, width, height)
rc = mock.MagicMock(key=0)
depth_data = jnp.zeros(
(nworld, width * height), dtype=jnp.float32
)
depth_data = jnp.zeros((nworld, width * height), dtype=jnp.float32)
with mock.patch.dict(
'mujoco.mjx.warp.render._MJX_RENDER_CONTEXT_BUFFERS',
{0: warp_rc},
):
depth = jax.jit(
jax.vmap(render_util.get_depth, in_axes=(None, 0, None, None)),
static_argnums=(0, 2, 3),
)(rc, depth_data, 0, 5.0)
jax.vmap(render_util.get_depth, in_axes=(None, None, 0, None)),
static_argnums=(0, 1, 3),
)(rc, 0, depth_data, 5.0)
self.assertEqual(depth.shape, (nworld, height, width))
+4 -2
View File
@@ -155,7 +155,7 @@ def benchmark(
def render_fn(mx, d, rc):
d = mjx.refit_bvh(mx, d, rc)
pixels, _ = mjx.render(mx, d, rc)
return render_util.get_rgb(rc, pixels, 0), d
return render_util.get_rgb(rc, 0, pixels), d
@jax_jit
def unroll(d):
@@ -232,7 +232,9 @@ def benchmark_raw_warp(
if render:
mjwarp.forward(mw, dw)
rc = mjwarp.create_render_context(
m, mw, dw,
m,
mw,
dw,
(_RENDER_WIDTH.value, _RENDER_HEIGHT.value),
[_RENDER_RGB.value] * ncam,
[_RENDER_DEPTH.value] * ncam,
+21 -43
View File
@@ -40,26 +40,18 @@ _MODELFILE = flags.DEFINE_string(
'humanoid/humanoid.xml',
'path to model',
)
_NWORLD = flags.DEFINE_integer(
'nworld', 4, 'number of worlds to render'
)
_NWORLD = flags.DEFINE_integer('nworld', 4, 'number of worlds to render')
_WIDTH = flags.DEFINE_integer('width', 512, 'image width')
_HEIGHT = flags.DEFINE_integer('height', 512, 'image height')
_CAMERA_ID = flags.DEFINE_integer(
'camera_id', 0, 'camera id to visualize'
)
_CAMERA_ID = flags.DEFINE_integer('camera_id', 0, 'camera id to visualize')
_OUTPUT_DIR = flags.DEFINE_string(
'output_dir', '/tmp/visualize_render', 'output directory'
)
_RANDOMIZE_QPOS = flags.DEFINE_boolean(
'randomize_qpos', False, 'randomize initial qpos'
)
_USE_TEXTURES = flags.DEFINE_boolean(
'use_textures', True, 'enable textures'
)
_USE_SHADOWS = flags.DEFINE_boolean(
'use_shadows', True, 'enable shadows'
)
_USE_TEXTURES = flags.DEFINE_boolean('use_textures', True, 'enable textures')
_USE_SHADOWS = flags.DEFINE_boolean('use_shadows', True, 'enable shadows')
_WP_KERNEL_CACHE_DIR = flags.DEFINE_string(
'wp_kernel_cache_dir',
'/tmp/wp_kernel_cache_dir_visualize_render',
@@ -67,9 +59,7 @@ _WP_KERNEL_CACHE_DIR = flags.DEFINE_string(
)
_COMPILER_OPTIONS = {'xla_gpu_graph_min_graph_size': 1}
jax_jit = functools.partial(
jax.jit, compiler_options=_COMPILER_OPTIONS
)
jax_jit = functools.partial(jax.jit, compiler_options=_COMPILER_OPTIONS)
def _save_single(rgb, out_path):
@@ -85,14 +75,10 @@ def _save_tiled(rgb, out_path):
nworld, height, width, _ = rgb.shape
cols = int(np.ceil(np.sqrt(nworld)))
rows = int(np.ceil(nworld / cols))
canvas = np.zeros(
(rows * height, cols * width, 3), dtype=np.uint8
)
canvas = np.zeros((rows * height, cols * width, 3), dtype=np.uint8)
for w in range(nworld):
img_uint8 = (np.asarray(rgb[w]) * 255).astype(
np.uint8
)
img_uint8 = (np.asarray(rgb[w]) * 255).astype(np.uint8)
r, c = w // cols, w % cols
y0, y1 = r * height, (r + 1) * height
x0, x1 = c * width, (c + 1) * width
@@ -136,18 +122,14 @@ def _main(_: Sequence[str]):
if _RANDOMIZE_QPOS.value:
# TODO(robotics-team): consider integrating velocity if there are free
# joints.
qpos = qpos0 + jax.random.uniform(
rng, (m.nq,), minval=-0.2, maxval=0.05
)
qpos = qpos0 + jax.random.uniform(rng, (m.nq,), minval=-0.2, maxval=0.05)
return dx.replace(qpos=qpos)
print('initializing data...')
dx_batch = jax_jit(init)(worldids)
print('running forward...')
dx_batch = jax_jit(
jax.vmap(forward.forward, in_axes=(None, 0))
)(mx, dx_batch)
dx_batch = jax_jit(jax.vmap(forward.forward, in_axes=(None, 0)))(mx, dx_batch)
print('creating render context...')
rc = io.create_render_context(
@@ -162,30 +144,26 @@ def _main(_: Sequence[str]):
)
print('rendering...')
dx_batch = jax_jit(
jax.vmap(
bvh.refit_bvh, in_axes=(None, 0, None)
)
)(mx, dx_batch, rc)
dx_batch = jax_jit(jax.vmap(bvh.refit_bvh, in_axes=(None, 0, None)))(
mx, dx_batch, rc
)
out_batch = jax_jit(
jax.vmap(
render.render, in_axes=(None, 0, None)
)
)(mx, dx_batch, rc)
out_batch = jax_jit(jax.vmap(render.render, in_axes=(None, 0, None)))(
mx, dx_batch, rc
)
rgb_packed = out_batch[0]
depth_packed = out_batch[1]
print(f' rgb shape: {rgb_packed.shape}')
print(f' depth shape: {depth_packed.shape}\n')
rgb = jax.vmap(
render_util.get_rgb, in_axes=(None, 0, None)
)(rc, rgb_packed, _CAMERA_ID.value)
rgb = jax.vmap(render_util.get_rgb, in_axes=(None, None, 0))(
rc, _CAMERA_ID.value, rgb_packed
)
depth = jax.vmap(
render_util.get_depth, in_axes=(None, 0, None, None)
)(rc, depth_packed, _CAMERA_ID.value, 10.0)
depth = jax.vmap(render_util.get_depth, in_axes=(None, None, 0, None))(
rc, _CAMERA_ID.value, depth_packed, 10.0
)
single_path = os.path.join(
_OUTPUT_DIR.value, f'camera_{_CAMERA_ID.value}.png'