From f6129596a6701678d4f6ac0711d87476657a3c72 Mon Sep 17 00:00:00 2001 From: Baruch Tabanpour Date: Tue, 17 Feb 2026 12:54:44 -0800 Subject: [PATCH] Save depth images in visualize_render.py PiperOrigin-RevId: 871454618 Change-Id: If844e2c4af10e271536c4e6a36e760e8d4d50374 --- mjx/mujoco/mjx/warp/visualize_render.py | 19 ++++++++++++++++++- 1 file changed, 18 insertions(+), 1 deletion(-) diff --git a/mjx/mujoco/mjx/warp/visualize_render.py b/mjx/mujoco/mjx/warp/visualize_render.py index 8e6515d2..6c71f299 100644 --- a/mjx/mujoco/mjx/warp/visualize_render.py +++ b/mjx/mujoco/mjx/warp/visualize_render.py @@ -175,23 +175,40 @@ def _main(_: Sequence[str]): )(mx, dx_batch, rc) rgb_packed = out_batch[0] - print(f' rgb shape: {rgb_packed.shape}\n') + 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) + depth = jax.vmap( + render_util.get_depth, in_axes=(None, 0, None, None) + )(rc, depth_packed, _CAMERA_ID.value, 10.0) + single_path = os.path.join( _OUTPUT_DIR.value, f'camera_{_CAMERA_ID.value}.png' ) _save_single(rgb, single_path) + depth_rgb = np.repeat(np.asarray(depth)[..., 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 _NWORLD.value > 1: tiled_path = os.path.join( _OUTPUT_DIR.value, f'tiled_{_CAMERA_ID.value}.png' ) _save_tiled(rgb, tiled_path) + depth_tiled_path = os.path.join( + _OUTPUT_DIR.value, f'depth_tiled_{_CAMERA_ID.value}.png' + ) + _save_tiled(depth_rgb, depth_tiled_path) + print('\ndone.')