Keep render_util arg order consistent between mujoco_warp and mjx-warp.
PiperOrigin-RevId: 872055868 Change-Id: Ic5b4bdedbdb353dc0b1385cf555261a86407cfdb
This commit is contained in:
committed by
Copybara-Service
parent
57c2316b9e
commit
3c858f871b
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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'
|
||||
|
||||
Reference in New Issue
Block a user