diff --git a/mjx/mujoco/mjx/_src/render_util.py b/mjx/mujoco/mjx/_src/render_util.py index e988243e..58e98ef4 100644 --- a/mjx/mujoco/mjx/_src/render_util.py +++ b/mjx/mujoco/mjx/_src/render_util.py @@ -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: diff --git a/mjx/mujoco/mjx/_src/render_util_test.py b/mjx/mujoco/mjx/_src/render_util_test.py index 56160d86..e62658ac 100644 --- a/mjx/mujoco/mjx/_src/render_util_test.py +++ b/mjx/mujoco/mjx/_src/render_util_test.py @@ -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)) diff --git a/mjx/mujoco/mjx/warp/testspeed.py b/mjx/mujoco/mjx/warp/testspeed.py index 48b2179d..e2525a8d 100644 --- a/mjx/mujoco/mjx/warp/testspeed.py +++ b/mjx/mujoco/mjx/warp/testspeed.py @@ -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, diff --git a/mjx/mujoco/mjx/warp/visualize_render.py b/mjx/mujoco/mjx/warp/visualize_render.py index 6c71f299..42ed02f1 100644 --- a/mjx/mujoco/mjx/warp/visualize_render.py +++ b/mjx/mujoco/mjx/warp/visualize_render.py @@ -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'